Merge pull request #217 from AAEE86/1233

perf: 优化请求鉴权链路并批量化统计/调度查询
This commit is contained in:
fawney19
2026-03-10 23:31:28 +08:00
committed by GitHub
44 changed files with 2745 additions and 571 deletions

View File

@@ -73,13 +73,7 @@ class AdminPercentilesAdapter(AdminApiAdapter):
)
return result
result = []
for local_date, day_start_utc, day_end_utc in time_range.get_local_day_hours():
percentiles = StatsAggregatorService.compute_daily_percentiles(
context.db, day_start_utc, day_end_utc
)
result.append({"date": local_date.isoformat(), **percentiles})
return result
return StatsAggregatorService.compute_percentiles_by_local_day(context.db, time_range)
@router.get("/performance/percentiles")

View File

@@ -10,7 +10,7 @@ from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import ValidationError
from sqlalchemy import case, func
from sqlalchemy.orm import Session
from sqlalchemy.orm import Session, selectinload
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
@@ -975,9 +975,6 @@ class AdminExportConfigAdapter(AdminApiAdapter):
from src.core.crypto import crypto_service
from src.models.database import (
GlobalModel,
Model,
ProviderAPIKey,
ProviderEndpoint,
ProxyNode,
)
@@ -991,22 +988,35 @@ class AdminExportConfigAdapter(AdminApiAdapter):
gm_name_map: dict[str, str] = {gm.id: gm.name for gm in global_models}
# 导出 Providers 及其关联数据
providers = db.query(Provider).all()
providers = (
db.query(Provider)
.options(
selectinload(Provider.endpoints),
selectinload(Provider.api_keys),
selectinload(Provider.models),
)
.all()
)
providers_data = []
def _normalize_created_at_for_sort(value: datetime | None) -> datetime:
if value is None:
return datetime.min.replace(tzinfo=timezone.utc)
return value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc)
for provider in providers:
# 导出 Endpoints
endpoints = (
db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
)
endpoints = list(provider.endpoints)
endpoints_data = [ep.to_export_dict() for ep in endpoints]
provider_endpoint_formats = self._collect_provider_endpoint_formats(endpoints)
# 导出 Provider Keys按 provider_id 归属,包含 api_formats
keys = (
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.provider_id == provider.id)
.order_by(ProviderAPIKey.internal_priority.asc(), ProviderAPIKey.created_at.asc())
.all()
keys = sorted(
provider.api_keys,
key=lambda key: (
key.internal_priority if key.internal_priority is not None else 0,
_normalize_created_at_for_sort(key.created_at),
),
)
keys_data = []
for key in keys:
@@ -1046,7 +1056,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
# 导出 Provider Models
# 注意提供商模型Model必须关联全局模型GlobalModel才能参与路由
# 导入时未关联 GlobalModel 的模型会被跳过,这是业务规则而非 bug
models = db.query(Model).filter(Model.provider_id == provider.id).all()
models = list(provider.models)
models_data = []
for model in models:
model_data = model.to_export_dict()

View File

@@ -224,10 +224,11 @@ async def get_usage_records(
**返回字段**:
- `records`: 使用记录列表,包含 id, user_id, user_email, username, api_key, provider, model, target_model,
input_tokens, output_tokens, cache_creation_input_tokens, cache_read_input_tokens, total_tokens,
cost, actual_cost, rate_multiplier, response_time_ms, first_byte_time_ms, created_at, is_stream,
input_price_per_1m, output_price_per_1m, cache_creation_price_per_1m, cache_read_price_per_1m,
status_code, error_message, status, has_fallback, has_retry, has_rectified, api_format, api_key_name, request_metadata
model_version, input_tokens, output_tokens, cache_creation_input_tokens, cache_read_input_tokens,
total_tokens, cost, actual_cost, rate_multiplier, response_time_ms, first_byte_time_ms, created_at,
is_stream, input_price_per_1m, output_price_per_1m, cache_creation_price_per_1m,
cache_read_price_per_1m, status_code, error_message, status, has_fallback, has_retry,
has_rectified, api_format, api_key_name
- `total`: 符合条件的总记录数
- `limit`: 当前分页限制
- `offset`: 当前分页偏移量
@@ -885,8 +886,12 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
count_query = count_query.outerjoin(ApiKey, Usage.api_key_id == ApiKey.id)
# -- 构建数据查询(完整 JOIN --
usage_model_version = Usage.request_metadata["model_version"].as_string().label(
"model_version"
)
query = (
db.query(Usage, User, ProviderEndpoint, ProviderAPIKey, ApiKey)
db.query(Usage, User, ProviderEndpoint, ProviderAPIKey, ApiKey, usage_model_version)
.outerjoin(User, Usage.user_id == User.id)
.outerjoin(ProviderEndpoint, Usage.provider_endpoint_id == ProviderEndpoint.id)
.outerjoin(ProviderAPIKey, Usage.provider_api_key_id == ProviderAPIKey.id)
@@ -1001,7 +1006,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
# Perf: count query uses fewer JOINs than the data query
total = int(count_query.scalar() or 0)
# Perf: do not load large request/response columns for list view
# Perf: do not load large request/response columns or full request_metadata for list view
query = query.options(
load_only(
Usage.id,
@@ -1032,7 +1037,6 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
Usage.api_format,
Usage.endpoint_api_format,
Usage.has_format_conversion,
Usage.request_metadata,
Usage.input_price_per_1m,
Usage.output_price_per_1m,
Usage.cache_creation_price_per_1m,
@@ -1047,7 +1051,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
)
request_ids = [usage.request_id for usage, _, _, _, _ in records if usage.request_id]
request_ids = [usage.request_id for usage, _, _, _, _, _ in records if usage.request_id]
fallback_map = {}
retry_map = {}
rectified_map = {}
@@ -1110,7 +1114,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
# 构建 provider_id -> Provider 名称的映射,避免 N+1 查询
provider_ids = list(
{usage.provider_id for usage, _, _, _, _ in records if usage.provider_id}
{usage.provider_id for usage, _, _, _, _, _ in records if usage.provider_id}
)
provider_map = {}
if provider_ids:
@@ -1121,7 +1125,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
data = []
api_key_display_cache: dict[str, str] = {}
for usage, user, endpoint, provider_api_key, user_api_key in records:
for usage, user, endpoint, provider_api_key, user_api_key, model_version in records:
actual_cost = (
float(usage.actual_total_cost_usd)
if usage.actual_total_cost_usd is not None
@@ -1198,7 +1202,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
"endpoint_api_format": endpoint_api_format,
"has_format_conversion": bool(has_format_conversion),
"api_key_name": provider_api_key.name if provider_api_key else None,
"request_metadata": usage.request_metadata, # Provider 响应元数据
"model_version": model_version, # Provider 返回的实际模型版本(轻量字段)
}
)
@@ -1227,7 +1231,10 @@ class AdminActiveRequestsAdapter(AdminApiAdapter):
return {"requests": []}
requests = UsageService.get_active_requests_status(
db=db, ids=id_list, include_admin_fields=True
db=db,
ids=id_list,
include_admin_fields=True,
maintain_status=True,
)
return {"requests": requests}

View File

@@ -18,7 +18,7 @@ from src.core.logger import logger
from src.database import get_db
from src.models.admin_requests import UpdateUserRequest
from src.models.api import CreateApiKeyRequest, CreateUserRequest
from src.models.database import ApiKey, User, UserRole
from src.models.database import ApiKey, User, UserRole, Wallet
from src.services.system.config import SystemConfigService
from src.services.user.apikey import ApiKeyService
from src.services.user.bulk_cleanup import pre_clean_api_key
@@ -30,8 +30,23 @@ router = APIRouter(prefix="/api/admin/users", tags=["Admin - Users"])
pipeline = ApiRequestPipeline()
def _serialize_user(db: Session, user: User) -> dict[str, Any]:
wallet = WalletService.get_wallet(db, user_id=user.id)
class _WalletSentinelType:
pass
_WALLET_SENTINEL = _WalletSentinelType()
def _serialize_user(
db: Session,
user: User,
wallet: Wallet | None | _WalletSentinelType = _WALLET_SENTINEL,
) -> dict[str, Any]:
resolved_wallet: Wallet | None
if wallet is _WALLET_SENTINEL:
resolved_wallet = WalletService.get_wallet(db, user_id=user.id)
else:
resolved_wallet = wallet
return {
"id": user.id,
"email": user.email,
@@ -40,7 +55,7 @@ def _serialize_user(db: Session, user: User) -> dict[str, Any]:
"allowed_providers": user.allowed_providers,
"allowed_api_formats": user.allowed_api_formats,
"allowed_models": user.allowed_models,
"unlimited": WalletService.is_unlimited_wallet(wallet),
"unlimited": WalletService.is_unlimited_wallet(resolved_wallet),
"is_active": user.is_active,
"created_at": user.created_at.isoformat(),
"updated_at": user.updated_at.isoformat() if user.updated_at else None,
@@ -334,7 +349,8 @@ class AdminListUsersAdapter(AdminApiAdapter):
except KeyError as exc:
raise InvalidRequestException("角色参数不合法") from exc
users = UserService.list_users(db, self.skip, self.limit, role_enum, self.is_active)
return [_serialize_user(db, u) for u in users]
wallets_by_user_id = WalletService.get_wallets_by_user_ids(db, [user.id for user in users])
return [_serialize_user(db, user, wallets_by_user_id.get(user.id)) for user in users]
class AdminGetUserAdapter(AdminApiAdapter):

View File

@@ -274,10 +274,10 @@ class AuthLoginAdapter(AuthPublicAdapter):
detail=f"登录请求过于频繁,请在 {reset_after} 秒后重试",
)
user = await AuthService.authenticate_user(
authenticated_user = await AuthService.authenticate_user_threadsafe(
db, login_request.email, login_request.password, login_request.auth_type
)
if not user:
if not authenticated_user:
AuditService.log_login_attempt(
db=db,
email=login_request.email,
@@ -296,22 +296,30 @@ class AuthLoginAdapter(AuthPublicAdapter):
success=True,
ip_address=client_ip,
user_agent=user_agent,
user_id=user.id,
user_id=authenticated_user.user_id,
)
db.commit()
context.request.state.tx_committed_by_route = True
access_token = AuthService.create_access_token(
data={
"user_id": user.id,
"role": user.role.value,
"created_at": user.created_at.isoformat() if user.created_at else None,
"user_id": authenticated_user.user_id,
"role": authenticated_user.role.value,
"created_at": (
authenticated_user.created_at.isoformat()
if authenticated_user.created_at
else None
),
}
)
refresh_token = AuthService.create_refresh_token(
data={
"user_id": user.id,
"created_at": user.created_at.isoformat() if user.created_at else None,
"user_id": authenticated_user.user_id,
"created_at": (
authenticated_user.created_at.isoformat()
if authenticated_user.created_at
else None
),
}
)
response = LoginResponse(
@@ -319,10 +327,10 @@ class AuthLoginAdapter(AuthPublicAdapter):
refresh_token=refresh_token,
token_type="bearer",
expires_in=86400,
user_id=user.id,
email=user.email,
username=user.username,
role=user.role.value,
user_id=authenticated_user.user_id,
email=authenticated_user.email,
username=authenticated_user.username,
role=authenticated_user.role.value,
)
return response.model_dump()
@@ -332,9 +340,6 @@ class AuthRefreshAdapter(AuthPublicAdapter):
db = context.db
payload = context.ensure_json_body()
refresh_request = RefreshTokenRequest.model_validate(payload)
client_ip = get_client_ip(context.request)
user_agent = get_user_agent(context.request)
try:
token_payload = await AuthService.verify_token(
refresh_request.refresh_token, token_type="refresh"
@@ -745,7 +750,6 @@ class AuthSendVerificationCodeAdapter(AuthPublicAdapter):
class AuthVerifyEmailAdapter(AuthPublicAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""验证邮箱验证码"""
db = context.db
payload = context.ensure_json_body()
try:

View File

@@ -25,6 +25,7 @@ class ApiAdapter(ABC):
audit_log_enabled: bool = True
audit_success_event = None
audit_failure_event = None
eager_request_body: bool = True
@abstractmethod
async def handle(self, context: ApiRequestContext) -> Response:

View File

@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import gzip
import json
import time
@@ -9,7 +10,9 @@ from typing import Any
from fastapi import HTTPException, Request
from sqlalchemy.orm import Session
from starlette.requests import ClientDisconnect
from src.config.settings import config
from src.core.api_format.headers import get_header_value
from src.core.http_compression import is_gzip_content_encoding, normalize_content_encoding
from src.core.logger import logger
@@ -53,6 +56,56 @@ class ApiRequestContext:
client_content_encoding: str | None = None
client_accept_encoding: str | None = None
async def ensure_raw_body_async(self) -> bytes:
"""按需读取原始请求体,避免所有请求都在 Pipeline 阶段预读。"""
if self.raw_body is not None:
return self.raw_body
perf_metrics = getattr(self.request.state, "perf_metrics", None)
perf_sampled = isinstance(perf_metrics, dict) and bool(perf_metrics)
body_start = PerfRecorder.start(force=perf_sampled)
body_size = 0
try:
self.raw_body = await asyncio.wait_for(
self.request.body(), timeout=config.request_body_timeout
)
body_size = len(self.raw_body or b"")
except TimeoutError as exc:
timeout_sec = int(config.request_body_timeout)
logger.error("读取请求体超时({}s),可能客户端未发送完整请求体", timeout_sec)
raise HTTPException(
status_code=408,
detail=f"Request timeout: body not received within {timeout_sec} seconds",
) from exc
except ClientDisconnect:
logger.warning(
"[Context] 客户端在读取请求体期间断开连接: {} {}",
self.request.method,
self.request.url.path,
)
raise HTTPException(
status_code=499,
detail="Client closed request",
)
finally:
body_duration = PerfRecorder.stop(
body_start,
"pipeline_body_read",
labels={"mode": self.mode},
log_hint=f"size={body_size}",
)
if isinstance(perf_metrics, dict):
pipeline_metrics = perf_metrics.setdefault("pipeline", {})
pipeline_metrics["body_read_ms"] = int((body_duration or 0) * 1000)
pipeline_metrics["body_bytes"] = int(body_size)
return self.raw_body or b""
async def ensure_json_body_async(self) -> dict[str, Any]:
"""异步懒加载 JSON 请求体。"""
await self.ensure_raw_body_async()
return self.ensure_json_body()
def ensure_json_body(self) -> dict[str, Any]:
"""确保请求体已解析为JSON并返回。"""
if self.json_body is not None:

View File

@@ -1,10 +1,12 @@
from __future__ import annotations
import time
from datetime import datetime, timezone
from enum import Enum
from typing import TYPE_CHECKING, Any
from fastapi import HTTPException, Request
from fastapi.concurrency import run_in_threadpool
from fastapi.responses import JSONResponse
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.orm import Session
@@ -14,6 +16,7 @@ from src.config.settings import config
from src.core.enums import UserRole
from src.core.exceptions import BalanceInsufficientException
from src.core.logger import logger
from src.database.database import create_session
from src.models.database import ApiKey, AuditEventType, User
from src.services.auth.service import AuthService
from src.services.system.audit import AuditService
@@ -108,14 +111,22 @@ class ApiRequestPipeline:
user, management_token = await self._authenticate_management(http_request, db)
api_key = None
else:
user, api_key = self._authenticate_client(http_request, db, adapter, quiet=is_quiet)
user, api_key = await self._authenticate_client(
http_request,
db,
adapter,
quiet=is_quiet,
)
management_token = None
finally:
auth_duration = PerfRecorder.stop(auth_start, "pipeline_auth", labels=perf_labels)
_record_perf_metric("auth_ms", auth_duration)
raw_body = None
if http_request.method in {"POST", "PUT", "PATCH"}:
should_eager_read_body = http_request.method in {"POST", "PUT", "PATCH"} and getattr(
adapter, "eager_request_body", True
)
if should_eager_read_body:
try:
import asyncio
@@ -141,7 +152,7 @@ class ApiRequestPipeline:
perf_metrics.setdefault("pipeline", {})["body_bytes"] = int(body_size)
except TimeoutError:
timeout_sec = int(config.request_body_timeout)
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
logger.error("读取请求体超时({}s),可能客户端未发送完整请求体", timeout_sec)
raise HTTPException(
status_code=408,
detail=f"Request timeout: body not received within {timeout_sec} seconds",
@@ -176,8 +187,11 @@ class ApiRequestPipeline:
context.management_token = management_token
# 存储 quiet 标志到 context用于审计日志判断
context.quiet_logging = is_quiet
if mode != ApiMode.ADMIN and user:
remaining = self._calculate_balance_remaining(db, user, api_key=api_key)
if mode in {ApiMode.STANDARD, ApiMode.PROXY, ApiMode.USER} and user:
if hasattr(http_request.state, "prefetched_balance_remaining"):
remaining = getattr(http_request.state, "prefetched_balance_remaining")
else:
remaining = await self._calculate_balance_remaining_async(user, api_key=api_key)
context.balance_remaining = remaining
# authorize 可能是异步的,需要检查并 await
authorize_start = PerfRecorder.start(force=perf_sampled)
@@ -238,7 +252,7 @@ class ApiRequestPipeline:
try:
context.db.rollback()
except Exception as rollback_exc:
logger.debug(f"[Pipeline] 回滚失败(可忽略): {rollback_exc}")
logger.debug("[Pipeline] 回滚失败(可忽略): {}", rollback_exc)
self._record_audit_event(
context,
adapter,
@@ -252,32 +266,52 @@ class ApiRequestPipeline:
# Internal helpers
# --------------------------------------------------------------------- #
def _authenticate_client(
async def _authenticate_client(
self, request: Request, db: Session, adapter: ApiAdapter, **_kw: object
) -> tuple[User, ApiKey]:
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
client_api_key = adapter.extract_api_key(request)
if not client_api_key:
raise HTTPException(status_code=401, detail="请提供API密钥")
auth_result = self.auth_service.authenticate_api_key(db, client_api_key)
auth_result = await self.auth_service.authenticate_api_key_threadsafe(client_api_key)
if not auth_result:
raise HTTPException(status_code=401, detail="无效的API密钥")
user, api_key = auth_result
user = auth_result.user
api_key = auth_result.api_key
if not user or not api_key:
raise HTTPException(status_code=401, detail="无效的API密钥")
request.state.user_id = user.id
request.state.api_key_id = api_key.id
# 线程池认证返回的是分离对象;重新绑定到路由会话,避免后续写入失效。
db_user = db.query(User).filter(User.id == user.id).first()
db_api_key = db.query(ApiKey).filter(ApiKey.id == api_key.id).first()
if not db_user or not db_api_key:
raise HTTPException(status_code=401, detail="无效的API密钥")
# 使用路由会话再核对一次状态,避免线程池认证结果与当前事务视图短暂不一致。
if not db_user.is_active or db_user.is_deleted:
raise HTTPException(status_code=401, detail="无效的API密钥")
if not db_api_key.is_active:
raise HTTPException(status_code=401, detail="无效的API密钥")
if db_api_key.is_locked and not db_api_key.is_standalone:
raise HTTPException(status_code=403, detail="该密钥已被管理员锁定,请联系管理员")
if db_api_key.user_id != db_user.id:
raise HTTPException(status_code=401, detail="无效的API密钥")
if db_api_key.expires_at:
expires_at = db_api_key.expires_at
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=timezone.utc)
if expires_at < datetime.now(timezone.utc):
raise HTTPException(status_code=401, detail="无效的API密钥")
# 检查余额(支持独立 Key
access_ok, _message = self.usage_service.check_request_balance(db, user, api_key=api_key)
if not access_ok:
remaining = self._calculate_balance_remaining(db, user, api_key=api_key)
request.state.user_id = db_user.id
request.state.api_key_id = db_api_key.id
request.state.prefetched_balance_remaining = auth_result.balance_remaining
if not auth_result.access_allowed:
remaining = auth_result.balance_remaining
raise BalanceInsufficientException(balance_type="USD", remaining=remaining)
return user, api_key
return db_user, db_api_key
async def _try_token_prefix_auth(
self, token: str, request: Request, db: Session
@@ -293,9 +327,7 @@ class ApiRequestPipeline:
from src.core.modules.hooks import AUTH_TOKEN_PREFIX_AUTHENTICATORS, get_hook_dispatcher
from src.utils.request_utils import get_client_ip
authenticators = await get_hook_dispatcher().dispatch(
AUTH_TOKEN_PREFIX_AUTHENTICATORS, db=db
)
authenticators = await get_hook_dispatcher().dispatch(AUTH_TOKEN_PREFIX_AUTHENTICATORS)
for auth_info in authenticators or []:
prefix = auth_info.get("prefix", "")
authenticate_fn = auth_info.get("authenticate")
@@ -304,111 +336,136 @@ class ApiRequestPipeline:
logger.warning("Token prefix '{}' has no authenticate callback", prefix)
raise HTTPException(status_code=401, detail="认证服务不可用")
client_ip = get_client_ip(request)
result = await authenticate_fn(db, token, client_ip)
if result:
return result
auth_db = create_session()
try:
result = await authenticate_fn(auth_db, token, client_ip)
if result:
for instance in result:
if instance is None:
continue
try:
auth_db.expunge(instance)
except Exception:
pass
return result
finally:
auth_db.close()
# 前缀匹配但认证失败
module_name = auth_info.get("module", "unknown")
raise HTTPException(status_code=401, detail=f"无效或过期的 Token ({module_name})")
return None # 无前缀匹配
def _reattach_token_auth_result(
self,
db: Session,
token_auth_result: tuple[User, Any],
) -> tuple[User, Any]:
"""将前缀认证返回对象绑定到当前请求会话,避免后续写入失效。"""
user, management_token = token_auth_result
db_user = db.query(User).filter(User.id == user.id).first()
if not db_user:
raise HTTPException(status_code=401, detail="无效或过期的 Token")
if management_token is None:
return db_user, None
token_id = getattr(management_token, "id", None)
token_model: Any = type(management_token)
if token_id is None or not hasattr(token_model, "id"):
return db_user, management_token
db_management_token = db.query(token_model).filter(token_model.id == token_id).first()
if not db_management_token:
raise HTTPException(status_code=401, detail="无效或过期的 Token")
return db_user, db_management_token
async def _authenticate_admin(
self, request: Request, db: Session
) -> tuple[User, ManagementToken | None]:
"""管理员认证,支持 JWT Management Token 两种方式"""
"""Admin auth supports JWT and Management Token."""
authorization = request.headers.get("authorization")
if not authorization or not authorization.lower().startswith("bearer "):
raise HTTPException(status_code=401, detail="缺少管理员凭证")
token = authorization[7:].strip()
# 通过钩子检查是否匹配模块注册的 token 前缀(如 ae_
token_auth_result = await self._try_token_prefix_auth(token, request, db)
if token_auth_result is not None:
user, management_token = token_auth_result
user, management_token = self._reattach_token_auth_result(db, token_auth_result)
# 检查管理员权限
if user.role != UserRole.ADMIN:
logger.warning(f"非管理员尝试通过 Management Token 访问管理端点: {user.email}")
logger.warning("非管理员尝试通过 Management Token 访问管理端点: {}", user.email)
raise HTTPException(status_code=403, detail="需要管理员权限")
# 存储到 request.state
request.state.user_id = user.id
request.state.management_token_id = management_token.id
request.state.management_token_id = management_token.id if management_token else None
return user, management_token
# JWT 认证
try:
payload = await self.auth_service.verify_token(token, token_type="access")
except HTTPException:
raise
except Exception as exc:
logger.error(f"Admin token 验证失败: {exc}")
logger.error("Admin token 验证失败: {}", exc)
raise HTTPException(status_code=401, detail="无效的管理员令牌")
user_id = payload.get("user_id")
if not user_id:
raise HTTPException(status_code=401, detail="无效的管理员令牌")
# 直接查询数据库,确保返回的是当前 Session 绑定的对象
user = db.query(User).filter(User.id == user_id).first()
if not user or not user.is_active or user.is_deleted:
db_user = db.query(User).filter(User.id == user_id).first()
if not db_user or not db_user.is_active or db_user.is_deleted:
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
if not self.auth_service.token_identity_matches_user(payload, user):
if not self.auth_service.token_identity_matches_user(payload, db_user):
raise HTTPException(status_code=403, detail="无效的管理员令牌")
# 检查管理员权限
if user.role != UserRole.ADMIN:
logger.warning(f"非管理员尝试通过 JWT 访问管理端点: {user.email}")
if db_user.role != UserRole.ADMIN:
logger.warning("非管理员尝试通过 JWT 访问管理端点: {}", db_user.email)
raise HTTPException(status_code=403, detail="需要管理员权限")
request.state.user_id = user.id
return user, None
request.state.user_id = db_user.id
return db_user, None
async def _authenticate_user(
self, request: Request, db: Session
) -> tuple[User, ManagementToken | None]:
"""用户认证,支持 JWT Management Token 两种方式"""
"""User auth supports JWT and Management Token."""
authorization = request.headers.get("authorization")
if not authorization or not authorization.lower().startswith("bearer "):
raise HTTPException(status_code=401, detail="缺少用户凭证")
token = authorization[7:].strip()
# 通过钩子检查是否匹配模块注册的 token 前缀(如 ae_
token_auth_result = await self._try_token_prefix_auth(token, request, db)
if token_auth_result is not None:
user, management_token = token_auth_result
user, management_token = self._reattach_token_auth_result(db, token_auth_result)
request.state.user_id = user.id
request.state.management_token_id = management_token.id
request.state.management_token_id = management_token.id if management_token else None
return user, management_token
# JWT 认证
try:
payload = await self.auth_service.verify_token(token, token_type="access")
except HTTPException:
raise
except Exception as exc:
logger.error(f"User token 验证失败: {exc}")
logger.error("User token 验证失败: {}", exc)
raise HTTPException(status_code=401, detail="无效的用户令牌")
user_id = payload.get("user_id")
if not user_id:
raise HTTPException(status_code=401, detail="无效的用户令牌")
user = db.query(User).filter(User.id == user_id).first()
if not user or not user.is_active or user.is_deleted:
db_user = db.query(User).filter(User.id == user_id).first()
if not db_user or not db_user.is_active or db_user.is_deleted:
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
if not self.auth_service.token_identity_matches_user(payload, user):
if not self.auth_service.token_identity_matches_user(payload, db_user):
raise HTTPException(status_code=403, detail="无效的用户令牌")
request.state.user_id = user.id
return user, None
request.state.user_id = db_user.id
return db_user, None
async def _authenticate_management(
self, request: Request, db: Session
@@ -424,11 +481,11 @@ class ApiRequestPipeline:
# _try_token_prefix_auth 会在前缀匹配但认证失败时直接抛 HTTPException
token_auth_result = await self._try_token_prefix_auth(token, request, db)
if token_auth_result is not None:
user, management_token = token_auth_result
user, management_token = self._reattach_token_auth_result(db, token_auth_result)
# 存储到 request.state
request.state.user_id = user.id
request.state.management_token_id = management_token.id
request.state.management_token_id = management_token.id if management_token else None
return user, management_token
@@ -437,15 +494,37 @@ class ApiRequestPipeline:
detail="无效的 Token 格式,需要 Management Token",
)
def _calculate_balance_remaining(
self, db: Session, user: User | None, api_key: ApiKey | None = None
async def _calculate_balance_remaining_async(
self, user: User | None, api_key: ApiKey | None = None
) -> float | None:
if not user:
return None
balance = WalletService.get_balance_snapshot(db, user=user, api_key=api_key)
if balance is None:
return None
return float(balance)
user_id = getattr(user, "id", None)
api_key_id = getattr(api_key, "id", None)
# API Key 链路通常已在认证阶段预取余额;这里只保留为无预取路径的兜底查询。
def _load_balance() -> float | None:
thread_db = create_session()
try:
db_user = (
thread_db.query(User).filter(User.id == user_id).first() if user_id else None
)
db_api_key = (
thread_db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
if api_key_id
else None
)
balance = WalletService.get_balance_snapshot(
thread_db,
user=db_user,
api_key=db_api_key,
)
return float(balance) if balance is not None else None
finally:
thread_db.close()
return await run_in_threadpool(_load_balance)
def _record_audit_event(
self,
@@ -502,7 +581,7 @@ class ApiRequestPipeline:
)
except Exception as exc:
# 审计失败不应影响主请求,仅记录警告
logger.warning(f"[Audit] Failed to record event for adapter={adapter.name}: {exc}")
logger.warning("[Audit] Failed to record event for adapter={}: {}", adapter.name, exc)
def _build_audit_metadata(
self,
@@ -568,7 +647,9 @@ class ApiRequestPipeline:
if adapter_details:
extra_details.update(adapter_details)
except Exception as exc:
logger.warning(f"[Audit] Adapter metadata failed: {adapter.__class__.__name__}: {exc}")
logger.warning(
"[Audit] Adapter metadata failed: {}: {}", adapter.__class__.__name__, exc
)
if extra_details:
metadata["details"] = extra_details
@@ -578,7 +659,7 @@ class ApiRequestPipeline:
return self._sanitize_metadata(metadata)
def _sanitize_metadata(self, value: Any, depth: int = 0) -> None:
def _sanitize_metadata(self, value: Any, depth: int = 0) -> Any:
if value is None:
return None
if depth > 5:

View File

@@ -56,6 +56,7 @@ class ChatAdapterBase(HandlerAdapterBase):
# 适配器配置
name: str = "chat.base"
mode = ApiMode.STANDARD
eager_request_body = False
async def handle(self, context: ApiRequestContext) -> Any:
"""处理 Chat API 请求"""
@@ -71,7 +72,7 @@ class ChatAdapterBase(HandlerAdapterBase):
original_headers = context.original_headers
query_params = context.query_params
original_request_body = context.ensure_json_body()
original_request_body = await context.ensure_json_body_async()
# 合并 path_params 到请求体(如 Gemini API 的 model 在 URL 路径中)
if context.path_params:

View File

@@ -20,7 +20,6 @@ from __future__ import annotations
from typing import Any
from fastapi import HTTPException
from fastapi.responses import JSONResponse
from src.api.base.adapter import ApiMode
from src.api.base.context import ApiRequestContext
@@ -58,6 +57,7 @@ class CliAdapterBase(HandlerAdapterBase):
# 适配器配置
name: str = "cli.base"
mode = ApiMode.PROXY
eager_request_body = False
async def handle(self, context: ApiRequestContext) -> Any:
"""处理 CLI API 请求"""
@@ -82,7 +82,7 @@ class CliAdapterBase(HandlerAdapterBase):
set_original_request_headers(original_headers)
original_request_body = context.ensure_json_body()
original_request_body = await context.ensure_json_body_async()
# 合并 path_params 到请求体(如 Gemini API 的 model 在 URL 路径中)
if context.path_params:

View File

@@ -35,6 +35,7 @@ class VideoAdapterBase(ApiAdapter):
name: str = "video.base"
mode = ApiMode.STANDARD
eager_request_body = False
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
@@ -103,7 +104,7 @@ class VideoAdapterBase(ApiAdapter):
# Remix task
if method == "POST" and path.endswith("/remix") and task_id:
original_request_body = context.ensure_json_body()
original_request_body = await context.ensure_json_body_async()
return await handler.handle_remix_task(
task_id=task_id,
http_request=http_request,
@@ -134,7 +135,7 @@ class VideoAdapterBase(ApiAdapter):
# Create task (default)
if method in {"POST", "PUT", "PATCH"}:
original_request_body = context.ensure_json_body()
original_request_body = await context.ensure_json_body_async()
return await handler.handle_create_task(
http_request=http_request,
original_headers=context.original_headers,

View File

@@ -184,6 +184,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
"thinking_enabled": bool(request_obj.thinking),
}
@classmethod
def build_endpoint_url(
cls,
base_url: str,
@@ -224,6 +225,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
name = "claude.token_count"
mode = ApiMode.STANDARD
eager_request_body = False
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
@@ -239,7 +241,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
return bearer_handler.extract_credentials(request)
async def handle(self, context: ApiRequestContext) -> Any:
payload = context.ensure_json_body()
payload = await context.ensure_json_body_async()
try:
request = ClaudeTokenCountRequest.model_validate(payload, strict=False)

View File

@@ -48,7 +48,7 @@ class OpenAICliAdapter(CliAdapterBase):
async def handle(self, context: ApiRequestContext) -> Any:
"""处理 CLI API 请求 -- compact 模式下注入标记并强制非流式"""
if self._compact:
body = context.ensure_json_body()
body = await context.ensure_json_body_async()
body["_aether_compact"] = True
# compact 端点永远非流式
body.pop("stream", None)

View File

@@ -231,7 +231,6 @@ async def get_my_usage(
- `total_tokens`: 总 Token 数
- `total_cost`: 总成本USD
- `summary_by_model`: 按模型分组统计
- `summary_by_provider`: 按提供商分组统计
- `records`: 详细使用记录列表
- `pagination`: 分页信息
"""
@@ -809,6 +808,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
user_id=user.id,
start_date=start_utc,
end_date=end_utc,
group_by=None,
)
# 过滤掉 unknown/pending provider 的记录(请求未到达任何提供商)
@@ -818,31 +818,23 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
if item.get("provider") not in ("unknown", "pending", None)
]
total_requests = sum(item["requests"] for item in filtered_summary)
total_input_tokens = (
sum(item["input_tokens"] for item in filtered_summary) if filtered_summary else 0
)
total_output_tokens = (
sum(item["output_tokens"] for item in filtered_summary) if filtered_summary else 0
)
total_tokens = (
sum(item["total_tokens"] for item in filtered_summary) if filtered_summary else 0
)
total_cost = (
sum(item["total_cost_usd"] for item in filtered_summary) if filtered_summary else 0.0
)
# 管理员可以看到真实成本
total_requests = 0
total_input_tokens = 0
total_output_tokens = 0
total_tokens = 0
total_cost = 0.0
total_actual_cost = 0.0
if user.role == UserRole.ADMIN:
total_actual_cost = (
sum(item.get("actual_total_cost_usd", 0.0) for item in filtered_summary)
if filtered_summary
else 0.0
)
model_summary = {}
provider_summary = {}
for item in filtered_summary:
total_requests += item["requests"]
total_input_tokens += item["input_tokens"]
total_output_tokens += item["output_tokens"]
total_tokens += item["total_tokens"]
total_cost += item["total_cost_usd"]
if user.role == UserRole.ADMIN:
total_actual_cost += item.get("actual_total_cost_usd", 0.0)
model_name = item["model"]
base_stats = {
"model": model_name,
@@ -866,13 +858,8 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
if user.role == UserRole.ADMIN:
stats["actual_total_cost_usd"] += item.get("actual_total_cost_usd", 0.0)
summary_by_model = sorted(model_summary.values(), key=lambda x: x["requests"], reverse=True)
# 按提供商汇总(用于 UsageProviderTable
provider_summary = {}
for item in filtered_summary:
provider_name = item["provider"]
base_stats = {
provider_base_stats = {
"provider": provider_name,
"requests": 0,
"total_tokens": 0,
@@ -881,32 +868,37 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
"total_response_time_ms": 0.0,
"response_time_count": 0,
}
stats = provider_summary.setdefault(provider_name, base_stats)
stats["requests"] += item["requests"]
stats["total_tokens"] += item["total_tokens"]
stats["total_cost_usd"] += item["total_cost_usd"]
# 假设 summary 中的都是成功的请求
stats["success_count"] += item["requests"]
if item.get("avg_response_time_ms") is not None:
stats["total_response_time_ms"] += item["avg_response_time_ms"] * item["requests"]
stats["response_time_count"] += item["requests"]
provider_stats = provider_summary.setdefault(provider_name, provider_base_stats)
provider_stats["requests"] += item["requests"]
provider_stats["total_tokens"] += item["total_tokens"]
provider_stats["total_cost_usd"] += item["total_cost_usd"]
provider_stats["success_count"] += int(item.get("success_count", 0) or 0)
success_response_time_count = int(item.get("success_response_time_count", 0) or 0)
if success_response_time_count > 0:
provider_stats["total_response_time_ms"] += float(
item.get("success_response_time_sum_ms", 0.0) or 0.0
)
provider_stats["response_time_count"] += success_response_time_count
summary_by_model = sorted(model_summary.values(), key=lambda x: x["requests"], reverse=True)
summary_by_provider = []
for stats in provider_summary.values():
for provider_stats in provider_summary.values():
avg_response_time_ms = (
stats["total_response_time_ms"] / stats["response_time_count"]
if stats["response_time_count"] > 0
provider_stats["total_response_time_ms"] / provider_stats["response_time_count"]
if provider_stats["response_time_count"] > 0
else 0
)
success_rate = (
(stats["success_count"] / stats["requests"] * 100) if stats["requests"] > 0 else 100
(provider_stats["success_count"] / provider_stats["requests"] * 100)
if provider_stats["requests"] > 0
else 100
)
summary_by_provider.append(
{
"provider": stats["provider"],
"requests": stats["requests"],
"total_tokens": stats["total_tokens"],
"total_cost_usd": stats["total_cost_usd"],
"provider": provider_stats["provider"],
"requests": provider_stats["requests"],
"total_tokens": provider_stats["total_tokens"],
"total_cost_usd": provider_stats["total_cost_usd"],
"success_rate": round(success_rate, 2),
"avg_response_time_ms": round(avg_response_time_ms, 2),
}
@@ -1003,6 +995,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
"avg_response_time": avg_response_time,
"billing": WalletService.serialize_wallet_summary(wallet),
"summary_by_model": summary_by_model,
"summary_by_provider": summary_by_provider,
# 分页信息
"pagination": {
"total": total_records,
@@ -1117,7 +1110,12 @@ class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
if not id_list:
return {"requests": []}
requests = UsageService.get_active_requests_status(db=db, ids=id_list, user_id=user.id)
requests = UsageService.get_active_requests_status(
db=db,
ids=id_list,
user_id=user.id,
maintain_status=True,
)
return {"requests": requests}

View File

@@ -9,6 +9,7 @@ import secrets
import time
import uuid
from collections import OrderedDict
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from threading import Lock
from typing import TYPE_CHECKING, Any
@@ -23,6 +24,7 @@ from src.config import config
from src.core.enums import AuthSource
from src.core.exceptions import ForbiddenException
from src.core.logger import logger
from src.database.database import create_session
from src.services.system.config import SystemConfigService
if TYPE_CHECKING:
@@ -32,6 +34,31 @@ from src.models.database import ApiKey, User, UserRole
from src.services.auth.jwt_blacklist import JWTBlacklistService
from src.services.cache.user_cache import UserCacheService
@dataclass
class AuthenticatedUserSnapshot:
user_id: str
email: str | None
username: str
role: UserRole
created_at: datetime | None
@dataclass
class ThreadsafeAPIKeyAuthResult:
user: User
api_key: ApiKey | None = None
balance_remaining: float | None = None
access_allowed: bool = True
access_message: str = "OK"
@property
def access_ok(self) -> bool:
return self.access_allowed
PipelineThreadsafeAuthResult = ThreadsafeAPIKeyAuthResult
# API Key last_used_at 更新节流配置
# 同一个 API Key 在此时间间隔内只会更新一次 last_used_at
_LAST_USED_UPDATE_INTERVAL = 60 # 秒
@@ -177,6 +204,192 @@ class AuthService:
except jwt.InvalidTokenError:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="无效的Token")
@staticmethod
def _authenticate_local_user_sync(
db: Session,
email: str,
password: str,
) -> User | None:
"""同步执行本地认证,供线程池隔离入口复用。"""
from sqlalchemy import or_
user = db.query(User).filter(or_(User.email == email, User.username == email)).first()
if not user:
logger.warning("登录失败 - 用户不存在: {}", email)
return None
if user.is_deleted:
logger.warning("登录失败 - 用户已删除: {}", email)
return None
from src.core.modules.hooks import AUTH_CHECK_EXCLUSIVE_MODE, get_hook_dispatcher
is_exclusive = get_hook_dispatcher().dispatch_sync(AUTH_CHECK_EXCLUSIVE_MODE, db=db)
if is_exclusive:
if user.role != UserRole.ADMIN or user.auth_source != AuthSource.LOCAL:
logger.warning("登录失败 - 排他登录模式下仅管理员可本地登录: {}", email)
return None
logger.warning("[EXCLUSIVE-MODE] 紧急恢复通道:本地管理员登录: {}", email)
if user.auth_source == AuthSource.LDAP:
logger.warning("登录失败 - 该用户使用 LDAP 认证: {}", email)
return None
if not user.verify_password(password):
logger.warning("登录失败 - 密码错误: {}", email)
return None
if not user.is_active:
logger.warning("登录失败 - 用户已禁用: {}", email)
return None
return user
@staticmethod
def _build_authenticated_snapshot(user: User) -> AuthenticatedUserSnapshot:
return AuthenticatedUserSnapshot(
user_id=user.id,
email=user.email,
username=user.username,
role=user.role,
created_at=user.created_at,
)
@staticmethod
def _detach_instance(db: Session, instance: User | ApiKey | None) -> None:
if instance is None:
return
try:
db.expunge(instance)
except Exception as exc:
logger.debug("expunge failed: {}", exc)
@staticmethod
def _load_user_for_token_sync(db: Session, user_id: str) -> User | None:
user = db.query(User).filter(User.id == user_id).first()
if not user or not user.is_active or user.is_deleted:
return None
return user
@staticmethod
async def load_user_for_token_threadsafe(user_id: str) -> User | None:
"""Load the JWT user in a threadpool and return a detached object."""
def _load_in_thread() -> User | None:
thread_db = create_session()
try:
user = AuthService._load_user_for_token_sync(thread_db, user_id)
if not user:
return None
AuthService._detach_instance(thread_db, user)
return user
finally:
thread_db.close()
return await run_in_threadpool(_load_in_thread)
@staticmethod
async def load_user_for_pipeline_threadsafe(
user_id: str,
*,
include_balance: bool = False,
) -> PipelineThreadsafeAuthResult | None:
"""Compatibility helper: load a user in a threadpool and optionally prefetch balance."""
def _load_in_thread() -> PipelineThreadsafeAuthResult | None:
from src.services.wallet import WalletService
thread_db = create_session()
try:
user = AuthService._load_user_for_token_sync(thread_db, user_id)
if not user:
return None
balance_remaining: float | None = None
if include_balance:
balance = WalletService.get_balance_snapshot(thread_db, user=user)
balance_remaining = float(balance) if balance is not None else None
AuthService._detach_instance(thread_db, user)
return PipelineThreadsafeAuthResult(
user=user,
balance_remaining=balance_remaining,
)
finally:
thread_db.close()
return await run_in_threadpool(_load_in_thread)
@staticmethod
async def authenticate_api_key_threadsafe(
api_key: str,
) -> ThreadsafeAPIKeyAuthResult | None:
"""Authenticate API key and check balance in a threadpool."""
def _authenticate_in_thread() -> ThreadsafeAPIKeyAuthResult | None:
from src.services.usage.service import UsageService
thread_db = create_session()
try:
auth_result = AuthService.authenticate_api_key(thread_db, api_key)
if not auth_result:
return None
user, key_record = auth_result
balance_result = UsageService.check_request_balance_details(
thread_db,
user,
api_key=key_record,
)
AuthService._detach_instance(thread_db, user)
AuthService._detach_instance(thread_db, key_record)
return ThreadsafeAPIKeyAuthResult(
user=user,
api_key=key_record,
balance_remaining=balance_result.remaining,
access_allowed=balance_result.allowed,
access_message=balance_result.message,
)
finally:
thread_db.close()
return await run_in_threadpool(_authenticate_in_thread)
@staticmethod
async def authenticate_user_threadsafe(
db: Session, email: str, password: str, auth_type: str = "local"
) -> AuthenticatedUserSnapshot | None:
"""为异步登录路由提供线程池隔离的本地认证入口。"""
if auth_type != "local":
user = await AuthService.authenticate_user(db, email, password, auth_type)
if not user:
return None
return AuthService._build_authenticated_snapshot(user)
def _authenticate_in_thread() -> AuthenticatedUserSnapshot | None:
thread_db = create_session()
try:
user = AuthService._authenticate_local_user_sync(thread_db, email, password)
if not user:
return None
user.last_login_at = datetime.now(timezone.utc)
thread_db.commit()
return AuthService._build_authenticated_snapshot(user)
finally:
thread_db.close()
snapshot = await run_in_threadpool(_authenticate_in_thread)
if not snapshot:
return None
await UserCacheService.invalidate_user_cache(snapshot.user_id, snapshot.email or "")
logger.info("用户登录成功: {} (ID: {})", email, snapshot.user_id)
return snapshot
@staticmethod
async def authenticate_user(
db: Session, email: str, password: str, auth_type: str = "local"
@@ -208,40 +421,8 @@ class AuthService:
# 本地认证
# 登录校验必须读取密码哈希,不能使用不包含 password_hash 的缓存对象
# 支持邮箱或用户名登录
from sqlalchemy import or_
user = db.query(User).filter(or_(User.email == email, User.username == email)).first()
user = AuthService._authenticate_local_user_sync(db, email, password)
if not user:
logger.warning(f"登录失败 - 用户不存在: {email}")
return None
if user.is_deleted:
logger.warning(f"登录失败 - 用户已删除: {email}")
return None
# 检查排他登录模式(如 LDAP exclusive仅允许本地管理员登录紧急恢复通道
from src.core.modules.hooks import AUTH_CHECK_EXCLUSIVE_MODE, get_hook_dispatcher
is_exclusive = get_hook_dispatcher().dispatch_sync(AUTH_CHECK_EXCLUSIVE_MODE, db=db)
if is_exclusive:
if user.role != UserRole.ADMIN or user.auth_source != AuthSource.LOCAL:
logger.warning(f"登录失败 - 排他登录模式下仅管理员可本地登录: {email}")
return None
logger.warning(f"[EXCLUSIVE-MODE] 紧急恢复通道:本地管理员登录: {email}")
# 检查用户认证来源
if user.auth_source == AuthSource.LDAP:
logger.warning(f"登录失败 - 该用户使用 LDAP 认证: {email}")
return None
# 在线程池中执行 bcrypt 密码验证,避免阻塞事件循环
if not await run_in_threadpool(user.verify_password, password):
logger.warning(f"登录失败 - 密码错误: {email}")
return None
if not user.is_active:
logger.warning(f"登录失败 - 用户已禁用: {email}")
return None
# 更新最后登录时间

View File

@@ -42,6 +42,8 @@ class ModelCacheService:
# 缓存 TTL- 使用统一常量
CACHE_TTL = CacheTTL.MODEL
PROVIDER_MAPPING_INDEX_CACHE_KEY = "global_model:resolve_index:provider_model_mappings"
MODEL_MAPPING_RULES_CACHE_KEY = "global_model:resolve_index:model_mappings"
@staticmethod
async def get_model_by_id(db: Session, model_id: str) -> Model | None:
@@ -238,6 +240,9 @@ class ModelCacheService:
if resolve_keys_to_clear:
logger.debug(f"Model resolve 缓存已清除: {resolve_keys_to_clear}")
# provider_model_mappings 更新后,需要重建映射索引缓存。
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
@staticmethod
async def invalidate_global_model_cache(global_model_id: str, name: str | None = None) -> None:
"""清除 GlobalModel 缓存"""
@@ -249,6 +254,8 @@ class ModelCacheService:
# 全量清除 resolve 缓存,确保映射规则变更后不命中旧缓存
try:
await CacheService.delete_pattern("global_model:resolve:*")
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
await CacheService.delete(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
except Exception as e:
logger.error(f"GlobalModel resolve 缓存清除失败,可能导致映射不一致: {e}")
logger.debug(f"GlobalModel 缓存已清除: {global_model_id}")
@@ -262,10 +269,108 @@ class ModelCacheService:
"""
try:
deleted = await CacheService.delete_pattern("global_model:resolve:*")
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
await CacheService.delete(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
logger.debug(f"已清除 {deleted} 个 GlobalModel resolve 缓存")
except Exception as e:
logger.error(f"GlobalModel resolve 缓存清除失败: {e}")
@staticmethod
async def _get_provider_mapping_index(
db: Session,
) -> dict[str, list[dict[str, object]]]:
cached_data = await CacheService.get(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
if isinstance(cached_data, dict):
return {
str(name): value
for name, value in cached_data.items()
if isinstance(name, str) and isinstance(value, list)
}
from src.models.database import Provider
rows = (
db.query(Model, GlobalModel)
.join(Provider, Model.provider_id == Provider.id)
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
.filter(
Provider.is_active == True,
Model.is_active == True,
GlobalModel.is_active == True,
Model.provider_model_mappings.isnot(None),
)
.all()
)
index: dict[str, list[dict[str, object]]] = {}
seen_pairs: set[tuple[str, str]] = set()
for model, global_model in rows:
raw_mappings = getattr(model, "provider_model_mappings", None)
if not isinstance(raw_mappings, list):
continue
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
for raw in raw_mappings:
if not isinstance(raw, dict):
continue
name = raw.get("name")
if not isinstance(name, str):
continue
normalized_name = name.strip()
if not normalized_name:
continue
pair_key = (normalized_name, str(global_model.id))
if pair_key in seen_pairs:
continue
seen_pairs.add(pair_key)
index.setdefault(normalized_name, []).append(global_model_dict)
await CacheService.set(
ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY,
index,
ttl_seconds=ModelCacheService.CACHE_TTL,
)
return index
@staticmethod
async def _get_model_mapping_rules(
db: Session,
) -> list[dict[str, object]]:
cached_data = await CacheService.get(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
if isinstance(cached_data, list):
return [entry for entry in cached_data if isinstance(entry, dict)]
rows = (
db.query(GlobalModel)
.filter(GlobalModel.is_active == True, GlobalModel.config.isnot(None))
.all()
)
rules: list[dict[str, object]] = []
for global_model in rows:
config = getattr(global_model, "config", None) or {}
mappings = config.get("model_mappings")
if not isinstance(mappings, list) or not mappings:
continue
patterns = [
pattern for pattern in mappings if isinstance(pattern, str) and pattern.strip()
]
if not patterns:
continue
rules.append(
{
"global_model": ModelCacheService._global_model_to_dict(global_model),
"patterns": patterns,
}
)
await CacheService.set(
ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY,
rules,
ttl_seconds=ModelCacheService.CACHE_TTL,
)
return rules
@staticmethod
async def resolve_global_model_by_name_or_mapping(
db: Session, model_name: str
@@ -391,41 +496,18 @@ class ModelCacheService:
return result_global_model
# 4. 通过 provider_model_mappings 匹配
models_with_mappings = (
db.query(Model, GlobalModel)
.join(Provider, Model.provider_id == Provider.id)
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
.filter(
Provider.is_active == True,
Model.is_active == True,
GlobalModel.is_active == True,
Model.provider_model_mappings.isnot(None),
)
.all()
)
provider_mapping_index = await ModelCacheService._get_provider_mapping_index(db)
mapping_matched_global_models = [
ModelCacheService._dict_to_global_model(global_model_dict)
for global_model_dict in provider_mapping_index.get(normalized_name, [])
if isinstance(global_model_dict, dict)
]
mapping_matched_global_models: list[GlobalModel] = []
mapping_seen_ids: set[str] = set()
for model, gm in models_with_mappings:
raw_mappings = model.provider_model_mappings
if not isinstance(raw_mappings, list):
continue
for raw in raw_mappings:
if not isinstance(raw, dict):
continue
name = raw.get("name")
if not isinstance(name, str):
continue
if name.strip() != normalized_name:
continue
if gm.id not in mapping_seen_ids:
mapping_seen_ids.add(gm.id)
mapping_matched_global_models.append(gm)
logger.debug(
f"模型名称 '{normalized_name}' 通过 provider_model_mappings 匹配到 "
f"GlobalModel: {gm.name} (Model: {model.id[:8]}...)"
)
break
for gm in mapping_matched_global_models:
logger.debug(
f"模型名称 '{normalized_name}' 通过 provider_model_mappings 匹配到 "
f"GlobalModel: {gm.name}"
)
if mapping_matched_global_models:
resolution_method = "provider_model_mappings"
@@ -453,32 +535,21 @@ class ModelCacheService:
return result_global_model
# 5. 通过 GlobalModel.config.model_mappings 匹配(支持正则)
from sqlalchemy import func
from src.core.model_permissions import match_model_with_pattern
mapping_rows = (
db.query(GlobalModel)
.filter(
GlobalModel.is_active == True,
GlobalModel.config.isnot(None),
GlobalModel.config["model_mappings"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_mappings"]) > 0,
)
.all()
)
mapping_matches: list[GlobalModel] = []
for gm in mapping_rows:
config = gm.config or {}
mappings = config.get("model_mappings")
if not isinstance(mappings, list):
for entry in await ModelCacheService._get_model_mapping_rules(db):
global_model_dict = entry.get("global_model")
patterns = entry.get("patterns")
if not isinstance(global_model_dict, dict) or not isinstance(patterns, list):
continue
for pattern in mappings:
for pattern in patterns:
if isinstance(pattern, str) and match_model_with_pattern(
pattern, normalized_name
):
mapping_matches.append(gm)
mapping_matches.append(
ModelCacheService._dict_to_global_model(global_model_dict)
)
break
if mapping_matches:

View File

@@ -15,6 +15,14 @@ from src.models.database import RequestCandidate
class RequestCandidateService:
"""请求候选记录服务"""
@staticmethod
def _persist_candidate_update(db: Session, *, immediate: bool) -> None:
if immediate:
db.commit()
return
db.flush()
get_batch_committer().mark_dirty(db)
@staticmethod
def create_candidate(
db: Session,
@@ -93,9 +101,9 @@ class RequestCandidateService:
if candidate:
candidate.status = "pending"
candidate.started_at = datetime.now(timezone.utc)
# 关键状态更新:立即提交,不使用批量提交
# 原因:前端需要实时看到请求开始执行
db.commit()
# 中间态改为 flush最终 success/failed 仍会立即提交,
# 但开始执行这一跳不再单独制造一次事务往返。
RequestCandidateService._persist_candidate_update(db, immediate=False)
@staticmethod
def update_candidate_status(db: Session, candidate_id: str, status: str) -> None:
@@ -113,8 +121,9 @@ class RequestCandidateService:
# 如果状态变更为 pending记录开始时间
if status == "pending" and not candidate.started_at:
candidate.started_at = datetime.now(timezone.utc)
# 立即提交,确保前端能实时看到状态变化
db.commit()
RequestCandidateService._persist_candidate_update(
db, immediate=status not in {"pending", "streaming"}
)
@staticmethod
def mark_candidate_streaming(
@@ -141,7 +150,7 @@ class RequestCandidateService:
candidate.status = "streaming"
candidate.concurrent_requests = concurrent_requests
# streaming 状态不设置 finished_at 和 status_code因为请求还在进行中
db.commit()
RequestCandidateService._persist_candidate_update(db, immediate=False)
@staticmethod
def mark_candidate_success(

View File

@@ -62,7 +62,6 @@ from src.services.rate_limit.adaptive_reservation import (
)
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.services.scheduling.affinity_manager import (
CacheAffinityManager,
get_affinity_manager,
)
from src.services.scheduling.candidate_builder import (
@@ -464,12 +463,41 @@ class CacheAwareScheduler:
return [], global_model_id, queried_provider_count
# 1. 查询 Providers委托给 CandidateBuilder
providers = self._candidate_builder._query_providers(
db=db,
provider_offset=provider_offset,
provider_limit=provider_limit,
)
queried_provider_count = len(providers)
providers = []
if allowed_providers is not None:
provider_refs = self._candidate_builder._query_provider_refs(
db=db,
provider_offset=provider_offset,
provider_limit=provider_limit,
)
queried_provider_count = len(provider_refs)
allowed_values = {value for value in allowed_providers if value}
matched_provider_ids = [
provider_id
for provider_id, provider_name in provider_refs
if provider_id in allowed_values or provider_name in allowed_values
]
if queried_provider_count != len(matched_provider_ids):
logger.debug(
"用户/API Key 过滤 Provider 预加载范围: {} -> {}",
queried_provider_count,
len(matched_provider_ids),
)
if matched_provider_ids:
providers = self._candidate_builder._query_providers(
db=db,
provider_ids=matched_provider_ids,
)
else:
providers = self._candidate_builder._query_providers(
db=db,
provider_offset=provider_offset,
provider_limit=provider_limit,
)
queried_provider_count = len(providers)
# Provider query starts a transaction; release connection before entering async candidate build.
release_db_connection_before_await(db)
@@ -480,19 +508,6 @@ class CacheAwareScheduler:
", ".join(p.name for p in providers),
)
if not providers:
return [], global_model_id, queried_provider_count
# 1.5 根据 allowed_providers 过滤(合并 ApiKey 和 User 的限制)
if allowed_providers is not None:
original_count = len(providers)
# 同时支持 provider id 和 name 匹配
providers = [
p for p in providers if p.id in allowed_providers or p.name in allowed_providers
]
if original_count != len(providers):
logger.debug("用户/API Key 过滤 Provider: {} -> {}", original_count, len(providers))
if not providers:
return [], global_model_id, queried_provider_count

View File

@@ -14,6 +14,7 @@ import re
from collections.abc import Sequence
from typing import TYPE_CHECKING
from sqlalchemy import or_
from sqlalchemy.orm import Session, selectinload
from src.core.api_format.conversion.compatibility import is_format_compatible
@@ -52,7 +53,7 @@ from src.services.cache.model_cache import ModelCacheService
def _get_pool_config(provider: Provider) -> "PoolConfig | None":
"""Return parsed PoolConfig if the provider has pool enabled, else None."""
from src.services.provider.pool.config import PoolConfig, parse_pool_config
from src.services.provider.pool.config import parse_pool_config
return parse_pool_config(getattr(provider, "config", None))
@@ -79,11 +80,36 @@ class CandidateBuilder:
def __init__(self, candidate_sorter: CandidateSorterProtocol) -> None:
self._sorter = candidate_sorter
def _query_provider_refs(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
) -> list[tuple[str, str]]:
"""仅查询当前分页内 Provider 的轻量引用信息。"""
provider_query = (
db.query(Provider.id, Provider.name)
.filter(Provider.is_active.is_(True))
.order_by(Provider.provider_priority.asc())
)
if provider_offset:
provider_query = provider_query.offset(provider_offset)
if provider_limit:
provider_query = provider_query.limit(provider_limit)
return [
(str(provider_id), str(provider_name))
for provider_id, provider_name in provider_query.all()
]
def _query_providers(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
allowed_providers: list[str] | None = None,
provider_ids: list[str] | None = None,
) -> list[Provider]:
"""
查询活跃的 Providers带预加载
@@ -132,12 +158,30 @@ class CandidateBuilder:
.order_by(Provider.provider_priority.asc())
)
if provider_offset:
if allowed_providers:
allowed_values = [value for value in allowed_providers if value]
if allowed_values:
provider_query = provider_query.filter(
or_(Provider.id.in_(allowed_values), Provider.name.in_(allowed_values))
)
if provider_ids is not None:
if not provider_ids:
return []
provider_query = provider_query.filter(Provider.id.in_(provider_ids))
if provider_ids is None and provider_offset:
provider_query = provider_query.offset(provider_offset)
if provider_limit:
if provider_ids is None and provider_limit:
provider_query = provider_query.limit(provider_limit)
return provider_query.all()
providers = provider_query.all()
if provider_ids is None:
return providers
order_map = {provider_id: index for index, provider_id in enumerate(provider_ids)}
providers.sort(key=lambda provider: order_map.get(str(provider.id), len(order_map)))
return providers
async def _check_model_support(
self,

View File

@@ -42,11 +42,20 @@ class CandidateSorterProtocol(Protocol):
class CandidateBuilderProtocol(Protocol):
def _query_provider_refs(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
) -> list[tuple[str, str]]: ...
def _query_providers(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
allowed_providers: list[str] | None = None,
provider_ids: list[str] | None = None,
) -> list[Provider]: ...
async def _build_candidates(

View File

@@ -12,7 +12,7 @@ from datetime import date, datetime, time, timedelta, timezone
from decimal import Decimal
from typing import Any
from sqlalchemy import Float, and_, case, cast, func
from sqlalchemy import Date, Float, and_, case, cast, func, text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
@@ -47,9 +47,64 @@ def _get_utc_day_range(value: datetime) -> tuple[datetime, datetime]:
return day_start, day_start + timedelta(days=1)
def _merge_consecutive_utc_days(days: list[date]) -> list[tuple[datetime, datetime]]:
"""将连续 UTC 日期合并为更少的 [start, end) 区间。"""
if not days:
return []
sorted_days = sorted(days)
ranges: list[tuple[datetime, datetime]] = []
start_day = sorted_days[0]
end_day = start_day
for current_day in sorted_days[1:]:
if current_day == end_day + timedelta(days=1):
end_day = current_day
continue
range_start = datetime.combine(start_day, time.min, tzinfo=timezone.utc)
range_end = datetime.combine(end_day + timedelta(days=1), time.min, tzinfo=timezone.utc)
ranges.append((range_start, range_end))
start_day = current_day
end_day = current_day
range_start = datetime.combine(start_day, time.min, tzinfo=timezone.utc)
range_end = datetime.combine(end_day + timedelta(days=1), time.min, tzinfo=timezone.utc)
ranges.append((range_start, range_end))
return ranges
class StatsAggregatorService:
"""统计数据聚合服务"""
@staticmethod
def _resolve_percentile_row(row: Any | None) -> tuple[int | None, int | None, int | None]:
if not row:
return None, None, None
count = int(getattr(row, "count", 0) or 0)
if count < MIN_PERCENTILE_SAMPLES:
return None, None, None
p50 = getattr(row, "p50", None)
p90 = getattr(row, "p90", None)
p99 = getattr(row, "p99", None)
return (
int(p50) if p50 is not None else None,
int(p90) if p90 is not None else None,
int(p99) if p99 is not None else None,
)
@staticmethod
def _build_local_day_expression(time_range: TimeRangeParams) -> Any:
if time_range.timezone and time_range.timezone != "UTC":
local_time_expr = func.timezone(time_range.timezone, Usage.created_at)
elif time_range.tz_offset_minutes:
interval_literal = text(f"INTERVAL '{int(time_range.tz_offset_minutes)} minutes'")
local_time_expr = Usage.created_at + interval_literal
else:
local_time_expr = Usage.created_at
return cast(func.date_trunc("day", local_time_expr), Date)
@staticmethod
def compute_daily_stats(db: Session, date: datetime) -> dict:
"""计算指定 UTC 日期的统计数据(不写入数据库)"""
@@ -196,23 +251,8 @@ class StatsAggregatorService:
.first()
)
def _resolve(row: Any | None) -> tuple[int | None, int | None, int | None]:
if not row:
return None, None, None
count = int(getattr(row, "count", 0) or 0)
if count < MIN_PERCENTILE_SAMPLES:
return None, None, None
p50 = getattr(row, "p50", None)
p90 = getattr(row, "p90", None)
p99 = getattr(row, "p99", None)
return (
int(p50) if p50 is not None else None,
int(p90) if p90 is not None else None,
int(p99) if p99 is not None else None,
)
p50_rt, p90_rt, p99_rt = _resolve(rt_row)
p50_ttfb, p90_ttfb, p99_ttfb = _resolve(ttfb_row)
p50_rt, p90_rt, p99_rt = StatsAggregatorService._resolve_percentile_row(rt_row)
p50_ttfb, p90_ttfb, p99_ttfb = StatsAggregatorService._resolve_percentile_row(ttfb_row)
return {
"p50_response_time_ms": p50_rt,
@@ -223,6 +263,103 @@ class StatsAggregatorService:
"p99_first_byte_time_ms": p99_ttfb,
}
@staticmethod
def compute_percentiles_by_local_day(
db: Session, time_range: TimeRangeParams
) -> list[dict[str, int | None | str]]:
"""按本地日期批量计算性能百分位,避免逐天 fan-out。"""
bind = db.bind
dialect = bind.dialect.name if bind is not None else "sqlite"
local_dates: list[date] = []
current_date = time_range.start_date
while current_date <= time_range.end_date:
local_dates.append(current_date)
current_date += timedelta(days=1)
if dialect != "postgresql":
return [
{
"date": local_date.isoformat(),
"p50_response_time_ms": None,
"p90_response_time_ms": None,
"p99_response_time_ms": None,
"p50_first_byte_time_ms": None,
"p90_first_byte_time_ms": None,
"p99_first_byte_time_ms": None,
}
for local_date in local_dates
]
start_utc, end_utc = time_range.to_utc_datetime_range()
local_day_expr = StatsAggregatorService._build_local_day_expression(time_range)
rt_rows = (
db.query(
local_day_expr.label("local_day"),
func.percentile_cont(0.5).within_group(Usage.response_time_ms).label("p50"),
func.percentile_cont(0.9).within_group(Usage.response_time_ms).label("p90"),
func.percentile_cont(0.99).within_group(Usage.response_time_ms).label("p99"),
func.count().label("count"),
)
.filter(
Usage.created_at >= start_utc,
Usage.created_at < end_utc,
Usage.status == "completed",
Usage.response_time_ms.isnot(None),
)
.group_by(local_day_expr)
.all()
)
ttfb_rows = (
db.query(
local_day_expr.label("local_day"),
func.percentile_cont(0.5).within_group(Usage.first_byte_time_ms).label("p50"),
func.percentile_cont(0.9).within_group(Usage.first_byte_time_ms).label("p90"),
func.percentile_cont(0.99).within_group(Usage.first_byte_time_ms).label("p99"),
func.count().label("count"),
)
.filter(
Usage.created_at >= start_utc,
Usage.created_at < end_utc,
Usage.status == "completed",
Usage.first_byte_time_ms.isnot(None),
)
.group_by(local_day_expr)
.all()
)
rt_by_day: dict[date, tuple[int | None, int | None, int | None]] = {}
for row in rt_rows:
local_day = getattr(row, "local_day", None)
if local_day is not None:
rt_by_day[local_day] = StatsAggregatorService._resolve_percentile_row(row)
ttfb_by_day: dict[date, tuple[int | None, int | None, int | None]] = {}
for row in ttfb_rows:
local_day = getattr(row, "local_day", None)
if local_day is not None:
ttfb_by_day[local_day] = StatsAggregatorService._resolve_percentile_row(row)
result: list[dict[str, int | None | str]] = []
for local_date in local_dates:
p50_rt, p90_rt, p99_rt = rt_by_day.get(local_date, (None, None, None))
p50_ttfb, p90_ttfb, p99_ttfb = ttfb_by_day.get(local_date, (None, None, None))
result.append(
{
"date": local_date.isoformat(),
"p50_response_time_ms": p50_rt,
"p90_response_time_ms": p90_rt,
"p99_response_time_ms": p99_rt,
"p50_first_byte_time_ms": p50_ttfb,
"p90_first_byte_time_ms": p90_ttfb,
"p99_first_byte_time_ms": p99_ttfb,
}
)
return result
@staticmethod
def aggregate_daily_stats(db: Session, date: datetime, commit: bool = True) -> StatsDaily:
"""聚合指定 UTC 日期的统计数据
@@ -656,6 +793,82 @@ class StatsAggregatorService:
db.commit()
return stats
@staticmethod
def aggregate_user_daily_stats_batch(
db: Session, date: datetime, user_ids: list[str], commit: bool = True
) -> list[StatsUserDaily]:
"""批量聚合单日用户统计,避免逐用户 fan-out 查询。"""
if not user_ids:
return []
day_start, day_end = _get_utc_day_range(date)
ordered_user_ids = list(dict.fromkeys(user_ids))
existing_rows = (
db.query(StatsUserDaily)
.filter(
and_(
StatsUserDaily.date == day_start,
StatsUserDaily.user_id.in_(ordered_user_ids),
)
)
.all()
)
existing_by_user = {row.user_id: row for row in existing_rows}
error_cond = (Usage.status_code >= 400) | (Usage.error_message.isnot(None))
aggregated_rows = (
db.query(
Usage.user_id.label("user_id"),
func.count(Usage.id).label("total_requests"),
func.sum(case((error_cond, 1), else_=0)).label("error_requests"),
func.sum(Usage.input_tokens).label("input_tokens"),
func.sum(Usage.output_tokens).label("output_tokens"),
func.sum(Usage.cache_creation_input_tokens).label("cache_creation_tokens"),
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
func.sum(Usage.total_cost_usd).label("total_cost"),
func.max(Usage.username).label("username"),
)
.filter(
and_(
Usage.user_id.in_(ordered_user_ids),
Usage.created_at >= day_start,
Usage.created_at < day_end,
)
)
.group_by(Usage.user_id)
.all()
)
aggregated_by_user = {row.user_id: row for row in aggregated_rows}
result: list[StatsUserDaily] = []
for user_id in ordered_user_ids:
stats = existing_by_user.get(user_id)
if stats is None:
stats = StatsUserDaily(id=str(uuid.uuid4()), user_id=user_id, date=day_start)
db.add(stats)
aggregated = aggregated_by_user.get(user_id)
if not stats.username and aggregated is not None:
username = getattr(aggregated, "username", None)
if username:
stats.username = username
total_requests = int(getattr(aggregated, "total_requests", 0) or 0)
error_requests = int(getattr(aggregated, "error_requests", 0) or 0)
stats.total_requests = total_requests
stats.success_requests = total_requests - error_requests
stats.error_requests = error_requests
stats.input_tokens = int(getattr(aggregated, "input_tokens", 0) or 0)
stats.output_tokens = int(getattr(aggregated, "output_tokens", 0) or 0)
stats.cache_creation_tokens = int(getattr(aggregated, "cache_creation_tokens", 0) or 0)
stats.cache_read_tokens = int(getattr(aggregated, "cache_read_tokens", 0) or 0)
stats.total_cost = float(getattr(aggregated, "total_cost", 0) or 0.0)
result.append(stats)
if commit:
db.commit()
return result
@staticmethod
def aggregate_daily_stats_bundle(
db: Session, date: datetime, user_ids: list[str] | None = None
@@ -668,8 +881,9 @@ class StatsAggregatorService:
StatsAggregatorService.aggregate_daily_error_stats(db, date, commit=False)
if user_ids:
for user_id in user_ids:
StatsAggregatorService.aggregate_user_daily_stats(db, user_id, date, commit=False)
StatsAggregatorService.aggregate_user_daily_stats_batch(
db, date, user_ids, commit=False
)
stats.is_complete = True
stats.aggregated_at = datetime.now(timezone.utc)
@@ -1317,28 +1531,32 @@ def query_stats_hybrid(
result = AggregatedStats()
preaggregate_dates: list[datetime] = []
realtime_dates: list[datetime] = []
for day in complete_dates:
if day >= today_utc:
realtime_dates.append(datetime.combine(day, time.min, tzinfo=timezone.utc))
continue
day_dt = datetime.combine(day, time.min, tzinfo=timezone.utc)
stats = (
db.query(StatsDaily)
.filter(StatsDaily.date == day_dt, StatsDaily.is_complete.is_(True))
.first()
)
if stats:
preaggregate_dates.append(day_dt)
else:
realtime_dates.append(day_dt)
historical_dates = [day for day in complete_dates if day < today_utc]
realtime_dates = [day for day in complete_dates if day >= today_utc]
# Pre-aggregated totals
for day_dt in preaggregate_dates:
stats = db.query(StatsDaily).filter(StatsDaily.date == day_dt).first()
if not stats:
continue
preaggregated_by_date: dict[date, StatsDaily] = {}
if historical_dates:
historical_start = datetime.combine(min(historical_dates), time.min, tzinfo=timezone.utc)
historical_end = datetime.combine(
max(historical_dates) + timedelta(days=1),
time.min,
tzinfo=timezone.utc,
)
historical_rows = (
db.query(StatsDaily)
.filter(
StatsDaily.date >= historical_start,
StatsDaily.date < historical_end,
StatsDaily.is_complete.is_(True),
)
.all()
)
preaggregated_by_date = {
row.date.astimezone(timezone.utc).date() if row.date.tzinfo else row.date.date(): row
for row in historical_rows
}
for stats in preaggregated_by_date.values():
result.total_requests += stats.total_requests
result.success_requests += stats.success_requests
result.error_requests += stats.error_requests
@@ -1352,9 +1570,10 @@ def query_stats_hybrid(
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
for day_dt in realtime_dates:
result.add(aggregate_usage_range(db, day_dt, day_dt + timedelta(days=1), filters=filters))
missing_historical_dates = [day for day in historical_dates if day not in preaggregated_by_date]
realtime_ranges = _merge_consecutive_utc_days(missing_historical_dates + realtime_dates)
for range_start, range_end in realtime_ranges:
result.add(aggregate_usage_range(db, range_start, range_end, filters=filters))
if head_boundary:
result.add(aggregate_usage_range(db, head_boundary[0], head_boundary[1], filters=filters))

View File

@@ -247,13 +247,14 @@ class UsageActiveRequestsMixin:
default_timeout_seconds: int = 300,
*,
include_admin_fields: bool = False,
maintain_status: bool | None = None,
) -> list[dict[str, Any]]:
"""
获取活跃请求状态(用于前端轮询),并自动清理超时的 pending/streaming 请求
获取活跃请求状态(用于前端轮询)
与 get_active_requests 不同,此方法:
1. 返回轻量级的状态字典而非完整 Usage 对象
2. 自动检测并清理超时的 pending/streaming 请求
2. 可选地检测并清理超时的 pending/streaming 请求
3. 支持按 ID 列表查询特定请求
Args:
@@ -261,6 +262,7 @@ class UsageActiveRequestsMixin:
ids: 指定要查询的请求 ID 列表(可选)
user_id: 限制只查询该用户的请求(可选,用于普通用户接口)
default_timeout_seconds: 默认超时时间(秒),当端点未配置时使用
maintain_status: 是否执行超时修复与状态回写;默认仅在全量活跃请求查询时执行
Returns:
请求状态列表
@@ -311,28 +313,26 @@ class UsageActiveRequestsMixin:
query = query.order_by(Usage.created_at.desc()).limit(50)
records = query.all()
should_maintain_status = maintain_status if maintain_status is not None else not ids
# 检查超时的 pending/streaming 请求
# 收集可能超时的 usage_id 列表
timeout_candidates: list[str] = []
for r in records:
if r.status in ("pending", "streaming") and r.created_at:
# 使用全局配置的超时时间
timeout_seconds = default_timeout_seconds
if should_maintain_status:
for r in records:
if r.status in ("pending", "streaming") and r.created_at:
timeout_seconds = default_timeout_seconds
# 处理时区:如果 created_at 没有时区信息,假定为 UTC
created_at = r.created_at
if created_at.tzinfo is None:
created_at = created_at.replace(tzinfo=timezone.utc)
elapsed = (now - created_at).total_seconds()
if elapsed > timeout_seconds:
# 需要获取 request_id 以便检查 RequestCandidate 表
# r.id 是 usage_id需要查询 request_id
timeout_candidates.append(r.id)
created_at = r.created_at
if created_at.tzinfo is None:
created_at = created_at.replace(tzinfo=timezone.utc)
elapsed = (now - created_at).total_seconds()
if elapsed > timeout_seconds:
timeout_candidates.append(r.id)
# 批量更新超时的请求(排除已有成功完成记录的请求)
timeout_ids = []
if timeout_candidates:
if should_maintain_status and timeout_candidates:
# 先获取这些 Usage 的 request_id
usage_request_ids = (
db.query(Usage.id, Usage.request_id).filter(Usage.id.in_(timeout_candidates)).all()
@@ -374,7 +374,8 @@ class UsageActiveRequestsMixin:
cls._sync_candidate_status_to_success(db, completed_request_ids)
db.commit()
logger.info(
f"[Usage] 恢复 {len(completed_usage_ids)} 个已完成请求的状态(遥测回调丢失)"
"[Usage] 恢复 {} 个已完成请求的状态(遥测回调丢失)",
len(completed_usage_ids),
)
result: list[dict[str, Any]] = []

View File

@@ -1,5 +1,6 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any
@@ -10,6 +11,13 @@ from src.core.logger import logger
from src.models.database import ApiKey, Usage, User
@dataclass(slots=True)
class RequestBalanceCheckResult:
allowed: bool
message: str
remaining: float | None
class UsageQueryMixin:
"""查询/统计相关方法"""
@@ -125,14 +133,14 @@ class UsageQueryMixin:
return result
@staticmethod
def check_request_balance(
def check_request_balance_details(
db: Session,
user: User,
estimated_tokens: int = 0,
estimated_cost: float = 0,
api_key: ApiKey | None = None,
) -> tuple[bool, str]:
"""检查请求是否满足余额条件(支持独立 Key"""
) -> RequestBalanceCheckResult:
"""Return a structured balance-check result."""
from src.services.wallet import WalletService
wallet_access = WalletService.check_request_allowed(
@@ -140,29 +148,52 @@ class UsageQueryMixin:
user=None if (api_key and api_key.is_standalone) else user,
api_key=api_key,
)
snapshot = wallet_access.balance_snapshot
if snapshot is None:
snapshot = wallet_access.remaining
remaining = float(snapshot) if snapshot is not None else None
if wallet_access.allowed:
return True, "OK"
return RequestBalanceCheckResult(True, "OK", remaining)
if wallet_access.message == "钱包欠费,请先充值":
if wallet_access.message in {"钱包欠费,请先充值", "账户欠费,请先充值"}:
if api_key and api_key.is_standalone:
return False, "Key欠费请先调账或充值"
return False, "账户欠费,请先充值"
return RequestBalanceCheckResult(False, "Key欠费请先调账或充值", remaining)
return RequestBalanceCheckResult(False, "账户欠费,请先充值", remaining)
if wallet_access.message == "钱包不可用":
if api_key and api_key.is_standalone:
return False, "Key钱包不可用"
return False, "钱包不可用"
return RequestBalanceCheckResult(False, "Key钱包不可用", remaining)
return RequestBalanceCheckResult(False, "钱包不可用", remaining)
remaining = float(wallet_access.remaining) if wallet_access.remaining is not None else None
if api_key and api_key.is_standalone:
if remaining is None:
return False, "Key余额不足"
return False, f"Key余额不足剩余: ${remaining:.2f}"
return RequestBalanceCheckResult(False, "Key余额不足", remaining)
return RequestBalanceCheckResult(
False, f"Key余额不足剩余: ${remaining:.2f}", remaining
)
# admin 已在 WalletService.check_request_allowed 中放行,此处不再重复检查
# Admin users are already allowed in WalletService.check_request_allowed.
if remaining is None:
return False, wallet_access.message or "余额不足"
return False, f"余额不足(剩余: ${remaining:.2f}"
return RequestBalanceCheckResult(False, wallet_access.message or "余额不足", remaining)
return RequestBalanceCheckResult(False, f"余额不足(剩余: ${remaining:.2f}", remaining)
@staticmethod
def check_request_balance(
db: Session,
user: User,
estimated_tokens: int = 0,
estimated_cost: float = 0,
api_key: ApiKey | None = None,
) -> tuple[bool, str]:
"""Check whether the request passes balance rules."""
result = UsageQueryMixin.check_request_balance_details(
db,
user,
estimated_tokens=estimated_tokens,
estimated_cost=estimated_cost,
api_key=api_key,
)
return result.allowed, result.message
@staticmethod
def get_usage_summary(
@@ -171,7 +202,7 @@ class UsageQueryMixin:
api_key_id: str | None = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
group_by: str = "day", # day, week, month
group_by: str | None = "day", # day, week, month, None(不按时间分桶)
) -> list[dict[str, Any]]:
"""获取使用汇总"""
@@ -188,34 +219,35 @@ class UsageQueryMixin:
if end_date:
query = query.filter(Usage.created_at < end_date)
# 使用跨数据库可用的日期函数
from src.utils.database_helpers import date_trunc_portable
select_columns = [Usage.provider_name, Usage.model]
group_columns = [Usage.provider_name, Usage.model]
# 检测数据库方言
bind = db.bind
dialect = bind.dialect.name if bind is not None else "sqlite"
if group_by is not None:
from src.utils.database_helpers import date_trunc_portable
# 根据分组类型选择日期函数(适配多种数据库)
if group_by == "day":
date_func = date_trunc_portable(dialect, "day", Usage.created_at)
elif group_by == "week":
date_func = date_trunc_portable(dialect, "week", Usage.created_at)
elif group_by == "month":
date_func = date_trunc_portable(dialect, "month", Usage.created_at)
else:
# 默认按天分组
date_func = date_trunc_portable(dialect, "day", Usage.created_at)
bind = db.bind
dialect = bind.dialect.name if bind is not None else "sqlite"
if group_by == "day":
date_func = date_trunc_portable(dialect, "day", Usage.created_at)
elif group_by == "week":
date_func = date_trunc_portable(dialect, "week", Usage.created_at)
elif group_by == "month":
date_func = date_trunc_portable(dialect, "month", Usage.created_at)
else:
date_func = date_trunc_portable(dialect, "day", Usage.created_at)
select_columns.insert(0, date_func.label("period"))
group_columns.insert(0, date_func)
# 汇总查询
summary = db.query(
date_func.label("period"),
Usage.provider_name,
Usage.model,
*select_columns,
func.count(Usage.id).label("requests"),
func.sum(Usage.input_tokens).label("input_tokens"),
func.sum(Usage.output_tokens).label("output_tokens"),
func.sum(Usage.total_tokens).label("total_tokens"),
func.sum(Usage.total_cost_usd).label("total_cost_usd"),
func.sum(Usage.actual_total_cost_usd).label("actual_total_cost_usd"),
func.sum(case((Usage.status_code == 200, 1), else_=0)).label("success_count"),
func.avg(Usage.response_time_ms).label("avg_response_time"),
func.sum(
case(
@@ -246,18 +278,20 @@ class UsageQueryMixin:
if end_date:
summary = summary.filter(Usage.created_at < end_date)
summary = summary.group_by(date_func, Usage.provider_name, Usage.model).all()
summary = summary.group_by(*group_columns).all()
return [
{
"period": row.period,
"period": getattr(row, "period", None),
"provider": row.provider_name,
"model": row.model,
"requests": row.requests,
"input_tokens": row.input_tokens,
"output_tokens": row.output_tokens,
"total_tokens": row.total_tokens,
"total_cost_usd": float(row.total_cost_usd),
"total_cost_usd": float(row.total_cost_usd or 0.0),
"actual_total_cost_usd": float(row.actual_total_cost_usd or 0.0),
"success_count": int(row.success_count or 0),
"avg_response_time_ms": (
float(row.avg_response_time) if row.avg_response_time else 0
),

View File

@@ -43,6 +43,7 @@ class WalletAccessResult:
remaining: Decimal | None
message: str
wallet: Wallet | None = None
balance_snapshot: Decimal | None = None
class WalletService:
@@ -229,6 +230,17 @@ class WalletService:
return db.query(Wallet).filter(Wallet.user_id == user_id).first()
return None
@classmethod
def get_wallets_by_user_ids(
cls,
db: Session,
user_ids: list[str],
) -> dict[str, Wallet]:
if not user_ids:
return {}
wallets = db.query(Wallet).filter(Wallet.user_id.in_(user_ids)).all()
return {wallet.user_id: wallet for wallet in wallets if wallet.user_id is not None}
@classmethod
def get_or_create_wallet(
cls,
@@ -291,6 +303,18 @@ class WalletService:
return wallet
raise
@classmethod
def _get_balance_snapshot_from_wallet(cls, wallet: Wallet | None) -> Decimal | None:
if wallet is None:
return None
recharge_balance = cls.get_recharge_balance_value(wallet)
if recharge_balance < Decimal("0"):
return recharge_balance
if cls.is_unlimited_wallet(wallet):
return None
return cls.get_spendable_balance_value(wallet)
@classmethod
def check_request_allowed(
cls,
@@ -299,25 +323,29 @@ class WalletService:
user: User | None,
api_key: ApiKey | None = None,
) -> WalletAccessResult:
if user and user.role == UserRole.ADMIN:
return WalletAccessResult(True, None, "OK", None)
wallet = cls.get_or_create_wallet(db, user=user, api_key=api_key)
balance_snapshot = cls._get_balance_snapshot_from_wallet(wallet)
if user and user.role == UserRole.ADMIN:
return WalletAccessResult(True, None, "OK", wallet, balance_snapshot)
if wallet is None:
return WalletAccessResult(False, Decimal("0"), "钱包不存在", None)
return WalletAccessResult(False, Decimal("0"), "钱包不存在", None, None)
remaining = cls.get_spendable_balance_value(wallet)
recharge_balance = cls.get_recharge_balance_value(wallet)
if wallet.status != "active":
return WalletAccessResult(False, remaining, "钱包不可用", wallet)
# 充值余额为负视为欠费,禁止继续消费(即使总可用余额仍为正)。
return WalletAccessResult(False, remaining, "钱包不可用", wallet, balance_snapshot)
# Negative recharge balance means overdue; block further spending.
if recharge_balance < Decimal("0"):
return WalletAccessResult(False, recharge_balance, "钱包欠费,请先充值", wallet)
return WalletAccessResult(
False, recharge_balance, "钱包欠费,请先充值", wallet, balance_snapshot
)
if cls.is_unlimited_wallet(wallet):
return WalletAccessResult(True, None, "OK", wallet)
return WalletAccessResult(True, None, "OK", wallet, balance_snapshot)
if remaining <= Decimal("0"):
return WalletAccessResult(False, remaining, "钱包余额不足", wallet)
return WalletAccessResult(True, remaining, "OK", wallet)
return WalletAccessResult(False, remaining, "钱包余额不足", wallet, balance_snapshot)
return WalletAccessResult(True, remaining, "OK", wallet, balance_snapshot)
@classmethod
def get_balance_snapshot(
@@ -328,14 +356,7 @@ class WalletService:
api_key: ApiKey | None = None,
) -> Decimal | None:
wallet = cls.get_or_create_wallet(db, user=user, api_key=api_key)
if wallet is None:
return None
recharge_balance = cls.get_recharge_balance_value(wallet)
if recharge_balance < Decimal("0"):
return recharge_balance
if cls.is_unlimited_wallet(wallet):
return None
return cls.get_spendable_balance_value(wallet)
return cls._get_balance_snapshot_from_wallet(wallet)
@classmethod
def _resolve_wallet_for_usage(