perf: 优化同步方法性能,避免事件循环开销

- pipeline.py: 日志格式改为 %-style 懒惰计算
- cost.py: get_model_price/get_tiered_pricing 改用直接数据库查询
- usage/service.py: 批量记录改为并行准备,提取 build_params 减少重复代码
This commit is contained in:
fawney19
2026-01-29 02:00:00 +08:00
parent b4bbf4eeb1
commit 06a46383a1
3 changed files with 431 additions and 183 deletions

View File

@@ -206,16 +206,56 @@ class ModelCostService:
return result
def get_tiered_pricing(self, provider: ProviderRef, model: str) -> Optional[dict]:
"""同步获取模型的阶梯计费配置"""
import asyncio
"""同步获取模型的阶梯计费配置(直接查缓存和数据库,避免事件循环开销)。
try:
loop = asyncio.get_event_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
Returns:
阶梯计费配置字典,如果未找到配置则返回 None。
注意:与 get_tiered_pricing_async 不同,此方法仅返回 pricing 部分,不含 source 字段。
"""
provider_name = self._provider_name(provider)
cache_key = f"{provider_name}:{model}:tiered_with_source"
return loop.run_until_complete(self.get_tiered_pricing_async(provider, model))
# 优先从缓存获取
if cache_key in self._tiered_pricing_cache:
result = self._tiered_pricing_cache[cache_key]
return result.get("pricing") if result else None
# 缓存未命中,同步查询数据库
provider_obj = self._resolve_provider(provider)
result = None
if provider_obj:
global_model = (
self.db.query(GlobalModel)
.filter(
GlobalModel.name == model,
GlobalModel.is_active == True,
)
.first()
)
if global_model:
model_obj = (
self.db.query(Model)
.filter(
Model.provider_id == provider_obj.id,
Model.global_model_id == global_model.id,
Model.is_active == True,
)
.first()
)
if model_obj:
if model_obj.tiered_pricing is not None:
result = {"pricing": model_obj.tiered_pricing, "source": "provider"}
elif global_model.default_tiered_pricing is not None:
result = {"pricing": global_model.default_tiered_pricing, "source": "global"}
else:
if global_model.default_tiered_pricing is not None:
result = {"pricing": global_model.default_tiered_pricing, "source": "global"}
self._tiered_pricing_cache[cache_key] = result
return result.get("pricing") if result else None
# ------------------------------------------------------------------
# 公共方法
@@ -315,11 +355,7 @@ class ModelCostService:
def get_model_price(self, provider: ProviderRef, model: str) -> Tuple[float, float]:
"""
返回给定 provider/model 的 (input_price, output_price)。
逻辑:
1. 直接通过 GlobalModel.name 匹配
2. 查找该 Provider 的 Model 实现
3. 获取价格配置
直接查缓存和数据库,避免事件循环开销。
Args:
provider: Provider 对象或提供商名称
@@ -328,16 +364,74 @@ class ModelCostService:
Returns:
(input_price, output_price) 元组
"""
import asyncio
provider_name = self._provider_name(provider)
cache_key = f"{provider_name}:{model}"
# 在同步上下文中调用异步方法
try:
loop = asyncio.get_event_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
# 优先从缓存获取
if cache_key in self._price_cache:
prices = self._price_cache[cache_key]
return prices["input"], prices["output"]
return loop.run_until_complete(self.get_model_price_async(provider, model))
# 缓存未命中,同步查询数据库
provider_obj = self._resolve_provider(provider)
input_price = None
output_price = None
if provider_obj:
global_model = (
self.db.query(GlobalModel)
.filter(
GlobalModel.name == model,
GlobalModel.is_active == True,
)
.first()
)
if global_model:
model_obj = (
self.db.query(Model)
.filter(
Model.provider_id == provider_obj.id,
Model.global_model_id == global_model.id,
Model.is_active == True,
)
.first()
)
if model_obj:
tiered = model_obj.get_effective_tiered_pricing()
if tiered and tiered.get("tiers"):
first_tier = tiered["tiers"][0]
input_price = first_tier.get("input_price_per_1m", 0)
output_price = first_tier.get("output_price_per_1m", 0)
else:
input_price = model_obj.get_effective_input_price()
output_price = model_obj.get_effective_output_price()
else:
tiered = global_model.default_tiered_pricing
if tiered and tiered.get("tiers"):
first_tier = tiered["tiers"][0]
input_price = first_tier.get("input_price_per_1m", 0)
output_price = first_tier.get("output_price_per_1m", 0)
else:
input_price = 0.0
output_price = 0.0
if input_price is None:
input_price = 0.0
if output_price is None:
output_price = 0.0
# 与异步版本保持一致:当 token 价格为 0 时,仅在没有按次计费配置时告警
if input_price == 0.0 and output_price == 0.0:
price_per_request = self.get_request_price(provider, model)
if price_per_request is None or price_per_request == 0.0:
logger.warning(
"未找到模型价格配置: %s/%s,请在 GlobalModel 中配置价格",
provider_name,
model,
)
self._price_cache[cache_key] = {"input": input_price, "output": output_price}
return input_price, output_price
async def get_cache_prices_async(
self, provider: ProviderRef, model: str, input_price: float
@@ -472,54 +566,122 @@ class ModelCostService:
def get_request_price(self, provider: ProviderRef, model: str) -> Optional[float]:
"""
返回按次计费价格(每次请求的固定费用)。
直接查数据库,避免事件循环开销。
Args:
provider: Provider 对象或提供商名称
model: 用户请求的模型名(可能是 GlobalModel.name 或映射名)
model: 用户请求的模型名
Returns:
按次计费价格,如果没有配置则返回 None
"""
import asyncio
provider_obj = self._resolve_provider(provider)
price_per_request = None
# 在同步上下文中调用异步方法
try:
loop = asyncio.get_event_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
if provider_obj:
global_model = (
self.db.query(GlobalModel)
.filter(
GlobalModel.name == model,
GlobalModel.is_active == True,
)
.first()
)
return loop.run_until_complete(self.get_request_price_async(provider, model))
if global_model:
model_obj = (
self.db.query(Model)
.filter(
Model.provider_id == provider_obj.id,
Model.global_model_id == global_model.id,
Model.is_active == True,
)
.first()
)
if model_obj:
price_per_request = model_obj.get_effective_price_per_request()
else:
price_per_request = global_model.default_price_per_request
return price_per_request
def get_cache_prices(
self, provider: ProviderRef, model: str, input_price: float
) -> Tuple[Optional[float], Optional[float]]:
"""
返回缓存创建/读取价格(每 1M tokens
逻辑:
1. 直接通过 GlobalModel.name 匹配
2. 查找该 Provider 的 Model 实现
3. 获取缓存价格配置
直接查缓存和数据库,避免事件循环开销。
Args:
provider: Provider 对象或提供商名称
model: 用户请求的模型名(必须是 GlobalModel.name
input_price: 基础输入价格(用于 Claude 模型的默认估算)
input_price: 基础输入价格(用于默认估算)
Returns:
(cache_creation_price, cache_read_price) 元组
"""
import asyncio
provider_name = self._provider_name(provider)
cache_key = f"{provider_name}:{model}"
# 在同步上下文中调用异步方法
try:
loop = asyncio.get_event_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
# 优先从缓存获取
if cache_key in self._cache_price_cache:
prices = self._cache_price_cache[cache_key]
return prices["creation"], prices["read"]
return loop.run_until_complete(self.get_cache_prices_async(provider, model, input_price))
# 缓存未命中,同步查询数据库
provider_obj = self._resolve_provider(provider)
cache_creation_price = None
cache_read_price = None
if provider_obj:
global_model = (
self.db.query(GlobalModel)
.filter(
GlobalModel.name == model,
GlobalModel.is_active == True,
)
.first()
)
if global_model:
model_obj = (
self.db.query(Model)
.filter(
Model.provider_id == provider_obj.id,
Model.global_model_id == global_model.id,
Model.is_active == True,
)
.first()
)
if model_obj:
tiered = model_obj.get_effective_tiered_pricing()
if tiered and tiered.get("tiers"):
first_tier = tiered["tiers"][0]
cache_creation_price = first_tier.get("cache_creation_price_per_1m")
cache_read_price = first_tier.get("cache_read_price_per_1m")
else:
cache_creation_price = model_obj.get_effective_cache_creation_price()
cache_read_price = model_obj.get_effective_cache_read_price()
else:
tiered = global_model.default_tiered_pricing
if tiered and tiered.get("tiers"):
first_tier = tiered["tiers"][0]
cache_creation_price = first_tier.get("cache_creation_price_per_1m")
cache_read_price = first_tier.get("cache_read_price_per_1m")
# 默认缓存价格估算
if cache_creation_price is None:
cache_creation_price = input_price * 1.25
if cache_read_price is None:
cache_read_price = input_price * 0.1
self._cache_price_cache[cache_key] = {
"creation": cache_creation_price,
"read": cache_read_price,
}
return cache_creation_price, cache_read_price
def calculate_cost(
self,

View File

@@ -422,12 +422,29 @@ class UsageService:
input_cost, output_cost, cache_creation_cost, cache_read_cost, cache_cost,
request_cost, total_cost, tier_index)
"""
# 获取模型价格信息
input_price, output_price = await cls.get_model_price_async(db, provider, model)
cache_creation_price, cache_read_price = await cls.get_cache_prices_async(
db, provider, model, input_price
import asyncio
service = ModelCostService(db)
# 并行获取模型价格、按次计费价格;阶梯计费时额外获取 tiered 配置
price_task = service.get_model_price_async(provider, model)
request_price_task = service.get_request_price_async(provider, model)
tiered_pricing: Optional[dict] = None
if use_tiered_pricing:
tiered_pricing_task = service.get_tiered_pricing_async(provider, model)
(input_price, output_price), request_price, tiered_pricing = await asyncio.gather(
price_task, request_price_task, tiered_pricing_task
)
else:
(input_price, output_price), request_price = await asyncio.gather(
price_task, request_price_task
)
# 缓存价格依赖 input_price需要串行获取
cache_creation_price, cache_read_price = await service.get_cache_prices_async(
provider, model, input_price
)
request_price = await cls.get_request_price_async(db, provider, model)
effective_request_price = None if is_failed_request else request_price
# 初始化成本变量
@@ -441,29 +458,62 @@ class UsageService:
tier_index = None
if use_tiered_pricing:
(
input_cost,
output_cost,
cache_creation_cost,
cache_read_cost,
cache_cost,
request_cost,
total_cost,
tier_index,
) = await cls.calculate_cost_with_strategy_async(
db=db,
provider=provider,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
api_format=api_format,
cache_ttl_minutes=cache_ttl_minutes,
)
if is_failed_request:
total_cost = total_cost - request_cost
request_cost = 0.0
# 使用与 ModelCostService.compute_cost_with_strategy_async 一致的 adapter 逻辑,
# 但复用本方法已获取的价格/配置,避免重复 I/O。
adapter = None
if api_format:
from src.api.handlers.base.chat_adapter_base import get_adapter_instance
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_instance
adapter = get_adapter_instance(api_format)
if adapter is None:
adapter = get_cli_adapter_instance(api_format)
if adapter:
result = adapter.compute_cost(
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
input_price_per_1m=input_price,
output_price_per_1m=output_price,
cache_creation_price_per_1m=cache_creation_price,
cache_read_price_per_1m=cache_read_price,
price_per_request=effective_request_price,
tiered_pricing=tiered_pricing,
cache_ttl_minutes=cache_ttl_minutes,
)
input_cost = result["input_cost"]
output_cost = result["output_cost"]
cache_creation_cost = result["cache_creation_cost"]
cache_read_cost = result["cache_read_cost"]
cache_cost = result["cache_cost"]
request_cost = result["request_cost"]
total_cost = result["total_cost"]
tier_index = result.get("tier_index")
else:
(
input_cost,
output_cost,
cache_creation_cost,
cache_read_cost,
cache_cost,
request_cost,
total_cost,
tier_index,
) = ModelCostService.compute_cost_with_tiered_pricing(
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
tiered_pricing=tiered_pricing,
cache_ttl_minutes=cache_ttl_minutes,
price_per_request=effective_request_price,
fallback_input_price_per_1m=input_price,
fallback_output_price_per_1m=output_price,
fallback_cache_creation_price_per_1m=cache_creation_price,
fallback_cache_read_price_per_1m=cache_read_price,
)
else:
(
input_cost,
@@ -764,6 +814,44 @@ class UsageService:
return usage_params, total_cost
@classmethod
async def _prepare_usage_records_batch(
cls,
params_list: List[UsageRecordParams],
) -> List[Tuple[Dict[str, Any], float, Optional[Exception]]]:
"""批量并行准备用量记录(性能优化)
并行调用 _prepare_usage_record提高批量处理效率。
Args:
params_list: 用量记录参数列表
Returns:
列表,每项为 (usage_params, total_cost, exception)
如果处理成功exception 为 None
"""
import asyncio
async def prepare_single(params: UsageRecordParams) -> Tuple[Dict[str, Any], float, Optional[Exception]]:
try:
usage_params, total_cost = await cls._prepare_usage_record(params)
return (usage_params, total_cost, None)
except Exception as e:
return ({}, 0.0, e)
if not params_list:
return []
# 避免一次性创建过多 task并且 _prepare_usage_record 内部也可能包含并行调用)
# 这里采用分批 gather 来限制并发量。
chunk_size = 50
results: List[Tuple[Dict[str, Any], float, Optional[Exception]]] = []
for i in range(0, len(params_list), chunk_size):
chunk = params_list[i : i + chunk_size]
chunk_results = await asyncio.gather(*(prepare_single(p) for p in chunk))
results.extend(chunk_results)
return results
@classmethod
async def record_usage_async(
cls,
@@ -1115,66 +1203,102 @@ class UsageService:
updated_count = 0
total_count = len(all_records)
# 1. 处理需要更新的记录pending/streaming -> completed/failed/cancelled
# 辅助函数:构建 UsageRecordParams
def build_params(record: Dict[str, Any], request_id: str) -> UsageRecordParams:
user_id = record.get("user_id")
api_key_id = record.get("api_key_id")
user = users_map.get(str(user_id)) if user_id else None
api_key = api_keys_map.get(str(api_key_id)) if api_key_id else None
return UsageRecordParams(
db=db,
user=user,
api_key=api_key,
provider=record.get("provider") or "unknown",
model=record.get("model") or "unknown",
input_tokens=int(record.get("input_tokens") or 0),
output_tokens=int(record.get("output_tokens") or 0),
cache_creation_input_tokens=int(record.get("cache_creation_input_tokens") or 0),
cache_read_input_tokens=int(record.get("cache_read_input_tokens") or 0),
request_type=record.get("request_type") or "chat",
api_format=record.get("api_format"),
endpoint_api_format=record.get("endpoint_api_format"),
has_format_conversion=bool(record.get("has_format_conversion")),
is_stream=bool(record.get("is_stream", True)),
response_time_ms=record.get("response_time_ms"),
first_byte_time_ms=record.get("first_byte_time_ms"),
status_code=int(record.get("status_code") or 200),
error_message=record.get("error_message"),
metadata=record.get("metadata"),
request_headers=record.get("request_headers"),
request_body=record.get("request_body"),
provider_request_headers=record.get("provider_request_headers"),
response_headers=record.get("response_headers"),
client_response_headers=record.get("client_response_headers"),
response_body=record.get("response_body"),
request_id=request_id,
provider_id=record.get("provider_id"),
provider_endpoint_id=record.get("provider_endpoint_id"),
provider_api_key_id=record.get("provider_api_key_id"),
status=record.get("status") or "completed",
cache_ttl_minutes=record.get("cache_ttl_minutes"),
use_tiered_pricing=record.get("use_tiered_pricing", True),
target_model=record.get("target_model"),
)
# 构建所有参数并并行准备
update_params_list: List[Tuple[Dict[str, Any], str, UsageRecordParams]] = []
for record in records_to_update:
request_id = record.get("request_id")
if request_id and request_id in existing_usages:
existing_usage = existing_usages[request_id]
if existing_usage:
try:
params = build_params(record, request_id)
update_params_list.append((record, request_id, params))
except Exception as e:
skipped_count += 1
logger.warning("批量记录中参数构建失败: %s, request_id=%s", e, request_id)
insert_params_list: List[Tuple[Dict[str, Any], str, UsageRecordParams]] = []
for record in records_to_insert:
request_id = record.get("request_id") or str(uuid.uuid4())[:8]
try:
request_id = record.get("request_id")
existing_usage = existing_usages.get(request_id)
if not existing_usage:
skipped_count += 1
continue
params = build_params(record, request_id)
insert_params_list.append((record, request_id, params))
except Exception as e:
skipped_count += 1
logger.warning("批量记录中参数构建失败: %s, request_id=%s", e, request_id)
# 从预取的 map 中获取 user 和 api_key 对象
user_id = record.get("user_id")
api_key_id = record.get("api_key_id")
user = users_map.get(str(user_id)) if user_id else None
api_key = api_keys_map.get(str(api_key_id)) if api_key_id else None
# 并行准备所有记录(性能优化)
all_params = [p for _, _, p in update_params_list] + [p for _, _, p in insert_params_list]
if all_params:
prepared_results = await cls._prepare_usage_records_batch(all_params)
else:
prepared_results = []
# 准备记录参数
params = UsageRecordParams(
db=db,
user=user,
api_key=api_key,
provider=record.get("provider") or "unknown",
model=record.get("model") or "unknown",
input_tokens=int(record.get("input_tokens") or 0),
output_tokens=int(record.get("output_tokens") or 0),
cache_creation_input_tokens=int(record.get("cache_creation_input_tokens") or 0),
cache_read_input_tokens=int(record.get("cache_read_input_tokens") or 0),
request_type=record.get("request_type") or "chat",
api_format=record.get("api_format"),
endpoint_api_format=record.get("endpoint_api_format"),
has_format_conversion=bool(record.get("has_format_conversion")),
is_stream=bool(record.get("is_stream", True)),
response_time_ms=record.get("response_time_ms"),
first_byte_time_ms=record.get("first_byte_time_ms"),
status_code=int(record.get("status_code") or 200),
error_message=record.get("error_message"),
metadata=record.get("metadata"),
request_headers=record.get("request_headers"),
request_body=record.get("request_body"),
provider_request_headers=record.get("provider_request_headers"),
response_headers=record.get("response_headers"),
client_response_headers=record.get("client_response_headers"),
response_body=record.get("response_body"),
request_id=request_id,
provider_id=record.get("provider_id"),
provider_endpoint_id=record.get("provider_endpoint_id"),
provider_api_key_id=record.get("provider_api_key_id"),
status=record.get("status") or "completed",
cache_ttl_minutes=record.get("cache_ttl_minutes"),
use_tiered_pricing=record.get("use_tiered_pricing", True),
target_model=record.get("target_model"),
)
# 分配准备结果
update_results = prepared_results[:len(update_params_list)]
insert_results = prepared_results[len(update_params_list):]
usage_params, total_cost = await cls._prepare_usage_record(params)
# 1. 处理需要更新的记录
for i, (record, request_id, params) in enumerate(update_params_list):
try:
usage_params, total_cost, exc = update_results[i]
if exc:
raise exc
# existing_usage 已在构建阶段验证存在
existing_usage = existing_usages[request_id]
user = params.user
api_key = params.api_key
# 更新已存在的 Usage 记录
cls._update_existing_usage(existing_usage, usage_params, record.get("target_model"))
usages.append(existing_usage)
updated_count += 1
# 聚合统计(更新记录也需要更新统计,因为 pending 状态没有计费)
# 聚合统计
model_name = record.get("model") or "unknown"
model_counts[model_name] += 1
@@ -1183,11 +1307,9 @@ class UsageService:
actual_cost = usage_params.get("actual_total_cost_usd", 0)
provider_costs[provider_id] += actual_cost
# 用户统计(独立 Key 不计入创建者)
if user and not (api_key and api_key.is_standalone):
user_costs[str(user.id)] += total_cost
# API Key 统计
if api_key:
key_id = str(api_key.id)
apikey_stats[key_id]["requests"] += 1
@@ -1196,57 +1318,18 @@ class UsageService:
except Exception as e:
skipped_count += 1
logger.warning(f"批量记录中更新失败: {e}, request_id={record.get('request_id')}")
logger.warning("批量记录中更新失败: %s, request_id=%s", e, request_id)
continue
# 2. 处理需要新建的记录
for record in records_to_insert:
for i, (record, request_id, params) in enumerate(insert_params_list):
try:
# 从预取的 map 中获取 user 和 api_key 对象
user_id = record.get("user_id")
api_key_id = record.get("api_key_id")
user = users_map.get(str(user_id)) if user_id else None
api_key = api_keys_map.get(str(api_key_id)) if api_key_id else None
usage_params, total_cost, exc = insert_results[i]
if exc:
raise exc
# 准备记录参数
request_id = record.get("request_id") or str(uuid.uuid4())[:8]
params = UsageRecordParams(
db=db,
user=user,
api_key=api_key,
provider=record.get("provider") or "unknown",
model=record.get("model") or "unknown",
input_tokens=int(record.get("input_tokens") or 0),
output_tokens=int(record.get("output_tokens") or 0),
cache_creation_input_tokens=int(record.get("cache_creation_input_tokens") or 0),
cache_read_input_tokens=int(record.get("cache_read_input_tokens") or 0),
request_type=record.get("request_type") or "chat",
api_format=record.get("api_format"),
endpoint_api_format=record.get("endpoint_api_format"),
has_format_conversion=bool(record.get("has_format_conversion")),
is_stream=bool(record.get("is_stream", True)),
response_time_ms=record.get("response_time_ms"),
first_byte_time_ms=record.get("first_byte_time_ms"),
status_code=int(record.get("status_code") or 200),
error_message=record.get("error_message"),
metadata=record.get("metadata"),
request_headers=record.get("request_headers"),
request_body=record.get("request_body"),
provider_request_headers=record.get("provider_request_headers"),
response_headers=record.get("response_headers"),
client_response_headers=record.get("client_response_headers"),
response_body=record.get("response_body"),
request_id=request_id,
provider_id=record.get("provider_id"),
provider_endpoint_id=record.get("provider_endpoint_id"),
provider_api_key_id=record.get("provider_api_key_id"),
status=record.get("status") or "completed",
cache_ttl_minutes=record.get("cache_ttl_minutes"),
use_tiered_pricing=record.get("use_tiered_pricing", True),
target_model=record.get("target_model"),
)
usage_params, total_cost = await cls._prepare_usage_record(params)
user = params.user
api_key = params.api_key
# 创建 Usage 记录
usage = Usage(**usage_params)
@@ -1275,7 +1358,7 @@ class UsageService:
except Exception as e:
skipped_count += 1
logger.warning(f"批量记录中跳过无效记录: {e}, request_id={record.get('request_id')}")
logger.warning("批量记录中跳过无效记录: %s, request_id=%s", e, request_id)
continue
# 统计跳过的记录,失败率超过 10% 时提升日志级别
@@ -1283,11 +1366,12 @@ class UsageService:
skip_ratio = skipped_count / total_count if total_count > 0 else 0
if skip_ratio > 0.1:
logger.error(
f"批量记录失败率过高: {skipped_count}/{total_count} ({skip_ratio:.1%}) 条记录被跳过"
"批量记录失败率过高: %d/%d (%.1f%%) 条记录被跳过",
skipped_count, total_count, skip_ratio * 100
)
else:
logger.warning(
f"批量记录部分失败: {skipped_count}/{total_count} 条记录被跳过"
"批量记录部分失败: %d/%d 条记录被跳过", skipped_count, total_count
)
# 批量更新 GlobalModel 使用计数