mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
# 更新最后登录时间
|
||||
|
||||
177
src/services/cache/model_cache.py
vendored
177
src/services/cache/model_cache.py
vendored
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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]] = []
|
||||
|
||||
@@ -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
|
||||
),
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user