mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 09:50:21 +08:00
perf: 优化请求鉴权链路并批量化统计/调度查询
- 为 Pipeline/Context 增加按需读取请求体能力,支持 async 懒加载 JSON body - 为 chat/cli/video/claude/openai-cli 适配器关闭默认预读,减少无效 body 读取与超时风险 - 将本地登录、JWT 用户加载、API Key 鉴权迁移到线程池隔离会话执行,避免阻塞事件循环 - 为 API Key 鉴权返回结构化余额结果,并在主请求会话中重新绑定 user/api_key 后再校验状态、过期和锁定信息 - 为 management/user token 前缀认证引入独立会话与结果回绑,避免跨会话对象写入失效 - 为 Usage 余额检查补充结构化返回,统一透出 remaining 与欠费/不可用文案映射 - 为用户与管理端活跃请求查询增加 maintain_status 开关,避免轮询指定 id 时误触发状态修复 - 重写 user_me usage 汇总逻辑,支持 group_by=None 的粗粒度聚合 - 修正 provider 维度成功率与平均响应时间统计,基于 success_count 和成功响应耗时汇总计算 - 前端 Usage 轮询由 setInterval 改为串行 setTimeout,避免并发轮询叠加 - 为 StatsAggregator 增加按本地日期批量计算百分位能力,替代逐天 fan-out 查询 - 为混合统计查询合并连续实时日期区间,并批量读取 StatsDaily,减少逐日查询次数 - 为用户日统计增加批量聚合入口,替代逐用户循环聚合 - 为系统配置导出改用 selectinload 预加载 provider 关联数据,减少 N+1 查询 - 为管理员用户列表增加钱包批量查询,避免逐用户回表 - 为调度器增加 provider 轻量引用预过滤,先按 allowed_providers 缩小范围再加载完整 provider 图 - 为 CandidateBuilder 增加 provider refs/provider_ids 查询能力,保留分页顺序 - 为模型缓存增加 provider_model_mappings 索引缓存与 model_mappings 规则缓存,减少重复全量扫描 - 为请求候选中间态改为 flush/batch commit,降低 pending/streaming 状态切换的事务往返 - 为钱包访问结果补充 balance_snapshot,并抽取余额快照复用逻辑 - 补充 pipeline、auth、admin users、user_me usage、stats aggregator、model cache、 scheduler、wallet、request candidate 等回归与契约测试
This commit is contained in:
@@ -245,7 +245,7 @@ const activeRequestIds = computed(() => {
|
|||||||
const hasActiveRequests = computed(() => activeRequestIds.value.length > 0)
|
const hasActiveRequests = computed(() => activeRequestIds.value.length > 0)
|
||||||
|
|
||||||
// 自动刷新定时器
|
// 自动刷新定时器
|
||||||
let autoRefreshTimer: ReturnType<typeof setInterval> | null = null
|
let autoRefreshTimer: ReturnType<typeof setTimeout> | null = null
|
||||||
let globalAutoRefreshTimer: ReturnType<typeof setInterval> | null = null
|
let globalAutoRefreshTimer: ReturnType<typeof setInterval> | null = null
|
||||||
let refreshInFlight: Promise<void> | null = null
|
let refreshInFlight: Promise<void> | null = null
|
||||||
const AUTO_REFRESH_INTERVAL = 1000 // 1秒刷新一次(用于活跃请求)
|
const AUTO_REFRESH_INTERVAL = 1000 // 1秒刷新一次(用于活跃请求)
|
||||||
@@ -271,8 +271,10 @@ async function pollActiveRequests() {
|
|||||||
|
|
||||||
let shouldRefresh = false
|
let shouldRefresh = false
|
||||||
|
|
||||||
|
const recordMap = new Map(currentRecords.value.map(record => [record.id, record]))
|
||||||
|
|
||||||
for (const update of requests) {
|
for (const update of requests) {
|
||||||
const record = currentRecords.value.find(r => r.id === update.id)
|
const record = recordMap.get(update.id)
|
||||||
if (!record) {
|
if (!record) {
|
||||||
// 后端返回了未知的活跃请求,触发刷新以获取完整数据
|
// 后端返回了未知的活跃请求,触发刷新以获取完整数据
|
||||||
shouldRefresh = true
|
shouldRefresh = true
|
||||||
@@ -339,17 +341,26 @@ async function pollActiveRequests() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function scheduleNextAutoRefresh() {
|
||||||
|
if (autoRefreshTimer) return
|
||||||
|
if (!isPageVisible.value || !hasActiveRequests.value) return
|
||||||
|
autoRefreshTimer = setTimeout(async () => {
|
||||||
|
autoRefreshTimer = null
|
||||||
|
await pollActiveRequests()
|
||||||
|
scheduleNextAutoRefresh()
|
||||||
|
}, AUTO_REFRESH_INTERVAL)
|
||||||
|
}
|
||||||
|
|
||||||
// 启动自动刷新
|
// 启动自动刷新
|
||||||
function startAutoRefresh() {
|
function startAutoRefresh() {
|
||||||
if (!isPageVisible.value) return
|
if (!isPageVisible.value) return
|
||||||
if (autoRefreshTimer) return
|
scheduleNextAutoRefresh()
|
||||||
autoRefreshTimer = setInterval(pollActiveRequests, AUTO_REFRESH_INTERVAL)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 停止自动刷新
|
// 停止自动刷新
|
||||||
function stopAutoRefresh() {
|
function stopAutoRefresh() {
|
||||||
if (autoRefreshTimer) {
|
if (autoRefreshTimer) {
|
||||||
clearInterval(autoRefreshTimer)
|
clearTimeout(autoRefreshTimer)
|
||||||
autoRefreshTimer = null
|
autoRefreshTimer = null
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -73,13 +73,7 @@ class AdminPercentilesAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
result = []
|
return StatsAggregatorService.compute_percentiles_by_local_day(context.db, time_range)
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
@router.get("/performance/percentiles")
|
@router.get("/performance/percentiles")
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ from typing import Any
|
|||||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
from sqlalchemy import case, func
|
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.admin_adapter import AdminApiAdapter
|
||||||
from src.api.base.context import ApiRequestContext
|
from src.api.base.context import ApiRequestContext
|
||||||
@@ -975,9 +975,6 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
|||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
from src.models.database import (
|
from src.models.database import (
|
||||||
GlobalModel,
|
GlobalModel,
|
||||||
Model,
|
|
||||||
ProviderAPIKey,
|
|
||||||
ProviderEndpoint,
|
|
||||||
ProxyNode,
|
ProxyNode,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -991,22 +988,35 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
|||||||
gm_name_map: dict[str, str] = {gm.id: gm.name for gm in global_models}
|
gm_name_map: dict[str, str] = {gm.id: gm.name for gm in global_models}
|
||||||
|
|
||||||
# 导出 Providers 及其关联数据
|
# 导出 Providers 及其关联数据
|
||||||
providers = db.query(Provider).all()
|
providers = (
|
||||||
|
db.query(Provider)
|
||||||
|
.options(
|
||||||
|
selectinload(Provider.endpoints),
|
||||||
|
selectinload(Provider.api_keys),
|
||||||
|
selectinload(Provider.models),
|
||||||
|
)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
providers_data = []
|
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:
|
for provider in providers:
|
||||||
# 导出 Endpoints
|
# 导出 Endpoints
|
||||||
endpoints = (
|
endpoints = list(provider.endpoints)
|
||||||
db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
|
|
||||||
)
|
|
||||||
endpoints_data = [ep.to_export_dict() for ep in endpoints]
|
endpoints_data = [ep.to_export_dict() for ep in endpoints]
|
||||||
provider_endpoint_formats = self._collect_provider_endpoint_formats(endpoints)
|
provider_endpoint_formats = self._collect_provider_endpoint_formats(endpoints)
|
||||||
|
|
||||||
# 导出 Provider Keys(按 provider_id 归属,包含 api_formats)
|
# 导出 Provider Keys(按 provider_id 归属,包含 api_formats)
|
||||||
keys = (
|
keys = sorted(
|
||||||
db.query(ProviderAPIKey)
|
provider.api_keys,
|
||||||
.filter(ProviderAPIKey.provider_id == provider.id)
|
key=lambda key: (
|
||||||
.order_by(ProviderAPIKey.internal_priority.asc(), ProviderAPIKey.created_at.asc())
|
key.internal_priority if key.internal_priority is not None else 0,
|
||||||
.all()
|
_normalize_created_at_for_sort(key.created_at),
|
||||||
|
),
|
||||||
)
|
)
|
||||||
keys_data = []
|
keys_data = []
|
||||||
for key in keys:
|
for key in keys:
|
||||||
@@ -1046,7 +1056,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
|||||||
# 导出 Provider Models
|
# 导出 Provider Models
|
||||||
# 注意:提供商模型(Model)必须关联全局模型(GlobalModel)才能参与路由
|
# 注意:提供商模型(Model)必须关联全局模型(GlobalModel)才能参与路由
|
||||||
# 导入时未关联 GlobalModel 的模型会被跳过,这是业务规则而非 bug
|
# 导入时未关联 GlobalModel 的模型会被跳过,这是业务规则而非 bug
|
||||||
models = db.query(Model).filter(Model.provider_id == provider.id).all()
|
models = list(provider.models)
|
||||||
models_data = []
|
models_data = []
|
||||||
for model in models:
|
for model in models:
|
||||||
model_data = model.to_export_dict()
|
model_data = model.to_export_dict()
|
||||||
|
|||||||
@@ -1227,7 +1227,10 @@ class AdminActiveRequestsAdapter(AdminApiAdapter):
|
|||||||
return {"requests": []}
|
return {"requests": []}
|
||||||
|
|
||||||
requests = UsageService.get_active_requests_status(
|
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}
|
return {"requests": requests}
|
||||||
|
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ from src.core.logger import logger
|
|||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
from src.models.admin_requests import UpdateUserRequest
|
from src.models.admin_requests import UpdateUserRequest
|
||||||
from src.models.api import CreateApiKeyRequest, CreateUserRequest
|
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.system.config import SystemConfigService
|
||||||
from src.services.user.apikey import ApiKeyService
|
from src.services.user.apikey import ApiKeyService
|
||||||
from src.services.user.bulk_cleanup import pre_clean_api_key
|
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()
|
pipeline = ApiRequestPipeline()
|
||||||
|
|
||||||
|
|
||||||
def _serialize_user(db: Session, user: User) -> dict[str, Any]:
|
class _WalletSentinelType:
|
||||||
wallet = WalletService.get_wallet(db, user_id=user.id)
|
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 {
|
return {
|
||||||
"id": user.id,
|
"id": user.id,
|
||||||
"email": user.email,
|
"email": user.email,
|
||||||
@@ -40,7 +55,7 @@ def _serialize_user(db: Session, user: User) -> dict[str, Any]:
|
|||||||
"allowed_providers": user.allowed_providers,
|
"allowed_providers": user.allowed_providers,
|
||||||
"allowed_api_formats": user.allowed_api_formats,
|
"allowed_api_formats": user.allowed_api_formats,
|
||||||
"allowed_models": user.allowed_models,
|
"allowed_models": user.allowed_models,
|
||||||
"unlimited": WalletService.is_unlimited_wallet(wallet),
|
"unlimited": WalletService.is_unlimited_wallet(resolved_wallet),
|
||||||
"is_active": user.is_active,
|
"is_active": user.is_active,
|
||||||
"created_at": user.created_at.isoformat(),
|
"created_at": user.created_at.isoformat(),
|
||||||
"updated_at": user.updated_at.isoformat() if user.updated_at else None,
|
"updated_at": user.updated_at.isoformat() if user.updated_at else None,
|
||||||
@@ -334,7 +349,8 @@ class AdminListUsersAdapter(AdminApiAdapter):
|
|||||||
except KeyError as exc:
|
except KeyError as exc:
|
||||||
raise InvalidRequestException("角色参数不合法") from exc
|
raise InvalidRequestException("角色参数不合法") from exc
|
||||||
users = UserService.list_users(db, self.skip, self.limit, role_enum, self.is_active)
|
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):
|
class AdminGetUserAdapter(AdminApiAdapter):
|
||||||
|
|||||||
@@ -274,10 +274,10 @@ class AuthLoginAdapter(AuthPublicAdapter):
|
|||||||
detail=f"登录请求过于频繁,请在 {reset_after} 秒后重试",
|
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
|
db, login_request.email, login_request.password, login_request.auth_type
|
||||||
)
|
)
|
||||||
if not user:
|
if not authenticated_user:
|
||||||
AuditService.log_login_attempt(
|
AuditService.log_login_attempt(
|
||||||
db=db,
|
db=db,
|
||||||
email=login_request.email,
|
email=login_request.email,
|
||||||
@@ -296,22 +296,30 @@ class AuthLoginAdapter(AuthPublicAdapter):
|
|||||||
success=True,
|
success=True,
|
||||||
ip_address=client_ip,
|
ip_address=client_ip,
|
||||||
user_agent=user_agent,
|
user_agent=user_agent,
|
||||||
user_id=user.id,
|
user_id=authenticated_user.user_id,
|
||||||
)
|
)
|
||||||
db.commit()
|
db.commit()
|
||||||
context.request.state.tx_committed_by_route = True
|
context.request.state.tx_committed_by_route = True
|
||||||
|
|
||||||
access_token = AuthService.create_access_token(
|
access_token = AuthService.create_access_token(
|
||||||
data={
|
data={
|
||||||
"user_id": user.id,
|
"user_id": authenticated_user.user_id,
|
||||||
"role": user.role.value,
|
"role": authenticated_user.role.value,
|
||||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
"created_at": (
|
||||||
|
authenticated_user.created_at.isoformat()
|
||||||
|
if authenticated_user.created_at
|
||||||
|
else None
|
||||||
|
),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
refresh_token = AuthService.create_refresh_token(
|
refresh_token = AuthService.create_refresh_token(
|
||||||
data={
|
data={
|
||||||
"user_id": user.id,
|
"user_id": authenticated_user.user_id,
|
||||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
"created_at": (
|
||||||
|
authenticated_user.created_at.isoformat()
|
||||||
|
if authenticated_user.created_at
|
||||||
|
else None
|
||||||
|
),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
response = LoginResponse(
|
response = LoginResponse(
|
||||||
@@ -319,10 +327,10 @@ class AuthLoginAdapter(AuthPublicAdapter):
|
|||||||
refresh_token=refresh_token,
|
refresh_token=refresh_token,
|
||||||
token_type="bearer",
|
token_type="bearer",
|
||||||
expires_in=86400,
|
expires_in=86400,
|
||||||
user_id=user.id,
|
user_id=authenticated_user.user_id,
|
||||||
email=user.email,
|
email=authenticated_user.email,
|
||||||
username=user.username,
|
username=authenticated_user.username,
|
||||||
role=user.role.value,
|
role=authenticated_user.role.value,
|
||||||
)
|
)
|
||||||
return response.model_dump()
|
return response.model_dump()
|
||||||
|
|
||||||
@@ -332,9 +340,6 @@ class AuthRefreshAdapter(AuthPublicAdapter):
|
|||||||
db = context.db
|
db = context.db
|
||||||
payload = context.ensure_json_body()
|
payload = context.ensure_json_body()
|
||||||
refresh_request = RefreshTokenRequest.model_validate(payload)
|
refresh_request = RefreshTokenRequest.model_validate(payload)
|
||||||
client_ip = get_client_ip(context.request)
|
|
||||||
user_agent = get_user_agent(context.request)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
token_payload = await AuthService.verify_token(
|
token_payload = await AuthService.verify_token(
|
||||||
refresh_request.refresh_token, token_type="refresh"
|
refresh_request.refresh_token, token_type="refresh"
|
||||||
@@ -745,7 +750,6 @@ class AuthSendVerificationCodeAdapter(AuthPublicAdapter):
|
|||||||
class AuthVerifyEmailAdapter(AuthPublicAdapter):
|
class AuthVerifyEmailAdapter(AuthPublicAdapter):
|
||||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||||
"""验证邮箱验证码"""
|
"""验证邮箱验证码"""
|
||||||
db = context.db
|
|
||||||
payload = context.ensure_json_body()
|
payload = context.ensure_json_body()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ class ApiAdapter(ABC):
|
|||||||
audit_log_enabled: bool = True
|
audit_log_enabled: bool = True
|
||||||
audit_success_event = None
|
audit_success_event = None
|
||||||
audit_failure_event = None
|
audit_failure_event = None
|
||||||
|
eager_request_body: bool = True
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def handle(self, context: ApiRequestContext) -> Response:
|
async def handle(self, context: ApiRequestContext) -> Response:
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import gzip
|
import gzip
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
@@ -9,7 +10,9 @@ from typing import Any
|
|||||||
|
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
from sqlalchemy.orm import Session
|
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.api_format.headers import get_header_value
|
||||||
from src.core.http_compression import is_gzip_content_encoding, normalize_content_encoding
|
from src.core.http_compression import is_gzip_content_encoding, normalize_content_encoding
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
@@ -53,6 +56,56 @@ class ApiRequestContext:
|
|||||||
client_content_encoding: str | None = None
|
client_content_encoding: str | None = None
|
||||||
client_accept_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]:
|
def ensure_json_body(self) -> dict[str, Any]:
|
||||||
"""确保请求体已解析为JSON并返回。"""
|
"""确保请求体已解析为JSON并返回。"""
|
||||||
if self.json_body is not None:
|
if self.json_body is not None:
|
||||||
|
|||||||
@@ -1,10 +1,12 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import time
|
import time
|
||||||
|
from datetime import datetime, timezone
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
|
from fastapi.concurrency import run_in_threadpool
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
from sqlalchemy.exc import SQLAlchemyError
|
from sqlalchemy.exc import SQLAlchemyError
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -14,6 +16,7 @@ from src.config.settings import config
|
|||||||
from src.core.enums import UserRole
|
from src.core.enums import UserRole
|
||||||
from src.core.exceptions import BalanceInsufficientException
|
from src.core.exceptions import BalanceInsufficientException
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
|
from src.database.database import create_session
|
||||||
from src.models.database import ApiKey, AuditEventType, User
|
from src.models.database import ApiKey, AuditEventType, User
|
||||||
from src.services.auth.service import AuthService
|
from src.services.auth.service import AuthService
|
||||||
from src.services.system.audit import AuditService
|
from src.services.system.audit import AuditService
|
||||||
@@ -108,14 +111,22 @@ class ApiRequestPipeline:
|
|||||||
user, management_token = await self._authenticate_management(http_request, db)
|
user, management_token = await self._authenticate_management(http_request, db)
|
||||||
api_key = None
|
api_key = None
|
||||||
else:
|
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
|
management_token = None
|
||||||
finally:
|
finally:
|
||||||
auth_duration = PerfRecorder.stop(auth_start, "pipeline_auth", labels=perf_labels)
|
auth_duration = PerfRecorder.stop(auth_start, "pipeline_auth", labels=perf_labels)
|
||||||
_record_perf_metric("auth_ms", auth_duration)
|
_record_perf_metric("auth_ms", auth_duration)
|
||||||
|
|
||||||
raw_body = None
|
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:
|
try:
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
@@ -141,7 +152,7 @@ class ApiRequestPipeline:
|
|||||||
perf_metrics.setdefault("pipeline", {})["body_bytes"] = int(body_size)
|
perf_metrics.setdefault("pipeline", {})["body_bytes"] = int(body_size)
|
||||||
except TimeoutError:
|
except TimeoutError:
|
||||||
timeout_sec = int(config.request_body_timeout)
|
timeout_sec = int(config.request_body_timeout)
|
||||||
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
|
logger.error("读取请求体超时({}s),可能客户端未发送完整请求体", timeout_sec)
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=408,
|
status_code=408,
|
||||||
detail=f"Request timeout: body not received within {timeout_sec} seconds",
|
detail=f"Request timeout: body not received within {timeout_sec} seconds",
|
||||||
@@ -176,8 +187,11 @@ class ApiRequestPipeline:
|
|||||||
context.management_token = management_token
|
context.management_token = management_token
|
||||||
# 存储 quiet 标志到 context,用于审计日志判断
|
# 存储 quiet 标志到 context,用于审计日志判断
|
||||||
context.quiet_logging = is_quiet
|
context.quiet_logging = is_quiet
|
||||||
if mode != ApiMode.ADMIN and user:
|
if mode in {ApiMode.STANDARD, ApiMode.PROXY, ApiMode.USER} and user:
|
||||||
remaining = self._calculate_balance_remaining(db, user, api_key=api_key)
|
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
|
context.balance_remaining = remaining
|
||||||
# authorize 可能是异步的,需要检查并 await
|
# authorize 可能是异步的,需要检查并 await
|
||||||
authorize_start = PerfRecorder.start(force=perf_sampled)
|
authorize_start = PerfRecorder.start(force=perf_sampled)
|
||||||
@@ -238,7 +252,7 @@ class ApiRequestPipeline:
|
|||||||
try:
|
try:
|
||||||
context.db.rollback()
|
context.db.rollback()
|
||||||
except Exception as rollback_exc:
|
except Exception as rollback_exc:
|
||||||
logger.debug(f"[Pipeline] 回滚失败(可忽略): {rollback_exc}")
|
logger.debug("[Pipeline] 回滚失败(可忽略): {}", rollback_exc)
|
||||||
self._record_audit_event(
|
self._record_audit_event(
|
||||||
context,
|
context,
|
||||||
adapter,
|
adapter,
|
||||||
@@ -252,32 +266,52 @@ class ApiRequestPipeline:
|
|||||||
# Internal helpers
|
# Internal helpers
|
||||||
# --------------------------------------------------------------------- #
|
# --------------------------------------------------------------------- #
|
||||||
|
|
||||||
def _authenticate_client(
|
async def _authenticate_client(
|
||||||
self, request: Request, db: Session, adapter: ApiAdapter, **_kw: object
|
self, request: Request, db: Session, adapter: ApiAdapter, **_kw: object
|
||||||
) -> tuple[User, ApiKey]:
|
) -> tuple[User, ApiKey]:
|
||||||
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
|
|
||||||
client_api_key = adapter.extract_api_key(request)
|
client_api_key = adapter.extract_api_key(request)
|
||||||
if not client_api_key:
|
if not client_api_key:
|
||||||
raise HTTPException(status_code=401, detail="请提供API密钥")
|
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:
|
if not auth_result:
|
||||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
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:
|
if not user or not api_key:
|
||||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
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)
|
request.state.user_id = db_user.id
|
||||||
access_ok, _message = self.usage_service.check_request_balance(db, user, api_key=api_key)
|
request.state.api_key_id = db_api_key.id
|
||||||
if not access_ok:
|
request.state.prefetched_balance_remaining = auth_result.balance_remaining
|
||||||
remaining = self._calculate_balance_remaining(db, user, api_key=api_key)
|
|
||||||
|
if not auth_result.access_allowed:
|
||||||
|
remaining = auth_result.balance_remaining
|
||||||
raise BalanceInsufficientException(balance_type="USD", remaining=remaining)
|
raise BalanceInsufficientException(balance_type="USD", remaining=remaining)
|
||||||
|
|
||||||
return user, api_key
|
return db_user, db_api_key
|
||||||
|
|
||||||
async def _try_token_prefix_auth(
|
async def _try_token_prefix_auth(
|
||||||
self, token: str, request: Request, db: Session
|
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.core.modules.hooks import AUTH_TOKEN_PREFIX_AUTHENTICATORS, get_hook_dispatcher
|
||||||
from src.utils.request_utils import get_client_ip
|
from src.utils.request_utils import get_client_ip
|
||||||
|
|
||||||
authenticators = await get_hook_dispatcher().dispatch(
|
authenticators = await get_hook_dispatcher().dispatch(AUTH_TOKEN_PREFIX_AUTHENTICATORS)
|
||||||
AUTH_TOKEN_PREFIX_AUTHENTICATORS, db=db
|
|
||||||
)
|
|
||||||
for auth_info in authenticators or []:
|
for auth_info in authenticators or []:
|
||||||
prefix = auth_info.get("prefix", "")
|
prefix = auth_info.get("prefix", "")
|
||||||
authenticate_fn = auth_info.get("authenticate")
|
authenticate_fn = auth_info.get("authenticate")
|
||||||
@@ -304,111 +336,136 @@ class ApiRequestPipeline:
|
|||||||
logger.warning("Token prefix '{}' has no authenticate callback", prefix)
|
logger.warning("Token prefix '{}' has no authenticate callback", prefix)
|
||||||
raise HTTPException(status_code=401, detail="认证服务不可用")
|
raise HTTPException(status_code=401, detail="认证服务不可用")
|
||||||
client_ip = get_client_ip(request)
|
client_ip = get_client_ip(request)
|
||||||
result = await authenticate_fn(db, token, client_ip)
|
auth_db = create_session()
|
||||||
if result:
|
try:
|
||||||
return result
|
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")
|
module_name = auth_info.get("module", "unknown")
|
||||||
raise HTTPException(status_code=401, detail=f"无效或过期的 Token ({module_name})")
|
raise HTTPException(status_code=401, detail=f"无效或过期的 Token ({module_name})")
|
||||||
return None # 无前缀匹配
|
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(
|
async def _authenticate_admin(
|
||||||
self, request: Request, db: Session
|
self, request: Request, db: Session
|
||||||
) -> tuple[User, ManagementToken | None]:
|
) -> tuple[User, ManagementToken | None]:
|
||||||
"""管理员认证,支持 JWT 和 Management Token 两种方式"""
|
"""Admin auth supports JWT and Management Token."""
|
||||||
authorization = request.headers.get("authorization")
|
authorization = request.headers.get("authorization")
|
||||||
if not authorization or not authorization.lower().startswith("bearer "):
|
if not authorization or not authorization.lower().startswith("bearer "):
|
||||||
raise HTTPException(status_code=401, detail="缺少管理员凭证")
|
raise HTTPException(status_code=401, detail="缺少管理员凭证")
|
||||||
|
|
||||||
token = authorization[7:].strip()
|
token = authorization[7:].strip()
|
||||||
|
|
||||||
# 通过钩子检查是否匹配模块注册的 token 前缀(如 ae_)
|
|
||||||
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
||||||
if token_auth_result is not None:
|
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:
|
if user.role != UserRole.ADMIN:
|
||||||
logger.warning(f"非管理员尝试通过 Management Token 访问管理端点: {user.email}")
|
logger.warning("非管理员尝试通过 Management Token 访问管理端点: {}", user.email)
|
||||||
raise HTTPException(status_code=403, detail="需要管理员权限")
|
raise HTTPException(status_code=403, detail="需要管理员权限")
|
||||||
|
|
||||||
# 存储到 request.state
|
|
||||||
request.state.user_id = user.id
|
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
|
return user, management_token
|
||||||
|
|
||||||
# JWT 认证
|
|
||||||
try:
|
try:
|
||||||
payload = await self.auth_service.verify_token(token, token_type="access")
|
payload = await self.auth_service.verify_token(token, token_type="access")
|
||||||
except HTTPException:
|
except HTTPException:
|
||||||
raise
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error(f"Admin token 验证失败: {exc}")
|
logger.error("Admin token 验证失败: {}", exc)
|
||||||
raise HTTPException(status_code=401, detail="无效的管理员令牌")
|
raise HTTPException(status_code=401, detail="无效的管理员令牌")
|
||||||
|
|
||||||
user_id = payload.get("user_id")
|
user_id = payload.get("user_id")
|
||||||
if not user_id:
|
if not user_id:
|
||||||
raise HTTPException(status_code=401, detail="无效的管理员令牌")
|
raise HTTPException(status_code=401, detail="无效的管理员令牌")
|
||||||
|
|
||||||
# 直接查询数据库,确保返回的是当前 Session 绑定的对象
|
db_user = db.query(User).filter(User.id == user_id).first()
|
||||||
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:
|
||||||
if not user or not user.is_active or user.is_deleted:
|
|
||||||
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
|
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="无效的管理员令牌")
|
raise HTTPException(status_code=403, detail="无效的管理员令牌")
|
||||||
|
|
||||||
# 检查管理员权限
|
if db_user.role != UserRole.ADMIN:
|
||||||
if user.role != UserRole.ADMIN:
|
logger.warning("非管理员尝试通过 JWT 访问管理端点: {}", db_user.email)
|
||||||
logger.warning(f"非管理员尝试通过 JWT 访问管理端点: {user.email}")
|
|
||||||
raise HTTPException(status_code=403, detail="需要管理员权限")
|
raise HTTPException(status_code=403, detail="需要管理员权限")
|
||||||
|
|
||||||
request.state.user_id = user.id
|
request.state.user_id = db_user.id
|
||||||
return user, None
|
return db_user, None
|
||||||
|
|
||||||
async def _authenticate_user(
|
async def _authenticate_user(
|
||||||
self, request: Request, db: Session
|
self, request: Request, db: Session
|
||||||
) -> tuple[User, ManagementToken | None]:
|
) -> tuple[User, ManagementToken | None]:
|
||||||
"""用户认证,支持 JWT 和 Management Token 两种方式"""
|
"""User auth supports JWT and Management Token."""
|
||||||
authorization = request.headers.get("authorization")
|
authorization = request.headers.get("authorization")
|
||||||
if not authorization or not authorization.lower().startswith("bearer "):
|
if not authorization or not authorization.lower().startswith("bearer "):
|
||||||
raise HTTPException(status_code=401, detail="缺少用户凭证")
|
raise HTTPException(status_code=401, detail="缺少用户凭证")
|
||||||
|
|
||||||
token = authorization[7:].strip()
|
token = authorization[7:].strip()
|
||||||
|
|
||||||
# 通过钩子检查是否匹配模块注册的 token 前缀(如 ae_)
|
|
||||||
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
||||||
if token_auth_result is not None:
|
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.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
|
return user, management_token
|
||||||
|
|
||||||
# JWT 认证
|
|
||||||
try:
|
try:
|
||||||
payload = await self.auth_service.verify_token(token, token_type="access")
|
payload = await self.auth_service.verify_token(token, token_type="access")
|
||||||
except HTTPException:
|
except HTTPException:
|
||||||
raise
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error(f"User token 验证失败: {exc}")
|
logger.error("User token 验证失败: {}", exc)
|
||||||
raise HTTPException(status_code=401, detail="无效的用户令牌")
|
raise HTTPException(status_code=401, detail="无效的用户令牌")
|
||||||
|
|
||||||
user_id = payload.get("user_id")
|
user_id = payload.get("user_id")
|
||||||
if not user_id:
|
if not user_id:
|
||||||
raise HTTPException(status_code=401, detail="无效的用户令牌")
|
raise HTTPException(status_code=401, detail="无效的用户令牌")
|
||||||
|
|
||||||
user = db.query(User).filter(User.id == user_id).first()
|
db_user = db.query(User).filter(User.id == user_id).first()
|
||||||
if not user or not user.is_active or user.is_deleted:
|
if not db_user or not db_user.is_active or db_user.is_deleted:
|
||||||
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
|
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="无效的用户令牌")
|
raise HTTPException(status_code=403, detail="无效的用户令牌")
|
||||||
|
|
||||||
request.state.user_id = user.id
|
request.state.user_id = db_user.id
|
||||||
return user, None
|
return db_user, None
|
||||||
|
|
||||||
async def _authenticate_management(
|
async def _authenticate_management(
|
||||||
self, request: Request, db: Session
|
self, request: Request, db: Session
|
||||||
@@ -424,11 +481,11 @@ class ApiRequestPipeline:
|
|||||||
# _try_token_prefix_auth 会在前缀匹配但认证失败时直接抛 HTTPException
|
# _try_token_prefix_auth 会在前缀匹配但认证失败时直接抛 HTTPException
|
||||||
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
||||||
if token_auth_result is not None:
|
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
|
||||||
request.state.user_id = user.id
|
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
|
return user, management_token
|
||||||
|
|
||||||
@@ -437,15 +494,37 @@ class ApiRequestPipeline:
|
|||||||
detail="无效的 Token 格式,需要 Management Token",
|
detail="无效的 Token 格式,需要 Management Token",
|
||||||
)
|
)
|
||||||
|
|
||||||
def _calculate_balance_remaining(
|
async def _calculate_balance_remaining_async(
|
||||||
self, db: Session, user: User | None, api_key: ApiKey | None = None
|
self, user: User | None, api_key: ApiKey | None = None
|
||||||
) -> float | None:
|
) -> float | None:
|
||||||
if not user:
|
if not user:
|
||||||
return None
|
return None
|
||||||
balance = WalletService.get_balance_snapshot(db, user=user, api_key=api_key)
|
|
||||||
if balance is None:
|
user_id = getattr(user, "id", None)
|
||||||
return None
|
api_key_id = getattr(api_key, "id", None)
|
||||||
return float(balance)
|
|
||||||
|
# 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(
|
def _record_audit_event(
|
||||||
self,
|
self,
|
||||||
@@ -502,7 +581,7 @@ class ApiRequestPipeline:
|
|||||||
)
|
)
|
||||||
except Exception as exc:
|
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(
|
def _build_audit_metadata(
|
||||||
self,
|
self,
|
||||||
@@ -568,7 +647,9 @@ class ApiRequestPipeline:
|
|||||||
if adapter_details:
|
if adapter_details:
|
||||||
extra_details.update(adapter_details)
|
extra_details.update(adapter_details)
|
||||||
except Exception as exc:
|
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:
|
if extra_details:
|
||||||
metadata["details"] = extra_details
|
metadata["details"] = extra_details
|
||||||
@@ -578,7 +659,7 @@ class ApiRequestPipeline:
|
|||||||
|
|
||||||
return self._sanitize_metadata(metadata)
|
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:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
if depth > 5:
|
if depth > 5:
|
||||||
|
|||||||
@@ -56,6 +56,7 @@ class ChatAdapterBase(HandlerAdapterBase):
|
|||||||
# 适配器配置
|
# 适配器配置
|
||||||
name: str = "chat.base"
|
name: str = "chat.base"
|
||||||
mode = ApiMode.STANDARD
|
mode = ApiMode.STANDARD
|
||||||
|
eager_request_body = False
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Any:
|
async def handle(self, context: ApiRequestContext) -> Any:
|
||||||
"""处理 Chat API 请求"""
|
"""处理 Chat API 请求"""
|
||||||
@@ -71,7 +72,7 @@ class ChatAdapterBase(HandlerAdapterBase):
|
|||||||
original_headers = context.original_headers
|
original_headers = context.original_headers
|
||||||
query_params = context.query_params
|
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 路径中)
|
# 合并 path_params 到请求体(如 Gemini API 的 model 在 URL 路径中)
|
||||||
if context.path_params:
|
if context.path_params:
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ from __future__ import annotations
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import HTTPException
|
from fastapi import HTTPException
|
||||||
from fastapi.responses import JSONResponse
|
|
||||||
|
|
||||||
from src.api.base.adapter import ApiMode
|
from src.api.base.adapter import ApiMode
|
||||||
from src.api.base.context import ApiRequestContext
|
from src.api.base.context import ApiRequestContext
|
||||||
@@ -58,6 +57,7 @@ class CliAdapterBase(HandlerAdapterBase):
|
|||||||
# 适配器配置
|
# 适配器配置
|
||||||
name: str = "cli.base"
|
name: str = "cli.base"
|
||||||
mode = ApiMode.PROXY
|
mode = ApiMode.PROXY
|
||||||
|
eager_request_body = False
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Any:
|
async def handle(self, context: ApiRequestContext) -> Any:
|
||||||
"""处理 CLI API 请求"""
|
"""处理 CLI API 请求"""
|
||||||
@@ -82,7 +82,7 @@ class CliAdapterBase(HandlerAdapterBase):
|
|||||||
|
|
||||||
set_original_request_headers(original_headers)
|
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 路径中)
|
# 合并 path_params 到请求体(如 Gemini API 的 model 在 URL 路径中)
|
||||||
if context.path_params:
|
if context.path_params:
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ class VideoAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
name: str = "video.base"
|
name: str = "video.base"
|
||||||
mode = ApiMode.STANDARD
|
mode = ApiMode.STANDARD
|
||||||
|
eager_request_body = False
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
||||||
@@ -103,7 +104,7 @@ class VideoAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
# Remix task
|
# Remix task
|
||||||
if method == "POST" and path.endswith("/remix") and task_id:
|
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(
|
return await handler.handle_remix_task(
|
||||||
task_id=task_id,
|
task_id=task_id,
|
||||||
http_request=http_request,
|
http_request=http_request,
|
||||||
@@ -134,7 +135,7 @@ class VideoAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
# Create task (default)
|
# Create task (default)
|
||||||
if method in {"POST", "PUT", "PATCH"}:
|
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(
|
return await handler.handle_create_task(
|
||||||
http_request=http_request,
|
http_request=http_request,
|
||||||
original_headers=context.original_headers,
|
original_headers=context.original_headers,
|
||||||
|
|||||||
@@ -184,6 +184,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
|||||||
"thinking_enabled": bool(request_obj.thinking),
|
"thinking_enabled": bool(request_obj.thinking),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
def build_endpoint_url(
|
def build_endpoint_url(
|
||||||
cls,
|
cls,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
@@ -224,6 +225,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
|
|||||||
|
|
||||||
name = "claude.token_count"
|
name = "claude.token_count"
|
||||||
mode = ApiMode.STANDARD
|
mode = ApiMode.STANDARD
|
||||||
|
eager_request_body = False
|
||||||
|
|
||||||
def extract_api_key(self, request: Request) -> str | None:
|
def extract_api_key(self, request: Request) -> str | None:
|
||||||
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
|
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
|
||||||
@@ -239,7 +241,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
|
|||||||
return bearer_handler.extract_credentials(request)
|
return bearer_handler.extract_credentials(request)
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Any:
|
async def handle(self, context: ApiRequestContext) -> Any:
|
||||||
payload = context.ensure_json_body()
|
payload = await context.ensure_json_body_async()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
request = ClaudeTokenCountRequest.model_validate(payload, strict=False)
|
request = ClaudeTokenCountRequest.model_validate(payload, strict=False)
|
||||||
|
|||||||
@@ -48,7 +48,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
|||||||
async def handle(self, context: ApiRequestContext) -> Any:
|
async def handle(self, context: ApiRequestContext) -> Any:
|
||||||
"""处理 CLI API 请求 -- compact 模式下注入标记并强制非流式"""
|
"""处理 CLI API 请求 -- compact 模式下注入标记并强制非流式"""
|
||||||
if self._compact:
|
if self._compact:
|
||||||
body = context.ensure_json_body()
|
body = await context.ensure_json_body_async()
|
||||||
body["_aether_compact"] = True
|
body["_aether_compact"] = True
|
||||||
# compact 端点永远非流式
|
# compact 端点永远非流式
|
||||||
body.pop("stream", None)
|
body.pop("stream", None)
|
||||||
|
|||||||
@@ -231,7 +231,6 @@ async def get_my_usage(
|
|||||||
- `total_tokens`: 总 Token 数
|
- `total_tokens`: 总 Token 数
|
||||||
- `total_cost`: 总成本(USD)
|
- `total_cost`: 总成本(USD)
|
||||||
- `summary_by_model`: 按模型分组统计
|
- `summary_by_model`: 按模型分组统计
|
||||||
- `summary_by_provider`: 按提供商分组统计
|
|
||||||
- `records`: 详细使用记录列表
|
- `records`: 详细使用记录列表
|
||||||
- `pagination`: 分页信息
|
- `pagination`: 分页信息
|
||||||
"""
|
"""
|
||||||
@@ -809,6 +808,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
|||||||
user_id=user.id,
|
user_id=user.id,
|
||||||
start_date=start_utc,
|
start_date=start_utc,
|
||||||
end_date=end_utc,
|
end_date=end_utc,
|
||||||
|
group_by=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 过滤掉 unknown/pending provider 的记录(请求未到达任何提供商)
|
# 过滤掉 unknown/pending provider 的记录(请求未到达任何提供商)
|
||||||
@@ -818,31 +818,23 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
|||||||
if item.get("provider") not in ("unknown", "pending", None)
|
if item.get("provider") not in ("unknown", "pending", None)
|
||||||
]
|
]
|
||||||
|
|
||||||
total_requests = sum(item["requests"] for item in filtered_summary)
|
total_requests = 0
|
||||||
total_input_tokens = (
|
total_input_tokens = 0
|
||||||
sum(item["input_tokens"] for item in filtered_summary) if filtered_summary else 0
|
total_output_tokens = 0
|
||||||
)
|
total_tokens = 0
|
||||||
total_output_tokens = (
|
total_cost = 0.0
|
||||||
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_actual_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 = {}
|
model_summary = {}
|
||||||
|
provider_summary = {}
|
||||||
for item in filtered_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"]
|
model_name = item["model"]
|
||||||
base_stats = {
|
base_stats = {
|
||||||
"model": model_name,
|
"model": model_name,
|
||||||
@@ -866,13 +858,8 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
|||||||
if user.role == UserRole.ADMIN:
|
if user.role == UserRole.ADMIN:
|
||||||
stats["actual_total_cost_usd"] += item.get("actual_total_cost_usd", 0.0)
|
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"]
|
provider_name = item["provider"]
|
||||||
base_stats = {
|
provider_base_stats = {
|
||||||
"provider": provider_name,
|
"provider": provider_name,
|
||||||
"requests": 0,
|
"requests": 0,
|
||||||
"total_tokens": 0,
|
"total_tokens": 0,
|
||||||
@@ -881,32 +868,37 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
|||||||
"total_response_time_ms": 0.0,
|
"total_response_time_ms": 0.0,
|
||||||
"response_time_count": 0,
|
"response_time_count": 0,
|
||||||
}
|
}
|
||||||
stats = provider_summary.setdefault(provider_name, base_stats)
|
provider_stats = provider_summary.setdefault(provider_name, provider_base_stats)
|
||||||
stats["requests"] += item["requests"]
|
provider_stats["requests"] += item["requests"]
|
||||||
stats["total_tokens"] += item["total_tokens"]
|
provider_stats["total_tokens"] += item["total_tokens"]
|
||||||
stats["total_cost_usd"] += item["total_cost_usd"]
|
provider_stats["total_cost_usd"] += item["total_cost_usd"]
|
||||||
# 假设 summary 中的都是成功的请求
|
provider_stats["success_count"] += int(item.get("success_count", 0) or 0)
|
||||||
stats["success_count"] += item["requests"]
|
success_response_time_count = int(item.get("success_response_time_count", 0) or 0)
|
||||||
if item.get("avg_response_time_ms") is not None:
|
if success_response_time_count > 0:
|
||||||
stats["total_response_time_ms"] += item["avg_response_time_ms"] * item["requests"]
|
provider_stats["total_response_time_ms"] += float(
|
||||||
stats["response_time_count"] += item["requests"]
|
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 = []
|
summary_by_provider = []
|
||||||
for stats in provider_summary.values():
|
for provider_stats in provider_summary.values():
|
||||||
avg_response_time_ms = (
|
avg_response_time_ms = (
|
||||||
stats["total_response_time_ms"] / stats["response_time_count"]
|
provider_stats["total_response_time_ms"] / provider_stats["response_time_count"]
|
||||||
if stats["response_time_count"] > 0
|
if provider_stats["response_time_count"] > 0
|
||||||
else 0
|
else 0
|
||||||
)
|
)
|
||||||
success_rate = (
|
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(
|
summary_by_provider.append(
|
||||||
{
|
{
|
||||||
"provider": stats["provider"],
|
"provider": provider_stats["provider"],
|
||||||
"requests": stats["requests"],
|
"requests": provider_stats["requests"],
|
||||||
"total_tokens": stats["total_tokens"],
|
"total_tokens": provider_stats["total_tokens"],
|
||||||
"total_cost_usd": stats["total_cost_usd"],
|
"total_cost_usd": provider_stats["total_cost_usd"],
|
||||||
"success_rate": round(success_rate, 2),
|
"success_rate": round(success_rate, 2),
|
||||||
"avg_response_time_ms": round(avg_response_time_ms, 2),
|
"avg_response_time_ms": round(avg_response_time_ms, 2),
|
||||||
}
|
}
|
||||||
@@ -1003,6 +995,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
|||||||
"avg_response_time": avg_response_time,
|
"avg_response_time": avg_response_time,
|
||||||
"billing": WalletService.serialize_wallet_summary(wallet),
|
"billing": WalletService.serialize_wallet_summary(wallet),
|
||||||
"summary_by_model": summary_by_model,
|
"summary_by_model": summary_by_model,
|
||||||
|
"summary_by_provider": summary_by_provider,
|
||||||
# 分页信息
|
# 分页信息
|
||||||
"pagination": {
|
"pagination": {
|
||||||
"total": total_records,
|
"total": total_records,
|
||||||
@@ -1117,7 +1110,12 @@ class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
|
|||||||
if not id_list:
|
if not id_list:
|
||||||
return {"requests": []}
|
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}
|
return {"requests": requests}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import secrets
|
|||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from threading import Lock
|
from threading import Lock
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
@@ -23,6 +24,7 @@ from src.config import config
|
|||||||
from src.core.enums import AuthSource
|
from src.core.enums import AuthSource
|
||||||
from src.core.exceptions import ForbiddenException
|
from src.core.exceptions import ForbiddenException
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
|
from src.database.database import create_session
|
||||||
from src.services.system.config import SystemConfigService
|
from src.services.system.config import SystemConfigService
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
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.auth.jwt_blacklist import JWTBlacklistService
|
||||||
from src.services.cache.user_cache import UserCacheService
|
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 更新节流配置
|
||||||
# 同一个 API Key 在此时间间隔内只会更新一次 last_used_at
|
# 同一个 API Key 在此时间间隔内只会更新一次 last_used_at
|
||||||
_LAST_USED_UPDATE_INTERVAL = 60 # 秒
|
_LAST_USED_UPDATE_INTERVAL = 60 # 秒
|
||||||
@@ -177,6 +204,192 @@ class AuthService:
|
|||||||
except jwt.InvalidTokenError:
|
except jwt.InvalidTokenError:
|
||||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="无效的Token")
|
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
|
@staticmethod
|
||||||
async def authenticate_user(
|
async def authenticate_user(
|
||||||
db: Session, email: str, password: str, auth_type: str = "local"
|
db: Session, email: str, password: str, auth_type: str = "local"
|
||||||
@@ -208,40 +421,8 @@ class AuthService:
|
|||||||
# 本地认证
|
# 本地认证
|
||||||
# 登录校验必须读取密码哈希,不能使用不包含 password_hash 的缓存对象
|
# 登录校验必须读取密码哈希,不能使用不包含 password_hash 的缓存对象
|
||||||
# 支持邮箱或用户名登录
|
# 支持邮箱或用户名登录
|
||||||
from sqlalchemy import or_
|
user = AuthService._authenticate_local_user_sync(db, email, password)
|
||||||
|
|
||||||
user = db.query(User).filter(or_(User.email == email, User.username == email)).first()
|
|
||||||
|
|
||||||
if not user:
|
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
|
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(秒)- 使用统一常量
|
# 缓存 TTL(秒)- 使用统一常量
|
||||||
CACHE_TTL = CacheTTL.MODEL
|
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
|
@staticmethod
|
||||||
async def get_model_by_id(db: Session, model_id: str) -> Model | None:
|
async def get_model_by_id(db: Session, model_id: str) -> Model | None:
|
||||||
@@ -238,6 +240,9 @@ class ModelCacheService:
|
|||||||
if resolve_keys_to_clear:
|
if resolve_keys_to_clear:
|
||||||
logger.debug(f"Model resolve 缓存已清除: {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
|
@staticmethod
|
||||||
async def invalidate_global_model_cache(global_model_id: str, name: str | None = None) -> None:
|
async def invalidate_global_model_cache(global_model_id: str, name: str | None = None) -> None:
|
||||||
"""清除 GlobalModel 缓存"""
|
"""清除 GlobalModel 缓存"""
|
||||||
@@ -249,6 +254,8 @@ class ModelCacheService:
|
|||||||
# 全量清除 resolve 缓存,确保映射规则变更后不命中旧缓存
|
# 全量清除 resolve 缓存,确保映射规则变更后不命中旧缓存
|
||||||
try:
|
try:
|
||||||
await CacheService.delete_pattern("global_model:resolve:*")
|
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:
|
except Exception as e:
|
||||||
logger.error(f"GlobalModel resolve 缓存清除失败,可能导致映射不一致: {e}")
|
logger.error(f"GlobalModel resolve 缓存清除失败,可能导致映射不一致: {e}")
|
||||||
logger.debug(f"GlobalModel 缓存已清除: {global_model_id}")
|
logger.debug(f"GlobalModel 缓存已清除: {global_model_id}")
|
||||||
@@ -262,10 +269,108 @@ class ModelCacheService:
|
|||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
deleted = await CacheService.delete_pattern("global_model:resolve:*")
|
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 缓存")
|
logger.debug(f"已清除 {deleted} 个 GlobalModel resolve 缓存")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"GlobalModel resolve 缓存清除失败: {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
|
@staticmethod
|
||||||
async def resolve_global_model_by_name_or_mapping(
|
async def resolve_global_model_by_name_or_mapping(
|
||||||
db: Session, model_name: str
|
db: Session, model_name: str
|
||||||
@@ -391,41 +496,18 @@ class ModelCacheService:
|
|||||||
return result_global_model
|
return result_global_model
|
||||||
|
|
||||||
# 4. 通过 provider_model_mappings 匹配
|
# 4. 通过 provider_model_mappings 匹配
|
||||||
models_with_mappings = (
|
provider_mapping_index = await ModelCacheService._get_provider_mapping_index(db)
|
||||||
db.query(Model, GlobalModel)
|
mapping_matched_global_models = [
|
||||||
.join(Provider, Model.provider_id == Provider.id)
|
ModelCacheService._dict_to_global_model(global_model_dict)
|
||||||
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
|
for global_model_dict in provider_mapping_index.get(normalized_name, [])
|
||||||
.filter(
|
if isinstance(global_model_dict, dict)
|
||||||
Provider.is_active == True,
|
]
|
||||||
Model.is_active == True,
|
|
||||||
GlobalModel.is_active == True,
|
|
||||||
Model.provider_model_mappings.isnot(None),
|
|
||||||
)
|
|
||||||
.all()
|
|
||||||
)
|
|
||||||
|
|
||||||
mapping_matched_global_models: list[GlobalModel] = []
|
for gm in mapping_matched_global_models:
|
||||||
mapping_seen_ids: set[str] = set()
|
logger.debug(
|
||||||
for model, gm in models_with_mappings:
|
f"模型名称 '{normalized_name}' 通过 provider_model_mappings 匹配到 "
|
||||||
raw_mappings = model.provider_model_mappings
|
f"GlobalModel: {gm.name}"
|
||||||
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
|
|
||||||
|
|
||||||
if mapping_matched_global_models:
|
if mapping_matched_global_models:
|
||||||
resolution_method = "provider_model_mappings"
|
resolution_method = "provider_model_mappings"
|
||||||
@@ -453,32 +535,21 @@ class ModelCacheService:
|
|||||||
return result_global_model
|
return result_global_model
|
||||||
|
|
||||||
# 5. 通过 GlobalModel.config.model_mappings 匹配(支持正则)
|
# 5. 通过 GlobalModel.config.model_mappings 匹配(支持正则)
|
||||||
from sqlalchemy import func
|
|
||||||
|
|
||||||
from src.core.model_permissions import match_model_with_pattern
|
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] = []
|
mapping_matches: list[GlobalModel] = []
|
||||||
for gm in mapping_rows:
|
for entry in await ModelCacheService._get_model_mapping_rules(db):
|
||||||
config = gm.config or {}
|
global_model_dict = entry.get("global_model")
|
||||||
mappings = config.get("model_mappings")
|
patterns = entry.get("patterns")
|
||||||
if not isinstance(mappings, list):
|
if not isinstance(global_model_dict, dict) or not isinstance(patterns, list):
|
||||||
continue
|
continue
|
||||||
for pattern in mappings:
|
for pattern in patterns:
|
||||||
if isinstance(pattern, str) and match_model_with_pattern(
|
if isinstance(pattern, str) and match_model_with_pattern(
|
||||||
pattern, normalized_name
|
pattern, normalized_name
|
||||||
):
|
):
|
||||||
mapping_matches.append(gm)
|
mapping_matches.append(
|
||||||
|
ModelCacheService._dict_to_global_model(global_model_dict)
|
||||||
|
)
|
||||||
break
|
break
|
||||||
|
|
||||||
if mapping_matches:
|
if mapping_matches:
|
||||||
|
|||||||
@@ -15,6 +15,14 @@ from src.models.database import RequestCandidate
|
|||||||
class RequestCandidateService:
|
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
|
@staticmethod
|
||||||
def create_candidate(
|
def create_candidate(
|
||||||
db: Session,
|
db: Session,
|
||||||
@@ -93,9 +101,9 @@ class RequestCandidateService:
|
|||||||
if candidate:
|
if candidate:
|
||||||
candidate.status = "pending"
|
candidate.status = "pending"
|
||||||
candidate.started_at = datetime.now(timezone.utc)
|
candidate.started_at = datetime.now(timezone.utc)
|
||||||
# 关键状态更新:立即提交,不使用批量提交
|
# 中间态改为 flush:最终 success/failed 仍会立即提交,
|
||||||
# 原因:前端需要实时看到请求开始执行
|
# 但开始执行这一跳不再单独制造一次事务往返。
|
||||||
db.commit()
|
RequestCandidateService._persist_candidate_update(db, immediate=False)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def update_candidate_status(db: Session, candidate_id: str, status: str) -> None:
|
def update_candidate_status(db: Session, candidate_id: str, status: str) -> None:
|
||||||
@@ -113,8 +121,9 @@ class RequestCandidateService:
|
|||||||
# 如果状态变更为 pending,记录开始时间
|
# 如果状态变更为 pending,记录开始时间
|
||||||
if status == "pending" and not candidate.started_at:
|
if status == "pending" and not candidate.started_at:
|
||||||
candidate.started_at = datetime.now(timezone.utc)
|
candidate.started_at = datetime.now(timezone.utc)
|
||||||
# 立即提交,确保前端能实时看到状态变化
|
RequestCandidateService._persist_candidate_update(
|
||||||
db.commit()
|
db, immediate=status not in {"pending", "streaming"}
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def mark_candidate_streaming(
|
def mark_candidate_streaming(
|
||||||
@@ -141,7 +150,7 @@ class RequestCandidateService:
|
|||||||
candidate.status = "streaming"
|
candidate.status = "streaming"
|
||||||
candidate.concurrent_requests = concurrent_requests
|
candidate.concurrent_requests = concurrent_requests
|
||||||
# streaming 状态不设置 finished_at 和 status_code,因为请求还在进行中
|
# streaming 状态不设置 finished_at 和 status_code,因为请求还在进行中
|
||||||
db.commit()
|
RequestCandidateService._persist_candidate_update(db, immediate=False)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def mark_candidate_success(
|
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.rate_limit.concurrency_manager import get_concurrency_manager
|
||||||
from src.services.scheduling.affinity_manager import (
|
from src.services.scheduling.affinity_manager import (
|
||||||
CacheAffinityManager,
|
|
||||||
get_affinity_manager,
|
get_affinity_manager,
|
||||||
)
|
)
|
||||||
from src.services.scheduling.candidate_builder import (
|
from src.services.scheduling.candidate_builder import (
|
||||||
@@ -465,12 +464,41 @@ class CacheAwareScheduler:
|
|||||||
return [], global_model_id, queried_provider_count
|
return [], global_model_id, queried_provider_count
|
||||||
|
|
||||||
# 1. 查询 Providers(委托给 CandidateBuilder)
|
# 1. 查询 Providers(委托给 CandidateBuilder)
|
||||||
providers = self._candidate_builder._query_providers(
|
providers = []
|
||||||
db=db,
|
if allowed_providers is not None:
|
||||||
provider_offset=provider_offset,
|
provider_refs = self._candidate_builder._query_provider_refs(
|
||||||
provider_limit=provider_limit,
|
db=db,
|
||||||
)
|
provider_offset=provider_offset,
|
||||||
queried_provider_count = len(providers)
|
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.
|
# Provider query starts a transaction; release connection before entering async candidate build.
|
||||||
release_db_connection_before_await(db)
|
release_db_connection_before_await(db)
|
||||||
@@ -481,19 +509,6 @@ class CacheAwareScheduler:
|
|||||||
", ".join(p.name for p in providers),
|
", ".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:
|
if not providers:
|
||||||
return [], global_model_id, queried_provider_count
|
return [], global_model_id, queried_provider_count
|
||||||
|
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import re
|
|||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from sqlalchemy import or_
|
||||||
from sqlalchemy.orm import Session, selectinload
|
from sqlalchemy.orm import Session, selectinload
|
||||||
|
|
||||||
from src.core.api_format.conversion.compatibility import is_format_compatible
|
from src.core.api_format.conversion.compatibility import is_format_compatible
|
||||||
@@ -49,7 +50,7 @@ from src.services.cache.model_cache import ModelCacheService
|
|||||||
|
|
||||||
def _get_pool_config(provider: Provider) -> "PoolConfig | None":
|
def _get_pool_config(provider: Provider) -> "PoolConfig | None":
|
||||||
"""Return parsed PoolConfig if the provider has pool enabled, else 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))
|
return parse_pool_config(getattr(provider, "config", None))
|
||||||
|
|
||||||
@@ -76,11 +77,36 @@ class CandidateBuilder:
|
|||||||
def __init__(self, candidate_sorter: CandidateSorterProtocol) -> None:
|
def __init__(self, candidate_sorter: CandidateSorterProtocol) -> None:
|
||||||
self._sorter = candidate_sorter
|
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(
|
def _query_providers(
|
||||||
self,
|
self,
|
||||||
db: Session,
|
db: Session,
|
||||||
provider_offset: int = 0,
|
provider_offset: int = 0,
|
||||||
provider_limit: int | None = None,
|
provider_limit: int | None = None,
|
||||||
|
allowed_providers: list[str] | None = None,
|
||||||
|
provider_ids: list[str] | None = None,
|
||||||
) -> list[Provider]:
|
) -> list[Provider]:
|
||||||
"""
|
"""
|
||||||
查询活跃的 Providers(带预加载)
|
查询活跃的 Providers(带预加载)
|
||||||
@@ -118,12 +144,30 @@ class CandidateBuilder:
|
|||||||
.order_by(Provider.provider_priority.asc())
|
.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)
|
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)
|
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(
|
async def _check_model_support(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -42,11 +42,20 @@ class CandidateSorterProtocol(Protocol):
|
|||||||
|
|
||||||
|
|
||||||
class CandidateBuilderProtocol(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(
|
def _query_providers(
|
||||||
self,
|
self,
|
||||||
db: Session,
|
db: Session,
|
||||||
provider_offset: int = 0,
|
provider_offset: int = 0,
|
||||||
provider_limit: int | None = None,
|
provider_limit: int | None = None,
|
||||||
|
allowed_providers: list[str] | None = None,
|
||||||
|
provider_ids: list[str] | None = None,
|
||||||
) -> list[Provider]: ...
|
) -> list[Provider]: ...
|
||||||
|
|
||||||
async def _build_candidates(
|
async def _build_candidates(
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ from datetime import date, datetime, time, timedelta, timezone
|
|||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
from typing import Any
|
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.exc import IntegrityError
|
||||||
from sqlalchemy.orm import Session
|
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)
|
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:
|
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
|
@staticmethod
|
||||||
def compute_daily_stats(db: Session, date: datetime) -> dict:
|
def compute_daily_stats(db: Session, date: datetime) -> dict:
|
||||||
"""计算指定 UTC 日期的统计数据(不写入数据库)"""
|
"""计算指定 UTC 日期的统计数据(不写入数据库)"""
|
||||||
@@ -196,23 +251,8 @@ class StatsAggregatorService:
|
|||||||
.first()
|
.first()
|
||||||
)
|
)
|
||||||
|
|
||||||
def _resolve(row: Any | None) -> tuple[int | None, int | None, int | None]:
|
p50_rt, p90_rt, p99_rt = StatsAggregatorService._resolve_percentile_row(rt_row)
|
||||||
if not row:
|
p50_ttfb, p90_ttfb, p99_ttfb = StatsAggregatorService._resolve_percentile_row(ttfb_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)
|
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"p50_response_time_ms": p50_rt,
|
"p50_response_time_ms": p50_rt,
|
||||||
@@ -223,6 +263,103 @@ class StatsAggregatorService:
|
|||||||
"p99_first_byte_time_ms": p99_ttfb,
|
"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
|
@staticmethod
|
||||||
def aggregate_daily_stats(db: Session, date: datetime, commit: bool = True) -> StatsDaily:
|
def aggregate_daily_stats(db: Session, date: datetime, commit: bool = True) -> StatsDaily:
|
||||||
"""聚合指定 UTC 日期的统计数据
|
"""聚合指定 UTC 日期的统计数据
|
||||||
@@ -656,6 +793,82 @@ class StatsAggregatorService:
|
|||||||
db.commit()
|
db.commit()
|
||||||
return stats
|
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
|
@staticmethod
|
||||||
def aggregate_daily_stats_bundle(
|
def aggregate_daily_stats_bundle(
|
||||||
db: Session, date: datetime, user_ids: list[str] | None = None
|
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)
|
StatsAggregatorService.aggregate_daily_error_stats(db, date, commit=False)
|
||||||
|
|
||||||
if user_ids:
|
if user_ids:
|
||||||
for user_id in user_ids:
|
StatsAggregatorService.aggregate_user_daily_stats_batch(
|
||||||
StatsAggregatorService.aggregate_user_daily_stats(db, user_id, date, commit=False)
|
db, date, user_ids, commit=False
|
||||||
|
)
|
||||||
|
|
||||||
stats.is_complete = True
|
stats.is_complete = True
|
||||||
stats.aggregated_at = datetime.now(timezone.utc)
|
stats.aggregated_at = datetime.now(timezone.utc)
|
||||||
@@ -1311,28 +1525,32 @@ def query_stats_hybrid(
|
|||||||
|
|
||||||
result = AggregatedStats()
|
result = AggregatedStats()
|
||||||
|
|
||||||
preaggregate_dates: list[datetime] = []
|
historical_dates = [day for day in complete_dates if day < today_utc]
|
||||||
realtime_dates: list[datetime] = []
|
realtime_dates = [day for day in complete_dates if day >= today_utc]
|
||||||
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)
|
|
||||||
|
|
||||||
# Pre-aggregated totals
|
preaggregated_by_date: dict[date, StatsDaily] = {}
|
||||||
for day_dt in preaggregate_dates:
|
if historical_dates:
|
||||||
stats = db.query(StatsDaily).filter(StatsDaily.date == day_dt).first()
|
historical_start = datetime.combine(min(historical_dates), time.min, tzinfo=timezone.utc)
|
||||||
if not stats:
|
historical_end = datetime.combine(
|
||||||
continue
|
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.total_requests += stats.total_requests
|
||||||
result.success_requests += stats.success_requests
|
result.success_requests += stats.success_requests
|
||||||
result.error_requests += stats.error_requests
|
result.error_requests += stats.error_requests
|
||||||
@@ -1346,9 +1564,10 @@ def query_stats_hybrid(
|
|||||||
result.actual_total_cost += float(stats.actual_total_cost or 0)
|
result.actual_total_cost += float(stats.actual_total_cost or 0)
|
||||||
result.total_response_time_ms += (stats.avg_response_time_ms or 0.0) * stats.total_requests
|
result.total_response_time_ms += (stats.avg_response_time_ms or 0.0) * stats.total_requests
|
||||||
|
|
||||||
# Realtime per day
|
missing_historical_dates = [day for day in historical_dates if day not in preaggregated_by_date]
|
||||||
for day_dt in realtime_dates:
|
realtime_ranges = _merge_consecutive_utc_days(missing_historical_dates + realtime_dates)
|
||||||
result.add(aggregate_usage_range(db, day_dt, day_dt + timedelta(days=1), filters=filters))
|
for range_start, range_end in realtime_ranges:
|
||||||
|
result.add(aggregate_usage_range(db, range_start, range_end, filters=filters))
|
||||||
|
|
||||||
if head_boundary:
|
if head_boundary:
|
||||||
result.add(aggregate_usage_range(db, head_boundary[0], head_boundary[1], filters=filters))
|
result.add(aggregate_usage_range(db, head_boundary[0], head_boundary[1], filters=filters))
|
||||||
|
|||||||
@@ -233,13 +233,14 @@ class UsageActiveRequestsMixin:
|
|||||||
default_timeout_seconds: int = 300,
|
default_timeout_seconds: int = 300,
|
||||||
*,
|
*,
|
||||||
include_admin_fields: bool = False,
|
include_admin_fields: bool = False,
|
||||||
|
maintain_status: bool | None = None,
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
获取活跃请求状态(用于前端轮询),并自动清理超时的 pending/streaming 请求
|
获取活跃请求状态(用于前端轮询)。
|
||||||
|
|
||||||
与 get_active_requests 不同,此方法:
|
与 get_active_requests 不同,此方法:
|
||||||
1. 返回轻量级的状态字典而非完整 Usage 对象
|
1. 返回轻量级的状态字典而非完整 Usage 对象
|
||||||
2. 自动检测并清理超时的 pending/streaming 请求
|
2. 可选地检测并清理超时的 pending/streaming 请求
|
||||||
3. 支持按 ID 列表查询特定请求
|
3. 支持按 ID 列表查询特定请求
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -247,6 +248,7 @@ class UsageActiveRequestsMixin:
|
|||||||
ids: 指定要查询的请求 ID 列表(可选)
|
ids: 指定要查询的请求 ID 列表(可选)
|
||||||
user_id: 限制只查询该用户的请求(可选,用于普通用户接口)
|
user_id: 限制只查询该用户的请求(可选,用于普通用户接口)
|
||||||
default_timeout_seconds: 默认超时时间(秒),当端点未配置时使用
|
default_timeout_seconds: 默认超时时间(秒),当端点未配置时使用
|
||||||
|
maintain_status: 是否执行超时修复与状态回写;默认仅在全量活跃请求查询时执行
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
请求状态列表
|
请求状态列表
|
||||||
@@ -297,28 +299,26 @@ class UsageActiveRequestsMixin:
|
|||||||
query = query.order_by(Usage.created_at.desc()).limit(50)
|
query = query.order_by(Usage.created_at.desc()).limit(50)
|
||||||
|
|
||||||
records = query.all()
|
records = query.all()
|
||||||
|
should_maintain_status = maintain_status if maintain_status is not None else not ids
|
||||||
|
|
||||||
# 检查超时的 pending/streaming 请求
|
# 检查超时的 pending/streaming 请求
|
||||||
# 收集可能超时的 usage_id 列表
|
# 收集可能超时的 usage_id 列表
|
||||||
timeout_candidates: list[str] = []
|
timeout_candidates: list[str] = []
|
||||||
for r in records:
|
if should_maintain_status:
|
||||||
if r.status in ("pending", "streaming") and r.created_at:
|
for r in records:
|
||||||
# 使用全局配置的超时时间
|
if r.status in ("pending", "streaming") and r.created_at:
|
||||||
timeout_seconds = default_timeout_seconds
|
timeout_seconds = default_timeout_seconds
|
||||||
|
|
||||||
# 处理时区:如果 created_at 没有时区信息,假定为 UTC
|
created_at = r.created_at
|
||||||
created_at = r.created_at
|
if created_at.tzinfo is None:
|
||||||
if created_at.tzinfo is None:
|
created_at = created_at.replace(tzinfo=timezone.utc)
|
||||||
created_at = created_at.replace(tzinfo=timezone.utc)
|
elapsed = (now - created_at).total_seconds()
|
||||||
elapsed = (now - created_at).total_seconds()
|
if elapsed > timeout_seconds:
|
||||||
if elapsed > timeout_seconds:
|
timeout_candidates.append(r.id)
|
||||||
# 需要获取 request_id 以便检查 RequestCandidate 表
|
|
||||||
# r.id 是 usage_id,需要查询 request_id
|
|
||||||
timeout_candidates.append(r.id)
|
|
||||||
|
|
||||||
# 批量更新超时的请求(排除已有成功完成记录的请求)
|
# 批量更新超时的请求(排除已有成功完成记录的请求)
|
||||||
timeout_ids = []
|
timeout_ids = []
|
||||||
if timeout_candidates:
|
if should_maintain_status and timeout_candidates:
|
||||||
# 先获取这些 Usage 的 request_id
|
# 先获取这些 Usage 的 request_id
|
||||||
usage_request_ids = (
|
usage_request_ids = (
|
||||||
db.query(Usage.id, Usage.request_id).filter(Usage.id.in_(timeout_candidates)).all()
|
db.query(Usage.id, Usage.request_id).filter(Usage.id.in_(timeout_candidates)).all()
|
||||||
@@ -360,7 +360,8 @@ class UsageActiveRequestsMixin:
|
|||||||
cls._sync_candidate_status_to_success(db, completed_request_ids)
|
cls._sync_candidate_status_to_success(db, completed_request_ids)
|
||||||
db.commit()
|
db.commit()
|
||||||
logger.info(
|
logger.info(
|
||||||
f"[Usage] 恢复 {len(completed_usage_ids)} 个已完成请求的状态(遥测回调丢失)"
|
"[Usage] 恢复 {} 个已完成请求的状态(遥测回调丢失)",
|
||||||
|
len(completed_usage_ids),
|
||||||
)
|
)
|
||||||
|
|
||||||
result: list[dict[str, Any]] = []
|
result: list[dict[str, Any]] = []
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -10,6 +11,13 @@ from src.core.logger import logger
|
|||||||
from src.models.database import ApiKey, Usage, User
|
from src.models.database import ApiKey, Usage, User
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(slots=True)
|
||||||
|
class RequestBalanceCheckResult:
|
||||||
|
allowed: bool
|
||||||
|
message: str
|
||||||
|
remaining: float | None
|
||||||
|
|
||||||
|
|
||||||
class UsageQueryMixin:
|
class UsageQueryMixin:
|
||||||
"""查询/统计相关方法"""
|
"""查询/统计相关方法"""
|
||||||
|
|
||||||
@@ -125,14 +133,14 @@ class UsageQueryMixin:
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def check_request_balance(
|
def check_request_balance_details(
|
||||||
db: Session,
|
db: Session,
|
||||||
user: User,
|
user: User,
|
||||||
estimated_tokens: int = 0,
|
estimated_tokens: int = 0,
|
||||||
estimated_cost: float = 0,
|
estimated_cost: float = 0,
|
||||||
api_key: ApiKey | None = None,
|
api_key: ApiKey | None = None,
|
||||||
) -> tuple[bool, str]:
|
) -> RequestBalanceCheckResult:
|
||||||
"""检查请求是否满足余额条件(支持独立 Key)。"""
|
"""Return a structured balance-check result."""
|
||||||
from src.services.wallet import WalletService
|
from src.services.wallet import WalletService
|
||||||
|
|
||||||
wallet_access = WalletService.check_request_allowed(
|
wallet_access = WalletService.check_request_allowed(
|
||||||
@@ -140,29 +148,52 @@ class UsageQueryMixin:
|
|||||||
user=None if (api_key and api_key.is_standalone) else user,
|
user=None if (api_key and api_key.is_standalone) else user,
|
||||||
api_key=api_key,
|
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:
|
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:
|
if api_key and api_key.is_standalone:
|
||||||
return False, "Key欠费,请先调账或充值"
|
return RequestBalanceCheckResult(False, "Key欠费,请先调账或充值", remaining)
|
||||||
return False, "账户欠费,请先充值"
|
return RequestBalanceCheckResult(False, "账户欠费,请先充值", remaining)
|
||||||
|
|
||||||
if wallet_access.message == "钱包不可用":
|
if wallet_access.message == "钱包不可用":
|
||||||
if api_key and api_key.is_standalone:
|
if api_key and api_key.is_standalone:
|
||||||
return False, "Key钱包不可用"
|
return RequestBalanceCheckResult(False, "Key钱包不可用", remaining)
|
||||||
return False, "钱包不可用"
|
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 api_key and api_key.is_standalone:
|
||||||
if remaining is None:
|
if remaining is None:
|
||||||
return False, "Key余额不足"
|
return RequestBalanceCheckResult(False, "Key余额不足", remaining)
|
||||||
return False, f"Key余额不足(剩余: ${remaining:.2f})"
|
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:
|
if remaining is None:
|
||||||
return False, wallet_access.message or "余额不足"
|
return RequestBalanceCheckResult(False, wallet_access.message or "余额不足", remaining)
|
||||||
return False, f"余额不足(剩余: ${remaining:.2f})"
|
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
|
@staticmethod
|
||||||
def get_usage_summary(
|
def get_usage_summary(
|
||||||
@@ -171,7 +202,7 @@ class UsageQueryMixin:
|
|||||||
api_key_id: str | None = None,
|
api_key_id: str | None = None,
|
||||||
start_date: datetime | None = None,
|
start_date: datetime | None = None,
|
||||||
end_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]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""获取使用汇总"""
|
"""获取使用汇总"""
|
||||||
|
|
||||||
@@ -188,34 +219,35 @@ class UsageQueryMixin:
|
|||||||
if end_date:
|
if end_date:
|
||||||
query = query.filter(Usage.created_at < end_date)
|
query = query.filter(Usage.created_at < end_date)
|
||||||
|
|
||||||
# 使用跨数据库可用的日期函数
|
select_columns = [Usage.provider_name, Usage.model]
|
||||||
from src.utils.database_helpers import date_trunc_portable
|
group_columns = [Usage.provider_name, Usage.model]
|
||||||
|
|
||||||
# 检测数据库方言
|
if group_by is not None:
|
||||||
bind = db.bind
|
from src.utils.database_helpers import date_trunc_portable
|
||||||
dialect = bind.dialect.name if bind is not None else "sqlite"
|
|
||||||
|
|
||||||
# 根据分组类型选择日期函数(适配多种数据库)
|
bind = db.bind
|
||||||
if group_by == "day":
|
dialect = bind.dialect.name if bind is not None else "sqlite"
|
||||||
date_func = date_trunc_portable(dialect, "day", Usage.created_at)
|
|
||||||
elif group_by == "week":
|
if group_by == "day":
|
||||||
date_func = date_trunc_portable(dialect, "week", Usage.created_at)
|
date_func = date_trunc_portable(dialect, "day", Usage.created_at)
|
||||||
elif group_by == "month":
|
elif group_by == "week":
|
||||||
date_func = date_trunc_portable(dialect, "month", Usage.created_at)
|
date_func = date_trunc_portable(dialect, "week", Usage.created_at)
|
||||||
else:
|
elif group_by == "month":
|
||||||
# 默认按天分组
|
date_func = date_trunc_portable(dialect, "month", Usage.created_at)
|
||||||
date_func = date_trunc_portable(dialect, "day", 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(
|
summary = db.query(
|
||||||
date_func.label("period"),
|
*select_columns,
|
||||||
Usage.provider_name,
|
|
||||||
Usage.model,
|
|
||||||
func.count(Usage.id).label("requests"),
|
func.count(Usage.id).label("requests"),
|
||||||
func.sum(Usage.input_tokens).label("input_tokens"),
|
func.sum(Usage.input_tokens).label("input_tokens"),
|
||||||
func.sum(Usage.output_tokens).label("output_tokens"),
|
func.sum(Usage.output_tokens).label("output_tokens"),
|
||||||
func.sum(Usage.total_tokens).label("total_tokens"),
|
func.sum(Usage.total_tokens).label("total_tokens"),
|
||||||
func.sum(Usage.total_cost_usd).label("total_cost_usd"),
|
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.avg(Usage.response_time_ms).label("avg_response_time"),
|
||||||
func.sum(
|
func.sum(
|
||||||
case(
|
case(
|
||||||
@@ -246,18 +278,20 @@ class UsageQueryMixin:
|
|||||||
if end_date:
|
if end_date:
|
||||||
summary = summary.filter(Usage.created_at < 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 [
|
return [
|
||||||
{
|
{
|
||||||
"period": row.period,
|
"period": getattr(row, "period", None),
|
||||||
"provider": row.provider_name,
|
"provider": row.provider_name,
|
||||||
"model": row.model,
|
"model": row.model,
|
||||||
"requests": row.requests,
|
"requests": row.requests,
|
||||||
"input_tokens": row.input_tokens,
|
"input_tokens": row.input_tokens,
|
||||||
"output_tokens": row.output_tokens,
|
"output_tokens": row.output_tokens,
|
||||||
"total_tokens": row.total_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": (
|
"avg_response_time_ms": (
|
||||||
float(row.avg_response_time) if row.avg_response_time else 0
|
float(row.avg_response_time) if row.avg_response_time else 0
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -43,6 +43,7 @@ class WalletAccessResult:
|
|||||||
remaining: Decimal | None
|
remaining: Decimal | None
|
||||||
message: str
|
message: str
|
||||||
wallet: Wallet | None = None
|
wallet: Wallet | None = None
|
||||||
|
balance_snapshot: Decimal | None = None
|
||||||
|
|
||||||
|
|
||||||
class WalletService:
|
class WalletService:
|
||||||
@@ -229,6 +230,17 @@ class WalletService:
|
|||||||
return db.query(Wallet).filter(Wallet.user_id == user_id).first()
|
return db.query(Wallet).filter(Wallet.user_id == user_id).first()
|
||||||
return None
|
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
|
@classmethod
|
||||||
def get_or_create_wallet(
|
def get_or_create_wallet(
|
||||||
cls,
|
cls,
|
||||||
@@ -291,6 +303,18 @@ class WalletService:
|
|||||||
return wallet
|
return wallet
|
||||||
raise
|
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
|
@classmethod
|
||||||
def check_request_allowed(
|
def check_request_allowed(
|
||||||
cls,
|
cls,
|
||||||
@@ -299,25 +323,29 @@ class WalletService:
|
|||||||
user: User | None,
|
user: User | None,
|
||||||
api_key: ApiKey | None = None,
|
api_key: ApiKey | None = None,
|
||||||
) -> WalletAccessResult:
|
) -> 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)
|
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:
|
if wallet is None:
|
||||||
return WalletAccessResult(False, Decimal("0"), "钱包不存在", None)
|
return WalletAccessResult(False, Decimal("0"), "钱包不存在", None, None)
|
||||||
|
|
||||||
remaining = cls.get_spendable_balance_value(wallet)
|
remaining = cls.get_spendable_balance_value(wallet)
|
||||||
recharge_balance = cls.get_recharge_balance_value(wallet)
|
recharge_balance = cls.get_recharge_balance_value(wallet)
|
||||||
if wallet.status != "active":
|
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"):
|
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):
|
if cls.is_unlimited_wallet(wallet):
|
||||||
return WalletAccessResult(True, None, "OK", wallet)
|
return WalletAccessResult(True, None, "OK", wallet, balance_snapshot)
|
||||||
if remaining <= Decimal("0"):
|
if remaining <= Decimal("0"):
|
||||||
return WalletAccessResult(False, remaining, "钱包余额不足", wallet)
|
return WalletAccessResult(False, remaining, "钱包余额不足", wallet, balance_snapshot)
|
||||||
return WalletAccessResult(True, remaining, "OK", wallet)
|
return WalletAccessResult(True, remaining, "OK", wallet, balance_snapshot)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_balance_snapshot(
|
def get_balance_snapshot(
|
||||||
@@ -328,14 +356,7 @@ class WalletService:
|
|||||||
api_key: ApiKey | None = None,
|
api_key: ApiKey | None = None,
|
||||||
) -> Decimal | None:
|
) -> Decimal | None:
|
||||||
wallet = cls.get_or_create_wallet(db, user=user, api_key=api_key)
|
wallet = cls.get_or_create_wallet(db, user=user, api_key=api_key)
|
||||||
if wallet is None:
|
return cls._get_balance_snapshot_from_wallet(wallet)
|
||||||
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
|
@classmethod
|
||||||
def _resolve_wallet_for_usage(
|
def _resolve_wallet_for_usage(
|
||||||
|
|||||||
93
tests/api/test_admin_user_routes.py
Normal file
93
tests/api/test_admin_user_routes.py
Normal file
@@ -0,0 +1,93 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from typing import Any
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi import FastAPI
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
from src.api.admin.users.routes import router as admin_users_router
|
||||||
|
from src.database import get_db
|
||||||
|
|
||||||
|
|
||||||
|
def _build_admin_users_app(db: MagicMock, monkeypatch: pytest.MonkeyPatch) -> TestClient:
|
||||||
|
app = FastAPI()
|
||||||
|
app.include_router(admin_users_router)
|
||||||
|
app.dependency_overrides[get_db] = lambda: db
|
||||||
|
|
||||||
|
async def _fake_pipeline_run(
|
||||||
|
*, adapter: Any, http_request: object, db: MagicMock, mode: object
|
||||||
|
) -> Any:
|
||||||
|
_ = http_request, mode
|
||||||
|
context = SimpleNamespace(
|
||||||
|
db=db,
|
||||||
|
request=SimpleNamespace(state=SimpleNamespace()),
|
||||||
|
user=SimpleNamespace(id="admin-1"),
|
||||||
|
ensure_json_body=lambda: {},
|
||||||
|
add_audit_metadata=lambda **_: None,
|
||||||
|
)
|
||||||
|
return await adapter.handle(context)
|
||||||
|
|
||||||
|
monkeypatch.setattr("src.api.admin.users.routes.pipeline.run", _fake_pipeline_run)
|
||||||
|
return TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
|
def test_list_users_uses_wallet_batch_lookup(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
db = MagicMock()
|
||||||
|
client = _build_admin_users_app(db, monkeypatch)
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
users = [
|
||||||
|
SimpleNamespace(
|
||||||
|
id="user-1",
|
||||||
|
email="u1@example.com",
|
||||||
|
username="user1",
|
||||||
|
role=SimpleNamespace(value="user"),
|
||||||
|
allowed_providers=None,
|
||||||
|
allowed_api_formats=None,
|
||||||
|
allowed_models=None,
|
||||||
|
is_active=True,
|
||||||
|
created_at=now,
|
||||||
|
updated_at=now,
|
||||||
|
last_login_at=None,
|
||||||
|
),
|
||||||
|
SimpleNamespace(
|
||||||
|
id="user-2",
|
||||||
|
email="u2@example.com",
|
||||||
|
username="user2",
|
||||||
|
role=SimpleNamespace(value="admin"),
|
||||||
|
allowed_providers=None,
|
||||||
|
allowed_api_formats=None,
|
||||||
|
allowed_models=None,
|
||||||
|
is_active=True,
|
||||||
|
created_at=now,
|
||||||
|
updated_at=None,
|
||||||
|
last_login_at=None,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
wallets_by_user_id = {
|
||||||
|
"user-1": SimpleNamespace(limit_mode="unlimited"),
|
||||||
|
}
|
||||||
|
batch_getter = MagicMock(return_value=wallets_by_user_id)
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.api.admin.users.routes.UserService.list_users", lambda *_a, **_k: users
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.api.admin.users.routes.WalletService.get_wallets_by_user_ids",
|
||||||
|
batch_getter,
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.api.admin.users.routes.WalletService.get_wallet",
|
||||||
|
lambda *_a, **_k: (_ for _ in ()).throw(AssertionError("不应回退到逐个钱包查询")),
|
||||||
|
)
|
||||||
|
|
||||||
|
response = client.get("/api/admin/users")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json()[0]["unlimited"] is True
|
||||||
|
assert response.json()[1]["unlimited"] is False
|
||||||
|
batch_getter.assert_called_once()
|
||||||
|
assert batch_getter.call_args.args[1] == ["user-1", "user-2"]
|
||||||
@@ -8,54 +8,139 @@ API Pipeline 测试
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from fastapi import HTTPException
|
from fastapi import HTTPException
|
||||||
|
|
||||||
|
from src.api.base.adapter import ApiMode
|
||||||
from src.api.base.pipeline import ApiRequestPipeline
|
from src.api.base.pipeline import ApiRequestPipeline
|
||||||
from src.core.enums import UserRole
|
from src.core.enums import UserRole
|
||||||
|
from src.core.modules.hooks import AUTH_TOKEN_PREFIX_AUTHENTICATORS
|
||||||
|
|
||||||
|
|
||||||
class TestPipelineBalanceCalculation:
|
class TestPipelineBalanceCalculation:
|
||||||
"""测试 Pipeline 余额计算"""
|
"""Balance calculation tests for Pipeline."""
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def pipeline(self) -> ApiRequestPipeline:
|
def pipeline(self) -> ApiRequestPipeline:
|
||||||
return ApiRequestPipeline()
|
return ApiRequestPipeline()
|
||||||
|
|
||||||
def test_calculate_balance_remaining_with_balance(self, pipeline: ApiRequestPipeline) -> None:
|
@pytest.mark.asyncio
|
||||||
"""测试有限制钱包时计算剩余余额"""
|
async def test_calculate_balance_remaining_with_balance(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
"""Returns remaining balance for limited wallets."""
|
||||||
mock_user = MagicMock()
|
mock_user = MagicMock()
|
||||||
mock_db = MagicMock()
|
mock_user.id = "user-123"
|
||||||
|
|
||||||
with patch(
|
thread_db = MagicMock()
|
||||||
"src.api.base.pipeline.WalletService.get_balance_snapshot",
|
db_user = MagicMock()
|
||||||
return_value=70.0,
|
thread_db.query.return_value.filter.return_value.first.return_value = db_user
|
||||||
):
|
|
||||||
remaining = pipeline._calculate_balance_remaining(mock_db, mock_user)
|
with patch("src.api.base.pipeline.create_session", return_value=thread_db):
|
||||||
|
with patch(
|
||||||
|
"src.api.base.pipeline.WalletService.get_balance_snapshot",
|
||||||
|
return_value=70.0,
|
||||||
|
):
|
||||||
|
remaining = await pipeline._calculate_balance_remaining_async(mock_user)
|
||||||
|
|
||||||
assert remaining == 70.0
|
assert remaining == 70.0
|
||||||
|
thread_db.close.assert_called_once()
|
||||||
|
|
||||||
def test_calculate_balance_remaining_unlimited(self, pipeline: ApiRequestPipeline) -> None:
|
@pytest.mark.asyncio
|
||||||
"""测试无限制钱包时返回 None"""
|
async def test_calculate_balance_remaining_unlimited(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
"""Returns None for unlimited wallets."""
|
||||||
mock_user = MagicMock()
|
mock_user = MagicMock()
|
||||||
mock_db = MagicMock()
|
mock_user.id = "user-123"
|
||||||
|
|
||||||
with patch(
|
thread_db = MagicMock()
|
||||||
"src.api.base.pipeline.WalletService.get_balance_snapshot",
|
db_user = MagicMock()
|
||||||
return_value=None,
|
thread_db.query.return_value.filter.return_value.first.return_value = db_user
|
||||||
|
|
||||||
|
with patch("src.api.base.pipeline.create_session", return_value=thread_db):
|
||||||
|
with patch(
|
||||||
|
"src.api.base.pipeline.WalletService.get_balance_snapshot",
|
||||||
|
return_value=None,
|
||||||
|
):
|
||||||
|
remaining = await pipeline._calculate_balance_remaining_async(mock_user)
|
||||||
|
|
||||||
|
assert remaining is None
|
||||||
|
thread_db.close.assert_called_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_calculate_balance_remaining_none_user(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
"""Returns None when user is missing."""
|
||||||
|
remaining = await pipeline._calculate_balance_remaining_async(None)
|
||||||
|
|
||||||
|
assert remaining is None
|
||||||
|
|
||||||
|
|
||||||
|
class TestPipelineRunModes:
|
||||||
|
"""Returns remaining balance for limited wallets."""
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def pipeline(self) -> ApiRequestPipeline:
|
||||||
|
return ApiRequestPipeline()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_management_mode_skips_balance_calculation(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
"""Management mode skips balance calculation."""
|
||||||
|
mock_request = MagicMock()
|
||||||
|
mock_request.method = "GET"
|
||||||
|
mock_request.url.path = "/api/admin/tokens"
|
||||||
|
mock_request.state = MagicMock()
|
||||||
|
|
||||||
|
mock_db = MagicMock()
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "admin-123"
|
||||||
|
mock_token = MagicMock()
|
||||||
|
mock_token.id = "mt-123"
|
||||||
|
|
||||||
|
mock_adapter = MagicMock()
|
||||||
|
mock_adapter.name = "test-adapter"
|
||||||
|
mock_adapter.authorize = MagicMock(return_value=None)
|
||||||
|
mock_response = MagicMock()
|
||||||
|
mock_response.status_code = 200
|
||||||
|
mock_adapter.handle = AsyncMock(return_value=mock_response)
|
||||||
|
|
||||||
|
mock_context = MagicMock()
|
||||||
|
mock_context.db = mock_db
|
||||||
|
mock_context.request = mock_request
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
pipeline,
|
||||||
|
"_authenticate_management",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value=(mock_user, mock_token),
|
||||||
):
|
):
|
||||||
remaining = pipeline._calculate_balance_remaining(mock_db, mock_user)
|
with patch(
|
||||||
|
"src.api.base.pipeline.ApiRequestContext.build",
|
||||||
|
return_value=mock_context,
|
||||||
|
):
|
||||||
|
with patch.object(
|
||||||
|
pipeline,
|
||||||
|
"_calculate_balance_remaining_async",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
) as mock_balance:
|
||||||
|
with patch.object(pipeline, "_record_audit_event"):
|
||||||
|
response = await pipeline.run(
|
||||||
|
mock_adapter,
|
||||||
|
mock_request,
|
||||||
|
mock_db,
|
||||||
|
mode=ApiMode.MANAGEMENT,
|
||||||
|
)
|
||||||
|
|
||||||
assert remaining is None
|
assert response == mock_response
|
||||||
|
assert mock_context.management_token == mock_token
|
||||||
def test_calculate_balance_remaining_none_user(self, pipeline: ApiRequestPipeline) -> None:
|
mock_balance.assert_not_called()
|
||||||
"""测试用户为 None 时返回 None"""
|
|
||||||
mock_db = MagicMock()
|
|
||||||
remaining = pipeline._calculate_balance_remaining(mock_db, None)
|
|
||||||
|
|
||||||
assert remaining is None
|
|
||||||
|
|
||||||
|
|
||||||
class TestPipelineAuditLogging:
|
class TestPipelineAuditLogging:
|
||||||
@@ -213,7 +298,8 @@ class TestPipelineAuthentication:
|
|||||||
def pipeline(self) -> ApiRequestPipeline:
|
def pipeline(self) -> ApiRequestPipeline:
|
||||||
return ApiRequestPipeline()
|
return ApiRequestPipeline()
|
||||||
|
|
||||||
def test_authenticate_client_missing_key(self, pipeline: ApiRequestPipeline) -> None:
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_client_missing_key(self, pipeline: ApiRequestPipeline) -> None:
|
||||||
"""测试缺少 API Key 时抛出异常"""
|
"""测试缺少 API Key 时抛出异常"""
|
||||||
mock_request = MagicMock()
|
mock_request = MagicMock()
|
||||||
mock_request.headers = {}
|
mock_request.headers = {}
|
||||||
@@ -226,12 +312,13 @@ class TestPipelineAuthentication:
|
|||||||
mock_adapter.extract_api_key = MagicMock(return_value=None)
|
mock_adapter.extract_api_key = MagicMock(return_value=None)
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
|
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
|
||||||
|
|
||||||
assert exc_info.value.status_code == 401
|
assert exc_info.value.status_code == 401
|
||||||
assert "API密钥" in exc_info.value.detail
|
assert "API密钥" in exc_info.value.detail
|
||||||
|
|
||||||
def test_authenticate_client_invalid_key(self, pipeline: ApiRequestPipeline) -> None:
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_client_invalid_key(self, pipeline: ApiRequestPipeline) -> None:
|
||||||
"""测试无效的 API Key"""
|
"""测试无效的 API Key"""
|
||||||
mock_request = MagicMock()
|
mock_request = MagicMock()
|
||||||
mock_request.headers = {"Authorization": "Bearer sk-invalid"}
|
mock_request.headers = {"Authorization": "Bearer sk-invalid"}
|
||||||
@@ -245,15 +332,17 @@ class TestPipelineAuthentication:
|
|||||||
|
|
||||||
with patch.object(
|
with patch.object(
|
||||||
pipeline.auth_service,
|
pipeline.auth_service,
|
||||||
"authenticate_api_key",
|
"authenticate_api_key_threadsafe",
|
||||||
|
new_callable=AsyncMock,
|
||||||
return_value=None,
|
return_value=None,
|
||||||
):
|
):
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
|
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
|
||||||
|
|
||||||
assert exc_info.value.status_code == 401
|
assert exc_info.value.status_code == 401
|
||||||
|
|
||||||
def test_authenticate_client_balance_exceeded(self, pipeline: ApiRequestPipeline) -> None:
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_client_balance_exceeded(self, pipeline: ApiRequestPipeline) -> None:
|
||||||
"""测试余额不足时抛出异常"""
|
"""测试余额不足时抛出异常"""
|
||||||
mock_user = MagicMock()
|
mock_user = MagicMock()
|
||||||
mock_user.id = "user-123"
|
mock_user.id = "user-123"
|
||||||
@@ -268,28 +357,198 @@ class TestPipelineAuthentication:
|
|||||||
mock_request.state = MagicMock()
|
mock_request.state = MagicMock()
|
||||||
|
|
||||||
mock_db = MagicMock()
|
mock_db = MagicMock()
|
||||||
|
db_user = MagicMock()
|
||||||
|
db_user.id = "user-123"
|
||||||
|
db_user.is_active = True
|
||||||
|
db_user.is_deleted = False
|
||||||
|
db_api_key = MagicMock()
|
||||||
|
db_api_key.id = "key-123"
|
||||||
|
db_api_key.user_id = "user-123"
|
||||||
|
db_api_key.is_active = True
|
||||||
|
db_api_key.is_locked = False
|
||||||
|
db_api_key.is_standalone = False
|
||||||
|
db_api_key.expires_at = None
|
||||||
|
user_query = MagicMock()
|
||||||
|
user_query.filter.return_value.first.return_value = db_user
|
||||||
|
api_key_query = MagicMock()
|
||||||
|
api_key_query.filter.return_value.first.return_value = db_api_key
|
||||||
|
mock_db.query.side_effect = [user_query, api_key_query]
|
||||||
|
|
||||||
mock_adapter = MagicMock()
|
mock_adapter = MagicMock()
|
||||||
mock_adapter.extract_api_key = MagicMock(return_value="sk-test")
|
mock_adapter.extract_api_key = MagicMock(return_value="sk-test")
|
||||||
|
|
||||||
with patch.object(
|
with patch.object(
|
||||||
pipeline.auth_service,
|
pipeline.auth_service,
|
||||||
"authenticate_api_key",
|
"authenticate_api_key_threadsafe",
|
||||||
return_value=(mock_user, mock_api_key),
|
new_callable=AsyncMock,
|
||||||
|
return_value=MagicMock(
|
||||||
|
user=mock_user,
|
||||||
|
api_key=mock_api_key,
|
||||||
|
access_allowed=False,
|
||||||
|
balance_remaining=0.0,
|
||||||
|
),
|
||||||
):
|
):
|
||||||
with patch.object(
|
from src.core.exceptions import BalanceInsufficientException
|
||||||
pipeline.usage_service,
|
|
||||||
"check_request_balance",
|
|
||||||
return_value=(False, "余额不足"),
|
|
||||||
):
|
|
||||||
with patch(
|
|
||||||
"src.api.base.pipeline.WalletService.get_balance_snapshot",
|
|
||||||
return_value=0.0,
|
|
||||||
):
|
|
||||||
from src.core.exceptions import BalanceInsufficientException
|
|
||||||
|
|
||||||
with pytest.raises(BalanceInsufficientException):
|
with pytest.raises(BalanceInsufficientException):
|
||||||
pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
|
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_client_requery_detects_inactive_user(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "user-123"
|
||||||
|
mock_api_key = MagicMock()
|
||||||
|
mock_api_key.id = "key-123"
|
||||||
|
|
||||||
|
mock_request = MagicMock()
|
||||||
|
mock_request.headers = {"Authorization": "Bearer sk-test"}
|
||||||
|
mock_request.url.path = "/v1/messages"
|
||||||
|
mock_request.state = MagicMock()
|
||||||
|
|
||||||
|
mock_adapter = MagicMock()
|
||||||
|
mock_adapter.extract_api_key = MagicMock(return_value="sk-test")
|
||||||
|
|
||||||
|
db_user = MagicMock()
|
||||||
|
db_user.id = "user-123"
|
||||||
|
db_user.is_active = False
|
||||||
|
db_user.is_deleted = False
|
||||||
|
db_api_key = MagicMock()
|
||||||
|
db_api_key.id = "key-123"
|
||||||
|
db_api_key.user_id = "user-123"
|
||||||
|
db_api_key.is_active = True
|
||||||
|
db_api_key.is_locked = False
|
||||||
|
db_api_key.is_standalone = False
|
||||||
|
db_api_key.expires_at = None
|
||||||
|
|
||||||
|
mock_db = MagicMock()
|
||||||
|
user_query = MagicMock()
|
||||||
|
user_query.filter.return_value.first.return_value = db_user
|
||||||
|
api_key_query = MagicMock()
|
||||||
|
api_key_query.filter.return_value.first.return_value = db_api_key
|
||||||
|
mock_db.query.side_effect = [user_query, api_key_query]
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
pipeline.auth_service,
|
||||||
|
"authenticate_api_key_threadsafe",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value=MagicMock(
|
||||||
|
user=mock_user,
|
||||||
|
api_key=mock_api_key,
|
||||||
|
access_allowed=True,
|
||||||
|
balance_remaining=10.0,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
|
||||||
|
|
||||||
|
assert exc_info.value.status_code == 401
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_client_requery_detects_locked_key(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "user-123"
|
||||||
|
mock_api_key = MagicMock()
|
||||||
|
mock_api_key.id = "key-123"
|
||||||
|
|
||||||
|
mock_request = MagicMock()
|
||||||
|
mock_request.headers = {"Authorization": "Bearer sk-test"}
|
||||||
|
mock_request.url.path = "/v1/messages"
|
||||||
|
mock_request.state = MagicMock()
|
||||||
|
|
||||||
|
mock_adapter = MagicMock()
|
||||||
|
mock_adapter.extract_api_key = MagicMock(return_value="sk-test")
|
||||||
|
|
||||||
|
db_user = MagicMock()
|
||||||
|
db_user.id = "user-123"
|
||||||
|
db_user.is_active = True
|
||||||
|
db_user.is_deleted = False
|
||||||
|
db_api_key = MagicMock()
|
||||||
|
db_api_key.id = "key-123"
|
||||||
|
db_api_key.user_id = "user-123"
|
||||||
|
db_api_key.is_active = True
|
||||||
|
db_api_key.is_locked = True
|
||||||
|
db_api_key.is_standalone = False
|
||||||
|
db_api_key.expires_at = None
|
||||||
|
|
||||||
|
mock_db = MagicMock()
|
||||||
|
user_query = MagicMock()
|
||||||
|
user_query.filter.return_value.first.return_value = db_user
|
||||||
|
api_key_query = MagicMock()
|
||||||
|
api_key_query.filter.return_value.first.return_value = db_api_key
|
||||||
|
mock_db.query.side_effect = [user_query, api_key_query]
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
pipeline.auth_service,
|
||||||
|
"authenticate_api_key_threadsafe",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value=MagicMock(
|
||||||
|
user=mock_user,
|
||||||
|
api_key=mock_api_key,
|
||||||
|
access_allowed=True,
|
||||||
|
balance_remaining=10.0,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
|
||||||
|
|
||||||
|
assert exc_info.value.status_code == 403
|
||||||
|
assert "锁定" in str(exc_info.value.detail)
|
||||||
|
|
||||||
|
|
||||||
|
class TestPipelineTokenPrefixAuth:
|
||||||
|
"""Tests token-prefix auth isolation."""
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def pipeline(self) -> ApiRequestPipeline:
|
||||||
|
return ApiRequestPipeline()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_try_token_prefix_auth_uses_isolated_session(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
mock_request = MagicMock()
|
||||||
|
mock_request.headers = {}
|
||||||
|
mock_request.client = MagicMock(host="127.0.0.1")
|
||||||
|
|
||||||
|
route_db = MagicMock()
|
||||||
|
auth_db = MagicMock()
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_token = MagicMock()
|
||||||
|
|
||||||
|
async def authenticate(db: Any, token: str, client_ip: str) -> tuple[Any, Any]:
|
||||||
|
assert db is auth_db
|
||||||
|
assert token == "ae_test"
|
||||||
|
assert client_ip == "127.0.0.1"
|
||||||
|
return mock_user, mock_token
|
||||||
|
|
||||||
|
with patch("src.api.base.pipeline.create_session", return_value=auth_db):
|
||||||
|
with patch("src.utils.request_utils.get_client_ip", return_value="127.0.0.1"):
|
||||||
|
with patch("src.core.modules.hooks.get_hook_dispatcher") as mock_get_dispatcher:
|
||||||
|
dispatcher = MagicMock()
|
||||||
|
dispatcher.dispatch = AsyncMock(
|
||||||
|
return_value=[
|
||||||
|
{
|
||||||
|
"prefix": "ae_",
|
||||||
|
"module": "management_tokens",
|
||||||
|
"authenticate": authenticate,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
)
|
||||||
|
mock_get_dispatcher.return_value = dispatcher
|
||||||
|
|
||||||
|
result = await pipeline._try_token_prefix_auth(
|
||||||
|
"ae_test", mock_request, route_db
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == (mock_user, mock_token)
|
||||||
|
dispatcher.dispatch.assert_awaited_once_with(AUTH_TOKEN_PREFIX_AUTHENTICATORS)
|
||||||
|
auth_db.expunge.assert_any_call(mock_user)
|
||||||
|
auth_db.expunge.assert_any_call(mock_token)
|
||||||
|
auth_db.close.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
class TestPipelineAdminAuth:
|
class TestPipelineAdminAuth:
|
||||||
|
|||||||
175
tests/api/test_user_me_usage_routes.py
Normal file
175
tests/api/test_user_me_usage_routes.py
Normal file
@@ -0,0 +1,175 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.api.user_me.routes import GetUsageAdapter
|
||||||
|
from src.core.enums import UserRole
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_usage_adapter_uses_coarse_summary_grouping(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
db = MagicMock()
|
||||||
|
query = MagicMock()
|
||||||
|
count_query = MagicMock()
|
||||||
|
count_query.scalar.return_value = 0
|
||||||
|
query.outerjoin.return_value = query
|
||||||
|
query.filter.return_value = query
|
||||||
|
query.with_entities.return_value = count_query
|
||||||
|
query.options.return_value = query
|
||||||
|
query.order_by.return_value = query
|
||||||
|
query.offset.return_value = query
|
||||||
|
query.limit.return_value = query
|
||||||
|
query.all.return_value = []
|
||||||
|
db.query.return_value = query
|
||||||
|
|
||||||
|
summary_getter = MagicMock(
|
||||||
|
return_value=[
|
||||||
|
{
|
||||||
|
"provider": "provider-a",
|
||||||
|
"model": "gpt-4o",
|
||||||
|
"requests": 2,
|
||||||
|
"input_tokens": 10,
|
||||||
|
"output_tokens": 5,
|
||||||
|
"total_tokens": 15,
|
||||||
|
"total_cost_usd": 1.5,
|
||||||
|
"actual_total_cost_usd": 1.2,
|
||||||
|
"success_count": 2,
|
||||||
|
"success_response_time_sum_ms": 1000.0,
|
||||||
|
"success_response_time_count": 2,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"provider": "pending",
|
||||||
|
"model": "gpt-4o",
|
||||||
|
"requests": 99,
|
||||||
|
"input_tokens": 999,
|
||||||
|
"output_tokens": 999,
|
||||||
|
"total_tokens": 1998,
|
||||||
|
"total_cost_usd": 9.9,
|
||||||
|
"actual_total_cost_usd": 9.9,
|
||||||
|
"success_count": 0,
|
||||||
|
"success_response_time_sum_ms": 0.0,
|
||||||
|
"success_response_time_count": 0,
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("src.api.user_me.routes.UsageService.get_usage_summary", summary_getter)
|
||||||
|
monkeypatch.setattr("src.api.user_me.routes.WalletService.get_wallet", lambda *_a, **_k: None)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.api.user_me.routes.WalletService.serialize_wallet_summary",
|
||||||
|
lambda _wallet: {"limit_mode": "finite"},
|
||||||
|
)
|
||||||
|
|
||||||
|
adapter = GetUsageAdapter(time_range=None, limit=20, offset=0)
|
||||||
|
context = SimpleNamespace(
|
||||||
|
db=db,
|
||||||
|
user=SimpleNamespace(id="user-1", role=UserRole.USER),
|
||||||
|
request=SimpleNamespace(state=SimpleNamespace()),
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await adapter.handle(context)
|
||||||
|
|
||||||
|
assert result["total_requests"] == 2
|
||||||
|
assert result["total_tokens"] == 15
|
||||||
|
assert result["summary_by_model"] == [
|
||||||
|
{
|
||||||
|
"model": "gpt-4o",
|
||||||
|
"requests": 2,
|
||||||
|
"input_tokens": 10,
|
||||||
|
"output_tokens": 5,
|
||||||
|
"total_tokens": 15,
|
||||||
|
"total_cost_usd": 1.5,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
assert "total_actual_cost" not in result
|
||||||
|
assert result["summary_by_provider"] == [
|
||||||
|
{
|
||||||
|
"provider": "provider-a",
|
||||||
|
"requests": 2,
|
||||||
|
"total_tokens": 15,
|
||||||
|
"total_cost_usd": 1.5,
|
||||||
|
"success_rate": 100.0,
|
||||||
|
"avg_response_time_ms": 500.0,
|
||||||
|
}
|
||||||
|
]
|
||||||
|
assert summary_getter.call_args.kwargs["group_by"] is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_usage_adapter_provider_success_rate_uses_success_count(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
db = MagicMock()
|
||||||
|
query = MagicMock()
|
||||||
|
count_query = MagicMock()
|
||||||
|
count_query.scalar.return_value = 0
|
||||||
|
query.outerjoin.return_value = query
|
||||||
|
query.filter.return_value = query
|
||||||
|
query.with_entities.return_value = count_query
|
||||||
|
query.options.return_value = query
|
||||||
|
query.order_by.return_value = query
|
||||||
|
query.offset.return_value = query
|
||||||
|
query.limit.return_value = query
|
||||||
|
query.all.return_value = []
|
||||||
|
db.query.return_value = query
|
||||||
|
|
||||||
|
summary_getter = MagicMock(
|
||||||
|
return_value=[
|
||||||
|
{
|
||||||
|
"provider": "provider-a",
|
||||||
|
"model": "gpt-4o",
|
||||||
|
"requests": 3,
|
||||||
|
"input_tokens": 30,
|
||||||
|
"output_tokens": 15,
|
||||||
|
"total_tokens": 45,
|
||||||
|
"total_cost_usd": 4.5,
|
||||||
|
"actual_total_cost_usd": 4.5,
|
||||||
|
"success_count": 2,
|
||||||
|
"success_response_time_sum_ms": 600.0,
|
||||||
|
"success_response_time_count": 2,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"provider": "provider-a",
|
||||||
|
"model": "gpt-4.1",
|
||||||
|
"requests": 1,
|
||||||
|
"input_tokens": 10,
|
||||||
|
"output_tokens": 5,
|
||||||
|
"total_tokens": 15,
|
||||||
|
"total_cost_usd": 1.5,
|
||||||
|
"actual_total_cost_usd": 1.5,
|
||||||
|
"success_count": 0,
|
||||||
|
"success_response_time_sum_ms": 0.0,
|
||||||
|
"success_response_time_count": 0,
|
||||||
|
},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("src.api.user_me.routes.UsageService.get_usage_summary", summary_getter)
|
||||||
|
monkeypatch.setattr("src.api.user_me.routes.WalletService.get_wallet", lambda *_a, **_k: None)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.api.user_me.routes.WalletService.serialize_wallet_summary",
|
||||||
|
lambda _wallet: {"limit_mode": "finite"},
|
||||||
|
)
|
||||||
|
|
||||||
|
adapter = GetUsageAdapter(time_range=None, limit=20, offset=0)
|
||||||
|
context = SimpleNamespace(
|
||||||
|
db=db,
|
||||||
|
user=SimpleNamespace(id="user-1", role=UserRole.USER),
|
||||||
|
request=SimpleNamespace(state=SimpleNamespace()),
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await adapter.handle(context)
|
||||||
|
|
||||||
|
assert result["summary_by_provider"] == [
|
||||||
|
{
|
||||||
|
"provider": "provider-a",
|
||||||
|
"requests": 4,
|
||||||
|
"total_tokens": 60,
|
||||||
|
"total_cost_usd": 6.0,
|
||||||
|
"success_rate": 50.0,
|
||||||
|
"avg_response_time_ms": 300.0,
|
||||||
|
}
|
||||||
|
]
|
||||||
@@ -69,26 +69,33 @@ async def test_list_all_candidates_returns_provider_batch_count_even_when_candid
|
|||||||
global_model = _make_global_model(gid="gm1", name="gpt-4o")
|
global_model = _make_global_model(gid="gm1", name="gpt-4o")
|
||||||
|
|
||||||
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
|
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
|
||||||
with patch.object(scheduler._candidate_builder, "_query_providers", return_value=providers):
|
with patch.object(
|
||||||
with patch(
|
scheduler._candidate_builder,
|
||||||
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
"_query_provider_refs",
|
||||||
new=AsyncMock(return_value=global_model),
|
return_value=[("p1", "p1"), ("p2", "p2")],
|
||||||
):
|
):
|
||||||
|
with patch.object(scheduler._candidate_builder, "_query_providers") as query_providers:
|
||||||
with patch(
|
with patch(
|
||||||
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||||
return_value=True,
|
new=AsyncMock(return_value=global_model),
|
||||||
):
|
):
|
||||||
candidates, global_model_id, provider_batch_count = (
|
with patch(
|
||||||
await scheduler.list_all_candidates(
|
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||||
db=db,
|
return_value=True,
|
||||||
api_format="openai:chat",
|
):
|
||||||
model_name="gpt-4o",
|
candidates, global_model_id, provider_batch_count = (
|
||||||
affinity_key=None,
|
await scheduler.list_all_candidates(
|
||||||
user_api_key=user_api_key, # type: ignore[arg-type]
|
db=db,
|
||||||
provider_offset=0,
|
api_format="openai:chat",
|
||||||
provider_limit=20,
|
model_name="gpt-4o",
|
||||||
|
affinity_key=None,
|
||||||
|
user_api_key=user_api_key, # type: ignore[arg-type]
|
||||||
|
provider_offset=0,
|
||||||
|
provider_limit=20,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
)
|
|
||||||
|
query_providers.assert_not_called()
|
||||||
|
|
||||||
assert candidates == []
|
assert candidates == []
|
||||||
assert global_model_id == "gm1"
|
assert global_model_id == "gm1"
|
||||||
|
|||||||
@@ -0,0 +1,93 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||||
|
|
||||||
|
|
||||||
|
def _make_db() -> MagicMock:
|
||||||
|
db = MagicMock()
|
||||||
|
db.new = []
|
||||||
|
db.dirty = []
|
||||||
|
db.deleted = []
|
||||||
|
db.in_transaction.return_value = False
|
||||||
|
return db
|
||||||
|
|
||||||
|
|
||||||
|
def _make_global_model() -> SimpleNamespace:
|
||||||
|
return SimpleNamespace(
|
||||||
|
id="gm1",
|
||||||
|
name="gpt-4o",
|
||||||
|
is_active=True,
|
||||||
|
config={},
|
||||||
|
supported_capabilities=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_list_all_candidates_prefilters_provider_graph_by_allowed_providers() -> None:
|
||||||
|
scheduler = CacheAwareScheduler()
|
||||||
|
scheduler.scheduling_mode = CacheAwareScheduler.SCHEDULING_MODE_FIXED_ORDER
|
||||||
|
|
||||||
|
db = _make_db()
|
||||||
|
global_model = _make_global_model()
|
||||||
|
user_api_key = SimpleNamespace(
|
||||||
|
id="ak1",
|
||||||
|
allowed_providers=["provider-b"],
|
||||||
|
allowed_models=None,
|
||||||
|
allowed_api_formats=None,
|
||||||
|
user=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
filtered_provider = SimpleNamespace(
|
||||||
|
id="provider-b",
|
||||||
|
name="provider-b",
|
||||||
|
endpoints=[],
|
||||||
|
models=[],
|
||||||
|
provider_priority=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
|
||||||
|
with patch.object(
|
||||||
|
scheduler._candidate_builder,
|
||||||
|
"_query_provider_refs",
|
||||||
|
return_value=[("provider-a", "provider-a"), ("provider-b", "provider-b")],
|
||||||
|
) as refs_mock:
|
||||||
|
with patch.object(
|
||||||
|
scheduler._candidate_builder,
|
||||||
|
"_query_providers",
|
||||||
|
return_value=[filtered_provider],
|
||||||
|
) as providers_mock:
|
||||||
|
with patch(
|
||||||
|
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||||
|
new=AsyncMock(return_value=global_model),
|
||||||
|
):
|
||||||
|
with patch(
|
||||||
|
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||||
|
return_value=True,
|
||||||
|
):
|
||||||
|
with patch.object(
|
||||||
|
scheduler._candidate_builder,
|
||||||
|
"_build_candidates",
|
||||||
|
new=AsyncMock(return_value=[]),
|
||||||
|
):
|
||||||
|
candidates, global_model_id, provider_batch_count = (
|
||||||
|
await scheduler.list_all_candidates(
|
||||||
|
db=db,
|
||||||
|
api_format="openai:chat",
|
||||||
|
model_name="gpt-4o",
|
||||||
|
affinity_key=None,
|
||||||
|
user_api_key=user_api_key, # type: ignore[arg-type]
|
||||||
|
provider_offset=0,
|
||||||
|
provider_limit=20,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert candidates == []
|
||||||
|
assert global_model_id == "gm1"
|
||||||
|
assert provider_batch_count == 2
|
||||||
|
refs_mock.assert_called_once()
|
||||||
|
providers_mock.assert_called_once_with(db=db, provider_ids=["provider-b"])
|
||||||
@@ -8,18 +8,20 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
|
from decimal import Decimal
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import jwt
|
import jwt
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from src.core.exceptions import ForbiddenException
|
|
||||||
from src.core.enums import AuthSource
|
from src.core.enums import AuthSource
|
||||||
|
from src.core.exceptions import ForbiddenException
|
||||||
from src.models.database import UserRole
|
from src.models.database import UserRole
|
||||||
from src.services.auth.service import (
|
from src.services.auth.service import (
|
||||||
JWT_ALGORITHM,
|
JWT_ALGORITHM,
|
||||||
JWT_EXPIRATION_HOURS,
|
JWT_EXPIRATION_HOURS,
|
||||||
JWT_SECRET_KEY,
|
JWT_SECRET_KEY,
|
||||||
|
AuthenticatedUserSnapshot,
|
||||||
AuthService,
|
AuthService,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -235,6 +237,115 @@ class TestUserAuthentication:
|
|||||||
|
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_user_threadsafe_uses_isolated_session_for_local_login(self) -> None:
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "user-123"
|
||||||
|
mock_user.email = "test@example.com"
|
||||||
|
mock_user.username = "tester"
|
||||||
|
mock_user.created_at = datetime.now(timezone.utc)
|
||||||
|
mock_user.is_deleted = False
|
||||||
|
mock_user.is_active = True
|
||||||
|
mock_user.auth_source = AuthSource.LOCAL
|
||||||
|
mock_user.role = UserRole.USER
|
||||||
|
mock_user.verify_password.return_value = True
|
||||||
|
|
||||||
|
thread_db = MagicMock()
|
||||||
|
thread_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||||
|
route_db = MagicMock()
|
||||||
|
|
||||||
|
with patch("src.services.auth.service.create_session", return_value=thread_db):
|
||||||
|
with patch(
|
||||||
|
"src.services.auth.service.UserCacheService.invalidate_user_cache",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
) as invalidate_cache:
|
||||||
|
result = await AuthService.authenticate_user_threadsafe(
|
||||||
|
route_db,
|
||||||
|
"test@example.com",
|
||||||
|
"password123",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert isinstance(result, AuthenticatedUserSnapshot)
|
||||||
|
assert result.user_id == "user-123"
|
||||||
|
assert result.username == "tester"
|
||||||
|
thread_db.commit.assert_called_once()
|
||||||
|
thread_db.close.assert_called_once()
|
||||||
|
route_db.commit.assert_not_called()
|
||||||
|
invalidate_cache.assert_awaited_once_with("user-123", "test@example.com")
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_load_user_for_pipeline_threadsafe_prefetches_balance(self) -> None:
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "user-123"
|
||||||
|
mock_user.is_active = True
|
||||||
|
mock_user.is_deleted = False
|
||||||
|
|
||||||
|
thread_db = MagicMock()
|
||||||
|
thread_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||||
|
|
||||||
|
with patch("src.services.auth.service.create_session", return_value=thread_db):
|
||||||
|
with patch(
|
||||||
|
"src.services.wallet.service.WalletService.get_balance_snapshot",
|
||||||
|
return_value=Decimal("7.5"),
|
||||||
|
):
|
||||||
|
result = await AuthService.load_user_for_pipeline_threadsafe(
|
||||||
|
"user-123",
|
||||||
|
include_balance=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert result.user == mock_user
|
||||||
|
assert result.balance_remaining == 7.5
|
||||||
|
thread_db.expunge.assert_called_with(mock_user)
|
||||||
|
thread_db.close.assert_called_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_api_key_threadsafe_returns_balance_and_access_result(self) -> None:
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "user-123"
|
||||||
|
mock_api_key = MagicMock()
|
||||||
|
mock_api_key.id = "key-123"
|
||||||
|
|
||||||
|
thread_db = MagicMock()
|
||||||
|
|
||||||
|
with patch("src.services.auth.service.create_session", return_value=thread_db):
|
||||||
|
with patch.object(
|
||||||
|
AuthService,
|
||||||
|
"authenticate_api_key",
|
||||||
|
return_value=(mock_user, mock_api_key),
|
||||||
|
):
|
||||||
|
with patch(
|
||||||
|
"src.services.usage.service.UsageService.check_request_balance_details",
|
||||||
|
return_value=MagicMock(allowed=False, message="????", remaining=0.0),
|
||||||
|
) as mock_balance_details:
|
||||||
|
with patch(
|
||||||
|
"src.services.wallet.service.WalletService.get_balance_snapshot"
|
||||||
|
) as mock_balance_snapshot:
|
||||||
|
result = await AuthService.authenticate_api_key_threadsafe("sk-test")
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert result.user == mock_user
|
||||||
|
assert result.api_key == mock_api_key
|
||||||
|
assert result.access_ok is False
|
||||||
|
assert result.balance_remaining == 0.0
|
||||||
|
assert result.access_message == "????"
|
||||||
|
mock_balance_details.assert_called_once()
|
||||||
|
mock_balance_snapshot.assert_not_called()
|
||||||
|
thread_db.expunge.assert_any_call(mock_user)
|
||||||
|
thread_db.expunge.assert_any_call(mock_api_key)
|
||||||
|
thread_db.close.assert_called_once()
|
||||||
|
|
||||||
|
def test_detach_instance_logs_debug_when_expunge_fails(self) -> None:
|
||||||
|
mock_db = MagicMock()
|
||||||
|
mock_db.expunge.side_effect = RuntimeError("expunge boom")
|
||||||
|
mock_instance = MagicMock()
|
||||||
|
|
||||||
|
with patch("src.services.auth.service.logger.debug") as mock_debug:
|
||||||
|
AuthService._detach_instance(mock_db, mock_instance)
|
||||||
|
|
||||||
|
mock_debug.assert_called_once()
|
||||||
|
assert "expunge failed" in mock_debug.call_args[0][0]
|
||||||
|
|
||||||
|
|
||||||
class TestAPIKeyAuthentication:
|
class TestAPIKeyAuthentication:
|
||||||
"""测试 API Key 认证"""
|
"""测试 API Key 认证"""
|
||||||
|
|||||||
@@ -1,36 +1,45 @@
|
|||||||
|
from types import SimpleNamespace
|
||||||
|
from typing import Any, Callable, cast
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from src.core.cache_service import CacheService
|
|
||||||
from src.models.database import GlobalModel, Model
|
from src.models.database import GlobalModel, Model
|
||||||
from src.services.cache.model_cache import ModelCacheService
|
from src.services.cache.model_cache import ModelCacheService
|
||||||
|
from src.core.cache_service import CacheService
|
||||||
|
|
||||||
|
|
||||||
class _FakeQuery:
|
class _FakeQuery:
|
||||||
def __init__(self, *, first_result=None, all_result=None, on_all=None):
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
first_result: Any = None,
|
||||||
|
all_result: list[Any] | None = None,
|
||||||
|
on_all: Callable[[], None] | None = None,
|
||||||
|
) -> None:
|
||||||
self._first_result = first_result
|
self._first_result = first_result
|
||||||
self._all_result = all_result if all_result is not None else []
|
self._all_result = all_result if all_result is not None else []
|
||||||
self._on_all = on_all
|
self._on_all = on_all
|
||||||
|
|
||||||
def join(self, *_args, **_kwargs):
|
def join(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def filter(self, *_args, **_kwargs):
|
def filter(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def first(self):
|
def first(self) -> Any:
|
||||||
return self._first_result
|
return self._first_result
|
||||||
|
|
||||||
def all(self):
|
def all(self) -> list[Any]:
|
||||||
if self._on_all:
|
if self._on_all:
|
||||||
self._on_all()
|
self._on_all()
|
||||||
return self._all_result
|
return self._all_result
|
||||||
|
|
||||||
|
|
||||||
class _FakeSession:
|
class _FakeSession:
|
||||||
def __init__(self, *, direct_match: GlobalModel):
|
def __init__(self, *, direct_match: GlobalModel) -> None:
|
||||||
self._direct_match = direct_match
|
self._direct_match = direct_match
|
||||||
|
|
||||||
def query(self, *entities):
|
def query(self, *entities: object) -> "_FakeQuery":
|
||||||
if entities == (GlobalModel,):
|
if entities == (GlobalModel,):
|
||||||
return _FakeQuery(first_result=self._direct_match)
|
return _FakeQuery(first_result=self._direct_match)
|
||||||
|
|
||||||
@@ -43,12 +52,43 @@ class _FakeSession:
|
|||||||
raise AssertionError(f"Unexpected query entities: {entities}")
|
raise AssertionError(f"Unexpected query entities: {entities}")
|
||||||
|
|
||||||
|
|
||||||
|
class _MappingIndexSession:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
provider_mapping_rows: list[tuple[object, GlobalModel]],
|
||||||
|
) -> None:
|
||||||
|
self._provider_mapping_rows = provider_mapping_rows
|
||||||
|
self.provider_mapping_scan_count = 0
|
||||||
|
self.model_global_query_count = 0
|
||||||
|
|
||||||
|
def query(self, *entities: object) -> "_FakeQuery":
|
||||||
|
if entities == (GlobalModel,):
|
||||||
|
return _FakeQuery(first_result=None, all_result=[])
|
||||||
|
|
||||||
|
if entities == (Model, GlobalModel):
|
||||||
|
self.model_global_query_count += 1
|
||||||
|
if self.model_global_query_count in {1, 3}:
|
||||||
|
return _FakeQuery(all_result=[])
|
||||||
|
if self.model_global_query_count == 2:
|
||||||
|
return _FakeQuery(
|
||||||
|
all_result=self._provider_mapping_rows,
|
||||||
|
on_all=self._record_provider_mapping_scan,
|
||||||
|
)
|
||||||
|
raise AssertionError("provider_model_mappings 全量扫描被重复触发")
|
||||||
|
|
||||||
|
raise AssertionError(f"Unexpected query entities: {entities}")
|
||||||
|
|
||||||
|
def _record_provider_mapping_scan(self) -> None:
|
||||||
|
self.provider_mapping_scan_count += 1
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_resolve_global_model_prefers_direct_match(monkeypatch) -> None:
|
async def test_resolve_global_model_prefers_direct_match(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
async def _fake_get(_key: str):
|
async def _fake_get(_key: str) -> None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _fake_set(_key: str, _value, ttl_seconds: int = 60): # noqa: ARG001
|
async def _fake_set(_key: str, _value: object, ttl_seconds: int = 60) -> bool: # noqa: ARG001
|
||||||
return True
|
return True
|
||||||
|
|
||||||
monkeypatch.setattr(CacheService, "get", staticmethod(_fake_get))
|
monkeypatch.setattr(CacheService, "get", staticmethod(_fake_get))
|
||||||
@@ -67,6 +107,98 @@ async def test_resolve_global_model_prefers_direct_match(monkeypatch) -> None:
|
|||||||
db = _FakeSession(direct_match=global_model)
|
db = _FakeSession(direct_match=global_model)
|
||||||
|
|
||||||
resolved = await ModelCacheService.resolve_global_model_by_name_or_mapping(
|
resolved = await ModelCacheService.resolve_global_model_by_name_or_mapping(
|
||||||
db, global_model.name
|
cast(Any, db),
|
||||||
|
cast(str, global_model.name),
|
||||||
)
|
)
|
||||||
assert resolved is global_model
|
assert resolved is global_model
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_resolve_global_model_reuses_provider_mapping_index_cache(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
cache_store: dict[str, object] = {}
|
||||||
|
|
||||||
|
async def _fake_get(key: str) -> object | None:
|
||||||
|
return cache_store.get(key)
|
||||||
|
|
||||||
|
async def _fake_set(key: str, value: object, ttl_seconds: int = 60) -> bool: # noqa: ARG001
|
||||||
|
cache_store[key] = value
|
||||||
|
return True
|
||||||
|
|
||||||
|
monkeypatch.setattr(CacheService, "get", staticmethod(_fake_get))
|
||||||
|
monkeypatch.setattr(CacheService, "set", staticmethod(_fake_set))
|
||||||
|
|
||||||
|
global_model_one = GlobalModel(
|
||||||
|
id="gm-1",
|
||||||
|
name="gpt-4o",
|
||||||
|
display_name="GPT-4o",
|
||||||
|
supported_capabilities=[],
|
||||||
|
config={},
|
||||||
|
default_tiered_pricing=None,
|
||||||
|
default_price_per_request=None,
|
||||||
|
is_active=True,
|
||||||
|
)
|
||||||
|
global_model_two = GlobalModel(
|
||||||
|
id="gm-2",
|
||||||
|
name="claude-3-7-sonnet",
|
||||||
|
display_name="Claude 3.7 Sonnet",
|
||||||
|
supported_capabilities=[],
|
||||||
|
config={},
|
||||||
|
default_tiered_pricing=None,
|
||||||
|
default_price_per_request=None,
|
||||||
|
is_active=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
model_one = SimpleNamespace(
|
||||||
|
id="m-1",
|
||||||
|
provider_model_mappings=[{"name": "mapped-one"}],
|
||||||
|
)
|
||||||
|
model_two = SimpleNamespace(
|
||||||
|
id="m-2",
|
||||||
|
provider_model_mappings=[{"name": "mapped-two"}],
|
||||||
|
)
|
||||||
|
|
||||||
|
db = _MappingIndexSession(
|
||||||
|
provider_mapping_rows=[
|
||||||
|
(model_one, global_model_one),
|
||||||
|
(model_two, global_model_two),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
resolved_one = await ModelCacheService.resolve_global_model_by_name_or_mapping(
|
||||||
|
cast(Any, db), "mapped-one"
|
||||||
|
)
|
||||||
|
resolved_two = await ModelCacheService.resolve_global_model_by_name_or_mapping(
|
||||||
|
cast(Any, db), "mapped-two"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resolved_one is not None
|
||||||
|
assert resolved_one.name == "gpt-4o"
|
||||||
|
assert resolved_two is not None
|
||||||
|
assert resolved_two.name == "claude-3-7-sonnet"
|
||||||
|
assert db.provider_mapping_scan_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_invalidate_model_cache_clears_provider_mapping_index(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
deleted_keys: list[str] = []
|
||||||
|
|
||||||
|
async def _fake_delete(key: str) -> bool:
|
||||||
|
deleted_keys.append(key)
|
||||||
|
return True
|
||||||
|
|
||||||
|
monkeypatch.setattr(CacheService, "delete", staticmethod(_fake_delete))
|
||||||
|
|
||||||
|
await ModelCacheService.invalidate_model_cache(
|
||||||
|
model_id="model-1",
|
||||||
|
provider_model_name="provider-model",
|
||||||
|
provider_model_mappings=[{"name": "alias-model"}],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "model:id:model-1" in deleted_keys
|
||||||
|
assert "global_model:resolve:provider-model" in deleted_keys
|
||||||
|
assert "global_model:resolve:alias-model" in deleted_keys
|
||||||
|
assert ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY in deleted_keys
|
||||||
|
|||||||
182
tests/services/test_stats_aggregator_optimization.py
Normal file
182
tests/services/test_stats_aggregator_optimization.py
Normal file
@@ -0,0 +1,182 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import date, datetime, time, timedelta, timezone
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.models.database import StatsDaily, StatsUserDaily
|
||||||
|
from src.services.system.stats_aggregator import (
|
||||||
|
AggregatedStats,
|
||||||
|
StatsAggregatorService,
|
||||||
|
query_stats_hybrid,
|
||||||
|
)
|
||||||
|
from src.services.system.time_range import TimeRangeParams
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeQuery:
|
||||||
|
def __init__(self, *, all_result: list[Any] | None = None) -> None:
|
||||||
|
self._all_result = all_result if all_result is not None else []
|
||||||
|
|
||||||
|
def filter(self, *_args: object, **_kwargs: object) -> _FakeQuery:
|
||||||
|
return self
|
||||||
|
|
||||||
|
def group_by(self, *_args: object, **_kwargs: object) -> _FakeQuery:
|
||||||
|
return self
|
||||||
|
|
||||||
|
def all(self) -> list[Any]:
|
||||||
|
return self._all_result
|
||||||
|
|
||||||
|
|
||||||
|
class _HybridQuerySession:
|
||||||
|
def __init__(self, stats_daily_rows: list[SimpleNamespace]) -> None:
|
||||||
|
self._stats_daily_rows = stats_daily_rows
|
||||||
|
self.stats_daily_query_count = 0
|
||||||
|
|
||||||
|
def query(self, entity: object) -> _FakeQuery:
|
||||||
|
if entity is StatsDaily:
|
||||||
|
self.stats_daily_query_count += 1
|
||||||
|
return _FakeQuery(all_result=self._stats_daily_rows)
|
||||||
|
raise AssertionError(f"Unexpected query entity: {entity}")
|
||||||
|
|
||||||
|
|
||||||
|
class _BatchUserStatsSession:
|
||||||
|
def __init__(
|
||||||
|
self, existing_rows: list[StatsUserDaily], aggregated_rows: list[SimpleNamespace]
|
||||||
|
) -> None:
|
||||||
|
self._responses: list[list[Any]] = [list(existing_rows), list(aggregated_rows)]
|
||||||
|
self.added: list[StatsUserDaily] = []
|
||||||
|
self.commit_count = 0
|
||||||
|
|
||||||
|
def query(self, *_entities: object) -> _FakeQuery:
|
||||||
|
if not self._responses:
|
||||||
|
raise AssertionError("Unexpected extra query")
|
||||||
|
return _FakeQuery(all_result=self._responses.pop(0))
|
||||||
|
|
||||||
|
def add(self, row: StatsUserDaily) -> None:
|
||||||
|
self.added.append(row)
|
||||||
|
|
||||||
|
def commit(self) -> None:
|
||||||
|
self.commit_count += 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_query_stats_hybrid_batches_statsdaily_lookup_and_merges_realtime_ranges(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
today = datetime.now(timezone.utc).date()
|
||||||
|
historical_cached_day = today - timedelta(days=4)
|
||||||
|
historical_missing_day = today - timedelta(days=3)
|
||||||
|
realtime_day = today
|
||||||
|
|
||||||
|
cached_row = SimpleNamespace(
|
||||||
|
date=datetime.combine(historical_cached_day, time.min, tzinfo=timezone.utc),
|
||||||
|
total_requests=10,
|
||||||
|
success_requests=9,
|
||||||
|
error_requests=1,
|
||||||
|
input_tokens=100,
|
||||||
|
output_tokens=50,
|
||||||
|
cache_creation_tokens=5,
|
||||||
|
cache_read_tokens=3,
|
||||||
|
cache_creation_cost=1.2,
|
||||||
|
cache_read_cost=0.8,
|
||||||
|
total_cost=3.5,
|
||||||
|
actual_total_cost=3.0,
|
||||||
|
avg_response_time_ms=200.0,
|
||||||
|
)
|
||||||
|
db = _HybridQuerySession(stats_daily_rows=[cached_row])
|
||||||
|
|
||||||
|
calls: list[tuple[datetime, datetime]] = []
|
||||||
|
|
||||||
|
def _fake_aggregate_usage_range(
|
||||||
|
_db: object,
|
||||||
|
start_utc: datetime,
|
||||||
|
end_utc: datetime,
|
||||||
|
filters: object | None = None, # noqa: ARG001
|
||||||
|
) -> AggregatedStats:
|
||||||
|
calls.append((start_utc, end_utc))
|
||||||
|
return AggregatedStats(total_requests=1, success_requests=1)
|
||||||
|
|
||||||
|
class _FakeParams:
|
||||||
|
def get_complete_utc_dates(self) -> tuple[list[date], None, None]:
|
||||||
|
return [historical_cached_day, historical_missing_day, realtime_day], None, None
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"src.services.system.stats_aggregator.aggregate_usage_range",
|
||||||
|
_fake_aggregate_usage_range,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = query_stats_hybrid(cast(Any, db), cast(Any, _FakeParams()))
|
||||||
|
|
||||||
|
assert db.stats_daily_query_count == 1
|
||||||
|
assert calls == [
|
||||||
|
(
|
||||||
|
datetime.combine(historical_missing_day, time.min, tzinfo=timezone.utc),
|
||||||
|
datetime.combine(
|
||||||
|
historical_missing_day + timedelta(days=1), time.min, tzinfo=timezone.utc
|
||||||
|
),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
datetime.combine(realtime_day, time.min, tzinfo=timezone.utc),
|
||||||
|
datetime.combine(realtime_day + timedelta(days=1), time.min, tzinfo=timezone.utc),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
assert result.total_requests == 12
|
||||||
|
assert result.success_requests == 11
|
||||||
|
|
||||||
|
|
||||||
|
def test_aggregate_user_daily_stats_batch_updates_all_users_in_two_queries() -> None:
|
||||||
|
target_day = datetime(2026, 3, 1, tzinfo=timezone.utc)
|
||||||
|
aggregated_rows = [
|
||||||
|
SimpleNamespace(
|
||||||
|
user_id="user-1",
|
||||||
|
username="alice",
|
||||||
|
total_requests=4,
|
||||||
|
error_requests=1,
|
||||||
|
input_tokens=20,
|
||||||
|
output_tokens=8,
|
||||||
|
cache_creation_tokens=2,
|
||||||
|
cache_read_tokens=1,
|
||||||
|
total_cost=1.5,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
db = _BatchUserStatsSession(existing_rows=[], aggregated_rows=aggregated_rows)
|
||||||
|
|
||||||
|
result = StatsAggregatorService.aggregate_user_daily_stats_batch(
|
||||||
|
cast(Any, db),
|
||||||
|
target_day,
|
||||||
|
["user-1", "user-2"],
|
||||||
|
commit=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(result) == 2
|
||||||
|
assert db.commit_count == 1
|
||||||
|
assert len(db.added) == 2
|
||||||
|
|
||||||
|
user_one = next(row for row in result if row.user_id == "user-1")
|
||||||
|
user_two = next(row for row in result if row.user_id == "user-2")
|
||||||
|
|
||||||
|
assert user_one.username == "alice"
|
||||||
|
assert user_one.total_requests == 4
|
||||||
|
assert user_one.success_requests == 3
|
||||||
|
assert user_one.total_cost == 1.5
|
||||||
|
|
||||||
|
assert user_two.total_requests == 0
|
||||||
|
assert user_two.success_requests == 0
|
||||||
|
assert user_two.error_requests == 0
|
||||||
|
assert user_two.total_cost == 0.0
|
||||||
|
|
||||||
|
|
||||||
|
def test_compute_percentiles_by_local_day_returns_sqlite_fallback_without_queries() -> None:
|
||||||
|
db = SimpleNamespace(bind=SimpleNamespace(dialect=SimpleNamespace(name="sqlite")))
|
||||||
|
time_range = TimeRangeParams(
|
||||||
|
start_date=date(2026, 3, 1),
|
||||||
|
end_date=date(2026, 3, 3),
|
||||||
|
timezone="Asia/Singapore",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = StatsAggregatorService.compute_percentiles_by_local_day(cast(Any, db), time_range)
|
||||||
|
|
||||||
|
assert [row["date"] for row in result] == ["2026-03-01", "2026-03-02", "2026-03-03"]
|
||||||
|
assert all(row["p50_response_time_ms"] is None for row in result)
|
||||||
|
assert all(row["p50_first_byte_time_ms"] is None for row in result)
|
||||||
@@ -8,7 +8,7 @@ UsageService 测试
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
@@ -146,6 +146,58 @@ class TestBalanceCheck:
|
|||||||
|
|
||||||
assert is_ok is True
|
assert is_ok is True
|
||||||
|
|
||||||
|
def test_check_request_balance_details_returns_remaining(self) -> None:
|
||||||
|
"""Balance detail helper returns remaining."""
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.role = MagicMock()
|
||||||
|
mock_user.role.value = "user"
|
||||||
|
|
||||||
|
mock_api_key = MagicMock()
|
||||||
|
mock_api_key.is_standalone = False
|
||||||
|
|
||||||
|
mock_db = MagicMock()
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"src.services.wallet.WalletService.check_request_allowed",
|
||||||
|
return_value=WalletAccessResult(
|
||||||
|
False, Decimal("12.5"), "\u94b1\u5305\u4f59\u989d\u4e0d\u8db3"
|
||||||
|
),
|
||||||
|
):
|
||||||
|
result = UsageService.check_request_balance_details(
|
||||||
|
mock_db, mock_user, api_key=mock_api_key
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.allowed is False
|
||||||
|
assert result.remaining == 12.5
|
||||||
|
assert "\u4f59\u989d\u4e0d\u8db3" in result.message
|
||||||
|
|
||||||
|
def test_check_request_balance_details_maps_overdue_message(self) -> None:
|
||||||
|
"""欠费状态应映射为对外统一文案。"""
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_api_key = MagicMock()
|
||||||
|
mock_api_key.is_standalone = False
|
||||||
|
mock_db = MagicMock()
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"src.services.wallet.WalletService.check_request_allowed",
|
||||||
|
return_value=WalletAccessResult(False, Decimal("-1"), "钱包欠费,请先充值"),
|
||||||
|
):
|
||||||
|
normal_result = UsageService.check_request_balance_details(
|
||||||
|
mock_db, mock_user, api_key=mock_api_key
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_api_key.is_standalone = True
|
||||||
|
with patch(
|
||||||
|
"src.services.wallet.WalletService.check_request_allowed",
|
||||||
|
return_value=WalletAccessResult(False, Decimal("-1"), "钱包欠费,请先充值"),
|
||||||
|
):
|
||||||
|
standalone_result = UsageService.check_request_balance_details(
|
||||||
|
mock_db, mock_user, api_key=mock_api_key
|
||||||
|
)
|
||||||
|
|
||||||
|
assert normal_result.message == "账户欠费,请先充值"
|
||||||
|
assert standalone_result.message == "Key欠费,请先调账或充值"
|
||||||
|
|
||||||
def test_check_request_balance_exceeded(self) -> None:
|
def test_check_request_balance_exceeded(self) -> None:
|
||||||
"""测试余额耗尽时拦截新请求"""
|
"""测试余额耗尽时拦截新请求"""
|
||||||
mock_user = MagicMock()
|
mock_user = MagicMock()
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from typing import Any, cast
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -52,7 +53,9 @@ def test_get_or_create_wallet_prefers_user_owner_for_non_standalone_key() -> Non
|
|||||||
user = SimpleNamespace(id="user-1")
|
user = SimpleNamespace(id="user-1")
|
||||||
api_key = SimpleNamespace(id="key-1", is_standalone=False)
|
api_key = SimpleNamespace(id="key-1", is_standalone=False)
|
||||||
|
|
||||||
wallet = WalletService.get_or_create_wallet(db, user=user, api_key=api_key)
|
wallet = WalletService.get_or_create_wallet(
|
||||||
|
db, user=cast(Any, user), api_key=cast(Any, api_key)
|
||||||
|
)
|
||||||
|
|
||||||
assert wallet is not None
|
assert wallet is not None
|
||||||
assert wallet.user_id == "user-1"
|
assert wallet.user_id == "user-1"
|
||||||
@@ -68,13 +71,35 @@ def test_get_or_create_wallet_uses_api_key_owner_for_standalone_key() -> None:
|
|||||||
user = SimpleNamespace(id="user-1")
|
user = SimpleNamespace(id="user-1")
|
||||||
api_key = SimpleNamespace(id="key-1", is_standalone=True)
|
api_key = SimpleNamespace(id="key-1", is_standalone=True)
|
||||||
|
|
||||||
wallet = WalletService.get_or_create_wallet(db, user=user, api_key=api_key)
|
wallet = WalletService.get_or_create_wallet(
|
||||||
|
db, user=cast(Any, user), api_key=cast(Any, api_key)
|
||||||
|
)
|
||||||
|
|
||||||
assert wallet is not None
|
assert wallet is not None
|
||||||
assert wallet.user_id is None
|
assert wallet.user_id is None
|
||||||
assert wallet.api_key_id == "key-1"
|
assert wallet.api_key_id == "key-1"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_wallets_by_user_ids_returns_mapping() -> None:
|
||||||
|
db = MagicMock()
|
||||||
|
wallet_1 = SimpleNamespace(user_id="user-1")
|
||||||
|
wallet_2 = SimpleNamespace(user_id="user-2")
|
||||||
|
db.query.return_value.filter.return_value.all.return_value = [wallet_1, wallet_2]
|
||||||
|
|
||||||
|
result = WalletService.get_wallets_by_user_ids(db, ["user-1", "user-2"])
|
||||||
|
|
||||||
|
assert result == {"user-1": wallet_1, "user-2": wallet_2}
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_wallets_by_user_ids_skips_query_for_empty_ids() -> None:
|
||||||
|
db = MagicMock()
|
||||||
|
|
||||||
|
result = WalletService.get_wallets_by_user_ids(db, [])
|
||||||
|
|
||||||
|
assert result == {}
|
||||||
|
db.query.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
def test_check_request_allowed_denies_when_recharge_negative_even_total_positive() -> None:
|
def test_check_request_allowed_denies_when_recharge_negative_even_total_positive() -> None:
|
||||||
wallet = _build_wallet(recharge="-1", gift="10", limit_mode="finite")
|
wallet = _build_wallet(recharge="-1", gift="10", limit_mode="finite")
|
||||||
db = MagicMock()
|
db = MagicMock()
|
||||||
@@ -104,7 +129,7 @@ def test_admin_adjust_balance_negative_from_gift_spills_to_recharge() -> None:
|
|||||||
|
|
||||||
tx = WalletService.admin_adjust_balance(
|
tx = WalletService.admin_adjust_balance(
|
||||||
db,
|
db,
|
||||||
wallet=wallet,
|
wallet=cast(Any, wallet),
|
||||||
amount_usd=Decimal("-10"),
|
amount_usd=Decimal("-10"),
|
||||||
balance_type="gift",
|
balance_type="gift",
|
||||||
operator_id="admin-1",
|
operator_id="admin-1",
|
||||||
@@ -127,7 +152,7 @@ def test_admin_adjust_balance_negative_from_recharge_then_gift() -> None:
|
|||||||
|
|
||||||
tx = WalletService.admin_adjust_balance(
|
tx = WalletService.admin_adjust_balance(
|
||||||
db,
|
db,
|
||||||
wallet=wallet,
|
wallet=cast(Any, wallet),
|
||||||
amount_usd=Decimal("-4"),
|
amount_usd=Decimal("-4"),
|
||||||
balance_type="recharge",
|
balance_type="recharge",
|
||||||
operator_id="admin-1",
|
operator_id="admin-1",
|
||||||
@@ -148,7 +173,7 @@ def test_admin_adjust_balance_positive_adds_to_selected_bucket_without_offset()
|
|||||||
|
|
||||||
tx = WalletService.admin_adjust_balance(
|
tx = WalletService.admin_adjust_balance(
|
||||||
db,
|
db,
|
||||||
wallet=wallet,
|
wallet=cast(Any, wallet),
|
||||||
amount_usd=Decimal("1"),
|
amount_usd=Decimal("1"),
|
||||||
balance_type="gift",
|
balance_type="gift",
|
||||||
operator_id="admin-1",
|
operator_id="admin-1",
|
||||||
@@ -181,7 +206,9 @@ def test_apply_usage_charge_prefers_gift_then_recharge() -> None:
|
|||||||
db = _build_locked_db(wallet)
|
db = _build_locked_db(wallet)
|
||||||
|
|
||||||
with patch.object(WalletService, "_resolve_wallet_for_usage", return_value=wallet):
|
with patch.object(WalletService, "_resolve_wallet_for_usage", return_value=wallet):
|
||||||
before, after = WalletService.apply_usage_charge(db, usage=usage, amount_usd=Decimal("6"))
|
before, after = WalletService.apply_usage_charge(
|
||||||
|
db, usage=cast(Any, usage), amount_usd=Decimal("6")
|
||||||
|
)
|
||||||
|
|
||||||
assert before == Decimal("8.00000000")
|
assert before == Decimal("8.00000000")
|
||||||
assert after == Decimal("2.00000000")
|
assert after == Decimal("2.00000000")
|
||||||
@@ -212,7 +239,9 @@ def test_apply_usage_charge_unlimited_wallet_keeps_balances() -> None:
|
|||||||
db = _build_locked_db(wallet)
|
db = _build_locked_db(wallet)
|
||||||
|
|
||||||
with patch.object(WalletService, "_resolve_wallet_for_usage", return_value=wallet):
|
with patch.object(WalletService, "_resolve_wallet_for_usage", return_value=wallet):
|
||||||
before, after = WalletService.apply_usage_charge(db, usage=usage, amount_usd=Decimal("4"))
|
before, after = WalletService.apply_usage_charge(
|
||||||
|
db, usage=cast(Any, usage), amount_usd=Decimal("4")
|
||||||
|
)
|
||||||
|
|
||||||
assert before == Decimal("8.00000000")
|
assert before == Decimal("8.00000000")
|
||||||
assert after == Decimal("8.00000000")
|
assert after == Decimal("8.00000000")
|
||||||
@@ -237,7 +266,7 @@ def test_complete_refund_requires_processing_status() -> None:
|
|||||||
db.query.return_value = query
|
db.query.return_value = query
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="processing"):
|
with pytest.raises(ValueError, match="processing"):
|
||||||
WalletService.complete_refund(db, refund=refund)
|
WalletService.complete_refund(db, refund=cast(Any, refund))
|
||||||
|
|
||||||
|
|
||||||
def test_get_or_create_wallet_reuses_existing_after_integrity_error() -> None:
|
def test_get_or_create_wallet_reuses_existing_after_integrity_error() -> None:
|
||||||
@@ -252,7 +281,7 @@ def test_get_or_create_wallet_reuses_existing_after_integrity_error() -> None:
|
|||||||
with patch.object(WalletService, "get_wallet", side_effect=[None, existing_wallet]):
|
with patch.object(WalletService, "get_wallet", side_effect=[None, existing_wallet]):
|
||||||
wallet = WalletService.get_or_create_wallet(
|
wallet = WalletService.get_or_create_wallet(
|
||||||
db,
|
db,
|
||||||
user=SimpleNamespace(id="user-1"),
|
user=cast(Any, SimpleNamespace(id="user-1")),
|
||||||
api_key=None,
|
api_key=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -287,18 +316,20 @@ def test_create_refund_request_rejects_uncredited_payment_order() -> None:
|
|||||||
|
|
||||||
db.query.side_effect = _query
|
db.query.side_effect = _query
|
||||||
|
|
||||||
with patch.object(WalletService, "_get_pending_refund_reserved_amount", return_value=Decimal("0")):
|
with patch.object(
|
||||||
|
WalletService, "_get_pending_refund_reserved_amount", return_value=Decimal("0")
|
||||||
|
):
|
||||||
with pytest.raises(ValueError, match="payment order is not refundable"):
|
with pytest.raises(ValueError, match="payment order is not refundable"):
|
||||||
WalletService.create_refund_request(
|
WalletService.create_refund_request(
|
||||||
db,
|
db,
|
||||||
wallet=wallet,
|
wallet=cast(Any, wallet),
|
||||||
user_id="user-1",
|
user_id="user-1",
|
||||||
amount_usd=Decimal("2"),
|
amount_usd=Decimal("2"),
|
||||||
refund_no="rf-1",
|
refund_no="rf-1",
|
||||||
source_type="payment_order",
|
source_type="payment_order",
|
||||||
source_id="order-1",
|
source_id="order-1",
|
||||||
refund_mode="original_channel",
|
refund_mode="original_channel",
|
||||||
payment_order=payment_order,
|
payment_order=cast(Any, payment_order),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -326,7 +357,7 @@ def test_create_refund_request_reserves_pending_wallet_amount() -> None:
|
|||||||
with pytest.raises(ValueError, match="available refundable recharge balance"):
|
with pytest.raises(ValueError, match="available refundable recharge balance"):
|
||||||
WalletService.create_refund_request(
|
WalletService.create_refund_request(
|
||||||
db,
|
db,
|
||||||
wallet=wallet,
|
wallet=cast(Any, wallet),
|
||||||
user_id="user-1",
|
user_id="user-1",
|
||||||
amount_usd=Decimal("2"),
|
amount_usd=Decimal("2"),
|
||||||
refund_no="rf-2",
|
refund_no="rf-2",
|
||||||
@@ -372,14 +403,14 @@ def test_create_refund_request_reserves_pending_order_amount() -> None:
|
|||||||
with pytest.raises(ValueError, match="available refundable amount"):
|
with pytest.raises(ValueError, match="available refundable amount"):
|
||||||
WalletService.create_refund_request(
|
WalletService.create_refund_request(
|
||||||
db,
|
db,
|
||||||
wallet=wallet,
|
wallet=cast(Any, wallet),
|
||||||
user_id="user-1",
|
user_id="user-1",
|
||||||
amount_usd=Decimal("2"),
|
amount_usd=Decimal("2"),
|
||||||
refund_no="rf-3",
|
refund_no="rf-3",
|
||||||
source_type="payment_order",
|
source_type="payment_order",
|
||||||
source_id="order-1",
|
source_id="order-1",
|
||||||
refund_mode="original_channel",
|
refund_mode="original_channel",
|
||||||
payment_order=payment_order,
|
payment_order=cast(Any, payment_order),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -430,9 +461,13 @@ def test_move_refund_to_processing_rejects_double_transition() -> None:
|
|||||||
tx = SimpleNamespace(id="tx-1")
|
tx = SimpleNamespace(id="tx-1")
|
||||||
|
|
||||||
with patch.object(WalletService, "create_wallet_transaction", return_value=tx) as create_tx:
|
with patch.object(WalletService, "create_wallet_transaction", return_value=tx) as create_tx:
|
||||||
first_tx = WalletService.move_refund_to_processing(db, refund=refund, operator_id="admin-1")
|
first_tx = WalletService.move_refund_to_processing(
|
||||||
|
db, refund=cast(Any, refund), operator_id="admin-1"
|
||||||
|
)
|
||||||
with pytest.raises(ValueError, match="not approvable"):
|
with pytest.raises(ValueError, match="not approvable"):
|
||||||
WalletService.move_refund_to_processing(db, refund=refund, operator_id="admin-1")
|
WalletService.move_refund_to_processing(
|
||||||
|
db, refund=cast(Any, refund), operator_id="admin-1"
|
||||||
|
)
|
||||||
|
|
||||||
assert first_tx is tx
|
assert first_tx is tx
|
||||||
assert create_tx.call_count == 1
|
assert create_tx.call_count == 1
|
||||||
@@ -488,7 +523,9 @@ def test_move_refund_to_processing_rechecks_payment_order_refundable_amount() ->
|
|||||||
|
|
||||||
with patch.object(WalletService, "create_wallet_transaction") as create_tx:
|
with patch.object(WalletService, "create_wallet_transaction") as create_tx:
|
||||||
with pytest.raises(ValueError, match="refund amount exceeds refundable amount"):
|
with pytest.raises(ValueError, match="refund amount exceeds refundable amount"):
|
||||||
WalletService.move_refund_to_processing(db, refund=refund, operator_id="admin-1")
|
WalletService.move_refund_to_processing(
|
||||||
|
db, refund=cast(Any, refund), operator_id="admin-1"
|
||||||
|
)
|
||||||
|
|
||||||
create_tx.assert_not_called()
|
create_tx.assert_not_called()
|
||||||
assert refund.status == "pending_approval"
|
assert refund.status == "pending_approval"
|
||||||
@@ -531,14 +568,14 @@ def test_fail_refund_rejects_invalid_status_after_first_failure() -> None:
|
|||||||
with patch.object(WalletService, "create_wallet_transaction", return_value=tx) as create_tx:
|
with patch.object(WalletService, "create_wallet_transaction", return_value=tx) as create_tx:
|
||||||
first_tx = WalletService.fail_refund(
|
first_tx = WalletService.fail_refund(
|
||||||
db,
|
db,
|
||||||
refund=refund,
|
refund=cast(Any, refund),
|
||||||
reason="first-failure",
|
reason="first-failure",
|
||||||
operator_id="admin-1",
|
operator_id="admin-1",
|
||||||
)
|
)
|
||||||
with pytest.raises(ValueError, match="cannot fail refund in status: failed"):
|
with pytest.raises(ValueError, match="cannot fail refund in status: failed"):
|
||||||
WalletService.fail_refund(
|
WalletService.fail_refund(
|
||||||
db,
|
db,
|
||||||
refund=refund,
|
refund=cast(Any, refund),
|
||||||
reason="retry-failure",
|
reason="retry-failure",
|
||||||
operator_id="admin-1",
|
operator_id="admin-1",
|
||||||
)
|
)
|
||||||
@@ -577,7 +614,7 @@ def test_fail_refund_rejects_succeeded_status() -> None:
|
|||||||
with pytest.raises(ValueError, match="cannot fail refund in status: succeeded"):
|
with pytest.raises(ValueError, match="cannot fail refund in status: succeeded"):
|
||||||
WalletService.fail_refund(
|
WalletService.fail_refund(
|
||||||
db,
|
db,
|
||||||
refund=refund,
|
refund=cast(Any, refund),
|
||||||
reason="should-not-override",
|
reason="should-not-override",
|
||||||
operator_id="admin-1",
|
operator_id="admin-1",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -11,6 +11,13 @@ from src.api.base.context import ApiRequestContext
|
|||||||
|
|
||||||
|
|
||||||
def _build_request(headers: dict[str, str] | None = None) -> Request:
|
def _build_request(headers: dict[str, str] | None = None) -> Request:
|
||||||
|
return _build_request_with_body(b"", headers=headers)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_request_with_body(
|
||||||
|
body: bytes,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
) -> Request:
|
||||||
header_items = [
|
header_items = [
|
||||||
(str(key).encode("latin-1"), str(value).encode("latin-1"))
|
(str(key).encode("latin-1"), str(value).encode("latin-1"))
|
||||||
for key, value in (headers or {}).items()
|
for key, value in (headers or {}).items()
|
||||||
@@ -28,8 +35,14 @@ def _build_request(headers: dict[str, str] | None = None) -> Request:
|
|||||||
"server": ("testserver", 80),
|
"server": ("testserver", 80),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
received = False
|
||||||
|
|
||||||
async def receive() -> dict[str, object]:
|
async def receive() -> dict[str, object]:
|
||||||
return {"type": "http.request", "body": b"", "more_body": False}
|
nonlocal received
|
||||||
|
if received:
|
||||||
|
return {"type": "http.request", "body": b"", "more_body": False}
|
||||||
|
received = True
|
||||||
|
return {"type": "http.request", "body": body, "more_body": False}
|
||||||
|
|
||||||
request = Request(scope, receive)
|
request = Request(scope, receive)
|
||||||
request.state.perf_metrics = {}
|
request.state.perf_metrics = {}
|
||||||
@@ -89,3 +102,22 @@ class TestApiRequestContextEnsureJsonBody:
|
|||||||
|
|
||||||
assert context.client_content_encoding == "gzip"
|
assert context.client_content_encoding == "gzip"
|
||||||
assert context.client_accept_encoding == "gzip, deflate"
|
assert context.client_accept_encoding == "gzip, deflate"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_ensure_json_body_async_loads_body_lazily(self) -> None:
|
||||||
|
payload = {"message": "hello", "count": 2}
|
||||||
|
request = _build_request_with_body(json.dumps(payload).encode("utf-8"))
|
||||||
|
context = ApiRequestContext.build(
|
||||||
|
request=request,
|
||||||
|
db=None, # type: ignore[arg-type]
|
||||||
|
user=None,
|
||||||
|
api_key=None,
|
||||||
|
raw_body=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert context.raw_body is None
|
||||||
|
|
||||||
|
result = await context.ensure_json_body_async()
|
||||||
|
|
||||||
|
assert result == payload
|
||||||
|
assert context.raw_body == json.dumps(payload).encode("utf-8")
|
||||||
|
|||||||
40
tests/unit/test_request_candidate_intermediate_status.py
Normal file
40
tests/unit/test_request_candidate_intermediate_status.py
Normal file
@@ -0,0 +1,40 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from src.services.request.candidate import RequestCandidateService
|
||||||
|
|
||||||
|
|
||||||
|
def _build_db_with_candidate(candidate: SimpleNamespace) -> MagicMock:
|
||||||
|
query = MagicMock()
|
||||||
|
query.filter.return_value.first.return_value = candidate
|
||||||
|
|
||||||
|
db = MagicMock()
|
||||||
|
db.query.return_value = query
|
||||||
|
db.info = {"managed_by_middleware": True}
|
||||||
|
return db
|
||||||
|
|
||||||
|
|
||||||
|
def test_mark_candidate_started_flushes_without_immediate_commit() -> None:
|
||||||
|
candidate = SimpleNamespace(status="available", started_at=None)
|
||||||
|
db = _build_db_with_candidate(candidate)
|
||||||
|
|
||||||
|
RequestCandidateService.mark_candidate_started(db, "candidate-1")
|
||||||
|
|
||||||
|
assert candidate.status == "pending"
|
||||||
|
assert candidate.started_at is not None
|
||||||
|
db.flush.assert_called_once()
|
||||||
|
db.commit.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
def test_mark_candidate_streaming_flushes_without_immediate_commit() -> None:
|
||||||
|
candidate = SimpleNamespace(status="pending", concurrent_requests=None)
|
||||||
|
db = _build_db_with_candidate(candidate)
|
||||||
|
|
||||||
|
RequestCandidateService.mark_candidate_streaming(db, "candidate-1", concurrent_requests=3)
|
||||||
|
|
||||||
|
assert candidate.status == "streaming"
|
||||||
|
assert candidate.concurrent_requests == 3
|
||||||
|
db.flush.assert_called_once()
|
||||||
|
db.commit.assert_not_called()
|
||||||
Reference in New Issue
Block a user