mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
refactor: 重构异步任务系统和计费服务架构
- 重构任务系统:新增 lifecycle (TaskStatus/BillingStatus)、context、application 模块 - 将 video tasks 泛化为 async tasks,支持更通用的异步任务管理 - 新增 Gemini Files 管理模块和管理界面 - 重构 billing 服务:拆分 schema.py 和 service.py - 新增 candidate 服务模块用于请求候选管理 - 数据库迁移:添加 billing_status、request_id、gemini_file_mappings 表和索引 - 移除废弃的 video_telemetry、task orchestrator 等模块
This commit is contained in:
543
src/api/admin/gemini_files.py
Normal file
543
src/api/admin/gemini_files.py
Normal file
@@ -0,0 +1,543 @@
|
||||
"""
|
||||
Gemini Files 管理 API
|
||||
|
||||
提供文件映射的管理功能:
|
||||
- 列出所有文件映射
|
||||
- 删除文件映射
|
||||
- 查看文件映射统计
|
||||
- 上传文件到 Gemini
|
||||
|
||||
优化:HTTP 上传期间不持有数据库连接,避免阻塞其他请求。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import delete, func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.database import create_session, get_db
|
||||
from src.models.database import GeminiFileMapping, ProviderAPIKey, User
|
||||
from src.services.gemini_files_mapping import delete_file_key_mapping, store_file_key_mapping
|
||||
|
||||
|
||||
@dataclass
|
||||
class KeyInfo:
|
||||
"""Key 信息(用于 HTTP 上传,不依赖数据库会话)"""
|
||||
|
||||
id: str
|
||||
name: str | None
|
||||
decrypted_api_key: str
|
||||
|
||||
|
||||
router = APIRouter(prefix="/api/admin/gemini-files", tags=["Gemini Files Management"])
|
||||
|
||||
# Gemini Files API 基础 URL
|
||||
GEMINI_FILES_BASE_URL = "https://generativelanguage.googleapis.com"
|
||||
|
||||
|
||||
# ============ Schema ============
|
||||
|
||||
|
||||
class FileMappingResponse(BaseModel):
|
||||
"""文件映射响应"""
|
||||
|
||||
id: str
|
||||
file_name: str
|
||||
key_id: str
|
||||
key_name: str | None = None
|
||||
user_id: str | None = None
|
||||
username: str | None = None
|
||||
display_name: str | None = None
|
||||
mime_type: str | None = None
|
||||
created_at: datetime
|
||||
expires_at: datetime
|
||||
is_expired: bool
|
||||
|
||||
|
||||
class FileMappingListResponse(BaseModel):
|
||||
"""文件映射列表响应"""
|
||||
|
||||
items: list[FileMappingResponse]
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
|
||||
|
||||
class FileMappingStatsResponse(BaseModel):
|
||||
"""文件映射统计响应"""
|
||||
|
||||
total_mappings: int
|
||||
active_mappings: int
|
||||
expired_mappings: int
|
||||
by_mime_type: dict[str, int]
|
||||
capable_keys_count: int
|
||||
|
||||
|
||||
# ============ Routes ============
|
||||
|
||||
|
||||
@router.get("/mappings", response_model=FileMappingListResponse)
|
||||
async def list_file_mappings(
|
||||
db: Session = Depends(get_db),
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(20, ge=1, le=100),
|
||||
include_expired: bool = Query(False),
|
||||
search: str | None = Query(None),
|
||||
) -> Any:
|
||||
"""
|
||||
列出所有文件映射
|
||||
|
||||
- **page**: 页码
|
||||
- **page_size**: 每页数量
|
||||
- **include_expired**: 是否包含已过期的映射
|
||||
- **search**: 搜索文件名或显示名
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
query = db.query(GeminiFileMapping)
|
||||
|
||||
# 过滤过期
|
||||
if not include_expired:
|
||||
query = query.filter(GeminiFileMapping.expires_at > now)
|
||||
|
||||
# 搜索
|
||||
if search:
|
||||
search_pattern = f"%{search}%"
|
||||
query = query.filter(
|
||||
(GeminiFileMapping.file_name.ilike(search_pattern))
|
||||
| (GeminiFileMapping.display_name.ilike(search_pattern))
|
||||
)
|
||||
|
||||
# 总数
|
||||
total = query.count()
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
mappings = (
|
||||
query.order_by(GeminiFileMapping.created_at.desc()).offset(offset).limit(page_size).all()
|
||||
)
|
||||
|
||||
# 获取关联的 Key 和 User 信息
|
||||
key_ids = {m.key_id for m in mappings}
|
||||
user_ids = {m.user_id for m in mappings if m.user_id}
|
||||
|
||||
keys_map = {}
|
||||
if key_ids:
|
||||
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.id.in_(key_ids)).all()
|
||||
keys_map = {str(k.id): k.name for k in keys}
|
||||
|
||||
users_map = {}
|
||||
if user_ids:
|
||||
users = db.query(User).filter(User.id.in_(user_ids)).all()
|
||||
users_map = {str(u.id): u.username for u in users}
|
||||
|
||||
items = []
|
||||
for m in mappings:
|
||||
items.append(
|
||||
FileMappingResponse(
|
||||
id=str(m.id),
|
||||
file_name=m.file_name,
|
||||
key_id=str(m.key_id),
|
||||
key_name=keys_map.get(str(m.key_id)),
|
||||
user_id=str(m.user_id) if m.user_id else None,
|
||||
username=users_map.get(str(m.user_id)) if m.user_id else None,
|
||||
display_name=m.display_name,
|
||||
mime_type=m.mime_type,
|
||||
created_at=m.created_at,
|
||||
expires_at=m.expires_at,
|
||||
is_expired=m.expires_at <= now,
|
||||
)
|
||||
)
|
||||
|
||||
return FileMappingListResponse(
|
||||
items=items,
|
||||
total=total,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/stats", response_model=FileMappingStatsResponse)
|
||||
async def get_file_mapping_stats(
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""获取文件映射统计信息"""
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# 总数
|
||||
total_mappings = db.query(func.count(GeminiFileMapping.id)).scalar() or 0
|
||||
|
||||
# 活跃数(未过期)
|
||||
active_mappings = (
|
||||
db.query(func.count(GeminiFileMapping.id))
|
||||
.filter(GeminiFileMapping.expires_at > now)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
|
||||
# 过期数
|
||||
expired_mappings = total_mappings - active_mappings
|
||||
|
||||
# 按 MIME 类型统计
|
||||
mime_stats = (
|
||||
db.query(GeminiFileMapping.mime_type, func.count(GeminiFileMapping.id))
|
||||
.filter(GeminiFileMapping.expires_at > now)
|
||||
.group_by(GeminiFileMapping.mime_type)
|
||||
.all()
|
||||
)
|
||||
by_mime_type = {(mt or "unknown"): count for mt, count in mime_stats}
|
||||
|
||||
# 有 gemini_files 能力的 Key 数量
|
||||
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
|
||||
capable_keys_count = sum(
|
||||
1 for key in keys if key.capabilities and key.capabilities.get("gemini_files", False)
|
||||
)
|
||||
|
||||
return FileMappingStatsResponse(
|
||||
total_mappings=total_mappings,
|
||||
active_mappings=active_mappings,
|
||||
expired_mappings=expired_mappings,
|
||||
by_mime_type=by_mime_type,
|
||||
capable_keys_count=capable_keys_count,
|
||||
)
|
||||
|
||||
|
||||
@router.delete("/mappings/{mapping_id}")
|
||||
async def delete_mapping(
|
||||
mapping_id: str,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
删除指定的文件映射
|
||||
|
||||
注意:这只删除映射记录,不会删除 Gemini 上的实际文件
|
||||
"""
|
||||
mapping = db.query(GeminiFileMapping).filter(GeminiFileMapping.id == mapping_id).first()
|
||||
|
||||
if not mapping:
|
||||
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||
|
||||
file_name = mapping.file_name
|
||||
|
||||
# 从数据库删除
|
||||
db.delete(mapping)
|
||||
db.commit()
|
||||
|
||||
# 同时从 Redis 删除
|
||||
await delete_file_key_mapping(file_name)
|
||||
|
||||
return {"message": "Mapping deleted successfully", "file_name": file_name}
|
||||
|
||||
|
||||
@router.delete("/mappings")
|
||||
async def cleanup_expired_mappings(
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""清理所有过期的文件映射"""
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
result = db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.expires_at <= now))
|
||||
db.commit()
|
||||
|
||||
deleted_count = result.rowcount
|
||||
|
||||
return {
|
||||
"message": f"Cleaned up {deleted_count} expired mappings",
|
||||
"deleted_count": deleted_count,
|
||||
}
|
||||
|
||||
|
||||
class CapableKeyResponse(BaseModel):
|
||||
"""可用 Key 响应"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
provider_name: str | None = None
|
||||
|
||||
|
||||
class UploadResultItem(BaseModel):
|
||||
"""单个 Key 的上传结果"""
|
||||
|
||||
key_id: str
|
||||
key_name: str | None = None
|
||||
success: bool
|
||||
file_name: str | None = None
|
||||
error: str | None = None
|
||||
|
||||
|
||||
class UploadResponse(BaseModel):
|
||||
"""上传响应"""
|
||||
|
||||
display_name: str
|
||||
mime_type: str
|
||||
size_bytes: int
|
||||
results: list[UploadResultItem]
|
||||
success_count: int
|
||||
fail_count: int
|
||||
|
||||
|
||||
@router.get("/capable-keys", response_model=list[CapableKeyResponse])
|
||||
async def list_capable_keys(
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""获取所有具有 gemini_files 能力的 Key 列表"""
|
||||
from src.models.database import Provider
|
||||
|
||||
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
|
||||
capable_keys = [
|
||||
key for key in keys if key.capabilities and key.capabilities.get("gemini_files", False)
|
||||
]
|
||||
|
||||
# 获取 Provider 名称
|
||||
provider_ids = {key.provider_id for key in capable_keys}
|
||||
providers = db.query(Provider).filter(Provider.id.in_(provider_ids)).all()
|
||||
provider_map = {str(p.id): p.name for p in providers}
|
||||
|
||||
return [
|
||||
CapableKeyResponse(
|
||||
id=str(key.id),
|
||||
name=key.name,
|
||||
provider_name=provider_map.get(str(key.provider_id)),
|
||||
)
|
||||
for key in capable_keys
|
||||
]
|
||||
|
||||
|
||||
async def _upload_to_key(
|
||||
key_info: KeyInfo,
|
||||
content: bytes,
|
||||
file_size: int,
|
||||
mime_type: str,
|
||||
display_name: str,
|
||||
source_hash: str,
|
||||
) -> UploadResultItem:
|
||||
"""上传文件到指定的 Key(不依赖数据库会话)"""
|
||||
api_key = key_info.decrypted_api_key
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
|
||||
try:
|
||||
# 第一步:初始化可恢复上传
|
||||
init_url = f"{GEMINI_FILES_BASE_URL}/upload/v1beta/files?key={api_key}"
|
||||
init_headers = {
|
||||
"X-Goog-Upload-Protocol": "resumable",
|
||||
"X-Goog-Upload-Command": "start",
|
||||
"X-Goog-Upload-Header-Content-Length": str(file_size),
|
||||
"X-Goog-Upload-Header-Content-Type": mime_type,
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
init_body = {"file": {"display_name": display_name}}
|
||||
|
||||
init_response = await client.post(init_url, headers=init_headers, json=init_body)
|
||||
|
||||
if init_response.status_code != 200:
|
||||
logger.error(
|
||||
f"Gemini upload init failed for key {key_info.id}: {init_response.status_code}"
|
||||
)
|
||||
return UploadResultItem(
|
||||
key_id=key_info.id,
|
||||
key_name=key_info.name,
|
||||
success=False,
|
||||
error=f"初始化失败: {init_response.status_code}",
|
||||
)
|
||||
|
||||
# 获取上传 URL
|
||||
upload_url = init_response.headers.get("X-Goog-Upload-URL")
|
||||
if not upload_url:
|
||||
return UploadResultItem(
|
||||
key_id=key_info.id,
|
||||
key_name=key_info.name,
|
||||
success=False,
|
||||
error="未获取到上传 URL",
|
||||
)
|
||||
|
||||
# 第二步:上传文件内容
|
||||
upload_headers = {
|
||||
"X-Goog-Upload-Command": "upload, finalize",
|
||||
"X-Goog-Upload-Offset": "0",
|
||||
"Content-Length": str(file_size),
|
||||
"Content-Type": mime_type,
|
||||
}
|
||||
|
||||
upload_response = await client.post(upload_url, headers=upload_headers, content=content)
|
||||
|
||||
if upload_response.status_code != 200:
|
||||
logger.error(
|
||||
f"Gemini upload failed for key {key_info.id}: {upload_response.status_code}"
|
||||
)
|
||||
return UploadResultItem(
|
||||
key_id=key_info.id,
|
||||
key_name=key_info.name,
|
||||
success=False,
|
||||
error=f"上传失败: {upload_response.status_code}",
|
||||
)
|
||||
|
||||
# 解析响应
|
||||
result = upload_response.json()
|
||||
file_info = result.get("file", {})
|
||||
file_name = file_info.get("name", "")
|
||||
response_display_name = file_info.get("displayName", display_name)
|
||||
response_mime_type = file_info.get("mimeType", mime_type)
|
||||
|
||||
# 存储文件映射(包含源文件哈希,用于关联相同源文件的不同上传)
|
||||
await store_file_key_mapping(
|
||||
file_name=file_name,
|
||||
key_id=key_info.id,
|
||||
user_id=None,
|
||||
display_name=response_display_name,
|
||||
mime_type=response_mime_type,
|
||||
source_hash=source_hash,
|
||||
)
|
||||
|
||||
return UploadResultItem(
|
||||
key_id=key_info.id,
|
||||
key_name=key_info.name,
|
||||
success=True,
|
||||
file_name=file_name,
|
||||
)
|
||||
|
||||
except Exception as exc:
|
||||
logger.error(f"Gemini upload error for key {key_info.id}: {exc}")
|
||||
return UploadResultItem(
|
||||
key_id=key_info.id,
|
||||
key_name=key_info.name,
|
||||
success=False,
|
||||
error=str(exc),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/upload", response_model=UploadResponse)
|
||||
async def upload_file(
|
||||
file: UploadFile = File(...),
|
||||
key_ids: str = Query(..., description="逗号分隔的 Key ID 列表"),
|
||||
) -> Any:
|
||||
"""
|
||||
上传文件到 Gemini Files API
|
||||
|
||||
- **file**: 要上传的文件
|
||||
- **key_ids**: 逗号分隔的 Key ID 列表,文件将上传到所有指定的 Key
|
||||
|
||||
支持的文件类型:视频、图片、音频、文档等
|
||||
文件大小限制:2GB
|
||||
文件有效期:48小时
|
||||
|
||||
优化:HTTP 上传期间不持有数据库连接
|
||||
"""
|
||||
# 解析 Key IDs
|
||||
key_id_list = [kid.strip() for kid in key_ids.split(",") if kid.strip()]
|
||||
if not key_id_list:
|
||||
raise HTTPException(status_code=400, detail="请至少选择一个 Key")
|
||||
|
||||
# ========== 阶段 1:读取文件内容并计算哈希 ==========
|
||||
content = await file.read()
|
||||
file_size = len(content)
|
||||
mime_type = file.content_type or "application/octet-stream"
|
||||
display_name = file.filename or "uploaded_file"
|
||||
|
||||
# 计算源文件哈希(用于重复检测和关联相同源文件的不同上传)
|
||||
source_hash = hashlib.sha256(content).hexdigest()
|
||||
|
||||
# ========== 阶段 2:查询数据库(短暂持有连接)==========
|
||||
key_infos: list[KeyInfo] = []
|
||||
existing_mappings: dict[str, str] = {} # key_id -> existing file_name
|
||||
|
||||
with create_session() as db:
|
||||
keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
.filter(
|
||||
ProviderAPIKey.id.in_(key_id_list),
|
||||
ProviderAPIKey.is_active.is_(True),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 过滤有 gemini_files 能力的 Key,并提取必要信息
|
||||
capable_key_ids = []
|
||||
for key in keys:
|
||||
if key.capabilities and key.capabilities.get("gemini_files", False):
|
||||
try:
|
||||
decrypted_key = crypto_service.decrypt(key.api_key)
|
||||
key_infos.append(
|
||||
KeyInfo(
|
||||
id=str(key.id),
|
||||
name=key.name,
|
||||
decrypted_api_key=decrypted_key,
|
||||
)
|
||||
)
|
||||
capable_key_ids.append(str(key.id))
|
||||
except Exception as exc:
|
||||
logger.error(f"Failed to decrypt provider key {key.id}: {exc}")
|
||||
|
||||
# 检查是否已存在相同 source_hash 的映射(重复检测)
|
||||
if capable_key_ids:
|
||||
now = datetime.now(timezone.utc)
|
||||
existing = (
|
||||
db.query(GeminiFileMapping)
|
||||
.filter(
|
||||
GeminiFileMapping.source_hash == source_hash,
|
||||
GeminiFileMapping.key_id.in_(capable_key_ids),
|
||||
GeminiFileMapping.expires_at > now, # 只查未过期的
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for mapping in existing:
|
||||
existing_mappings[str(mapping.key_id)] = mapping.file_name
|
||||
|
||||
if not key_infos:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="选中的 Key 都没有「Gemini 文件 API」能力或解密失败",
|
||||
)
|
||||
|
||||
# ========== 阶段 3:并发上传(跳过已有相同文件的 Key)==========
|
||||
results: list[UploadResultItem] = []
|
||||
keys_to_upload: list[KeyInfo] = []
|
||||
|
||||
for key_info in key_infos:
|
||||
if key_info.id in existing_mappings:
|
||||
# 该 Key 已有相同文件,跳过上传
|
||||
results.append(
|
||||
UploadResultItem(
|
||||
key_id=key_info.id,
|
||||
key_name=key_info.name,
|
||||
success=True,
|
||||
file_name=existing_mappings[key_info.id],
|
||||
error=None,
|
||||
)
|
||||
)
|
||||
logger.info(
|
||||
f"跳过重复上传: Key {key_info.id} 已有文件 {existing_mappings[key_info.id]}"
|
||||
)
|
||||
else:
|
||||
keys_to_upload.append(key_info)
|
||||
|
||||
# 只对需要上传的 Key 执行上传
|
||||
if keys_to_upload:
|
||||
tasks = [
|
||||
_upload_to_key(key_info, content, file_size, mime_type, display_name, source_hash)
|
||||
for key_info in keys_to_upload
|
||||
]
|
||||
upload_results = await asyncio.gather(*tasks)
|
||||
results.extend(upload_results)
|
||||
|
||||
success_count = sum(1 for r in results if r.success)
|
||||
fail_count = len(results) - success_count
|
||||
|
||||
return UploadResponse(
|
||||
display_name=display_name,
|
||||
mime_type=mime_type,
|
||||
size_bytes=file_size,
|
||||
results=results,
|
||||
success_count=success_count,
|
||||
fail_count=fail_count,
|
||||
)
|
||||
@@ -20,6 +20,7 @@ from src.core.logger import logger
|
||||
from src.database.database import get_db
|
||||
from src.models.database import Provider, ProviderEndpoint, User
|
||||
from src.services.model.fetch_scheduler import (
|
||||
MODEL_FETCH_HTTP_TIMEOUT,
|
||||
get_upstream_models_from_cache,
|
||||
set_upstream_models_to_cache,
|
||||
)
|
||||
@@ -135,7 +136,9 @@ async def query_available_models(
|
||||
return [], f"Key {api_key.name or api_key.id}: decrypt failed", False
|
||||
|
||||
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint)
|
||||
models, errors, has_success = await fetch_models_from_endpoints(endpoint_configs)
|
||||
models, errors, has_success = await fetch_models_from_endpoints(
|
||||
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
|
||||
)
|
||||
|
||||
# 写入缓存
|
||||
if models:
|
||||
@@ -274,7 +277,9 @@ async def _fetch_models_for_single_key(
|
||||
raise HTTPException(status_code=500, detail="Failed to decrypt API key")
|
||||
|
||||
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint)
|
||||
all_models, errors, has_success = await fetch_models_from_endpoints(endpoint_configs)
|
||||
all_models, errors, has_success = await fetch_models_from_endpoints(
|
||||
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
|
||||
)
|
||||
|
||||
# 按 model id 聚合,合并所有 api_format
|
||||
unique_models = _aggregate_models_by_id(all_models)
|
||||
|
||||
@@ -309,6 +309,7 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
|
||||
website=provider.website,
|
||||
provider_priority=provider.provider_priority,
|
||||
keep_priority_on_conversion=provider.keep_priority_on_conversion,
|
||||
enable_format_conversion=provider.enable_format_conversion,
|
||||
is_active=provider.is_active,
|
||||
billing_type=provider.billing_type.value if provider.billing_type else None,
|
||||
monthly_quota_usd=provider.monthly_quota_usd,
|
||||
|
||||
@@ -271,6 +271,9 @@ class AdminUsageStatsAdapter(AdminApiAdapter):
|
||||
self.end_date = end_date
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
# Perf: use a single aggregate query (avoid 3 full scans).
|
||||
from sqlalchemy import case
|
||||
|
||||
db = context.db
|
||||
query = db.query(Usage)
|
||||
if self.start_date:
|
||||
@@ -278,56 +281,53 @@ class AdminUsageStatsAdapter(AdminApiAdapter):
|
||||
if self.end_date:
|
||||
query = query.filter(Usage.created_at <= self.end_date)
|
||||
|
||||
total_stats = query.with_entities(
|
||||
stats = query.with_entities(
|
||||
func.count(Usage.id).label("total_requests"),
|
||||
func.sum(Usage.total_tokens).label("total_tokens"),
|
||||
func.sum(Usage.total_cost_usd).label("total_cost"),
|
||||
func.sum(Usage.actual_total_cost_usd).label("total_actual_cost"),
|
||||
func.avg(Usage.response_time_ms).label("avg_response_time_ms"),
|
||||
).first()
|
||||
|
||||
# 缓存统计
|
||||
cache_stats = query.with_entities(
|
||||
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.cache_creation_cost_usd).label("cache_creation_cost"),
|
||||
func.sum(Usage.cache_read_cost_usd).label("cache_read_cost"),
|
||||
func.sum(
|
||||
case(
|
||||
(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None)),
|
||||
1,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
).label("error_count"),
|
||||
).first()
|
||||
|
||||
# 错误统计
|
||||
error_count = query.filter(
|
||||
(Usage.status_code >= 400) | (Usage.error_message.isnot(None))
|
||||
).count()
|
||||
|
||||
context.add_audit_metadata(
|
||||
action="usage_stats",
|
||||
start_date=self.start_date.isoformat() if self.start_date else None,
|
||||
end_date=self.end_date.isoformat() if self.end_date else None,
|
||||
)
|
||||
|
||||
total_requests = total_stats.total_requests if total_stats else 0
|
||||
avg_response_time_ms = float(total_stats.avg_response_time_ms or 0) if total_stats else 0
|
||||
total_requests = int(stats.total_requests or 0) if stats else 0
|
||||
avg_response_time_ms = float(stats.avg_response_time_ms or 0) if stats else 0
|
||||
avg_response_time = avg_response_time_ms / 1000.0
|
||||
error_count = int(stats.error_count or 0) if stats else 0
|
||||
|
||||
return {
|
||||
"total_requests": total_requests,
|
||||
"total_tokens": int(total_stats.total_tokens or 0),
|
||||
"total_cost": float(total_stats.total_cost or 0),
|
||||
"total_actual_cost": float(total_stats.total_actual_cost or 0),
|
||||
"total_tokens": int(stats.total_tokens or 0) if stats else 0,
|
||||
"total_cost": float(stats.total_cost or 0) if stats else 0,
|
||||
"total_actual_cost": float(stats.total_actual_cost or 0) if stats else 0,
|
||||
"avg_response_time": round(avg_response_time, 2),
|
||||
"error_count": error_count,
|
||||
"error_rate": (
|
||||
round((error_count / total_requests) * 100, 2) if total_requests > 0 else 0
|
||||
),
|
||||
"cache_stats": {
|
||||
"cache_creation_tokens": (
|
||||
int(cache_stats.cache_creation_tokens or 0) if cache_stats else 0
|
||||
),
|
||||
"cache_read_tokens": int(cache_stats.cache_read_tokens or 0) if cache_stats else 0,
|
||||
"cache_creation_cost": (
|
||||
float(cache_stats.cache_creation_cost or 0) if cache_stats else 0
|
||||
),
|
||||
"cache_read_cost": float(cache_stats.cache_read_cost or 0) if cache_stats else 0,
|
||||
"cache_creation_tokens": (int(stats.cache_creation_tokens or 0) if stats else 0),
|
||||
"cache_read_tokens": int(stats.cache_read_tokens or 0) if stats else 0,
|
||||
"cache_creation_cost": (float(stats.cache_creation_cost or 0) if stats else 0),
|
||||
"cache_read_cost": float(stats.cache_read_cost or 0) if stats else 0,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -637,6 +637,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from sqlalchemy import or_
|
||||
from sqlalchemy.orm import load_only
|
||||
|
||||
from src.utils.database_helpers import escape_like_pattern, safe_truncate_escaped
|
||||
|
||||
@@ -677,13 +678,13 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
escaped = escape_like_pattern(self.username)
|
||||
query = query.filter(User.username.ilike(f"%{escaped}%", escape="\\"))
|
||||
if self.model:
|
||||
# 支持模型名模糊搜索
|
||||
escaped = escape_like_pattern(self.model)
|
||||
query = query.filter(Usage.model.ilike(f"%{escaped}%", escape="\\"))
|
||||
# 模型筛选:前端为下拉框精确值,使用精确匹配以启用索引
|
||||
# 如需模糊搜索,请使用 search 参数。
|
||||
query = query.filter(Usage.model == self.model)
|
||||
if self.provider:
|
||||
# 支持提供商名称搜索
|
||||
escaped = escape_like_pattern(self.provider)
|
||||
query = query.filter(Provider.name.ilike(f"%{escaped}%", escape="\\"))
|
||||
# 提供商筛选:前端为下拉框精确值,使用精确匹配以启用索引
|
||||
# 如需模糊搜索,请使用 search 参数。
|
||||
query = query.filter(Provider.name == self.provider)
|
||||
if self.status:
|
||||
# 状态筛选
|
||||
# 旧的筛选值(基于 is_stream 和 status_code):stream, standard, error
|
||||
@@ -714,7 +715,51 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
if self.end_date:
|
||||
query = query.filter(Usage.created_at <= self.end_date)
|
||||
|
||||
total = query.count()
|
||||
# Perf: avoid Query.count() building a subquery selecting many columns
|
||||
total = int(query.with_entities(func.count(Usage.id)).scalar() or 0)
|
||||
|
||||
# Perf: do not load large request/response columns for list view
|
||||
query = query.options(
|
||||
load_only(
|
||||
Usage.id,
|
||||
Usage.request_id,
|
||||
Usage.user_id,
|
||||
Usage.api_key_id,
|
||||
Usage.provider_name,
|
||||
Usage.provider_id,
|
||||
Usage.provider_endpoint_id,
|
||||
Usage.provider_api_key_id,
|
||||
Usage.model,
|
||||
Usage.target_model,
|
||||
Usage.input_tokens,
|
||||
Usage.output_tokens,
|
||||
Usage.cache_creation_input_tokens,
|
||||
Usage.cache_read_input_tokens,
|
||||
Usage.total_tokens,
|
||||
Usage.total_cost_usd,
|
||||
Usage.actual_total_cost_usd,
|
||||
Usage.rate_multiplier,
|
||||
Usage.response_time_ms,
|
||||
Usage.first_byte_time_ms,
|
||||
Usage.created_at,
|
||||
Usage.is_stream,
|
||||
Usage.status_code,
|
||||
Usage.error_message,
|
||||
Usage.status,
|
||||
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,
|
||||
Usage.cache_read_price_per_1m,
|
||||
),
|
||||
load_only(User.id, User.email, User.username),
|
||||
load_only(ProviderEndpoint.id, ProviderEndpoint.api_format),
|
||||
load_only(ProviderAPIKey.id, ProviderAPIKey.name),
|
||||
load_only(ApiKey.id, ApiKey.name, ApiKey.key_encrypted),
|
||||
)
|
||||
records = (
|
||||
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
|
||||
)
|
||||
@@ -779,7 +824,9 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 构建 provider_id -> Provider 名称的映射,避免 N+1 查询
|
||||
provider_ids = [usage.provider_id for usage, _, _, _, _ in records if usage.provider_id]
|
||||
provider_ids = list(
|
||||
{usage.provider_id for usage, _, _, _, _ in records if usage.provider_id}
|
||||
)
|
||||
provider_map = {}
|
||||
if provider_ids:
|
||||
providers_data = (
|
||||
@@ -788,6 +835,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
provider_map = {str(p.id): p.name for p in providers_data}
|
||||
|
||||
data = []
|
||||
api_key_display_cache: dict[str, str] = {}
|
||||
for usage, user, endpoint, provider_api_key, user_api_key in records:
|
||||
actual_cost = (
|
||||
float(usage.actual_total_cost_usd)
|
||||
@@ -829,7 +877,9 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
{
|
||||
"id": user_api_key.id,
|
||||
"name": user_api_key.name,
|
||||
"display": user_api_key.get_display_key(),
|
||||
"display": api_key_display_cache.setdefault(
|
||||
user_api_key.id, user_api_key.get_display_key()
|
||||
),
|
||||
}
|
||||
if user_api_key
|
||||
else None
|
||||
|
||||
Reference in New Issue
Block a user