refactor(cost,perf): 成本字段 Float 改 Numeric 并优化多处查询性能

- 数据库所有 cost/price 字段从 Float 改为 Numeric(20,8),解决浮点精度问题
- API 响应中 Decimal 值统一用 float() 转换确保 JSON 序列化
- 路由中手动 commit 后标记 tx_committed_by_route 防止中间件重复提交
- 多处查询优化:SQL 聚合替代 Python 遍历、load_only/defer 减少字段加载、
  批量 DELETE 替代逐条 ORM 删除、N+1 查询消除、UNION ALL 合并多表日期查询
- 新增 provider_api_keys (provider_id, is_active) 复合索引
- 候选构建热路径 defer 冷字段,钱包扣费合并解析与加锁查询
This commit is contained in:
fawney19
2026-03-08 16:44:16 +08:00
parent bd3f73c2fc
commit 48d13762d9
25 changed files with 474 additions and 171 deletions

View File

@@ -16,7 +16,7 @@ from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import case, func
from sqlalchemy.orm import Session
from sqlalchemy.orm import Session, load_only
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, RequestCandidate
@@ -75,11 +75,21 @@ class EndpointHealthService:
# 收集所有 provider_ids
all_provider_ids = list({ep.provider_id for ep in endpoints})
# 批量查询所有密钥(通过 provider_id 关联)
# 批量查询所有密钥(通过 provider_id 关联,只加载分组和统计所需列
all_keys = (
(
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.provider_id.in_(all_provider_ids))
.options(
load_only(
ProviderAPIKey.id,
ProviderAPIKey.provider_id,
ProviderAPIKey.is_active,
ProviderAPIKey.api_formats,
ProviderAPIKey.health_by_format,
ProviderAPIKey.circuit_breaker_by_format,
)
)
.all()
)
if all_provider_ids

View File

@@ -810,16 +810,22 @@ class HealthMonitor:
),
).first()
# 统计 Key需要遍历 JSON 字段计算熔断状态
keys = db.query(ProviderAPIKey).all()
total_keys = len(keys)
active_keys = sum(1 for k in keys if k.is_active)
# 统计 Key只加载必要列,避免全字段全表扫描
key_rows = db.query(
ProviderAPIKey.is_active,
ProviderAPIKey.health_by_format,
ProviderAPIKey.circuit_breaker_by_format,
).all()
total_keys = len(key_rows)
active_keys = 0
unhealthy_keys = 0
circuit_open_keys = 0
for key in keys:
health_by_format = key.health_by_format or {}
circuit_by_format = key.circuit_breaker_by_format or {}
for is_active, health_by_format, circuit_by_format in key_rows:
if is_active:
active_keys += 1
health_by_format = health_by_format or {}
circuit_by_format = circuit_by_format or {}
# 检查是否有任何格式健康度低于 0.5
for fmt, health_data in health_by_format.items():

View File

@@ -8,6 +8,8 @@ from __future__ import annotations
from typing import cast
from sqlalchemy import delete as sa_delete
from sqlalchemy import func
from sqlalchemy.orm import Session, joinedload, load_only
from src.core.exceptions import InvalidRequestException, NotFoundException
@@ -199,21 +201,15 @@ class GlobalModelService:
"""
global_model = GlobalModelService.get_global_model(db, global_model_id)
# 查找所有关联的 Model使用 global_model.id预加载 provider 关联)
associated_models = (
db.query(Model)
.options(joinedload(Model.provider))
.filter(Model.global_model_id == global_model.id)
.all()
# 批量删除所有关联的 Provider 模型实现
assoc_count = (
db.query(func.count(Model.id)).filter(Model.global_model_id == global_model.id).scalar()
)
# 级联删除所有关联的 Provider 模型实现
if associated_models:
if assoc_count:
logger.info(
f"删除 GlobalModel {global_model.name}{len(associated_models)} 个关联 Provider 模型"
f"删除 GlobalModel {global_model.name}{assoc_count} 个关联 Provider 模型"
)
for model in associated_models:
db.delete(model)
db.execute(sa_delete(Model).where(Model.global_model_id == global_model.id))
# 删除 GlobalModel
db.delete(global_model)

View File

@@ -11,6 +11,7 @@ from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import delete as sa_delete
from sqlalchemy.orm import Session
from src.core.crypto import crypto_service
@@ -514,13 +515,12 @@ async def batch_delete_endpoint_keys_response(db: Session, key_ids: list[str]) -
# 收集受影响的 provider_id
affected_provider_ids = {key.provider_id for key in keys if key.provider_id}
# 批量删除,一次提交
# 批量 SQL DELETE,一次提交
success_count = 0
try:
for key in keys:
db.delete(key)
db.execute(sa_delete(ProviderAPIKey).where(ProviderAPIKey.id.in_(list(found_ids))))
db.commit()
success_count = len(keys)
success_count = len(found_ids)
except Exception as exc:
db.rollback()
logger.error("批量删除 Key 提交失败: {}", exc)

View File

@@ -254,6 +254,9 @@ class PoolQuotaProbeScheduler:
db = create_session()
try:
providers = db.query(Provider).filter(Provider.is_active == True).all() # noqa: E712
# 先筛选出符合条件的 provider收集其 ID 和配置
eligible_providers: list[tuple[str, str, int]] = [] # (id, type, interval_seconds)
for provider in providers:
provider_id = str(getattr(provider, "id", "") or "")
provider_type = normalize_provider_type(getattr(provider, "provider_type", ""))
@@ -267,16 +270,29 @@ class PoolQuotaProbeScheduler:
interval_minutes = _normalize_probe_interval_minutes(
pool_cfg.probing_interval_minutes
)
interval_seconds = interval_minutes * 60
eligible_providers.append((provider_id, provider_type, interval_minutes * 60))
keys = (
# 批量查询所有符合条件的 provider 的活跃 keys避免 N+1
eligible_ids = [p[0] for p in eligible_providers]
all_keys: list[ProviderAPIKey] = []
if eligible_ids:
all_keys = (
db.query(ProviderAPIKey)
.filter(
ProviderAPIKey.provider_id == provider_id,
ProviderAPIKey.provider_id.in_(eligible_ids),
ProviderAPIKey.is_active == True, # noqa: E712
)
.all()
)
# 按 provider_id 分组
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
for key in all_keys:
pid = str(key.provider_id)
keys_by_provider.setdefault(pid, []).append(key)
for provider_id, provider_type, interval_seconds in eligible_providers:
keys = keys_by_provider.get(provider_id, [])
if not keys:
continue

View File

@@ -97,7 +97,19 @@ class CandidateBuilder:
db.query(Provider)
.options(
# 预加载 Provider 级别的 api_keys
selectinload(Provider.api_keys),
# defer 排除仅后台管理/模型获取用的冷字段,热路径字段全部加载
selectinload(Provider.api_keys).defer(
ProviderAPIKey.note,
ProviderAPIKey.last_error_msg,
ProviderAPIKey.auto_fetch_models,
ProviderAPIKey.locked_models,
ProviderAPIKey.model_include_patterns,
ProviderAPIKey.model_exclude_patterns,
ProviderAPIKey.last_models_fetch_at,
ProviderAPIKey.last_models_fetch_error,
ProviderAPIKey.max_probe_interval_minutes,
ProviderAPIKey.expires_at,
),
# 预加载 endpoints用于按 api_format 选择请求配置)
selectinload(Provider.endpoints),
# 同时加载 models 和 global_model 关系

View File

@@ -18,10 +18,10 @@
from __future__ import annotations
import asyncio
from datetime import datetime, timedelta, timezone
from datetime import date, datetime, timedelta, timezone
from typing import Any
from sqlalchemy import delete, text
from sqlalchemy import delete, literal_column, text
from src.core.logger import logger
from src.database import create_session
@@ -393,40 +393,43 @@ class MaintenanceScheduler:
check_start_date, datetime.min.time(), tzinfo=timezone.utc
)
# 获取 StatsDaily 和 StatsDailyModel 中已有数据的日期集合
existing_daily_dates = set()
existing_model_dates = set()
existing_provider_dates = set()
# 单次查询获取三张统计表中已有数据的日期UNION ALL 合并)
existing_daily_dates: set[date] = set()
existing_model_dates: set[date] = set()
existing_provider_dates: set[date] = set()
daily_stats = (
db.query(StatsDaily.date).filter(StatsDaily.date >= check_start_dt).all()
)
for (stat_date,) in daily_stats:
if stat_date.tzinfo is None:
stat_date = stat_date.replace(tzinfo=timezone.utc)
existing_daily_dates.add(stat_date.date())
model_stats = (
db.query(StatsDailyModel.date)
q_daily = db.query(
StatsDaily.date.label("dt"),
literal_column("'daily'").label("src"),
).filter(StatsDaily.date >= check_start_dt)
q_model = (
db.query(
StatsDailyModel.date.label("dt"),
literal_column("'model'").label("src"),
)
.filter(StatsDailyModel.date >= check_start_dt)
.distinct()
.all()
)
for (stat_date,) in model_stats:
if stat_date.tzinfo is None:
stat_date = stat_date.replace(tzinfo=timezone.utc)
existing_model_dates.add(stat_date.date())
provider_stats = (
db.query(StatsDailyProvider.date)
q_provider = (
db.query(
StatsDailyProvider.date.label("dt"),
literal_column("'provider'").label("src"),
)
.filter(StatsDailyProvider.date >= check_start_dt)
.distinct()
.all()
)
for (stat_date,) in provider_stats:
combined = q_daily.union_all(q_model).union_all(q_provider).all()
for stat_date, src in combined:
if stat_date.tzinfo is None:
stat_date = stat_date.replace(tzinfo=timezone.utc)
existing_provider_dates.add(stat_date.date())
d = stat_date.date()
if src == "daily":
existing_daily_dates.add(d)
elif src == "model":
existing_model_dates.add(d)
else:
existing_provider_dates.add(d)
# 找出需要回填的日期
all_dates = set()

View File

@@ -1339,10 +1339,10 @@ def query_stats_hybrid(
result.output_tokens += stats.output_tokens
result.cache_creation_tokens += stats.cache_creation_tokens
result.cache_read_tokens += stats.cache_read_tokens
result.cache_creation_cost += stats.cache_creation_cost
result.cache_read_cost += stats.cache_read_cost
result.total_cost += stats.total_cost
result.actual_total_cost += stats.actual_total_cost
result.cache_creation_cost += float(stats.cache_creation_cost or 0)
result.cache_read_cost += float(stats.cache_read_cost or 0)
result.total_cost += float(stats.total_cost or 0)
result.actual_total_cost += float(stats.actual_total_cost or 0)
result.total_response_time_ms += (stats.avg_response_time_ms or 0.0) * stats.total_requests
# Realtime per day

View File

@@ -37,7 +37,8 @@ class SyncStatsService:
try:
# 获取要同步的API密钥使用分页避免大数据量问题
if api_key_id:
api_keys = db.query(ApiKey).filter(ApiKey.id == api_key_id).all()
single_key = db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
api_keys = [single_key] if single_key else []
else:
# 分页处理,避免一次加载所有数据
offset = 0
@@ -106,7 +107,7 @@ class SyncStatsService:
api_key.total_requests = actual_requests
needs_update = True
if abs(api_key.total_cost_usd - actual_cost) > 0.0001:
if abs(float(api_key.total_cost_usd or 0) - actual_cost) > 0.0001:
logger.info(
f"API密钥 {api_key.id} 费用不一致: {api_key.total_cost_usd} -> {actual_cost}"
)

View File

@@ -800,7 +800,7 @@ class StreamUsageTracker:
if usage_record:
try:
# 在 usage_record 仍在会话中时,立即获取所需属性
total_cost = usage_record.total_cost_usd or 0.0
total_cost = float(usage_record.total_cost_usd or 0)
except Exception as e:
logger.warning(f"Failed to access total_cost_usd from usage_record: {e}")
total_cost = 0.0

View File

@@ -9,7 +9,7 @@ from datetime import datetime, timezone
from typing import Any
from sqlalchemy import and_, func, or_
from sqlalchemy.orm import Session
from sqlalchemy.orm import Session, contains_eager
from src.core.logger import logger
from src.core.validators import EmailValidator, PasswordValidator, UsernameValidator
@@ -481,10 +481,11 @@ class UserService:
if not all_active_provider_ids:
return []
# 查询所有活跃的 Model关联 GlobalModel
# 查询所有活跃的 Model关联 GlobalModelcontains_eager 避免循环中懒加载
all_models = (
db.query(Model)
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
.options(contains_eager(Model.global_model))
.filter(
and_(
Model.provider_id.in_(all_active_provider_ids),

View File

@@ -338,20 +338,32 @@ class WalletService:
return cls.get_spendable_balance_value(wallet)
@classmethod
def _resolve_wallet_for_usage(cls, db: Session, usage: Usage) -> Wallet | None:
def _resolve_wallet_for_usage(
cls, db: Session, usage: Usage, *, for_update: bool = False
) -> Wallet | None:
"""解析 Usage 对应的钱包。for_update=True 时直接返回加锁的钱包,避免二次查询。"""
wallet_id = None
if usage.wallet_id:
wallet = db.query(Wallet).filter(Wallet.id == usage.wallet_id).first()
if wallet:
return wallet
api_key = None
if usage.api_key_id:
api_key = db.query(ApiKey).filter(ApiKey.id == usage.api_key_id).first()
if api_key and api_key.is_standalone:
return cls.get_or_create_wallet(db, api_key=api_key)
if usage.user_id:
user = db.query(User).filter(User.id == usage.user_id).first()
return cls.get_or_create_wallet(db, user=user, api_key=api_key)
return None
wallet_id = usage.wallet_id
else:
api_key = None
if usage.api_key_id:
api_key = db.query(ApiKey).filter(ApiKey.id == usage.api_key_id).first()
if api_key and api_key.is_standalone:
wallet = cls.get_or_create_wallet(db, api_key=api_key)
wallet_id = wallet.id if wallet else None
if wallet_id is None and usage.user_id:
user = db.query(User).filter(User.id == usage.user_id).first()
wallet = cls.get_or_create_wallet(db, user=user, api_key=api_key)
wallet_id = wallet.id if wallet else None
if wallet_id is None:
return None
query = db.query(Wallet).filter(Wallet.id == wallet_id)
if for_update:
query = query.with_for_update()
return query.first()
@classmethod
def apply_usage_charge(
@@ -365,13 +377,7 @@ class WalletService:
if amount <= Decimal("0"):
return None, None
wallet = cls._resolve_wallet_for_usage(db, usage)
if wallet is None:
return None, None
locked_wallet = (
db.query(Wallet).filter(Wallet.id == wallet.id).with_for_update().one_or_none()
)
locked_wallet = cls._resolve_wallet_for_usage(db, usage, for_update=True)
if locked_wallet is None:
return None, None