mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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 关系
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(关联 GlobalModel,contains_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),
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user