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:
fawney19
2026-02-02 03:16:52 +08:00
parent feb7484fda
commit 9e31efe26c
75 changed files with 7511 additions and 2068 deletions

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

View File

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

View File

@@ -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,

View File

@@ -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_codestream, 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