Merge pull request #217 from AAEE86/1233

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

View File

@@ -46,6 +46,13 @@ interface Props {
const chartRef = ref<HTMLCanvasElement>()
let chart: ChartJS<'line'> | null = null
function buildChartOptions(): ChartOptions<'line'> {
return {
...defaultOptions,
...props.options
}
}
const defaultOptions: ChartOptions<'line'> = {
responsive: true,
maintainAspectRatio: false,
@@ -89,10 +96,7 @@ function createChart() {
chart = new ChartJS(chartRef.value, {
type: 'line',
data: props.data,
options: {
...defaultOptions,
...props.options
}
options: buildChartOptions()
})
}
@@ -115,15 +119,12 @@ onUnmounted(() => {
}
})
// 监听数据变化
watch(() => props.data, updateChart, { deep: true })
// 监听引用变化,避免深监听触发整图重算
watch(() => props.data, updateChart)
watch(() => props.options, () => {
if (chart) {
chart.options = {
...defaultOptions,
...props.options
chart.options = buildChartOptions()
chart.update('none')
}
chart.update()
}
}, { deep: true })
})
</script>

View File

@@ -119,6 +119,11 @@ let chart: ChartJS<'scatter'> | null = null
const crosshairY = ref<number | null>(null)
const gapInfoList = ref<GapInfo[]>([])
interface PreparedRenderData {
chartData: ChartData<'scatter'>
gaps: GapInfo[]
}
const crosshairStats = computed<CrosshairStats | null>(() => {
if (crosshairY.value === null || !props.data.datasets) return null
@@ -294,6 +299,22 @@ function transformData(data: ChartData<'scatter'>): ChartData<'scatter'> {
}
}
function prepareRenderData(): PreparedRenderData {
let dataToUse = props.data
let gaps: GapInfo[] = []
if (props.compressGaps) {
const compressedResult = compressTimeGaps(props.data)
dataToUse = compressedResult.data
gaps = compressedResult.gaps
}
return {
chartData: transformData(dataToUse),
gaps
}
}
// 格式化时长
function formatDuration(ms: number): string {
const hours = Math.floor(ms / (1000 * 60 * 60))
@@ -516,22 +537,12 @@ function handleMouseLeave() {
function createChart() {
if (!chartRef.value) return
let dataToUse = props.data
gapInfoList.value = []
// 如果启用间隙压缩
if (props.compressGaps) {
const { data: compressedData, gaps } = compressTimeGaps(props.data)
dataToUse = compressedData
const { chartData, gaps } = prepareRenderData()
gapInfoList.value = gaps
}
// 转换数据
const transformedData = transformData(dataToUse)
chart = new ChartJS(chartRef.value, {
type: 'scatter',
data: transformedData,
data: chartData,
options: {
...defaultOptions,
...props.options
@@ -544,16 +555,9 @@ function createChart() {
function updateChart() {
if (chart) {
let dataToUse = props.data
gapInfoList.value = []
if (props.compressGaps) {
const { data: compressedData, gaps } = compressTimeGaps(props.data)
dataToUse = compressedData
const { chartData, gaps } = prepareRenderData()
gapInfoList.value = gaps
}
chart.data = transformData(dataToUse)
chart.data = chartData
chart.update('none')
}
}
@@ -573,21 +577,22 @@ onUnmounted(() => {
}
})
watch(() => props.data, updateChart, { deep: true })
watch(() => props.compressGaps, () => {
if (chart) {
chart.destroy()
chart = null
}
createChart()
})
watch(
[
() => props.data,
() => props.compressGaps,
() => props.gapThreshold,
() => props.compressedGapSize
],
updateChart
)
watch(() => props.options, () => {
if (chart) {
chart.options = {
...defaultOptions,
...props.options
}
chart.update()
chart.update('none')
}
}, { deep: true })
})
</script>

View File

@@ -809,7 +809,7 @@ function getApiFormatTooltip(record: UsageRecord): string {
return record.api_format
}
// 获取实际使用的模型(优先 target_model其次 model_version
// 获取实际使用的模型(优先 target_model其次列表接口下发的 model_version
// 只有当实际模型与请求模型不同时才返回,用于显示映射箭头
function getActualModel(record: UsageRecord): string | null {
// 优先显示模型映射
@@ -817,8 +817,8 @@ function getActualModel(record: UsageRecord): string | null {
return record.target_model
}
// 其次显示 Provider 返回的实际版本(如 Gemini 的 modelVersion
if (record.request_metadata?.model_version && record.request_metadata.model_version !== record.model) {
return record.request_metadata.model_version
if (record.model_version && record.model_version !== record.model) {
return record.model_version
}
return null
}

View File

@@ -71,6 +71,7 @@ export interface UsageRecord {
rate_multiplier?: number
model: string
target_model?: string | null // 映射后的目标模型名(若无映射则为空)
model_version?: string | null // Provider 返回的实际模型版本(列表轻量字段)
api_format?: string
endpoint_api_format?: string // 端点原生格式
has_format_conversion?: boolean // 是否发生了格式转换
@@ -90,10 +91,6 @@ export interface UsageRecord {
created_at: string
has_fallback?: boolean
has_retry?: boolean
request_metadata?: {
model_version?: string // Provider 返回的实际模型版本(如 Gemini 的 modelVersion
[key: string]: unknown
}
}
// 日期范围参数

View File

@@ -392,7 +392,7 @@ function generateMockUsageRecords(count: number = 100) {
status,
created_at: createdAt.toISOString(),
has_fallback: Math.random() > 0.9,
request_metadata: model.provider === 'google' ? { model_version: 'gemini-3-pro-preview-2025-01' } : undefined
model_version: model.provider === 'google' ? 'gemini-3-pro-preview-2025-01' : undefined
})
}

View File

@@ -49,6 +49,7 @@ const currentPage = ref(1)
const pageSize = ref(20)
const currentTime = ref(Math.floor(Date.now() / 1000))
const isPageVisible = ref(typeof document === 'undefined' ? true : !document.hidden)
const nextExpireAt = ref<number | null>(null)
// ==================== 模型映射缓存 ====================
@@ -121,6 +122,8 @@ async function fetchAffinityList(keyword?: string) {
const response = await cacheApi.listAffinities(keyword)
affinityList.value = response.items
matchedUserId.value = response.matched_user_id ?? null
currentTime.value = Math.floor(Date.now() / 1000)
pruneExpiredAffinities(currentTime.value, true)
if (keyword && response.total === 0) {
showInfo('未找到匹配的缓存记录')
@@ -238,6 +241,38 @@ function handlePageChange() {
window.scrollTo({ top: 0, behavior: 'smooth' })
}
function recalculateNextExpireAt(now: number = currentTime.value) {
let nearestExpireAt: number | null = null
for (const item of affinityList.value) {
if (!item.expire_at || item.expire_at <= now) continue
if (nearestExpireAt === null || item.expire_at < nearestExpireAt) {
nearestExpireAt = item.expire_at
}
}
nextExpireAt.value = nearestExpireAt
}
function pruneExpiredAffinities(now: number, silent = false) {
const beforeCount = affinityList.value.length
const activeItems = affinityList.value.filter(
item => item.expire_at && item.expire_at > now
)
if (activeItems.length === beforeCount) {
recalculateNextExpireAt(now)
return
}
affinityList.value = activeItems
recalculateNextExpireAt(now)
if (!silent) {
showInfo(`${beforeCount - activeItems.length} 个缓存已自动过期移除`)
}
}
// ==================== 定时器管理 ====================
function startCountdown() {
@@ -247,14 +282,8 @@ function startCountdown() {
countdownTimer = setInterval(() => {
currentTime.value = Math.floor(Date.now() / 1000)
const beforeCount = affinityList.value.length
affinityList.value = affinityList.value.filter(
item => item.expire_at && item.expire_at > currentTime.value
)
if (beforeCount > affinityList.value.length) {
const removedCount = beforeCount - affinityList.value.length
showInfo(`${removedCount} 个缓存已自动过期移除`)
if (nextExpireAt.value !== null && currentTime.value >= nextExpireAt.value) {
pruneExpiredAffinities(currentTime.value)
}
}, 1000)
}
@@ -273,6 +302,9 @@ function handleVisibilityChange() {
return
}
currentTime.value = Math.floor(Date.now() / 1000)
if (nextExpireAt.value !== null && currentTime.value >= nextExpireAt.value) {
pruneExpiredAffinities(currentTime.value)
}
startCountdown()
}

View File

@@ -245,7 +245,7 @@ const activeRequestIds = computed(() => {
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 refreshInFlight: Promise<void> | null = null
const AUTO_REFRESH_INTERVAL = 1000 // 1秒刷新一次用于活跃请求
@@ -271,8 +271,10 @@ async function pollActiveRequests() {
let shouldRefresh = false
const recordMap = new Map(currentRecords.value.map(record => [record.id, record]))
for (const update of requests) {
const record = currentRecords.value.find(r => r.id === update.id)
const record = recordMap.get(update.id)
if (!record) {
// 后端返回了未知的活跃请求,触发刷新以获取完整数据
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() {
if (!isPageVisible.value) return
if (autoRefreshTimer) return
autoRefreshTimer = setInterval(pollActiveRequests, AUTO_REFRESH_INTERVAL)
scheduleNextAutoRefresh()
}
// 停止自动刷新
function stopAutoRefresh() {
if (autoRefreshTimer) {
clearInterval(autoRefreshTimer)
clearTimeout(autoRefreshTimer)
autoRefreshTimer = null
}
}

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -0,0 +1,143 @@
from __future__ import annotations
from datetime import datetime, timezone
from decimal import Decimal
from types import SimpleNamespace
from typing import Any
import pytest
from src.api.admin.usage.routes import AdminUsageRecordsAdapter
class _FakeQuery:
def __init__(
self,
*,
scalar_result: int | None = None,
all_result: list[Any] | None = None,
) -> None:
self.scalar_result = scalar_result
self.all_result = all_result or []
self.options_args: tuple[Any, ...] = ()
def outerjoin(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def join(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def options(self, *args: Any) -> _FakeQuery:
self.options_args = args
return self
def order_by(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def offset(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def limit(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def scalar(self) -> int | None:
return self.scalar_result
def all(self) -> list[Any]:
return self.all_result
class _FakeDb:
def __init__(self, queries: list[_FakeQuery]) -> None:
self._queries = queries
self.query_calls: list[tuple[Any, ...]] = []
def query(self, *args: Any) -> _FakeQuery:
self.query_calls.append(args)
return self._queries[len(self.query_calls) - 1]
@pytest.mark.asyncio
async def test_admin_usage_records_returns_model_version_without_request_metadata(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr("src.utils.cache_decorator.get_redis_client_sync", lambda: None)
usage = SimpleNamespace(
id="usage-1",
request_id=None,
user_id="user-1",
api_key_id=None,
provider_name="google",
provider_id=None,
provider_endpoint_id=None,
provider_api_key_id=None,
model="gemini-2.5-pro",
target_model=None,
input_tokens=120,
output_tokens=80,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
total_tokens=200,
total_cost_usd=Decimal("1.25"),
actual_total_cost_usd=Decimal("1.25"),
rate_multiplier=Decimal("1.0"),
response_time_ms=850,
first_byte_time_ms=230,
created_at=datetime(2026, 3, 9, 8, 30, tzinfo=timezone.utc),
is_stream=False,
status_code=200,
error_message=None,
status="completed",
api_format="gemini:chat",
endpoint_api_format=None,
has_format_conversion=False,
input_price_per_1m=Decimal("0.10"),
output_price_per_1m=Decimal("0.30"),
cache_creation_price_per_1m=None,
cache_read_price_per_1m=None,
)
user = SimpleNamespace(id="user-1", email="user@example.com", username="tester")
count_query = _FakeQuery(scalar_result=1)
data_query = _FakeQuery(
all_result=[
(usage, user, None, None, None, "gemini-2.5-pro-001"),
]
)
db = _FakeDb([count_query, data_query])
context = SimpleNamespace(
db=db,
user=SimpleNamespace(id="admin-1"),
add_audit_metadata=lambda **_: None,
)
adapter = AdminUsageRecordsAdapter(
time_range=None,
search=None,
user_id=None,
username=None,
model=None,
provider=None,
api_format=None,
status=None,
limit=100,
offset=0,
)
result = await adapter.handle(context)
assert len(db.query_calls) == 2
assert len(db.query_calls[1]) == 6
assert getattr(db.query_calls[1][-1], "name", None) == "model_version"
record = result["records"][0]
assert record["model_version"] == "gemini-2.5-pro-001"
assert "request_metadata" not in record
usage_load_only = data_query.options_args[0]
usage_paths = {str(option.path) for option in usage_load_only.context}
assert "ORM Path[Mapper[Usage(usage)] -> Usage.request_metadata]" not in usage_paths

View 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"]

View File

@@ -8,54 +8,139 @@ API Pipeline 测试
"""
from datetime import datetime, timezone
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from src.api.base.adapter import ApiMode
from src.api.base.pipeline import ApiRequestPipeline
from src.core.enums import UserRole
from src.core.modules.hooks import AUTH_TOKEN_PREFIX_AUTHENTICATORS
class TestPipelineBalanceCalculation:
"""测试 Pipeline 余额计算"""
"""Balance calculation tests for Pipeline."""
@pytest.fixture
def pipeline(self) -> 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_db = MagicMock()
mock_user.id = "user-123"
thread_db = MagicMock()
db_user = MagicMock()
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=70.0,
):
remaining = pipeline._calculate_balance_remaining(mock_db, mock_user)
remaining = await pipeline._calculate_balance_remaining_async(mock_user)
assert remaining == 70.0
thread_db.close.assert_called_once()
def test_calculate_balance_remaining_unlimited(self, pipeline: ApiRequestPipeline) -> None:
"""测试无限制钱包时返回 None"""
@pytest.mark.asyncio
async def test_calculate_balance_remaining_unlimited(
self, pipeline: ApiRequestPipeline
) -> None:
"""Returns None for unlimited wallets."""
mock_user = MagicMock()
mock_db = MagicMock()
mock_user.id = "user-123"
thread_db = MagicMock()
db_user = MagicMock()
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 = pipeline._calculate_balance_remaining(mock_db, mock_user)
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
def test_calculate_balance_remaining_none_user(self, pipeline: ApiRequestPipeline) -> None:
"""测试用户为 None 时返回 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()
remaining = pipeline._calculate_balance_remaining(mock_db, None)
mock_user = MagicMock()
mock_user.id = "admin-123"
mock_token = MagicMock()
mock_token.id = "mt-123"
assert remaining is None
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),
):
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 response == mock_response
assert mock_context.management_token == mock_token
mock_balance.assert_not_called()
class TestPipelineAuditLogging:
@@ -213,7 +298,8 @@ class TestPipelineAuthentication:
def pipeline(self) -> 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 时抛出异常"""
mock_request = MagicMock()
mock_request.headers = {}
@@ -226,12 +312,13 @@ class TestPipelineAuthentication:
mock_adapter.extract_api_key = MagicMock(return_value=None)
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 "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"""
mock_request = MagicMock()
mock_request.headers = {"Authorization": "Bearer sk-invalid"}
@@ -245,15 +332,17 @@ class TestPipelineAuthentication:
with patch.object(
pipeline.auth_service,
"authenticate_api_key",
"authenticate_api_key_threadsafe",
new_callable=AsyncMock,
return_value=None,
):
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
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.id = "user-123"
@@ -268,28 +357,198 @@ class TestPipelineAuthentication:
mock_request.state = 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.extract_api_key = MagicMock(return_value="sk-test")
with patch.object(
pipeline.auth_service,
"authenticate_api_key",
return_value=(mock_user, mock_api_key),
):
with patch.object(
pipeline.usage_service,
"check_request_balance",
return_value=(False, "余额不足"),
):
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=0.0,
"authenticate_api_key_threadsafe",
new_callable=AsyncMock,
return_value=MagicMock(
user=mock_user,
api_key=mock_api_key,
access_allowed=False,
balance_remaining=0.0,
),
):
from src.core.exceptions import 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:

View 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,
}
]

View File

@@ -69,7 +69,12 @@ async def test_list_all_candidates_returns_provider_batch_count_even_when_candid
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._candidate_builder, "_query_providers", return_value=providers):
with patch.object(
scheduler._candidate_builder,
"_query_provider_refs",
return_value=[("p1", "p1"), ("p2", "p2")],
):
with patch.object(scheduler._candidate_builder, "_query_providers") as query_providers:
with patch(
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
new=AsyncMock(return_value=global_model),
@@ -90,6 +95,8 @@ async def test_list_all_candidates_returns_provider_batch_count_even_when_candid
)
)
query_providers.assert_not_called()
assert candidates == []
assert global_model_id == "gm1"
assert provider_batch_count == 2

View File

@@ -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"])

View File

@@ -8,18 +8,20 @@
"""
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from unittest.mock import AsyncMock, MagicMock, patch
import jwt
import pytest
from src.core.exceptions import ForbiddenException
from src.core.enums import AuthSource
from src.core.exceptions import ForbiddenException
from src.models.database import UserRole
from src.services.auth.service import (
JWT_ALGORITHM,
JWT_EXPIRATION_HOURS,
JWT_SECRET_KEY,
AuthenticatedUserSnapshot,
AuthService,
)
@@ -235,6 +237,115 @@ class TestUserAuthentication:
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:
"""测试 API Key 认证"""

View File

@@ -1,36 +1,45 @@
from types import SimpleNamespace
from typing import Any, Callable, cast
import pytest
from src.core.cache_service import CacheService
from src.models.database import GlobalModel, Model
from src.services.cache.model_cache import ModelCacheService
from src.core.cache_service import CacheService
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._all_result = all_result if all_result is not None else []
self._on_all = on_all
def join(self, *_args, **_kwargs):
def join(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
return self
def filter(self, *_args, **_kwargs):
def filter(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
return self
def first(self):
def first(self) -> Any:
return self._first_result
def all(self):
def all(self) -> list[Any]:
if self._on_all:
self._on_all()
return self._all_result
class _FakeSession:
def __init__(self, *, direct_match: GlobalModel):
def __init__(self, *, direct_match: GlobalModel) -> None:
self._direct_match = direct_match
def query(self, *entities):
def query(self, *entities: object) -> "_FakeQuery":
if entities == (GlobalModel,):
return _FakeQuery(first_result=self._direct_match)
@@ -43,12 +52,43 @@ class _FakeSession:
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
async def test_resolve_global_model_prefers_direct_match(monkeypatch) -> None:
async def _fake_get(_key: str):
async def test_resolve_global_model_prefers_direct_match(monkeypatch: pytest.MonkeyPatch) -> None:
async def _fake_get(_key: str) -> 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
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)
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
@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

View 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)

View File

@@ -8,7 +8,7 @@ UsageService 测试
"""
from decimal import Decimal
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import MagicMock, patch
import pytest
@@ -146,6 +146,58 @@ class TestBalanceCheck:
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:
"""测试余额耗尽时拦截新请求"""
mock_user = MagicMock()

View File

@@ -1,5 +1,6 @@
from decimal import Decimal
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock, patch
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")
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.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")
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.user_id is None
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:
wallet = _build_wallet(recharge="-1", gift="10", limit_mode="finite")
db = MagicMock()
@@ -104,7 +129,7 @@ def test_admin_adjust_balance_negative_from_gift_spills_to_recharge() -> None:
tx = WalletService.admin_adjust_balance(
db,
wallet=wallet,
wallet=cast(Any, wallet),
amount_usd=Decimal("-10"),
balance_type="gift",
operator_id="admin-1",
@@ -127,7 +152,7 @@ def test_admin_adjust_balance_negative_from_recharge_then_gift() -> None:
tx = WalletService.admin_adjust_balance(
db,
wallet=wallet,
wallet=cast(Any, wallet),
amount_usd=Decimal("-4"),
balance_type="recharge",
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(
db,
wallet=wallet,
wallet=cast(Any, wallet),
amount_usd=Decimal("1"),
balance_type="gift",
operator_id="admin-1",
@@ -181,7 +206,9 @@ def test_apply_usage_charge_prefers_gift_then_recharge() -> None:
db = _build_locked_db(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 after == Decimal("2.00000000")
@@ -212,7 +239,9 @@ def test_apply_usage_charge_unlimited_wallet_keeps_balances() -> None:
db = _build_locked_db(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 after == Decimal("8.00000000")
@@ -237,7 +266,7 @@ def test_complete_refund_requires_processing_status() -> None:
db.query.return_value = query
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:
@@ -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]):
wallet = WalletService.get_or_create_wallet(
db,
user=SimpleNamespace(id="user-1"),
user=cast(Any, SimpleNamespace(id="user-1")),
api_key=None,
)
@@ -287,18 +316,20 @@ def test_create_refund_request_rejects_uncredited_payment_order() -> None:
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"):
WalletService.create_refund_request(
db,
wallet=wallet,
wallet=cast(Any, wallet),
user_id="user-1",
amount_usd=Decimal("2"),
refund_no="rf-1",
source_type="payment_order",
source_id="order-1",
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"):
WalletService.create_refund_request(
db,
wallet=wallet,
wallet=cast(Any, wallet),
user_id="user-1",
amount_usd=Decimal("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"):
WalletService.create_refund_request(
db,
wallet=wallet,
wallet=cast(Any, wallet),
user_id="user-1",
amount_usd=Decimal("2"),
refund_no="rf-3",
source_type="payment_order",
source_id="order-1",
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")
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"):
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 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 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()
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:
first_tx = WalletService.fail_refund(
db,
refund=refund,
refund=cast(Any, refund),
reason="first-failure",
operator_id="admin-1",
)
with pytest.raises(ValueError, match="cannot fail refund in status: failed"):
WalletService.fail_refund(
db,
refund=refund,
refund=cast(Any, refund),
reason="retry-failure",
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"):
WalletService.fail_refund(
db,
refund=refund,
refund=cast(Any, refund),
reason="should-not-override",
operator_id="admin-1",
)

View File

@@ -11,6 +11,13 @@ from src.api.base.context import ApiRequestContext
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 = [
(str(key).encode("latin-1"), str(value).encode("latin-1"))
for key, value in (headers or {}).items()
@@ -28,8 +35,14 @@ def _build_request(headers: dict[str, str] | None = None) -> Request:
"server": ("testserver", 80),
}
received = False
async def receive() -> dict[str, object]:
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.state.perf_metrics = {}
@@ -89,3 +102,22 @@ class TestApiRequestContextEnsureJsonBody:
assert context.client_content_encoding == "gzip"
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")

View 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()