mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
@@ -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>
|
||||
@@ -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>
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
// 日期范围参数
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -73,13 +73,7 @@ class AdminPercentilesAdapter(AdminApiAdapter):
|
||||
)
|
||||
return result
|
||||
|
||||
result = []
|
||||
for local_date, day_start_utc, day_end_utc in time_range.get_local_day_hours():
|
||||
percentiles = StatsAggregatorService.compute_daily_percentiles(
|
||||
context.db, day_start_utc, day_end_utc
|
||||
)
|
||||
result.append({"date": local_date.isoformat(), **percentiles})
|
||||
return result
|
||||
return StatsAggregatorService.compute_percentiles_by_local_day(context.db, time_range)
|
||||
|
||||
|
||||
@router.get("/performance/percentiles")
|
||||
|
||||
@@ -10,7 +10,7 @@ from typing import Any
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy import case, func
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
@@ -975,9 +975,6 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
from src.core.crypto import crypto_service
|
||||
from src.models.database import (
|
||||
GlobalModel,
|
||||
Model,
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
ProxyNode,
|
||||
)
|
||||
|
||||
@@ -991,22 +988,35 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
gm_name_map: dict[str, str] = {gm.id: gm.name for gm in global_models}
|
||||
|
||||
# 导出 Providers 及其关联数据
|
||||
providers = db.query(Provider).all()
|
||||
providers = (
|
||||
db.query(Provider)
|
||||
.options(
|
||||
selectinload(Provider.endpoints),
|
||||
selectinload(Provider.api_keys),
|
||||
selectinload(Provider.models),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
providers_data = []
|
||||
|
||||
def _normalize_created_at_for_sort(value: datetime | None) -> datetime:
|
||||
if value is None:
|
||||
return datetime.min.replace(tzinfo=timezone.utc)
|
||||
return value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc)
|
||||
|
||||
for provider in providers:
|
||||
# 导出 Endpoints
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint).filter(ProviderEndpoint.provider_id == provider.id).all()
|
||||
)
|
||||
endpoints = list(provider.endpoints)
|
||||
endpoints_data = [ep.to_export_dict() for ep in endpoints]
|
||||
provider_endpoint_formats = self._collect_provider_endpoint_formats(endpoints)
|
||||
|
||||
# 导出 Provider Keys(按 provider_id 归属,包含 api_formats)
|
||||
keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
.filter(ProviderAPIKey.provider_id == provider.id)
|
||||
.order_by(ProviderAPIKey.internal_priority.asc(), ProviderAPIKey.created_at.asc())
|
||||
.all()
|
||||
keys = sorted(
|
||||
provider.api_keys,
|
||||
key=lambda key: (
|
||||
key.internal_priority if key.internal_priority is not None else 0,
|
||||
_normalize_created_at_for_sort(key.created_at),
|
||||
),
|
||||
)
|
||||
keys_data = []
|
||||
for key in keys:
|
||||
@@ -1046,7 +1056,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
|
||||
# 导出 Provider Models
|
||||
# 注意:提供商模型(Model)必须关联全局模型(GlobalModel)才能参与路由
|
||||
# 导入时未关联 GlobalModel 的模型会被跳过,这是业务规则而非 bug
|
||||
models = db.query(Model).filter(Model.provider_id == provider.id).all()
|
||||
models = list(provider.models)
|
||||
models_data = []
|
||||
for model in models:
|
||||
model_data = model.to_export_dict()
|
||||
|
||||
@@ -224,10 +224,11 @@ async def get_usage_records(
|
||||
|
||||
**返回字段**:
|
||||
- `records`: 使用记录列表,包含 id, user_id, user_email, username, api_key, provider, model, target_model,
|
||||
input_tokens, output_tokens, cache_creation_input_tokens, cache_read_input_tokens, total_tokens,
|
||||
cost, actual_cost, rate_multiplier, response_time_ms, first_byte_time_ms, created_at, is_stream,
|
||||
input_price_per_1m, output_price_per_1m, cache_creation_price_per_1m, cache_read_price_per_1m,
|
||||
status_code, error_message, status, has_fallback, has_retry, has_rectified, api_format, api_key_name, request_metadata
|
||||
model_version, input_tokens, output_tokens, cache_creation_input_tokens, cache_read_input_tokens,
|
||||
total_tokens, cost, actual_cost, rate_multiplier, response_time_ms, first_byte_time_ms, created_at,
|
||||
is_stream, input_price_per_1m, output_price_per_1m, cache_creation_price_per_1m,
|
||||
cache_read_price_per_1m, status_code, error_message, status, has_fallback, has_retry,
|
||||
has_rectified, api_format, api_key_name
|
||||
- `total`: 符合条件的总记录数
|
||||
- `limit`: 当前分页限制
|
||||
- `offset`: 当前分页偏移量
|
||||
@@ -885,8 +886,12 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
count_query = count_query.outerjoin(ApiKey, Usage.api_key_id == ApiKey.id)
|
||||
|
||||
# -- 构建数据查询(完整 JOIN) --
|
||||
usage_model_version = Usage.request_metadata["model_version"].as_string().label(
|
||||
"model_version"
|
||||
)
|
||||
|
||||
query = (
|
||||
db.query(Usage, User, ProviderEndpoint, ProviderAPIKey, ApiKey)
|
||||
db.query(Usage, User, ProviderEndpoint, ProviderAPIKey, ApiKey, usage_model_version)
|
||||
.outerjoin(User, Usage.user_id == User.id)
|
||||
.outerjoin(ProviderEndpoint, Usage.provider_endpoint_id == ProviderEndpoint.id)
|
||||
.outerjoin(ProviderAPIKey, Usage.provider_api_key_id == ProviderAPIKey.id)
|
||||
@@ -1001,7 +1006,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
# Perf: count query uses fewer JOINs than the data query
|
||||
total = int(count_query.scalar() or 0)
|
||||
|
||||
# Perf: do not load large request/response columns for list view
|
||||
# Perf: do not load large request/response columns or full request_metadata for list view
|
||||
query = query.options(
|
||||
load_only(
|
||||
Usage.id,
|
||||
@@ -1032,7 +1037,6 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
Usage.api_format,
|
||||
Usage.endpoint_api_format,
|
||||
Usage.has_format_conversion,
|
||||
Usage.request_metadata,
|
||||
Usage.input_price_per_1m,
|
||||
Usage.output_price_per_1m,
|
||||
Usage.cache_creation_price_per_1m,
|
||||
@@ -1047,7 +1051,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
|
||||
)
|
||||
|
||||
request_ids = [usage.request_id for usage, _, _, _, _ in records if usage.request_id]
|
||||
request_ids = [usage.request_id for usage, _, _, _, _, _ in records if usage.request_id]
|
||||
fallback_map = {}
|
||||
retry_map = {}
|
||||
rectified_map = {}
|
||||
@@ -1110,7 +1114,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
|
||||
# 构建 provider_id -> Provider 名称的映射,避免 N+1 查询
|
||||
provider_ids = list(
|
||||
{usage.provider_id for usage, _, _, _, _ in records if usage.provider_id}
|
||||
{usage.provider_id for usage, _, _, _, _, _ in records if usage.provider_id}
|
||||
)
|
||||
provider_map = {}
|
||||
if provider_ids:
|
||||
@@ -1121,7 +1125,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
|
||||
data = []
|
||||
api_key_display_cache: dict[str, str] = {}
|
||||
for usage, user, endpoint, provider_api_key, user_api_key in records:
|
||||
for usage, user, endpoint, provider_api_key, user_api_key, model_version in records:
|
||||
actual_cost = (
|
||||
float(usage.actual_total_cost_usd)
|
||||
if usage.actual_total_cost_usd is not None
|
||||
@@ -1198,7 +1202,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
"endpoint_api_format": endpoint_api_format,
|
||||
"has_format_conversion": bool(has_format_conversion),
|
||||
"api_key_name": provider_api_key.name if provider_api_key else None,
|
||||
"request_metadata": usage.request_metadata, # Provider 响应元数据
|
||||
"model_version": model_version, # Provider 返回的实际模型版本(轻量字段)
|
||||
}
|
||||
)
|
||||
|
||||
@@ -1227,7 +1231,10 @@ class AdminActiveRequestsAdapter(AdminApiAdapter):
|
||||
return {"requests": []}
|
||||
|
||||
requests = UsageService.get_active_requests_status(
|
||||
db=db, ids=id_list, include_admin_fields=True
|
||||
db=db,
|
||||
ids=id_list,
|
||||
include_admin_fields=True,
|
||||
maintain_status=True,
|
||||
)
|
||||
return {"requests": requests}
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.admin_requests import UpdateUserRequest
|
||||
from src.models.api import CreateApiKeyRequest, CreateUserRequest
|
||||
from src.models.database import ApiKey, User, UserRole
|
||||
from src.models.database import ApiKey, User, UserRole, Wallet
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.user.apikey import ApiKeyService
|
||||
from src.services.user.bulk_cleanup import pre_clean_api_key
|
||||
@@ -30,8 +30,23 @@ router = APIRouter(prefix="/api/admin/users", tags=["Admin - Users"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def _serialize_user(db: Session, user: User) -> dict[str, Any]:
|
||||
wallet = WalletService.get_wallet(db, user_id=user.id)
|
||||
class _WalletSentinelType:
|
||||
pass
|
||||
|
||||
|
||||
_WALLET_SENTINEL = _WalletSentinelType()
|
||||
|
||||
|
||||
def _serialize_user(
|
||||
db: Session,
|
||||
user: User,
|
||||
wallet: Wallet | None | _WalletSentinelType = _WALLET_SENTINEL,
|
||||
) -> dict[str, Any]:
|
||||
resolved_wallet: Wallet | None
|
||||
if wallet is _WALLET_SENTINEL:
|
||||
resolved_wallet = WalletService.get_wallet(db, user_id=user.id)
|
||||
else:
|
||||
resolved_wallet = wallet
|
||||
return {
|
||||
"id": user.id,
|
||||
"email": user.email,
|
||||
@@ -40,7 +55,7 @@ def _serialize_user(db: Session, user: User) -> dict[str, Any]:
|
||||
"allowed_providers": user.allowed_providers,
|
||||
"allowed_api_formats": user.allowed_api_formats,
|
||||
"allowed_models": user.allowed_models,
|
||||
"unlimited": WalletService.is_unlimited_wallet(wallet),
|
||||
"unlimited": WalletService.is_unlimited_wallet(resolved_wallet),
|
||||
"is_active": user.is_active,
|
||||
"created_at": user.created_at.isoformat(),
|
||||
"updated_at": user.updated_at.isoformat() if user.updated_at else None,
|
||||
@@ -334,7 +349,8 @@ class AdminListUsersAdapter(AdminApiAdapter):
|
||||
except KeyError as exc:
|
||||
raise InvalidRequestException("角色参数不合法") from exc
|
||||
users = UserService.list_users(db, self.skip, self.limit, role_enum, self.is_active)
|
||||
return [_serialize_user(db, u) for u in users]
|
||||
wallets_by_user_id = WalletService.get_wallets_by_user_ids(db, [user.id for user in users])
|
||||
return [_serialize_user(db, user, wallets_by_user_id.get(user.id)) for user in users]
|
||||
|
||||
|
||||
class AdminGetUserAdapter(AdminApiAdapter):
|
||||
|
||||
@@ -274,10 +274,10 @@ class AuthLoginAdapter(AuthPublicAdapter):
|
||||
detail=f"登录请求过于频繁,请在 {reset_after} 秒后重试",
|
||||
)
|
||||
|
||||
user = await AuthService.authenticate_user(
|
||||
authenticated_user = await AuthService.authenticate_user_threadsafe(
|
||||
db, login_request.email, login_request.password, login_request.auth_type
|
||||
)
|
||||
if not user:
|
||||
if not authenticated_user:
|
||||
AuditService.log_login_attempt(
|
||||
db=db,
|
||||
email=login_request.email,
|
||||
@@ -296,22 +296,30 @@ class AuthLoginAdapter(AuthPublicAdapter):
|
||||
success=True,
|
||||
ip_address=client_ip,
|
||||
user_agent=user_agent,
|
||||
user_id=user.id,
|
||||
user_id=authenticated_user.user_id,
|
||||
)
|
||||
db.commit()
|
||||
context.request.state.tx_committed_by_route = True
|
||||
|
||||
access_token = AuthService.create_access_token(
|
||||
data={
|
||||
"user_id": user.id,
|
||||
"role": user.role.value,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
"user_id": authenticated_user.user_id,
|
||||
"role": authenticated_user.role.value,
|
||||
"created_at": (
|
||||
authenticated_user.created_at.isoformat()
|
||||
if authenticated_user.created_at
|
||||
else None
|
||||
),
|
||||
}
|
||||
)
|
||||
refresh_token = AuthService.create_refresh_token(
|
||||
data={
|
||||
"user_id": user.id,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
"user_id": authenticated_user.user_id,
|
||||
"created_at": (
|
||||
authenticated_user.created_at.isoformat()
|
||||
if authenticated_user.created_at
|
||||
else None
|
||||
),
|
||||
}
|
||||
)
|
||||
response = LoginResponse(
|
||||
@@ -319,10 +327,10 @@ class AuthLoginAdapter(AuthPublicAdapter):
|
||||
refresh_token=refresh_token,
|
||||
token_type="bearer",
|
||||
expires_in=86400,
|
||||
user_id=user.id,
|
||||
email=user.email,
|
||||
username=user.username,
|
||||
role=user.role.value,
|
||||
user_id=authenticated_user.user_id,
|
||||
email=authenticated_user.email,
|
||||
username=authenticated_user.username,
|
||||
role=authenticated_user.role.value,
|
||||
)
|
||||
return response.model_dump()
|
||||
|
||||
@@ -332,9 +340,6 @@ class AuthRefreshAdapter(AuthPublicAdapter):
|
||||
db = context.db
|
||||
payload = context.ensure_json_body()
|
||||
refresh_request = RefreshTokenRequest.model_validate(payload)
|
||||
client_ip = get_client_ip(context.request)
|
||||
user_agent = get_user_agent(context.request)
|
||||
|
||||
try:
|
||||
token_payload = await AuthService.verify_token(
|
||||
refresh_request.refresh_token, token_type="refresh"
|
||||
@@ -745,7 +750,6 @@ class AuthSendVerificationCodeAdapter(AuthPublicAdapter):
|
||||
class AuthVerifyEmailAdapter(AuthPublicAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
"""验证邮箱验证码"""
|
||||
db = context.db
|
||||
payload = context.ensure_json_body()
|
||||
|
||||
try:
|
||||
|
||||
@@ -25,6 +25,7 @@ class ApiAdapter(ABC):
|
||||
audit_log_enabled: bool = True
|
||||
audit_success_event = None
|
||||
audit_failure_event = None
|
||||
eager_request_body: bool = True
|
||||
|
||||
@abstractmethod
|
||||
async def handle(self, context: ApiRequestContext) -> Response:
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import gzip
|
||||
import json
|
||||
import time
|
||||
@@ -9,7 +10,9 @@ from typing import Any
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.requests import ClientDisconnect
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.api_format.headers import get_header_value
|
||||
from src.core.http_compression import is_gzip_content_encoding, normalize_content_encoding
|
||||
from src.core.logger import logger
|
||||
@@ -53,6 +56,56 @@ class ApiRequestContext:
|
||||
client_content_encoding: str | None = None
|
||||
client_accept_encoding: str | None = None
|
||||
|
||||
async def ensure_raw_body_async(self) -> bytes:
|
||||
"""按需读取原始请求体,避免所有请求都在 Pipeline 阶段预读。"""
|
||||
if self.raw_body is not None:
|
||||
return self.raw_body
|
||||
|
||||
perf_metrics = getattr(self.request.state, "perf_metrics", None)
|
||||
perf_sampled = isinstance(perf_metrics, dict) and bool(perf_metrics)
|
||||
body_start = PerfRecorder.start(force=perf_sampled)
|
||||
body_size = 0
|
||||
try:
|
||||
self.raw_body = await asyncio.wait_for(
|
||||
self.request.body(), timeout=config.request_body_timeout
|
||||
)
|
||||
body_size = len(self.raw_body or b"")
|
||||
except TimeoutError as exc:
|
||||
timeout_sec = int(config.request_body_timeout)
|
||||
logger.error("读取请求体超时({}s),可能客户端未发送完整请求体", timeout_sec)
|
||||
raise HTTPException(
|
||||
status_code=408,
|
||||
detail=f"Request timeout: body not received within {timeout_sec} seconds",
|
||||
) from exc
|
||||
except ClientDisconnect:
|
||||
logger.warning(
|
||||
"[Context] 客户端在读取请求体期间断开连接: {} {}",
|
||||
self.request.method,
|
||||
self.request.url.path,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=499,
|
||||
detail="Client closed request",
|
||||
)
|
||||
finally:
|
||||
body_duration = PerfRecorder.stop(
|
||||
body_start,
|
||||
"pipeline_body_read",
|
||||
labels={"mode": self.mode},
|
||||
log_hint=f"size={body_size}",
|
||||
)
|
||||
if isinstance(perf_metrics, dict):
|
||||
pipeline_metrics = perf_metrics.setdefault("pipeline", {})
|
||||
pipeline_metrics["body_read_ms"] = int((body_duration or 0) * 1000)
|
||||
pipeline_metrics["body_bytes"] = int(body_size)
|
||||
|
||||
return self.raw_body or b""
|
||||
|
||||
async def ensure_json_body_async(self) -> dict[str, Any]:
|
||||
"""异步懒加载 JSON 请求体。"""
|
||||
await self.ensure_raw_body_async()
|
||||
return self.ensure_json_body()
|
||||
|
||||
def ensure_json_body(self) -> dict[str, Any]:
|
||||
"""确保请求体已解析为JSON并返回。"""
|
||||
if self.json_body is not None:
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from fastapi.concurrency import run_in_threadpool
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -14,6 +16,7 @@ from src.config.settings import config
|
||||
from src.core.enums import UserRole
|
||||
from src.core.exceptions import BalanceInsufficientException
|
||||
from src.core.logger import logger
|
||||
from src.database.database import create_session
|
||||
from src.models.database import ApiKey, AuditEventType, User
|
||||
from src.services.auth.service import AuthService
|
||||
from src.services.system.audit import AuditService
|
||||
@@ -108,14 +111,22 @@ class ApiRequestPipeline:
|
||||
user, management_token = await self._authenticate_management(http_request, db)
|
||||
api_key = None
|
||||
else:
|
||||
user, api_key = self._authenticate_client(http_request, db, adapter, quiet=is_quiet)
|
||||
user, api_key = await self._authenticate_client(
|
||||
http_request,
|
||||
db,
|
||||
adapter,
|
||||
quiet=is_quiet,
|
||||
)
|
||||
management_token = None
|
||||
finally:
|
||||
auth_duration = PerfRecorder.stop(auth_start, "pipeline_auth", labels=perf_labels)
|
||||
_record_perf_metric("auth_ms", auth_duration)
|
||||
|
||||
raw_body = None
|
||||
if http_request.method in {"POST", "PUT", "PATCH"}:
|
||||
should_eager_read_body = http_request.method in {"POST", "PUT", "PATCH"} and getattr(
|
||||
adapter, "eager_request_body", True
|
||||
)
|
||||
if should_eager_read_body:
|
||||
try:
|
||||
import asyncio
|
||||
|
||||
@@ -141,7 +152,7 @@ class ApiRequestPipeline:
|
||||
perf_metrics.setdefault("pipeline", {})["body_bytes"] = int(body_size)
|
||||
except TimeoutError:
|
||||
timeout_sec = int(config.request_body_timeout)
|
||||
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
|
||||
logger.error("读取请求体超时({}s),可能客户端未发送完整请求体", timeout_sec)
|
||||
raise HTTPException(
|
||||
status_code=408,
|
||||
detail=f"Request timeout: body not received within {timeout_sec} seconds",
|
||||
@@ -176,8 +187,11 @@ class ApiRequestPipeline:
|
||||
context.management_token = management_token
|
||||
# 存储 quiet 标志到 context,用于审计日志判断
|
||||
context.quiet_logging = is_quiet
|
||||
if mode != ApiMode.ADMIN and user:
|
||||
remaining = self._calculate_balance_remaining(db, user, api_key=api_key)
|
||||
if mode in {ApiMode.STANDARD, ApiMode.PROXY, ApiMode.USER} and user:
|
||||
if hasattr(http_request.state, "prefetched_balance_remaining"):
|
||||
remaining = getattr(http_request.state, "prefetched_balance_remaining")
|
||||
else:
|
||||
remaining = await self._calculate_balance_remaining_async(user, api_key=api_key)
|
||||
context.balance_remaining = remaining
|
||||
# authorize 可能是异步的,需要检查并 await
|
||||
authorize_start = PerfRecorder.start(force=perf_sampled)
|
||||
@@ -238,7 +252,7 @@ class ApiRequestPipeline:
|
||||
try:
|
||||
context.db.rollback()
|
||||
except Exception as rollback_exc:
|
||||
logger.debug(f"[Pipeline] 回滚失败(可忽略): {rollback_exc}")
|
||||
logger.debug("[Pipeline] 回滚失败(可忽略): {}", rollback_exc)
|
||||
self._record_audit_event(
|
||||
context,
|
||||
adapter,
|
||||
@@ -252,32 +266,52 @@ class ApiRequestPipeline:
|
||||
# Internal helpers
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
def _authenticate_client(
|
||||
async def _authenticate_client(
|
||||
self, request: Request, db: Session, adapter: ApiAdapter, **_kw: object
|
||||
) -> tuple[User, ApiKey]:
|
||||
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
|
||||
client_api_key = adapter.extract_api_key(request)
|
||||
if not client_api_key:
|
||||
raise HTTPException(status_code=401, detail="请提供API密钥")
|
||||
|
||||
auth_result = self.auth_service.authenticate_api_key(db, client_api_key)
|
||||
auth_result = await self.auth_service.authenticate_api_key_threadsafe(client_api_key)
|
||||
if not auth_result:
|
||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
||||
|
||||
user, api_key = auth_result
|
||||
user = auth_result.user
|
||||
api_key = auth_result.api_key
|
||||
if not user or not api_key:
|
||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
||||
|
||||
request.state.user_id = user.id
|
||||
request.state.api_key_id = api_key.id
|
||||
# 线程池认证返回的是分离对象;重新绑定到路由会话,避免后续写入失效。
|
||||
db_user = db.query(User).filter(User.id == user.id).first()
|
||||
db_api_key = db.query(ApiKey).filter(ApiKey.id == api_key.id).first()
|
||||
if not db_user or not db_api_key:
|
||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
||||
# 使用路由会话再核对一次状态,避免线程池认证结果与当前事务视图短暂不一致。
|
||||
if not db_user.is_active or db_user.is_deleted:
|
||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
||||
if not db_api_key.is_active:
|
||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
||||
if db_api_key.is_locked and not db_api_key.is_standalone:
|
||||
raise HTTPException(status_code=403, detail="该密钥已被管理员锁定,请联系管理员")
|
||||
if db_api_key.user_id != db_user.id:
|
||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
||||
if db_api_key.expires_at:
|
||||
expires_at = db_api_key.expires_at
|
||||
if expires_at.tzinfo is None:
|
||||
expires_at = expires_at.replace(tzinfo=timezone.utc)
|
||||
if expires_at < datetime.now(timezone.utc):
|
||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
||||
|
||||
# 检查余额(支持独立 Key)
|
||||
access_ok, _message = self.usage_service.check_request_balance(db, user, api_key=api_key)
|
||||
if not access_ok:
|
||||
remaining = self._calculate_balance_remaining(db, user, api_key=api_key)
|
||||
request.state.user_id = db_user.id
|
||||
request.state.api_key_id = db_api_key.id
|
||||
request.state.prefetched_balance_remaining = auth_result.balance_remaining
|
||||
|
||||
if not auth_result.access_allowed:
|
||||
remaining = auth_result.balance_remaining
|
||||
raise BalanceInsufficientException(balance_type="USD", remaining=remaining)
|
||||
|
||||
return user, api_key
|
||||
return db_user, db_api_key
|
||||
|
||||
async def _try_token_prefix_auth(
|
||||
self, token: str, request: Request, db: Session
|
||||
@@ -293,9 +327,7 @@ class ApiRequestPipeline:
|
||||
from src.core.modules.hooks import AUTH_TOKEN_PREFIX_AUTHENTICATORS, get_hook_dispatcher
|
||||
from src.utils.request_utils import get_client_ip
|
||||
|
||||
authenticators = await get_hook_dispatcher().dispatch(
|
||||
AUTH_TOKEN_PREFIX_AUTHENTICATORS, db=db
|
||||
)
|
||||
authenticators = await get_hook_dispatcher().dispatch(AUTH_TOKEN_PREFIX_AUTHENTICATORS)
|
||||
for auth_info in authenticators or []:
|
||||
prefix = auth_info.get("prefix", "")
|
||||
authenticate_fn = auth_info.get("authenticate")
|
||||
@@ -304,111 +336,136 @@ class ApiRequestPipeline:
|
||||
logger.warning("Token prefix '{}' has no authenticate callback", prefix)
|
||||
raise HTTPException(status_code=401, detail="认证服务不可用")
|
||||
client_ip = get_client_ip(request)
|
||||
result = await authenticate_fn(db, token, client_ip)
|
||||
auth_db = create_session()
|
||||
try:
|
||||
result = await authenticate_fn(auth_db, token, client_ip)
|
||||
if result:
|
||||
for instance in result:
|
||||
if instance is None:
|
||||
continue
|
||||
try:
|
||||
auth_db.expunge(instance)
|
||||
except Exception:
|
||||
pass
|
||||
return result
|
||||
finally:
|
||||
auth_db.close()
|
||||
# 前缀匹配但认证失败
|
||||
module_name = auth_info.get("module", "unknown")
|
||||
raise HTTPException(status_code=401, detail=f"无效或过期的 Token ({module_name})")
|
||||
return None # 无前缀匹配
|
||||
|
||||
def _reattach_token_auth_result(
|
||||
self,
|
||||
db: Session,
|
||||
token_auth_result: tuple[User, Any],
|
||||
) -> tuple[User, Any]:
|
||||
"""将前缀认证返回对象绑定到当前请求会话,避免后续写入失效。"""
|
||||
user, management_token = token_auth_result
|
||||
|
||||
db_user = db.query(User).filter(User.id == user.id).first()
|
||||
if not db_user:
|
||||
raise HTTPException(status_code=401, detail="无效或过期的 Token")
|
||||
|
||||
if management_token is None:
|
||||
return db_user, None
|
||||
|
||||
token_id = getattr(management_token, "id", None)
|
||||
token_model: Any = type(management_token)
|
||||
if token_id is None or not hasattr(token_model, "id"):
|
||||
return db_user, management_token
|
||||
|
||||
db_management_token = db.query(token_model).filter(token_model.id == token_id).first()
|
||||
if not db_management_token:
|
||||
raise HTTPException(status_code=401, detail="无效或过期的 Token")
|
||||
return db_user, db_management_token
|
||||
|
||||
async def _authenticate_admin(
|
||||
self, request: Request, db: Session
|
||||
) -> tuple[User, ManagementToken | None]:
|
||||
"""管理员认证,支持 JWT 和 Management Token 两种方式"""
|
||||
"""Admin auth supports JWT and Management Token."""
|
||||
authorization = request.headers.get("authorization")
|
||||
if not authorization or not authorization.lower().startswith("bearer "):
|
||||
raise HTTPException(status_code=401, detail="缺少管理员凭证")
|
||||
|
||||
token = authorization[7:].strip()
|
||||
|
||||
# 通过钩子检查是否匹配模块注册的 token 前缀(如 ae_)
|
||||
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
||||
if token_auth_result is not None:
|
||||
user, management_token = token_auth_result
|
||||
user, management_token = self._reattach_token_auth_result(db, token_auth_result)
|
||||
|
||||
# 检查管理员权限
|
||||
if user.role != UserRole.ADMIN:
|
||||
logger.warning(f"非管理员尝试通过 Management Token 访问管理端点: {user.email}")
|
||||
logger.warning("非管理员尝试通过 Management Token 访问管理端点: {}", user.email)
|
||||
raise HTTPException(status_code=403, detail="需要管理员权限")
|
||||
|
||||
# 存储到 request.state
|
||||
request.state.user_id = user.id
|
||||
request.state.management_token_id = management_token.id
|
||||
|
||||
request.state.management_token_id = management_token.id if management_token else None
|
||||
return user, management_token
|
||||
|
||||
# JWT 认证
|
||||
try:
|
||||
payload = await self.auth_service.verify_token(token, token_type="access")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error(f"Admin token 验证失败: {exc}")
|
||||
logger.error("Admin token 验证失败: {}", exc)
|
||||
raise HTTPException(status_code=401, detail="无效的管理员令牌")
|
||||
|
||||
user_id = payload.get("user_id")
|
||||
if not user_id:
|
||||
raise HTTPException(status_code=401, detail="无效的管理员令牌")
|
||||
|
||||
# 直接查询数据库,确保返回的是当前 Session 绑定的对象
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user or not user.is_active or user.is_deleted:
|
||||
db_user = db.query(User).filter(User.id == user_id).first()
|
||||
if not db_user or not db_user.is_active or db_user.is_deleted:
|
||||
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
|
||||
|
||||
if not self.auth_service.token_identity_matches_user(payload, user):
|
||||
if not self.auth_service.token_identity_matches_user(payload, db_user):
|
||||
raise HTTPException(status_code=403, detail="无效的管理员令牌")
|
||||
|
||||
# 检查管理员权限
|
||||
if user.role != UserRole.ADMIN:
|
||||
logger.warning(f"非管理员尝试通过 JWT 访问管理端点: {user.email}")
|
||||
if db_user.role != UserRole.ADMIN:
|
||||
logger.warning("非管理员尝试通过 JWT 访问管理端点: {}", db_user.email)
|
||||
raise HTTPException(status_code=403, detail="需要管理员权限")
|
||||
|
||||
request.state.user_id = user.id
|
||||
return user, None
|
||||
request.state.user_id = db_user.id
|
||||
return db_user, None
|
||||
|
||||
async def _authenticate_user(
|
||||
self, request: Request, db: Session
|
||||
) -> tuple[User, ManagementToken | None]:
|
||||
"""用户认证,支持 JWT 和 Management Token 两种方式"""
|
||||
"""User auth supports JWT and Management Token."""
|
||||
authorization = request.headers.get("authorization")
|
||||
if not authorization or not authorization.lower().startswith("bearer "):
|
||||
raise HTTPException(status_code=401, detail="缺少用户凭证")
|
||||
|
||||
token = authorization[7:].strip()
|
||||
|
||||
# 通过钩子检查是否匹配模块注册的 token 前缀(如 ae_)
|
||||
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
||||
if token_auth_result is not None:
|
||||
user, management_token = token_auth_result
|
||||
|
||||
user, management_token = self._reattach_token_auth_result(db, token_auth_result)
|
||||
request.state.user_id = user.id
|
||||
request.state.management_token_id = management_token.id
|
||||
|
||||
request.state.management_token_id = management_token.id if management_token else None
|
||||
return user, management_token
|
||||
|
||||
# JWT 认证
|
||||
try:
|
||||
payload = await self.auth_service.verify_token(token, token_type="access")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error(f"User token 验证失败: {exc}")
|
||||
logger.error("User token 验证失败: {}", exc)
|
||||
raise HTTPException(status_code=401, detail="无效的用户令牌")
|
||||
|
||||
user_id = payload.get("user_id")
|
||||
if not user_id:
|
||||
raise HTTPException(status_code=401, detail="无效的用户令牌")
|
||||
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user or not user.is_active or user.is_deleted:
|
||||
db_user = db.query(User).filter(User.id == user_id).first()
|
||||
if not db_user or not db_user.is_active or db_user.is_deleted:
|
||||
raise HTTPException(status_code=403, detail="用户不存在或已禁用")
|
||||
|
||||
if not self.auth_service.token_identity_matches_user(payload, user):
|
||||
if not self.auth_service.token_identity_matches_user(payload, db_user):
|
||||
raise HTTPException(status_code=403, detail="无效的用户令牌")
|
||||
|
||||
request.state.user_id = user.id
|
||||
return user, None
|
||||
request.state.user_id = db_user.id
|
||||
return db_user, None
|
||||
|
||||
async def _authenticate_management(
|
||||
self, request: Request, db: Session
|
||||
@@ -424,11 +481,11 @@ class ApiRequestPipeline:
|
||||
# _try_token_prefix_auth 会在前缀匹配但认证失败时直接抛 HTTPException
|
||||
token_auth_result = await self._try_token_prefix_auth(token, request, db)
|
||||
if token_auth_result is not None:
|
||||
user, management_token = token_auth_result
|
||||
user, management_token = self._reattach_token_auth_result(db, token_auth_result)
|
||||
|
||||
# 存储到 request.state
|
||||
request.state.user_id = user.id
|
||||
request.state.management_token_id = management_token.id
|
||||
request.state.management_token_id = management_token.id if management_token else None
|
||||
|
||||
return user, management_token
|
||||
|
||||
@@ -437,15 +494,37 @@ class ApiRequestPipeline:
|
||||
detail="无效的 Token 格式,需要 Management Token",
|
||||
)
|
||||
|
||||
def _calculate_balance_remaining(
|
||||
self, db: Session, user: User | None, api_key: ApiKey | None = None
|
||||
async def _calculate_balance_remaining_async(
|
||||
self, user: User | None, api_key: ApiKey | None = None
|
||||
) -> float | None:
|
||||
if not user:
|
||||
return None
|
||||
balance = WalletService.get_balance_snapshot(db, user=user, api_key=api_key)
|
||||
if balance is None:
|
||||
return None
|
||||
return float(balance)
|
||||
|
||||
user_id = getattr(user, "id", None)
|
||||
api_key_id = getattr(api_key, "id", None)
|
||||
|
||||
# API Key 链路通常已在认证阶段预取余额;这里只保留为无预取路径的兜底查询。
|
||||
def _load_balance() -> float | None:
|
||||
thread_db = create_session()
|
||||
try:
|
||||
db_user = (
|
||||
thread_db.query(User).filter(User.id == user_id).first() if user_id else None
|
||||
)
|
||||
db_api_key = (
|
||||
thread_db.query(ApiKey).filter(ApiKey.id == api_key_id).first()
|
||||
if api_key_id
|
||||
else None
|
||||
)
|
||||
balance = WalletService.get_balance_snapshot(
|
||||
thread_db,
|
||||
user=db_user,
|
||||
api_key=db_api_key,
|
||||
)
|
||||
return float(balance) if balance is not None else None
|
||||
finally:
|
||||
thread_db.close()
|
||||
|
||||
return await run_in_threadpool(_load_balance)
|
||||
|
||||
def _record_audit_event(
|
||||
self,
|
||||
@@ -502,7 +581,7 @@ class ApiRequestPipeline:
|
||||
)
|
||||
except Exception as exc:
|
||||
# 审计失败不应影响主请求,仅记录警告
|
||||
logger.warning(f"[Audit] Failed to record event for adapter={adapter.name}: {exc}")
|
||||
logger.warning("[Audit] Failed to record event for adapter={}: {}", adapter.name, exc)
|
||||
|
||||
def _build_audit_metadata(
|
||||
self,
|
||||
@@ -568,7 +647,9 @@ class ApiRequestPipeline:
|
||||
if adapter_details:
|
||||
extra_details.update(adapter_details)
|
||||
except Exception as exc:
|
||||
logger.warning(f"[Audit] Adapter metadata failed: {adapter.__class__.__name__}: {exc}")
|
||||
logger.warning(
|
||||
"[Audit] Adapter metadata failed: {}: {}", adapter.__class__.__name__, exc
|
||||
)
|
||||
|
||||
if extra_details:
|
||||
metadata["details"] = extra_details
|
||||
@@ -578,7 +659,7 @@ class ApiRequestPipeline:
|
||||
|
||||
return self._sanitize_metadata(metadata)
|
||||
|
||||
def _sanitize_metadata(self, value: Any, depth: int = 0) -> None:
|
||||
def _sanitize_metadata(self, value: Any, depth: int = 0) -> Any:
|
||||
if value is None:
|
||||
return None
|
||||
if depth > 5:
|
||||
|
||||
@@ -56,6 +56,7 @@ class ChatAdapterBase(HandlerAdapterBase):
|
||||
# 适配器配置
|
||||
name: str = "chat.base"
|
||||
mode = ApiMode.STANDARD
|
||||
eager_request_body = False
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
"""处理 Chat API 请求"""
|
||||
@@ -71,7 +72,7 @@ class ChatAdapterBase(HandlerAdapterBase):
|
||||
original_headers = context.original_headers
|
||||
query_params = context.query_params
|
||||
|
||||
original_request_body = context.ensure_json_body()
|
||||
original_request_body = await context.ensure_json_body_async()
|
||||
|
||||
# 合并 path_params 到请求体(如 Gemini API 的 model 在 URL 路径中)
|
||||
if context.path_params:
|
||||
|
||||
@@ -20,7 +20,6 @@ from __future__ import annotations
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from src.api.base.adapter import ApiMode
|
||||
from src.api.base.context import ApiRequestContext
|
||||
@@ -58,6 +57,7 @@ class CliAdapterBase(HandlerAdapterBase):
|
||||
# 适配器配置
|
||||
name: str = "cli.base"
|
||||
mode = ApiMode.PROXY
|
||||
eager_request_body = False
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
"""处理 CLI API 请求"""
|
||||
@@ -82,7 +82,7 @@ class CliAdapterBase(HandlerAdapterBase):
|
||||
|
||||
set_original_request_headers(original_headers)
|
||||
|
||||
original_request_body = context.ensure_json_body()
|
||||
original_request_body = await context.ensure_json_body_async()
|
||||
|
||||
# 合并 path_params 到请求体(如 Gemini API 的 model 在 URL 路径中)
|
||||
if context.path_params:
|
||||
|
||||
@@ -35,6 +35,7 @@ class VideoAdapterBase(ApiAdapter):
|
||||
|
||||
name: str = "video.base"
|
||||
mode = ApiMode.STANDARD
|
||||
eager_request_body = False
|
||||
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
||||
@@ -103,7 +104,7 @@ class VideoAdapterBase(ApiAdapter):
|
||||
|
||||
# Remix task
|
||||
if method == "POST" and path.endswith("/remix") and task_id:
|
||||
original_request_body = context.ensure_json_body()
|
||||
original_request_body = await context.ensure_json_body_async()
|
||||
return await handler.handle_remix_task(
|
||||
task_id=task_id,
|
||||
http_request=http_request,
|
||||
@@ -134,7 +135,7 @@ class VideoAdapterBase(ApiAdapter):
|
||||
|
||||
# Create task (default)
|
||||
if method in {"POST", "PUT", "PATCH"}:
|
||||
original_request_body = context.ensure_json_body()
|
||||
original_request_body = await context.ensure_json_body_async()
|
||||
return await handler.handle_create_task(
|
||||
http_request=http_request,
|
||||
original_headers=context.original_headers,
|
||||
|
||||
@@ -184,6 +184,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
"thinking_enabled": bool(request_obj.thinking),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(
|
||||
cls,
|
||||
base_url: str,
|
||||
@@ -224,6 +225,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
|
||||
|
||||
name = "claude.token_count"
|
||||
mode = ApiMode.STANDARD
|
||||
eager_request_body = False
|
||||
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
|
||||
@@ -239,7 +241,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
|
||||
return bearer_handler.extract_credentials(request)
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
payload = context.ensure_json_body()
|
||||
payload = await context.ensure_json_body_async()
|
||||
|
||||
try:
|
||||
request = ClaudeTokenCountRequest.model_validate(payload, strict=False)
|
||||
|
||||
@@ -48,7 +48,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
async def handle(self, context: ApiRequestContext) -> Any:
|
||||
"""处理 CLI API 请求 -- compact 模式下注入标记并强制非流式"""
|
||||
if self._compact:
|
||||
body = context.ensure_json_body()
|
||||
body = await context.ensure_json_body_async()
|
||||
body["_aether_compact"] = True
|
||||
# compact 端点永远非流式
|
||||
body.pop("stream", None)
|
||||
|
||||
@@ -231,7 +231,6 @@ async def get_my_usage(
|
||||
- `total_tokens`: 总 Token 数
|
||||
- `total_cost`: 总成本(USD)
|
||||
- `summary_by_model`: 按模型分组统计
|
||||
- `summary_by_provider`: 按提供商分组统计
|
||||
- `records`: 详细使用记录列表
|
||||
- `pagination`: 分页信息
|
||||
"""
|
||||
@@ -809,6 +808,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
user_id=user.id,
|
||||
start_date=start_utc,
|
||||
end_date=end_utc,
|
||||
group_by=None,
|
||||
)
|
||||
|
||||
# 过滤掉 unknown/pending provider 的记录(请求未到达任何提供商)
|
||||
@@ -818,31 +818,23 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
if item.get("provider") not in ("unknown", "pending", None)
|
||||
]
|
||||
|
||||
total_requests = sum(item["requests"] for item in filtered_summary)
|
||||
total_input_tokens = (
|
||||
sum(item["input_tokens"] for item in filtered_summary) if filtered_summary else 0
|
||||
)
|
||||
total_output_tokens = (
|
||||
sum(item["output_tokens"] for item in filtered_summary) if filtered_summary else 0
|
||||
)
|
||||
total_tokens = (
|
||||
sum(item["total_tokens"] for item in filtered_summary) if filtered_summary else 0
|
||||
)
|
||||
total_cost = (
|
||||
sum(item["total_cost_usd"] for item in filtered_summary) if filtered_summary else 0.0
|
||||
)
|
||||
|
||||
# 管理员可以看到真实成本
|
||||
total_requests = 0
|
||||
total_input_tokens = 0
|
||||
total_output_tokens = 0
|
||||
total_tokens = 0
|
||||
total_cost = 0.0
|
||||
total_actual_cost = 0.0
|
||||
if user.role == UserRole.ADMIN:
|
||||
total_actual_cost = (
|
||||
sum(item.get("actual_total_cost_usd", 0.0) for item in filtered_summary)
|
||||
if filtered_summary
|
||||
else 0.0
|
||||
)
|
||||
|
||||
model_summary = {}
|
||||
provider_summary = {}
|
||||
for item in filtered_summary:
|
||||
total_requests += item["requests"]
|
||||
total_input_tokens += item["input_tokens"]
|
||||
total_output_tokens += item["output_tokens"]
|
||||
total_tokens += item["total_tokens"]
|
||||
total_cost += item["total_cost_usd"]
|
||||
if user.role == UserRole.ADMIN:
|
||||
total_actual_cost += item.get("actual_total_cost_usd", 0.0)
|
||||
|
||||
model_name = item["model"]
|
||||
base_stats = {
|
||||
"model": model_name,
|
||||
@@ -866,13 +858,8 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
if user.role == UserRole.ADMIN:
|
||||
stats["actual_total_cost_usd"] += item.get("actual_total_cost_usd", 0.0)
|
||||
|
||||
summary_by_model = sorted(model_summary.values(), key=lambda x: x["requests"], reverse=True)
|
||||
|
||||
# 按提供商汇总(用于 UsageProviderTable)
|
||||
provider_summary = {}
|
||||
for item in filtered_summary:
|
||||
provider_name = item["provider"]
|
||||
base_stats = {
|
||||
provider_base_stats = {
|
||||
"provider": provider_name,
|
||||
"requests": 0,
|
||||
"total_tokens": 0,
|
||||
@@ -881,32 +868,37 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
"total_response_time_ms": 0.0,
|
||||
"response_time_count": 0,
|
||||
}
|
||||
stats = provider_summary.setdefault(provider_name, base_stats)
|
||||
stats["requests"] += item["requests"]
|
||||
stats["total_tokens"] += item["total_tokens"]
|
||||
stats["total_cost_usd"] += item["total_cost_usd"]
|
||||
# 假设 summary 中的都是成功的请求
|
||||
stats["success_count"] += item["requests"]
|
||||
if item.get("avg_response_time_ms") is not None:
|
||||
stats["total_response_time_ms"] += item["avg_response_time_ms"] * item["requests"]
|
||||
stats["response_time_count"] += item["requests"]
|
||||
provider_stats = provider_summary.setdefault(provider_name, provider_base_stats)
|
||||
provider_stats["requests"] += item["requests"]
|
||||
provider_stats["total_tokens"] += item["total_tokens"]
|
||||
provider_stats["total_cost_usd"] += item["total_cost_usd"]
|
||||
provider_stats["success_count"] += int(item.get("success_count", 0) or 0)
|
||||
success_response_time_count = int(item.get("success_response_time_count", 0) or 0)
|
||||
if success_response_time_count > 0:
|
||||
provider_stats["total_response_time_ms"] += float(
|
||||
item.get("success_response_time_sum_ms", 0.0) or 0.0
|
||||
)
|
||||
provider_stats["response_time_count"] += success_response_time_count
|
||||
|
||||
summary_by_model = sorted(model_summary.values(), key=lambda x: x["requests"], reverse=True)
|
||||
summary_by_provider = []
|
||||
for stats in provider_summary.values():
|
||||
for provider_stats in provider_summary.values():
|
||||
avg_response_time_ms = (
|
||||
stats["total_response_time_ms"] / stats["response_time_count"]
|
||||
if stats["response_time_count"] > 0
|
||||
provider_stats["total_response_time_ms"] / provider_stats["response_time_count"]
|
||||
if provider_stats["response_time_count"] > 0
|
||||
else 0
|
||||
)
|
||||
success_rate = (
|
||||
(stats["success_count"] / stats["requests"] * 100) if stats["requests"] > 0 else 100
|
||||
(provider_stats["success_count"] / provider_stats["requests"] * 100)
|
||||
if provider_stats["requests"] > 0
|
||||
else 100
|
||||
)
|
||||
summary_by_provider.append(
|
||||
{
|
||||
"provider": stats["provider"],
|
||||
"requests": stats["requests"],
|
||||
"total_tokens": stats["total_tokens"],
|
||||
"total_cost_usd": stats["total_cost_usd"],
|
||||
"provider": provider_stats["provider"],
|
||||
"requests": provider_stats["requests"],
|
||||
"total_tokens": provider_stats["total_tokens"],
|
||||
"total_cost_usd": provider_stats["total_cost_usd"],
|
||||
"success_rate": round(success_rate, 2),
|
||||
"avg_response_time_ms": round(avg_response_time_ms, 2),
|
||||
}
|
||||
@@ -1003,6 +995,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
"avg_response_time": avg_response_time,
|
||||
"billing": WalletService.serialize_wallet_summary(wallet),
|
||||
"summary_by_model": summary_by_model,
|
||||
"summary_by_provider": summary_by_provider,
|
||||
# 分页信息
|
||||
"pagination": {
|
||||
"total": total_records,
|
||||
@@ -1117,7 +1110,12 @@ class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
|
||||
if not id_list:
|
||||
return {"requests": []}
|
||||
|
||||
requests = UsageService.get_active_requests_status(db=db, ids=id_list, user_id=user.id)
|
||||
requests = UsageService.get_active_requests_status(
|
||||
db=db,
|
||||
ids=id_list,
|
||||
user_id=user.id,
|
||||
maintain_status=True,
|
||||
)
|
||||
return {"requests": requests}
|
||||
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ import secrets
|
||||
import time
|
||||
import uuid
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from threading import Lock
|
||||
from typing import TYPE_CHECKING, Any
|
||||
@@ -23,6 +24,7 @@ from src.config import config
|
||||
from src.core.enums import AuthSource
|
||||
from src.core.exceptions import ForbiddenException
|
||||
from src.core.logger import logger
|
||||
from src.database.database import create_session
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -32,6 +34,31 @@ from src.models.database import ApiKey, User, UserRole
|
||||
from src.services.auth.jwt_blacklist import JWTBlacklistService
|
||||
from src.services.cache.user_cache import UserCacheService
|
||||
|
||||
|
||||
@dataclass
|
||||
class AuthenticatedUserSnapshot:
|
||||
user_id: str
|
||||
email: str | None
|
||||
username: str
|
||||
role: UserRole
|
||||
created_at: datetime | None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ThreadsafeAPIKeyAuthResult:
|
||||
user: User
|
||||
api_key: ApiKey | None = None
|
||||
balance_remaining: float | None = None
|
||||
access_allowed: bool = True
|
||||
access_message: str = "OK"
|
||||
|
||||
@property
|
||||
def access_ok(self) -> bool:
|
||||
return self.access_allowed
|
||||
|
||||
|
||||
PipelineThreadsafeAuthResult = ThreadsafeAPIKeyAuthResult
|
||||
|
||||
# API Key last_used_at 更新节流配置
|
||||
# 同一个 API Key 在此时间间隔内只会更新一次 last_used_at
|
||||
_LAST_USED_UPDATE_INTERVAL = 60 # 秒
|
||||
@@ -177,6 +204,192 @@ class AuthService:
|
||||
except jwt.InvalidTokenError:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="无效的Token")
|
||||
|
||||
@staticmethod
|
||||
def _authenticate_local_user_sync(
|
||||
db: Session,
|
||||
email: str,
|
||||
password: str,
|
||||
) -> User | None:
|
||||
"""同步执行本地认证,供线程池隔离入口复用。"""
|
||||
from sqlalchemy import or_
|
||||
|
||||
user = db.query(User).filter(or_(User.email == email, User.username == email)).first()
|
||||
|
||||
if not user:
|
||||
logger.warning("登录失败 - 用户不存在: {}", email)
|
||||
return None
|
||||
|
||||
if user.is_deleted:
|
||||
logger.warning("登录失败 - 用户已删除: {}", email)
|
||||
return None
|
||||
|
||||
from src.core.modules.hooks import AUTH_CHECK_EXCLUSIVE_MODE, get_hook_dispatcher
|
||||
|
||||
is_exclusive = get_hook_dispatcher().dispatch_sync(AUTH_CHECK_EXCLUSIVE_MODE, db=db)
|
||||
if is_exclusive:
|
||||
if user.role != UserRole.ADMIN or user.auth_source != AuthSource.LOCAL:
|
||||
logger.warning("登录失败 - 排他登录模式下仅管理员可本地登录: {}", email)
|
||||
return None
|
||||
logger.warning("[EXCLUSIVE-MODE] 紧急恢复通道:本地管理员登录: {}", email)
|
||||
|
||||
if user.auth_source == AuthSource.LDAP:
|
||||
logger.warning("登录失败 - 该用户使用 LDAP 认证: {}", email)
|
||||
return None
|
||||
|
||||
if not user.verify_password(password):
|
||||
logger.warning("登录失败 - 密码错误: {}", email)
|
||||
return None
|
||||
|
||||
if not user.is_active:
|
||||
logger.warning("登录失败 - 用户已禁用: {}", email)
|
||||
return None
|
||||
|
||||
return user
|
||||
|
||||
@staticmethod
|
||||
def _build_authenticated_snapshot(user: User) -> AuthenticatedUserSnapshot:
|
||||
return AuthenticatedUserSnapshot(
|
||||
user_id=user.id,
|
||||
email=user.email,
|
||||
username=user.username,
|
||||
role=user.role,
|
||||
created_at=user.created_at,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _detach_instance(db: Session, instance: User | ApiKey | None) -> None:
|
||||
if instance is None:
|
||||
return
|
||||
try:
|
||||
db.expunge(instance)
|
||||
except Exception as exc:
|
||||
logger.debug("expunge failed: {}", exc)
|
||||
|
||||
@staticmethod
|
||||
def _load_user_for_token_sync(db: Session, user_id: str) -> User | None:
|
||||
user = db.query(User).filter(User.id == user_id).first()
|
||||
if not user or not user.is_active or user.is_deleted:
|
||||
return None
|
||||
return user
|
||||
|
||||
@staticmethod
|
||||
async def load_user_for_token_threadsafe(user_id: str) -> User | None:
|
||||
"""Load the JWT user in a threadpool and return a detached object."""
|
||||
|
||||
def _load_in_thread() -> User | None:
|
||||
thread_db = create_session()
|
||||
try:
|
||||
user = AuthService._load_user_for_token_sync(thread_db, user_id)
|
||||
if not user:
|
||||
return None
|
||||
|
||||
AuthService._detach_instance(thread_db, user)
|
||||
return user
|
||||
finally:
|
||||
thread_db.close()
|
||||
|
||||
return await run_in_threadpool(_load_in_thread)
|
||||
|
||||
@staticmethod
|
||||
async def load_user_for_pipeline_threadsafe(
|
||||
user_id: str,
|
||||
*,
|
||||
include_balance: bool = False,
|
||||
) -> PipelineThreadsafeAuthResult | None:
|
||||
"""Compatibility helper: load a user in a threadpool and optionally prefetch balance."""
|
||||
|
||||
def _load_in_thread() -> PipelineThreadsafeAuthResult | None:
|
||||
from src.services.wallet import WalletService
|
||||
|
||||
thread_db = create_session()
|
||||
try:
|
||||
user = AuthService._load_user_for_token_sync(thread_db, user_id)
|
||||
if not user:
|
||||
return None
|
||||
|
||||
balance_remaining: float | None = None
|
||||
if include_balance:
|
||||
balance = WalletService.get_balance_snapshot(thread_db, user=user)
|
||||
balance_remaining = float(balance) if balance is not None else None
|
||||
|
||||
AuthService._detach_instance(thread_db, user)
|
||||
return PipelineThreadsafeAuthResult(
|
||||
user=user,
|
||||
balance_remaining=balance_remaining,
|
||||
)
|
||||
finally:
|
||||
thread_db.close()
|
||||
|
||||
return await run_in_threadpool(_load_in_thread)
|
||||
|
||||
@staticmethod
|
||||
async def authenticate_api_key_threadsafe(
|
||||
api_key: str,
|
||||
) -> ThreadsafeAPIKeyAuthResult | None:
|
||||
"""Authenticate API key and check balance in a threadpool."""
|
||||
|
||||
def _authenticate_in_thread() -> ThreadsafeAPIKeyAuthResult | None:
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
thread_db = create_session()
|
||||
try:
|
||||
auth_result = AuthService.authenticate_api_key(thread_db, api_key)
|
||||
if not auth_result:
|
||||
return None
|
||||
|
||||
user, key_record = auth_result
|
||||
balance_result = UsageService.check_request_balance_details(
|
||||
thread_db,
|
||||
user,
|
||||
api_key=key_record,
|
||||
)
|
||||
|
||||
AuthService._detach_instance(thread_db, user)
|
||||
AuthService._detach_instance(thread_db, key_record)
|
||||
return ThreadsafeAPIKeyAuthResult(
|
||||
user=user,
|
||||
api_key=key_record,
|
||||
balance_remaining=balance_result.remaining,
|
||||
access_allowed=balance_result.allowed,
|
||||
access_message=balance_result.message,
|
||||
)
|
||||
finally:
|
||||
thread_db.close()
|
||||
|
||||
return await run_in_threadpool(_authenticate_in_thread)
|
||||
|
||||
@staticmethod
|
||||
async def authenticate_user_threadsafe(
|
||||
db: Session, email: str, password: str, auth_type: str = "local"
|
||||
) -> AuthenticatedUserSnapshot | None:
|
||||
"""为异步登录路由提供线程池隔离的本地认证入口。"""
|
||||
if auth_type != "local":
|
||||
user = await AuthService.authenticate_user(db, email, password, auth_type)
|
||||
if not user:
|
||||
return None
|
||||
return AuthService._build_authenticated_snapshot(user)
|
||||
|
||||
def _authenticate_in_thread() -> AuthenticatedUserSnapshot | None:
|
||||
thread_db = create_session()
|
||||
try:
|
||||
user = AuthService._authenticate_local_user_sync(thread_db, email, password)
|
||||
if not user:
|
||||
return None
|
||||
|
||||
user.last_login_at = datetime.now(timezone.utc)
|
||||
thread_db.commit()
|
||||
return AuthService._build_authenticated_snapshot(user)
|
||||
finally:
|
||||
thread_db.close()
|
||||
|
||||
snapshot = await run_in_threadpool(_authenticate_in_thread)
|
||||
if not snapshot:
|
||||
return None
|
||||
|
||||
await UserCacheService.invalidate_user_cache(snapshot.user_id, snapshot.email or "")
|
||||
logger.info("用户登录成功: {} (ID: {})", email, snapshot.user_id)
|
||||
return snapshot
|
||||
|
||||
@staticmethod
|
||||
async def authenticate_user(
|
||||
db: Session, email: str, password: str, auth_type: str = "local"
|
||||
@@ -208,40 +421,8 @@ class AuthService:
|
||||
# 本地认证
|
||||
# 登录校验必须读取密码哈希,不能使用不包含 password_hash 的缓存对象
|
||||
# 支持邮箱或用户名登录
|
||||
from sqlalchemy import or_
|
||||
|
||||
user = db.query(User).filter(or_(User.email == email, User.username == email)).first()
|
||||
|
||||
user = AuthService._authenticate_local_user_sync(db, email, password)
|
||||
if not user:
|
||||
logger.warning(f"登录失败 - 用户不存在: {email}")
|
||||
return None
|
||||
|
||||
if user.is_deleted:
|
||||
logger.warning(f"登录失败 - 用户已删除: {email}")
|
||||
return None
|
||||
|
||||
# 检查排他登录模式(如 LDAP exclusive):仅允许本地管理员登录(紧急恢复通道)
|
||||
from src.core.modules.hooks import AUTH_CHECK_EXCLUSIVE_MODE, get_hook_dispatcher
|
||||
|
||||
is_exclusive = get_hook_dispatcher().dispatch_sync(AUTH_CHECK_EXCLUSIVE_MODE, db=db)
|
||||
if is_exclusive:
|
||||
if user.role != UserRole.ADMIN or user.auth_source != AuthSource.LOCAL:
|
||||
logger.warning(f"登录失败 - 排他登录模式下仅管理员可本地登录: {email}")
|
||||
return None
|
||||
logger.warning(f"[EXCLUSIVE-MODE] 紧急恢复通道:本地管理员登录: {email}")
|
||||
|
||||
# 检查用户认证来源
|
||||
if user.auth_source == AuthSource.LDAP:
|
||||
logger.warning(f"登录失败 - 该用户使用 LDAP 认证: {email}")
|
||||
return None
|
||||
|
||||
# 在线程池中执行 bcrypt 密码验证,避免阻塞事件循环
|
||||
if not await run_in_threadpool(user.verify_password, password):
|
||||
logger.warning(f"登录失败 - 密码错误: {email}")
|
||||
return None
|
||||
|
||||
if not user.is_active:
|
||||
logger.warning(f"登录失败 - 用户已禁用: {email}")
|
||||
return None
|
||||
|
||||
# 更新最后登录时间
|
||||
|
||||
171
src/services/cache/model_cache.py
vendored
171
src/services/cache/model_cache.py
vendored
@@ -42,6 +42,8 @@ class ModelCacheService:
|
||||
|
||||
# 缓存 TTL(秒)- 使用统一常量
|
||||
CACHE_TTL = CacheTTL.MODEL
|
||||
PROVIDER_MAPPING_INDEX_CACHE_KEY = "global_model:resolve_index:provider_model_mappings"
|
||||
MODEL_MAPPING_RULES_CACHE_KEY = "global_model:resolve_index:model_mappings"
|
||||
|
||||
@staticmethod
|
||||
async def get_model_by_id(db: Session, model_id: str) -> Model | None:
|
||||
@@ -238,6 +240,9 @@ class ModelCacheService:
|
||||
if resolve_keys_to_clear:
|
||||
logger.debug(f"Model resolve 缓存已清除: {resolve_keys_to_clear}")
|
||||
|
||||
# provider_model_mappings 更新后,需要重建映射索引缓存。
|
||||
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
|
||||
|
||||
@staticmethod
|
||||
async def invalidate_global_model_cache(global_model_id: str, name: str | None = None) -> None:
|
||||
"""清除 GlobalModel 缓存"""
|
||||
@@ -249,6 +254,8 @@ class ModelCacheService:
|
||||
# 全量清除 resolve 缓存,确保映射规则变更后不命中旧缓存
|
||||
try:
|
||||
await CacheService.delete_pattern("global_model:resolve:*")
|
||||
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
|
||||
await CacheService.delete(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
|
||||
except Exception as e:
|
||||
logger.error(f"GlobalModel resolve 缓存清除失败,可能导致映射不一致: {e}")
|
||||
logger.debug(f"GlobalModel 缓存已清除: {global_model_id}")
|
||||
@@ -262,10 +269,108 @@ class ModelCacheService:
|
||||
"""
|
||||
try:
|
||||
deleted = await CacheService.delete_pattern("global_model:resolve:*")
|
||||
await CacheService.delete(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
|
||||
await CacheService.delete(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
|
||||
logger.debug(f"已清除 {deleted} 个 GlobalModel resolve 缓存")
|
||||
except Exception as e:
|
||||
logger.error(f"GlobalModel resolve 缓存清除失败: {e}")
|
||||
|
||||
@staticmethod
|
||||
async def _get_provider_mapping_index(
|
||||
db: Session,
|
||||
) -> dict[str, list[dict[str, object]]]:
|
||||
cached_data = await CacheService.get(ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY)
|
||||
if isinstance(cached_data, dict):
|
||||
return {
|
||||
str(name): value
|
||||
for name, value in cached_data.items()
|
||||
if isinstance(name, str) and isinstance(value, list)
|
||||
}
|
||||
|
||||
from src.models.database import Provider
|
||||
|
||||
rows = (
|
||||
db.query(Model, GlobalModel)
|
||||
.join(Provider, Model.provider_id == Provider.id)
|
||||
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
|
||||
.filter(
|
||||
Provider.is_active == True,
|
||||
Model.is_active == True,
|
||||
GlobalModel.is_active == True,
|
||||
Model.provider_model_mappings.isnot(None),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
index: dict[str, list[dict[str, object]]] = {}
|
||||
seen_pairs: set[tuple[str, str]] = set()
|
||||
for model, global_model in rows:
|
||||
raw_mappings = getattr(model, "provider_model_mappings", None)
|
||||
if not isinstance(raw_mappings, list):
|
||||
continue
|
||||
|
||||
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
|
||||
for raw in raw_mappings:
|
||||
if not isinstance(raw, dict):
|
||||
continue
|
||||
name = raw.get("name")
|
||||
if not isinstance(name, str):
|
||||
continue
|
||||
normalized_name = name.strip()
|
||||
if not normalized_name:
|
||||
continue
|
||||
pair_key = (normalized_name, str(global_model.id))
|
||||
if pair_key in seen_pairs:
|
||||
continue
|
||||
seen_pairs.add(pair_key)
|
||||
index.setdefault(normalized_name, []).append(global_model_dict)
|
||||
|
||||
await CacheService.set(
|
||||
ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY,
|
||||
index,
|
||||
ttl_seconds=ModelCacheService.CACHE_TTL,
|
||||
)
|
||||
return index
|
||||
|
||||
@staticmethod
|
||||
async def _get_model_mapping_rules(
|
||||
db: Session,
|
||||
) -> list[dict[str, object]]:
|
||||
cached_data = await CacheService.get(ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY)
|
||||
if isinstance(cached_data, list):
|
||||
return [entry for entry in cached_data if isinstance(entry, dict)]
|
||||
|
||||
rows = (
|
||||
db.query(GlobalModel)
|
||||
.filter(GlobalModel.is_active == True, GlobalModel.config.isnot(None))
|
||||
.all()
|
||||
)
|
||||
|
||||
rules: list[dict[str, object]] = []
|
||||
for global_model in rows:
|
||||
config = getattr(global_model, "config", None) or {}
|
||||
mappings = config.get("model_mappings")
|
||||
if not isinstance(mappings, list) or not mappings:
|
||||
continue
|
||||
patterns = [
|
||||
pattern for pattern in mappings if isinstance(pattern, str) and pattern.strip()
|
||||
]
|
||||
if not patterns:
|
||||
continue
|
||||
rules.append(
|
||||
{
|
||||
"global_model": ModelCacheService._global_model_to_dict(global_model),
|
||||
"patterns": patterns,
|
||||
}
|
||||
)
|
||||
|
||||
await CacheService.set(
|
||||
ModelCacheService.MODEL_MAPPING_RULES_CACHE_KEY,
|
||||
rules,
|
||||
ttl_seconds=ModelCacheService.CACHE_TTL,
|
||||
)
|
||||
return rules
|
||||
|
||||
@staticmethod
|
||||
async def resolve_global_model_by_name_or_mapping(
|
||||
db: Session, model_name: str
|
||||
@@ -391,41 +496,18 @@ class ModelCacheService:
|
||||
return result_global_model
|
||||
|
||||
# 4. 通过 provider_model_mappings 匹配
|
||||
models_with_mappings = (
|
||||
db.query(Model, GlobalModel)
|
||||
.join(Provider, Model.provider_id == Provider.id)
|
||||
.join(GlobalModel, Model.global_model_id == GlobalModel.id)
|
||||
.filter(
|
||||
Provider.is_active == True,
|
||||
Model.is_active == True,
|
||||
GlobalModel.is_active == True,
|
||||
Model.provider_model_mappings.isnot(None),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
provider_mapping_index = await ModelCacheService._get_provider_mapping_index(db)
|
||||
mapping_matched_global_models = [
|
||||
ModelCacheService._dict_to_global_model(global_model_dict)
|
||||
for global_model_dict in provider_mapping_index.get(normalized_name, [])
|
||||
if isinstance(global_model_dict, dict)
|
||||
]
|
||||
|
||||
mapping_matched_global_models: list[GlobalModel] = []
|
||||
mapping_seen_ids: set[str] = set()
|
||||
for model, gm in models_with_mappings:
|
||||
raw_mappings = model.provider_model_mappings
|
||||
if not isinstance(raw_mappings, list):
|
||||
continue
|
||||
for raw in raw_mappings:
|
||||
if not isinstance(raw, dict):
|
||||
continue
|
||||
name = raw.get("name")
|
||||
if not isinstance(name, str):
|
||||
continue
|
||||
if name.strip() != normalized_name:
|
||||
continue
|
||||
if gm.id not in mapping_seen_ids:
|
||||
mapping_seen_ids.add(gm.id)
|
||||
mapping_matched_global_models.append(gm)
|
||||
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:
|
||||
|
||||
@@ -15,6 +15,14 @@ from src.models.database import RequestCandidate
|
||||
class RequestCandidateService:
|
||||
"""请求候选记录服务"""
|
||||
|
||||
@staticmethod
|
||||
def _persist_candidate_update(db: Session, *, immediate: bool) -> None:
|
||||
if immediate:
|
||||
db.commit()
|
||||
return
|
||||
db.flush()
|
||||
get_batch_committer().mark_dirty(db)
|
||||
|
||||
@staticmethod
|
||||
def create_candidate(
|
||||
db: Session,
|
||||
@@ -93,9 +101,9 @@ class RequestCandidateService:
|
||||
if candidate:
|
||||
candidate.status = "pending"
|
||||
candidate.started_at = datetime.now(timezone.utc)
|
||||
# 关键状态更新:立即提交,不使用批量提交
|
||||
# 原因:前端需要实时看到请求开始执行
|
||||
db.commit()
|
||||
# 中间态改为 flush:最终 success/failed 仍会立即提交,
|
||||
# 但开始执行这一跳不再单独制造一次事务往返。
|
||||
RequestCandidateService._persist_candidate_update(db, immediate=False)
|
||||
|
||||
@staticmethod
|
||||
def update_candidate_status(db: Session, candidate_id: str, status: str) -> None:
|
||||
@@ -113,8 +121,9 @@ class RequestCandidateService:
|
||||
# 如果状态变更为 pending,记录开始时间
|
||||
if status == "pending" and not candidate.started_at:
|
||||
candidate.started_at = datetime.now(timezone.utc)
|
||||
# 立即提交,确保前端能实时看到状态变化
|
||||
db.commit()
|
||||
RequestCandidateService._persist_candidate_update(
|
||||
db, immediate=status not in {"pending", "streaming"}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def mark_candidate_streaming(
|
||||
@@ -141,7 +150,7 @@ class RequestCandidateService:
|
||||
candidate.status = "streaming"
|
||||
candidate.concurrent_requests = concurrent_requests
|
||||
# streaming 状态不设置 finished_at 和 status_code,因为请求还在进行中
|
||||
db.commit()
|
||||
RequestCandidateService._persist_candidate_update(db, immediate=False)
|
||||
|
||||
@staticmethod
|
||||
def mark_candidate_success(
|
||||
|
||||
@@ -62,7 +62,6 @@ from src.services.rate_limit.adaptive_reservation import (
|
||||
)
|
||||
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
|
||||
from src.services.scheduling.affinity_manager import (
|
||||
CacheAffinityManager,
|
||||
get_affinity_manager,
|
||||
)
|
||||
from src.services.scheduling.candidate_builder import (
|
||||
@@ -464,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
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ import re
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from sqlalchemy import or_
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from src.core.api_format.conversion.compatibility import is_format_compatible
|
||||
@@ -52,7 +53,7 @@ from src.services.cache.model_cache import ModelCacheService
|
||||
|
||||
def _get_pool_config(provider: Provider) -> "PoolConfig | None":
|
||||
"""Return parsed PoolConfig if the provider has pool enabled, else None."""
|
||||
from src.services.provider.pool.config import PoolConfig, parse_pool_config
|
||||
from src.services.provider.pool.config import parse_pool_config
|
||||
|
||||
return parse_pool_config(getattr(provider, "config", None))
|
||||
|
||||
@@ -79,11 +80,36 @@ class CandidateBuilder:
|
||||
def __init__(self, candidate_sorter: CandidateSorterProtocol) -> None:
|
||||
self._sorter = candidate_sorter
|
||||
|
||||
def _query_provider_refs(
|
||||
self,
|
||||
db: Session,
|
||||
provider_offset: int = 0,
|
||||
provider_limit: int | None = None,
|
||||
) -> list[tuple[str, str]]:
|
||||
"""仅查询当前分页内 Provider 的轻量引用信息。"""
|
||||
provider_query = (
|
||||
db.query(Provider.id, Provider.name)
|
||||
.filter(Provider.is_active.is_(True))
|
||||
.order_by(Provider.provider_priority.asc())
|
||||
)
|
||||
|
||||
if provider_offset:
|
||||
provider_query = provider_query.offset(provider_offset)
|
||||
if provider_limit:
|
||||
provider_query = provider_query.limit(provider_limit)
|
||||
|
||||
return [
|
||||
(str(provider_id), str(provider_name))
|
||||
for provider_id, provider_name in provider_query.all()
|
||||
]
|
||||
|
||||
def _query_providers(
|
||||
self,
|
||||
db: Session,
|
||||
provider_offset: int = 0,
|
||||
provider_limit: int | None = None,
|
||||
allowed_providers: list[str] | None = None,
|
||||
provider_ids: list[str] | None = None,
|
||||
) -> list[Provider]:
|
||||
"""
|
||||
查询活跃的 Providers(带预加载)
|
||||
@@ -132,12 +158,30 @@ class CandidateBuilder:
|
||||
.order_by(Provider.provider_priority.asc())
|
||||
)
|
||||
|
||||
if provider_offset:
|
||||
if allowed_providers:
|
||||
allowed_values = [value for value in allowed_providers if value]
|
||||
if allowed_values:
|
||||
provider_query = provider_query.filter(
|
||||
or_(Provider.id.in_(allowed_values), Provider.name.in_(allowed_values))
|
||||
)
|
||||
|
||||
if provider_ids is not None:
|
||||
if not provider_ids:
|
||||
return []
|
||||
provider_query = provider_query.filter(Provider.id.in_(provider_ids))
|
||||
|
||||
if provider_ids is None and provider_offset:
|
||||
provider_query = provider_query.offset(provider_offset)
|
||||
if provider_limit:
|
||||
if provider_ids is None and provider_limit:
|
||||
provider_query = provider_query.limit(provider_limit)
|
||||
|
||||
return provider_query.all()
|
||||
providers = provider_query.all()
|
||||
if provider_ids is None:
|
||||
return providers
|
||||
|
||||
order_map = {provider_id: index for index, provider_id in enumerate(provider_ids)}
|
||||
providers.sort(key=lambda provider: order_map.get(str(provider.id), len(order_map)))
|
||||
return providers
|
||||
|
||||
async def _check_model_support(
|
||||
self,
|
||||
|
||||
@@ -42,11 +42,20 @@ class CandidateSorterProtocol(Protocol):
|
||||
|
||||
|
||||
class CandidateBuilderProtocol(Protocol):
|
||||
def _query_provider_refs(
|
||||
self,
|
||||
db: Session,
|
||||
provider_offset: int = 0,
|
||||
provider_limit: int | None = None,
|
||||
) -> list[tuple[str, str]]: ...
|
||||
|
||||
def _query_providers(
|
||||
self,
|
||||
db: Session,
|
||||
provider_offset: int = 0,
|
||||
provider_limit: int | None = None,
|
||||
allowed_providers: list[str] | None = None,
|
||||
provider_ids: list[str] | None = None,
|
||||
) -> list[Provider]: ...
|
||||
|
||||
async def _build_candidates(
|
||||
|
||||
@@ -12,7 +12,7 @@ from datetime import date, datetime, time, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import Float, and_, case, cast, func
|
||||
from sqlalchemy import Date, Float, and_, case, cast, func, text
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -47,9 +47,64 @@ def _get_utc_day_range(value: datetime) -> tuple[datetime, datetime]:
|
||||
return day_start, day_start + timedelta(days=1)
|
||||
|
||||
|
||||
def _merge_consecutive_utc_days(days: list[date]) -> list[tuple[datetime, datetime]]:
|
||||
"""将连续 UTC 日期合并为更少的 [start, end) 区间。"""
|
||||
if not days:
|
||||
return []
|
||||
|
||||
sorted_days = sorted(days)
|
||||
ranges: list[tuple[datetime, datetime]] = []
|
||||
start_day = sorted_days[0]
|
||||
end_day = start_day
|
||||
|
||||
for current_day in sorted_days[1:]:
|
||||
if current_day == end_day + timedelta(days=1):
|
||||
end_day = current_day
|
||||
continue
|
||||
|
||||
range_start = datetime.combine(start_day, time.min, tzinfo=timezone.utc)
|
||||
range_end = datetime.combine(end_day + timedelta(days=1), time.min, tzinfo=timezone.utc)
|
||||
ranges.append((range_start, range_end))
|
||||
start_day = current_day
|
||||
end_day = current_day
|
||||
|
||||
range_start = datetime.combine(start_day, time.min, tzinfo=timezone.utc)
|
||||
range_end = datetime.combine(end_day + timedelta(days=1), time.min, tzinfo=timezone.utc)
|
||||
ranges.append((range_start, range_end))
|
||||
return ranges
|
||||
|
||||
|
||||
class StatsAggregatorService:
|
||||
"""统计数据聚合服务"""
|
||||
|
||||
@staticmethod
|
||||
def _resolve_percentile_row(row: Any | None) -> tuple[int | None, int | None, int | None]:
|
||||
if not row:
|
||||
return None, None, None
|
||||
count = int(getattr(row, "count", 0) or 0)
|
||||
if count < MIN_PERCENTILE_SAMPLES:
|
||||
return None, None, None
|
||||
p50 = getattr(row, "p50", None)
|
||||
p90 = getattr(row, "p90", None)
|
||||
p99 = getattr(row, "p99", None)
|
||||
return (
|
||||
int(p50) if p50 is not None else None,
|
||||
int(p90) if p90 is not None else None,
|
||||
int(p99) if p99 is not None else None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_local_day_expression(time_range: TimeRangeParams) -> Any:
|
||||
if time_range.timezone and time_range.timezone != "UTC":
|
||||
local_time_expr = func.timezone(time_range.timezone, Usage.created_at)
|
||||
elif time_range.tz_offset_minutes:
|
||||
interval_literal = text(f"INTERVAL '{int(time_range.tz_offset_minutes)} minutes'")
|
||||
local_time_expr = Usage.created_at + interval_literal
|
||||
else:
|
||||
local_time_expr = Usage.created_at
|
||||
|
||||
return cast(func.date_trunc("day", local_time_expr), Date)
|
||||
|
||||
@staticmethod
|
||||
def compute_daily_stats(db: Session, date: datetime) -> dict:
|
||||
"""计算指定 UTC 日期的统计数据(不写入数据库)"""
|
||||
@@ -196,23 +251,8 @@ class StatsAggregatorService:
|
||||
.first()
|
||||
)
|
||||
|
||||
def _resolve(row: Any | None) -> tuple[int | None, int | None, int | None]:
|
||||
if not row:
|
||||
return None, None, None
|
||||
count = int(getattr(row, "count", 0) or 0)
|
||||
if count < MIN_PERCENTILE_SAMPLES:
|
||||
return None, None, None
|
||||
p50 = getattr(row, "p50", None)
|
||||
p90 = getattr(row, "p90", None)
|
||||
p99 = getattr(row, "p99", None)
|
||||
return (
|
||||
int(p50) if p50 is not None else None,
|
||||
int(p90) if p90 is not None else None,
|
||||
int(p99) if p99 is not None else None,
|
||||
)
|
||||
|
||||
p50_rt, p90_rt, p99_rt = _resolve(rt_row)
|
||||
p50_ttfb, p90_ttfb, p99_ttfb = _resolve(ttfb_row)
|
||||
p50_rt, p90_rt, p99_rt = StatsAggregatorService._resolve_percentile_row(rt_row)
|
||||
p50_ttfb, p90_ttfb, p99_ttfb = StatsAggregatorService._resolve_percentile_row(ttfb_row)
|
||||
|
||||
return {
|
||||
"p50_response_time_ms": p50_rt,
|
||||
@@ -223,6 +263,103 @@ class StatsAggregatorService:
|
||||
"p99_first_byte_time_ms": p99_ttfb,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def compute_percentiles_by_local_day(
|
||||
db: Session, time_range: TimeRangeParams
|
||||
) -> list[dict[str, int | None | str]]:
|
||||
"""按本地日期批量计算性能百分位,避免逐天 fan-out。"""
|
||||
bind = db.bind
|
||||
dialect = bind.dialect.name if bind is not None else "sqlite"
|
||||
|
||||
local_dates: list[date] = []
|
||||
current_date = time_range.start_date
|
||||
while current_date <= time_range.end_date:
|
||||
local_dates.append(current_date)
|
||||
current_date += timedelta(days=1)
|
||||
|
||||
if dialect != "postgresql":
|
||||
return [
|
||||
{
|
||||
"date": local_date.isoformat(),
|
||||
"p50_response_time_ms": None,
|
||||
"p90_response_time_ms": None,
|
||||
"p99_response_time_ms": None,
|
||||
"p50_first_byte_time_ms": None,
|
||||
"p90_first_byte_time_ms": None,
|
||||
"p99_first_byte_time_ms": None,
|
||||
}
|
||||
for local_date in local_dates
|
||||
]
|
||||
|
||||
start_utc, end_utc = time_range.to_utc_datetime_range()
|
||||
local_day_expr = StatsAggregatorService._build_local_day_expression(time_range)
|
||||
|
||||
rt_rows = (
|
||||
db.query(
|
||||
local_day_expr.label("local_day"),
|
||||
func.percentile_cont(0.5).within_group(Usage.response_time_ms).label("p50"),
|
||||
func.percentile_cont(0.9).within_group(Usage.response_time_ms).label("p90"),
|
||||
func.percentile_cont(0.99).within_group(Usage.response_time_ms).label("p99"),
|
||||
func.count().label("count"),
|
||||
)
|
||||
.filter(
|
||||
Usage.created_at >= start_utc,
|
||||
Usage.created_at < end_utc,
|
||||
Usage.status == "completed",
|
||||
Usage.response_time_ms.isnot(None),
|
||||
)
|
||||
.group_by(local_day_expr)
|
||||
.all()
|
||||
)
|
||||
|
||||
ttfb_rows = (
|
||||
db.query(
|
||||
local_day_expr.label("local_day"),
|
||||
func.percentile_cont(0.5).within_group(Usage.first_byte_time_ms).label("p50"),
|
||||
func.percentile_cont(0.9).within_group(Usage.first_byte_time_ms).label("p90"),
|
||||
func.percentile_cont(0.99).within_group(Usage.first_byte_time_ms).label("p99"),
|
||||
func.count().label("count"),
|
||||
)
|
||||
.filter(
|
||||
Usage.created_at >= start_utc,
|
||||
Usage.created_at < end_utc,
|
||||
Usage.status == "completed",
|
||||
Usage.first_byte_time_ms.isnot(None),
|
||||
)
|
||||
.group_by(local_day_expr)
|
||||
.all()
|
||||
)
|
||||
|
||||
rt_by_day: dict[date, tuple[int | None, int | None, int | None]] = {}
|
||||
for row in rt_rows:
|
||||
local_day = getattr(row, "local_day", None)
|
||||
if local_day is not None:
|
||||
rt_by_day[local_day] = StatsAggregatorService._resolve_percentile_row(row)
|
||||
|
||||
ttfb_by_day: dict[date, tuple[int | None, int | None, int | None]] = {}
|
||||
for row in ttfb_rows:
|
||||
local_day = getattr(row, "local_day", None)
|
||||
if local_day is not None:
|
||||
ttfb_by_day[local_day] = StatsAggregatorService._resolve_percentile_row(row)
|
||||
|
||||
result: list[dict[str, int | None | str]] = []
|
||||
for local_date in local_dates:
|
||||
p50_rt, p90_rt, p99_rt = rt_by_day.get(local_date, (None, None, None))
|
||||
p50_ttfb, p90_ttfb, p99_ttfb = ttfb_by_day.get(local_date, (None, None, None))
|
||||
result.append(
|
||||
{
|
||||
"date": local_date.isoformat(),
|
||||
"p50_response_time_ms": p50_rt,
|
||||
"p90_response_time_ms": p90_rt,
|
||||
"p99_response_time_ms": p99_rt,
|
||||
"p50_first_byte_time_ms": p50_ttfb,
|
||||
"p90_first_byte_time_ms": p90_ttfb,
|
||||
"p99_first_byte_time_ms": p99_ttfb,
|
||||
}
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def aggregate_daily_stats(db: Session, date: datetime, commit: bool = True) -> StatsDaily:
|
||||
"""聚合指定 UTC 日期的统计数据
|
||||
@@ -656,6 +793,82 @@ class StatsAggregatorService:
|
||||
db.commit()
|
||||
return stats
|
||||
|
||||
@staticmethod
|
||||
def aggregate_user_daily_stats_batch(
|
||||
db: Session, date: datetime, user_ids: list[str], commit: bool = True
|
||||
) -> list[StatsUserDaily]:
|
||||
"""批量聚合单日用户统计,避免逐用户 fan-out 查询。"""
|
||||
if not user_ids:
|
||||
return []
|
||||
|
||||
day_start, day_end = _get_utc_day_range(date)
|
||||
ordered_user_ids = list(dict.fromkeys(user_ids))
|
||||
existing_rows = (
|
||||
db.query(StatsUserDaily)
|
||||
.filter(
|
||||
and_(
|
||||
StatsUserDaily.date == day_start,
|
||||
StatsUserDaily.user_id.in_(ordered_user_ids),
|
||||
)
|
||||
)
|
||||
.all()
|
||||
)
|
||||
existing_by_user = {row.user_id: row for row in existing_rows}
|
||||
|
||||
error_cond = (Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
aggregated_rows = (
|
||||
db.query(
|
||||
Usage.user_id.label("user_id"),
|
||||
func.count(Usage.id).label("total_requests"),
|
||||
func.sum(case((error_cond, 1), else_=0)).label("error_requests"),
|
||||
func.sum(Usage.input_tokens).label("input_tokens"),
|
||||
func.sum(Usage.output_tokens).label("output_tokens"),
|
||||
func.sum(Usage.cache_creation_input_tokens).label("cache_creation_tokens"),
|
||||
func.sum(Usage.cache_read_input_tokens).label("cache_read_tokens"),
|
||||
func.sum(Usage.total_cost_usd).label("total_cost"),
|
||||
func.max(Usage.username).label("username"),
|
||||
)
|
||||
.filter(
|
||||
and_(
|
||||
Usage.user_id.in_(ordered_user_ids),
|
||||
Usage.created_at >= day_start,
|
||||
Usage.created_at < day_end,
|
||||
)
|
||||
)
|
||||
.group_by(Usage.user_id)
|
||||
.all()
|
||||
)
|
||||
aggregated_by_user = {row.user_id: row for row in aggregated_rows}
|
||||
|
||||
result: list[StatsUserDaily] = []
|
||||
for user_id in ordered_user_ids:
|
||||
stats = existing_by_user.get(user_id)
|
||||
if stats is None:
|
||||
stats = StatsUserDaily(id=str(uuid.uuid4()), user_id=user_id, date=day_start)
|
||||
db.add(stats)
|
||||
|
||||
aggregated = aggregated_by_user.get(user_id)
|
||||
if not stats.username and aggregated is not None:
|
||||
username = getattr(aggregated, "username", None)
|
||||
if username:
|
||||
stats.username = username
|
||||
|
||||
total_requests = int(getattr(aggregated, "total_requests", 0) or 0)
|
||||
error_requests = int(getattr(aggregated, "error_requests", 0) or 0)
|
||||
stats.total_requests = total_requests
|
||||
stats.success_requests = total_requests - error_requests
|
||||
stats.error_requests = error_requests
|
||||
stats.input_tokens = int(getattr(aggregated, "input_tokens", 0) or 0)
|
||||
stats.output_tokens = int(getattr(aggregated, "output_tokens", 0) or 0)
|
||||
stats.cache_creation_tokens = int(getattr(aggregated, "cache_creation_tokens", 0) or 0)
|
||||
stats.cache_read_tokens = int(getattr(aggregated, "cache_read_tokens", 0) or 0)
|
||||
stats.total_cost = float(getattr(aggregated, "total_cost", 0) or 0.0)
|
||||
result.append(stats)
|
||||
|
||||
if commit:
|
||||
db.commit()
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def aggregate_daily_stats_bundle(
|
||||
db: Session, date: datetime, user_ids: list[str] | None = None
|
||||
@@ -668,8 +881,9 @@ class StatsAggregatorService:
|
||||
StatsAggregatorService.aggregate_daily_error_stats(db, date, commit=False)
|
||||
|
||||
if user_ids:
|
||||
for user_id in user_ids:
|
||||
StatsAggregatorService.aggregate_user_daily_stats(db, user_id, date, commit=False)
|
||||
StatsAggregatorService.aggregate_user_daily_stats_batch(
|
||||
db, date, user_ids, commit=False
|
||||
)
|
||||
|
||||
stats.is_complete = True
|
||||
stats.aggregated_at = datetime.now(timezone.utc)
|
||||
@@ -1317,28 +1531,32 @@ def query_stats_hybrid(
|
||||
|
||||
result = AggregatedStats()
|
||||
|
||||
preaggregate_dates: list[datetime] = []
|
||||
realtime_dates: list[datetime] = []
|
||||
for day in complete_dates:
|
||||
if day >= today_utc:
|
||||
realtime_dates.append(datetime.combine(day, time.min, tzinfo=timezone.utc))
|
||||
continue
|
||||
day_dt = datetime.combine(day, time.min, tzinfo=timezone.utc)
|
||||
stats = (
|
||||
db.query(StatsDaily)
|
||||
.filter(StatsDaily.date == day_dt, StatsDaily.is_complete.is_(True))
|
||||
.first()
|
||||
)
|
||||
if stats:
|
||||
preaggregate_dates.append(day_dt)
|
||||
else:
|
||||
realtime_dates.append(day_dt)
|
||||
historical_dates = [day for day in complete_dates if day < today_utc]
|
||||
realtime_dates = [day for day in complete_dates if day >= today_utc]
|
||||
|
||||
# Pre-aggregated totals
|
||||
for day_dt in preaggregate_dates:
|
||||
stats = db.query(StatsDaily).filter(StatsDaily.date == day_dt).first()
|
||||
if not stats:
|
||||
continue
|
||||
preaggregated_by_date: dict[date, StatsDaily] = {}
|
||||
if historical_dates:
|
||||
historical_start = datetime.combine(min(historical_dates), time.min, tzinfo=timezone.utc)
|
||||
historical_end = datetime.combine(
|
||||
max(historical_dates) + timedelta(days=1),
|
||||
time.min,
|
||||
tzinfo=timezone.utc,
|
||||
)
|
||||
historical_rows = (
|
||||
db.query(StatsDaily)
|
||||
.filter(
|
||||
StatsDaily.date >= historical_start,
|
||||
StatsDaily.date < historical_end,
|
||||
StatsDaily.is_complete.is_(True),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
preaggregated_by_date = {
|
||||
row.date.astimezone(timezone.utc).date() if row.date.tzinfo else row.date.date(): row
|
||||
for row in historical_rows
|
||||
}
|
||||
|
||||
for stats in preaggregated_by_date.values():
|
||||
result.total_requests += stats.total_requests
|
||||
result.success_requests += stats.success_requests
|
||||
result.error_requests += stats.error_requests
|
||||
@@ -1352,9 +1570,10 @@ def query_stats_hybrid(
|
||||
result.actual_total_cost += float(stats.actual_total_cost or 0)
|
||||
result.total_response_time_ms += (stats.avg_response_time_ms or 0.0) * stats.total_requests
|
||||
|
||||
# Realtime per day
|
||||
for day_dt in realtime_dates:
|
||||
result.add(aggregate_usage_range(db, day_dt, day_dt + timedelta(days=1), filters=filters))
|
||||
missing_historical_dates = [day for day in historical_dates if day not in preaggregated_by_date]
|
||||
realtime_ranges = _merge_consecutive_utc_days(missing_historical_dates + realtime_dates)
|
||||
for range_start, range_end in realtime_ranges:
|
||||
result.add(aggregate_usage_range(db, range_start, range_end, filters=filters))
|
||||
|
||||
if head_boundary:
|
||||
result.add(aggregate_usage_range(db, head_boundary[0], head_boundary[1], filters=filters))
|
||||
|
||||
@@ -247,13 +247,14 @@ class UsageActiveRequestsMixin:
|
||||
default_timeout_seconds: int = 300,
|
||||
*,
|
||||
include_admin_fields: bool = False,
|
||||
maintain_status: bool | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
获取活跃请求状态(用于前端轮询),并自动清理超时的 pending/streaming 请求
|
||||
获取活跃请求状态(用于前端轮询)。
|
||||
|
||||
与 get_active_requests 不同,此方法:
|
||||
1. 返回轻量级的状态字典而非完整 Usage 对象
|
||||
2. 自动检测并清理超时的 pending/streaming 请求
|
||||
2. 可选地检测并清理超时的 pending/streaming 请求
|
||||
3. 支持按 ID 列表查询特定请求
|
||||
|
||||
Args:
|
||||
@@ -261,6 +262,7 @@ class UsageActiveRequestsMixin:
|
||||
ids: 指定要查询的请求 ID 列表(可选)
|
||||
user_id: 限制只查询该用户的请求(可选,用于普通用户接口)
|
||||
default_timeout_seconds: 默认超时时间(秒),当端点未配置时使用
|
||||
maintain_status: 是否执行超时修复与状态回写;默认仅在全量活跃请求查询时执行
|
||||
|
||||
Returns:
|
||||
请求状态列表
|
||||
@@ -311,28 +313,26 @@ class UsageActiveRequestsMixin:
|
||||
query = query.order_by(Usage.created_at.desc()).limit(50)
|
||||
|
||||
records = query.all()
|
||||
should_maintain_status = maintain_status if maintain_status is not None else not ids
|
||||
|
||||
# 检查超时的 pending/streaming 请求
|
||||
# 收集可能超时的 usage_id 列表
|
||||
timeout_candidates: list[str] = []
|
||||
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]] = []
|
||||
|
||||
@@ -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
|
||||
),
|
||||
|
||||
@@ -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(
|
||||
|
||||
143
tests/api/test_admin_usage_routes.py
Normal file
143
tests/api/test_admin_usage_routes.py
Normal 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
|
||||
93
tests/api/test_admin_user_routes.py
Normal file
93
tests/api/test_admin_user_routes.py
Normal file
@@ -0,0 +1,93 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from src.api.admin.users.routes import router as admin_users_router
|
||||
from src.database import get_db
|
||||
|
||||
|
||||
def _build_admin_users_app(db: MagicMock, monkeypatch: pytest.MonkeyPatch) -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(admin_users_router)
|
||||
app.dependency_overrides[get_db] = lambda: db
|
||||
|
||||
async def _fake_pipeline_run(
|
||||
*, adapter: Any, http_request: object, db: MagicMock, mode: object
|
||||
) -> Any:
|
||||
_ = http_request, mode
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
user=SimpleNamespace(id="admin-1"),
|
||||
ensure_json_body=lambda: {},
|
||||
add_audit_metadata=lambda **_: None,
|
||||
)
|
||||
return await adapter.handle(context)
|
||||
|
||||
monkeypatch.setattr("src.api.admin.users.routes.pipeline.run", _fake_pipeline_run)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_list_users_uses_wallet_batch_lookup(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_admin_users_app(db, monkeypatch)
|
||||
now = datetime.now(timezone.utc)
|
||||
users = [
|
||||
SimpleNamespace(
|
||||
id="user-1",
|
||||
email="u1@example.com",
|
||||
username="user1",
|
||||
role=SimpleNamespace(value="user"),
|
||||
allowed_providers=None,
|
||||
allowed_api_formats=None,
|
||||
allowed_models=None,
|
||||
is_active=True,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
last_login_at=None,
|
||||
),
|
||||
SimpleNamespace(
|
||||
id="user-2",
|
||||
email="u2@example.com",
|
||||
username="user2",
|
||||
role=SimpleNamespace(value="admin"),
|
||||
allowed_providers=None,
|
||||
allowed_api_formats=None,
|
||||
allowed_models=None,
|
||||
is_active=True,
|
||||
created_at=now,
|
||||
updated_at=None,
|
||||
last_login_at=None,
|
||||
),
|
||||
]
|
||||
wallets_by_user_id = {
|
||||
"user-1": SimpleNamespace(limit_mode="unlimited"),
|
||||
}
|
||||
batch_getter = MagicMock(return_value=wallets_by_user_id)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.users.routes.UserService.list_users", lambda *_a, **_k: users
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.users.routes.WalletService.get_wallets_by_user_ids",
|
||||
batch_getter,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.users.routes.WalletService.get_wallet",
|
||||
lambda *_a, **_k: (_ for _ in ()).throw(AssertionError("不应回退到逐个钱包查询")),
|
||||
)
|
||||
|
||||
response = client.get("/api/admin/users")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()[0]["unlimited"] is True
|
||||
assert response.json()[1]["unlimited"] is False
|
||||
batch_getter.assert_called_once()
|
||||
assert batch_getter.call_args.args[1] == ["user-1", "user-2"]
|
||||
@@ -8,54 +8,139 @@ API Pipeline 测试
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from 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:
|
||||
|
||||
175
tests/api/test_user_me_usage_routes.py
Normal file
175
tests/api/test_user_me_usage_routes.py
Normal file
@@ -0,0 +1,175 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.api.user_me.routes import GetUsageAdapter
|
||||
from src.core.enums import UserRole
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_usage_adapter_uses_coarse_summary_grouping(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
query = MagicMock()
|
||||
count_query = MagicMock()
|
||||
count_query.scalar.return_value = 0
|
||||
query.outerjoin.return_value = query
|
||||
query.filter.return_value = query
|
||||
query.with_entities.return_value = count_query
|
||||
query.options.return_value = query
|
||||
query.order_by.return_value = query
|
||||
query.offset.return_value = query
|
||||
query.limit.return_value = query
|
||||
query.all.return_value = []
|
||||
db.query.return_value = query
|
||||
|
||||
summary_getter = MagicMock(
|
||||
return_value=[
|
||||
{
|
||||
"provider": "provider-a",
|
||||
"model": "gpt-4o",
|
||||
"requests": 2,
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"total_cost_usd": 1.5,
|
||||
"actual_total_cost_usd": 1.2,
|
||||
"success_count": 2,
|
||||
"success_response_time_sum_ms": 1000.0,
|
||||
"success_response_time_count": 2,
|
||||
},
|
||||
{
|
||||
"provider": "pending",
|
||||
"model": "gpt-4o",
|
||||
"requests": 99,
|
||||
"input_tokens": 999,
|
||||
"output_tokens": 999,
|
||||
"total_tokens": 1998,
|
||||
"total_cost_usd": 9.9,
|
||||
"actual_total_cost_usd": 9.9,
|
||||
"success_count": 0,
|
||||
"success_response_time_sum_ms": 0.0,
|
||||
"success_response_time_count": 0,
|
||||
},
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr("src.api.user_me.routes.UsageService.get_usage_summary", summary_getter)
|
||||
monkeypatch.setattr("src.api.user_me.routes.WalletService.get_wallet", lambda *_a, **_k: None)
|
||||
monkeypatch.setattr(
|
||||
"src.api.user_me.routes.WalletService.serialize_wallet_summary",
|
||||
lambda _wallet: {"limit_mode": "finite"},
|
||||
)
|
||||
|
||||
adapter = GetUsageAdapter(time_range=None, limit=20, offset=0)
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
user=SimpleNamespace(id="user-1", role=UserRole.USER),
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
)
|
||||
|
||||
result = await adapter.handle(context)
|
||||
|
||||
assert result["total_requests"] == 2
|
||||
assert result["total_tokens"] == 15
|
||||
assert result["summary_by_model"] == [
|
||||
{
|
||||
"model": "gpt-4o",
|
||||
"requests": 2,
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"total_cost_usd": 1.5,
|
||||
}
|
||||
]
|
||||
assert "total_actual_cost" not in result
|
||||
assert result["summary_by_provider"] == [
|
||||
{
|
||||
"provider": "provider-a",
|
||||
"requests": 2,
|
||||
"total_tokens": 15,
|
||||
"total_cost_usd": 1.5,
|
||||
"success_rate": 100.0,
|
||||
"avg_response_time_ms": 500.0,
|
||||
}
|
||||
]
|
||||
assert summary_getter.call_args.kwargs["group_by"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_usage_adapter_provider_success_rate_uses_success_count(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
query = MagicMock()
|
||||
count_query = MagicMock()
|
||||
count_query.scalar.return_value = 0
|
||||
query.outerjoin.return_value = query
|
||||
query.filter.return_value = query
|
||||
query.with_entities.return_value = count_query
|
||||
query.options.return_value = query
|
||||
query.order_by.return_value = query
|
||||
query.offset.return_value = query
|
||||
query.limit.return_value = query
|
||||
query.all.return_value = []
|
||||
db.query.return_value = query
|
||||
|
||||
summary_getter = MagicMock(
|
||||
return_value=[
|
||||
{
|
||||
"provider": "provider-a",
|
||||
"model": "gpt-4o",
|
||||
"requests": 3,
|
||||
"input_tokens": 30,
|
||||
"output_tokens": 15,
|
||||
"total_tokens": 45,
|
||||
"total_cost_usd": 4.5,
|
||||
"actual_total_cost_usd": 4.5,
|
||||
"success_count": 2,
|
||||
"success_response_time_sum_ms": 600.0,
|
||||
"success_response_time_count": 2,
|
||||
},
|
||||
{
|
||||
"provider": "provider-a",
|
||||
"model": "gpt-4.1",
|
||||
"requests": 1,
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"total_cost_usd": 1.5,
|
||||
"actual_total_cost_usd": 1.5,
|
||||
"success_count": 0,
|
||||
"success_response_time_sum_ms": 0.0,
|
||||
"success_response_time_count": 0,
|
||||
},
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr("src.api.user_me.routes.UsageService.get_usage_summary", summary_getter)
|
||||
monkeypatch.setattr("src.api.user_me.routes.WalletService.get_wallet", lambda *_a, **_k: None)
|
||||
monkeypatch.setattr(
|
||||
"src.api.user_me.routes.WalletService.serialize_wallet_summary",
|
||||
lambda _wallet: {"limit_mode": "finite"},
|
||||
)
|
||||
|
||||
adapter = GetUsageAdapter(time_range=None, limit=20, offset=0)
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
user=SimpleNamespace(id="user-1", role=UserRole.USER),
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
)
|
||||
|
||||
result = await adapter.handle(context)
|
||||
|
||||
assert result["summary_by_provider"] == [
|
||||
{
|
||||
"provider": "provider-a",
|
||||
"requests": 4,
|
||||
"total_tokens": 60,
|
||||
"total_cost_usd": 6.0,
|
||||
"success_rate": 50.0,
|
||||
"avg_response_time_ms": 300.0,
|
||||
}
|
||||
]
|
||||
@@ -69,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
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
|
||||
|
||||
def _make_db() -> MagicMock:
|
||||
db = MagicMock()
|
||||
db.new = []
|
||||
db.dirty = []
|
||||
db.deleted = []
|
||||
db.in_transaction.return_value = False
|
||||
return db
|
||||
|
||||
|
||||
def _make_global_model() -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
id="gm1",
|
||||
name="gpt-4o",
|
||||
is_active=True,
|
||||
config={},
|
||||
supported_capabilities=[],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_all_candidates_prefilters_provider_graph_by_allowed_providers() -> None:
|
||||
scheduler = CacheAwareScheduler()
|
||||
scheduler.scheduling_mode = CacheAwareScheduler.SCHEDULING_MODE_FIXED_ORDER
|
||||
|
||||
db = _make_db()
|
||||
global_model = _make_global_model()
|
||||
user_api_key = SimpleNamespace(
|
||||
id="ak1",
|
||||
allowed_providers=["provider-b"],
|
||||
allowed_models=None,
|
||||
allowed_api_formats=None,
|
||||
user=None,
|
||||
)
|
||||
|
||||
filtered_provider = SimpleNamespace(
|
||||
id="provider-b",
|
||||
name="provider-b",
|
||||
endpoints=[],
|
||||
models=[],
|
||||
provider_priority=2,
|
||||
)
|
||||
|
||||
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
|
||||
with patch.object(
|
||||
scheduler._candidate_builder,
|
||||
"_query_provider_refs",
|
||||
return_value=[("provider-a", "provider-a"), ("provider-b", "provider-b")],
|
||||
) as refs_mock:
|
||||
with patch.object(
|
||||
scheduler._candidate_builder,
|
||||
"_query_providers",
|
||||
return_value=[filtered_provider],
|
||||
) as providers_mock:
|
||||
with patch(
|
||||
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
|
||||
new=AsyncMock(return_value=global_model),
|
||||
):
|
||||
with patch(
|
||||
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
|
||||
return_value=True,
|
||||
):
|
||||
with patch.object(
|
||||
scheduler._candidate_builder,
|
||||
"_build_candidates",
|
||||
new=AsyncMock(return_value=[]),
|
||||
):
|
||||
candidates, global_model_id, provider_batch_count = (
|
||||
await scheduler.list_all_candidates(
|
||||
db=db,
|
||||
api_format="openai:chat",
|
||||
model_name="gpt-4o",
|
||||
affinity_key=None,
|
||||
user_api_key=user_api_key, # type: ignore[arg-type]
|
||||
provider_offset=0,
|
||||
provider_limit=20,
|
||||
)
|
||||
)
|
||||
|
||||
assert candidates == []
|
||||
assert global_model_id == "gm1"
|
||||
assert provider_batch_count == 2
|
||||
refs_mock.assert_called_once()
|
||||
providers_mock.assert_called_once_with(db=db, provider_ids=["provider-b"])
|
||||
@@ -8,18 +8,20 @@
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from 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 认证"""
|
||||
|
||||
@@ -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
|
||||
|
||||
182
tests/services/test_stats_aggregator_optimization.py
Normal file
182
tests/services/test_stats_aggregator_optimization.py
Normal file
@@ -0,0 +1,182 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime, time, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from src.models.database import StatsDaily, StatsUserDaily
|
||||
from src.services.system.stats_aggregator import (
|
||||
AggregatedStats,
|
||||
StatsAggregatorService,
|
||||
query_stats_hybrid,
|
||||
)
|
||||
from src.services.system.time_range import TimeRangeParams
|
||||
|
||||
|
||||
class _FakeQuery:
|
||||
def __init__(self, *, all_result: list[Any] | None = None) -> None:
|
||||
self._all_result = all_result if all_result is not None else []
|
||||
|
||||
def filter(self, *_args: object, **_kwargs: object) -> _FakeQuery:
|
||||
return self
|
||||
|
||||
def group_by(self, *_args: object, **_kwargs: object) -> _FakeQuery:
|
||||
return self
|
||||
|
||||
def all(self) -> list[Any]:
|
||||
return self._all_result
|
||||
|
||||
|
||||
class _HybridQuerySession:
|
||||
def __init__(self, stats_daily_rows: list[SimpleNamespace]) -> None:
|
||||
self._stats_daily_rows = stats_daily_rows
|
||||
self.stats_daily_query_count = 0
|
||||
|
||||
def query(self, entity: object) -> _FakeQuery:
|
||||
if entity is StatsDaily:
|
||||
self.stats_daily_query_count += 1
|
||||
return _FakeQuery(all_result=self._stats_daily_rows)
|
||||
raise AssertionError(f"Unexpected query entity: {entity}")
|
||||
|
||||
|
||||
class _BatchUserStatsSession:
|
||||
def __init__(
|
||||
self, existing_rows: list[StatsUserDaily], aggregated_rows: list[SimpleNamespace]
|
||||
) -> None:
|
||||
self._responses: list[list[Any]] = [list(existing_rows), list(aggregated_rows)]
|
||||
self.added: list[StatsUserDaily] = []
|
||||
self.commit_count = 0
|
||||
|
||||
def query(self, *_entities: object) -> _FakeQuery:
|
||||
if not self._responses:
|
||||
raise AssertionError("Unexpected extra query")
|
||||
return _FakeQuery(all_result=self._responses.pop(0))
|
||||
|
||||
def add(self, row: StatsUserDaily) -> None:
|
||||
self.added.append(row)
|
||||
|
||||
def commit(self) -> None:
|
||||
self.commit_count += 1
|
||||
|
||||
|
||||
def test_query_stats_hybrid_batches_statsdaily_lookup_and_merges_realtime_ranges(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
today = datetime.now(timezone.utc).date()
|
||||
historical_cached_day = today - timedelta(days=4)
|
||||
historical_missing_day = today - timedelta(days=3)
|
||||
realtime_day = today
|
||||
|
||||
cached_row = SimpleNamespace(
|
||||
date=datetime.combine(historical_cached_day, time.min, tzinfo=timezone.utc),
|
||||
total_requests=10,
|
||||
success_requests=9,
|
||||
error_requests=1,
|
||||
input_tokens=100,
|
||||
output_tokens=50,
|
||||
cache_creation_tokens=5,
|
||||
cache_read_tokens=3,
|
||||
cache_creation_cost=1.2,
|
||||
cache_read_cost=0.8,
|
||||
total_cost=3.5,
|
||||
actual_total_cost=3.0,
|
||||
avg_response_time_ms=200.0,
|
||||
)
|
||||
db = _HybridQuerySession(stats_daily_rows=[cached_row])
|
||||
|
||||
calls: list[tuple[datetime, datetime]] = []
|
||||
|
||||
def _fake_aggregate_usage_range(
|
||||
_db: object,
|
||||
start_utc: datetime,
|
||||
end_utc: datetime,
|
||||
filters: object | None = None, # noqa: ARG001
|
||||
) -> AggregatedStats:
|
||||
calls.append((start_utc, end_utc))
|
||||
return AggregatedStats(total_requests=1, success_requests=1)
|
||||
|
||||
class _FakeParams:
|
||||
def get_complete_utc_dates(self) -> tuple[list[date], None, None]:
|
||||
return [historical_cached_day, historical_missing_day, realtime_day], None, None
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.services.system.stats_aggregator.aggregate_usage_range",
|
||||
_fake_aggregate_usage_range,
|
||||
)
|
||||
|
||||
result = query_stats_hybrid(cast(Any, db), cast(Any, _FakeParams()))
|
||||
|
||||
assert db.stats_daily_query_count == 1
|
||||
assert calls == [
|
||||
(
|
||||
datetime.combine(historical_missing_day, time.min, tzinfo=timezone.utc),
|
||||
datetime.combine(
|
||||
historical_missing_day + timedelta(days=1), time.min, tzinfo=timezone.utc
|
||||
),
|
||||
),
|
||||
(
|
||||
datetime.combine(realtime_day, time.min, tzinfo=timezone.utc),
|
||||
datetime.combine(realtime_day + timedelta(days=1), time.min, tzinfo=timezone.utc),
|
||||
),
|
||||
]
|
||||
assert result.total_requests == 12
|
||||
assert result.success_requests == 11
|
||||
|
||||
|
||||
def test_aggregate_user_daily_stats_batch_updates_all_users_in_two_queries() -> None:
|
||||
target_day = datetime(2026, 3, 1, tzinfo=timezone.utc)
|
||||
aggregated_rows = [
|
||||
SimpleNamespace(
|
||||
user_id="user-1",
|
||||
username="alice",
|
||||
total_requests=4,
|
||||
error_requests=1,
|
||||
input_tokens=20,
|
||||
output_tokens=8,
|
||||
cache_creation_tokens=2,
|
||||
cache_read_tokens=1,
|
||||
total_cost=1.5,
|
||||
)
|
||||
]
|
||||
db = _BatchUserStatsSession(existing_rows=[], aggregated_rows=aggregated_rows)
|
||||
|
||||
result = StatsAggregatorService.aggregate_user_daily_stats_batch(
|
||||
cast(Any, db),
|
||||
target_day,
|
||||
["user-1", "user-2"],
|
||||
commit=True,
|
||||
)
|
||||
|
||||
assert len(result) == 2
|
||||
assert db.commit_count == 1
|
||||
assert len(db.added) == 2
|
||||
|
||||
user_one = next(row for row in result if row.user_id == "user-1")
|
||||
user_two = next(row for row in result if row.user_id == "user-2")
|
||||
|
||||
assert user_one.username == "alice"
|
||||
assert user_one.total_requests == 4
|
||||
assert user_one.success_requests == 3
|
||||
assert user_one.total_cost == 1.5
|
||||
|
||||
assert user_two.total_requests == 0
|
||||
assert user_two.success_requests == 0
|
||||
assert user_two.error_requests == 0
|
||||
assert user_two.total_cost == 0.0
|
||||
|
||||
|
||||
def test_compute_percentiles_by_local_day_returns_sqlite_fallback_without_queries() -> None:
|
||||
db = SimpleNamespace(bind=SimpleNamespace(dialect=SimpleNamespace(name="sqlite")))
|
||||
time_range = TimeRangeParams(
|
||||
start_date=date(2026, 3, 1),
|
||||
end_date=date(2026, 3, 3),
|
||||
timezone="Asia/Singapore",
|
||||
)
|
||||
|
||||
result = StatsAggregatorService.compute_percentiles_by_local_day(cast(Any, db), time_range)
|
||||
|
||||
assert [row["date"] for row in result] == ["2026-03-01", "2026-03-02", "2026-03-03"]
|
||||
assert all(row["p50_response_time_ms"] is None for row in result)
|
||||
assert all(row["p50_first_byte_time_ms"] is None for row in result)
|
||||
@@ -8,7 +8,7 @@ UsageService 测试
|
||||
"""
|
||||
|
||||
from decimal import Decimal
|
||||
from 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()
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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")
|
||||
|
||||
40
tests/unit/test_request_candidate_intermediate_status.py
Normal file
40
tests/unit/test_request_candidate_intermediate_status.py
Normal file
@@ -0,0 +1,40 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
|
||||
|
||||
def _build_db_with_candidate(candidate: SimpleNamespace) -> MagicMock:
|
||||
query = MagicMock()
|
||||
query.filter.return_value.first.return_value = candidate
|
||||
|
||||
db = MagicMock()
|
||||
db.query.return_value = query
|
||||
db.info = {"managed_by_middleware": True}
|
||||
return db
|
||||
|
||||
|
||||
def test_mark_candidate_started_flushes_without_immediate_commit() -> None:
|
||||
candidate = SimpleNamespace(status="available", started_at=None)
|
||||
db = _build_db_with_candidate(candidate)
|
||||
|
||||
RequestCandidateService.mark_candidate_started(db, "candidate-1")
|
||||
|
||||
assert candidate.status == "pending"
|
||||
assert candidate.started_at is not None
|
||||
db.flush.assert_called_once()
|
||||
db.commit.assert_not_called()
|
||||
|
||||
|
||||
def test_mark_candidate_streaming_flushes_without_immediate_commit() -> None:
|
||||
candidate = SimpleNamespace(status="pending", concurrent_requests=None)
|
||||
db = _build_db_with_candidate(candidate)
|
||||
|
||||
RequestCandidateService.mark_candidate_streaming(db, "candidate-1", concurrent_requests=3)
|
||||
|
||||
assert candidate.status == "streaming"
|
||||
assert candidate.concurrent_requests == 3
|
||||
db.flush.assert_called_once()
|
||||
db.commit.assert_not_called()
|
||||
Reference in New Issue
Block a user