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

View File

@@ -311,7 +311,7 @@ def get_compatible_provider_formats(
endpoint_format,
format_acceptance_config,
is_stream=False,
global_conversion_enabled=global_conversion_enabled,
effective_conversion_enabled=global_conversion_enabled,
)
if not is_compatible:
continue

View File

@@ -813,23 +813,31 @@ class DashboardRecentRequestsAdapter(DashboardAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
query = db.query(Usage)
# Perf: select only required columns (avoid loading large JSON/BLOB fields).
query = db.query(
Usage.id,
Usage.user_id,
Usage.model,
Usage.total_tokens,
Usage.created_at,
Usage.is_stream,
DBUser.username,
).outerjoin(DBUser, DBUser.id == Usage.user_id)
if user.role != UserRole.ADMIN:
query = query.filter(Usage.user_id == user.id)
recent_requests = query.order_by(Usage.created_at.desc()).limit(self.limit).all()
rows = query.order_by(Usage.created_at.desc()).limit(self.limit).all()
results = []
for req in recent_requests:
owner = db.query(DBUser).filter(DBUser.id == req.user_id).first()
for req_id, _user_id, model, total_tokens, created_at, is_stream, username in rows:
results.append(
{
"id": req.id,
"user": owner.username if owner else "Unknown",
"model": req.model or "N/A",
"tokens": req.total_tokens,
"time": req.created_at.strftime("%H:%M") if req.created_at else None,
"is_stream": req.is_stream,
"id": req_id,
"user": username or "Unknown",
"model": model or "N/A",
"tokens": int(total_tokens or 0),
"time": created_at.strftime("%H:%M") if created_at else None,
"is_stream": bool(is_stream),
}
)
@@ -848,18 +856,25 @@ class DashboardProviderStatusAdapter(DashboardAdapter):
providers = db.query(Provider).filter(Provider.is_active.is_(True)).all()
since = datetime.now(timezone.utc) - timedelta(days=1)
# Avoid N+1: compute 24h request counts for all providers in one GROUP BY query.
provider_names = [p.name for p in providers if p and p.name]
counts: dict[str, int] = {}
if provider_names:
rows = (
db.query(Usage.provider_name, func.count(Usage.id))
.filter(and_(Usage.created_at >= since, Usage.provider_name.in_(provider_names)))
.group_by(Usage.provider_name)
.all()
)
counts = {str(name): int(cnt or 0) for name, cnt in rows if name}
entries = []
for provider in providers:
count = (
db.query(func.count(Usage.id))
.filter(and_(Usage.provider_name == provider.name, Usage.created_at >= since))
.scalar()
)
entries.append(
{
"name": provider.name,
"status": "active" if provider.is_active else "inactive",
"requests": count,
"requests": int(counts.get(provider.name, 0)),
}
)

View File

@@ -79,10 +79,8 @@ class VideoAdapterBase(ApiAdapter):
path_params=path_params,
)
# Cancel task
if method in {"DELETE", "POST"} and (
path.endswith("/cancel") or path_params.get("action") == "cancel"
):
# Cancel task (POST /videos/{id}/cancel or explicit action=cancel)
if (method == "POST" and path.endswith("/cancel")) or path_params.get("action") == "cancel":
if not task_id:
raise HTTPException(
status_code=400, detail="Task ID is required for cancel operation"
@@ -95,6 +93,16 @@ class VideoAdapterBase(ApiAdapter):
path_params=path_params,
)
# Delete task (DELETE /videos/{id})
if method == "DELETE" and task_id:
return await handler.handle_delete_task(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# Remix task
if method == "POST" and path.endswith("/remix") and task_id:
return await handler.handle_remix_task(

View File

@@ -26,7 +26,7 @@ from src.services.cache.aware_scheduler import ProviderCandidate
if TYPE_CHECKING:
import httpx
from src.services.task.orchestrator import SubmitOutcome
from src.services.candidate.submit import SubmitOutcome
# 敏感信息匹配正则(预编译提升性能)
_SENSITIVE_PATTERN = re.compile(
@@ -53,21 +53,39 @@ def sanitize_error_message(message: str, max_length: int = 200) -> str:
return sanitized[:max_length]
def extract_short_id_from_operation(operation_id: str) -> str:
"""
从 operation ID 中提取短 ID
我们对外暴露的 operation name 格式是:
- models/{model}/operations/{short_id}
此函数提取最后一部分作为 short_id用于在数据库中查找任务。
Args:
operation_id: 原始 operation ID"models/veo-3.1/operations/abc123"
Returns:
short_id"abc123"
"""
# 格式: models/{model}/operations/{short_id}
# 或者直接是 short_id
if "/" in operation_id:
# 提取最后一部分
return operation_id.rsplit("/", 1)[-1]
return operation_id
def normalize_gemini_operation_id(operation_id: str) -> str:
"""
规范化 Gemini operation ID,确保以 "operations/" 开头
Gemini API 返回的任务 ID 格式可能是 "operations/xxx""xxx"
此函数统一规范化为 "operations/xxx" 格式。
规范化 Gemini operation ID(保留用于向后兼容)
Args:
operation_id: 原始 operation ID
Returns:
规范化后的 operation ID
规范化后的 operation ID(原样返回)
"""
if not operation_id.startswith("operations/"):
return f"operations/{operation_id}"
return operation_id
@@ -143,6 +161,18 @@ class VideoHandlerBase(ABC):
) -> JSONResponse:
"""取消任务"""
async def handle_delete_task(
self,
*,
task_id: str,
http_request: Request,
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""删除已完成或失败的视频任务 - 可选实现"""
raise HTTPException(status_code=501, detail="Delete not supported for this provider")
async def handle_remix_task(
self,
*,
@@ -205,6 +235,7 @@ class VideoHandlerBase(ABC):
}
def _get_task(self, task_id: str) -> VideoTask:
"""通过 UUID 查找任务OpenAI Sora 风格)"""
task = (
self.db.query(VideoTask)
.filter(VideoTask.id == task_id, VideoTask.user_id == self.user.id)
@@ -229,7 +260,7 @@ class VideoHandlerBase(ABC):
except ValueError:
status = VideoStatus.PENDING
return InternalVideoTask(
id=task.id,
id=task.id, # OpenAI Sora 使用 UUID
external_id=task.external_task_id,
status=status,
progress_percent=task.progress_percent or 0,
@@ -243,6 +274,52 @@ class VideoHandlerBase(ABC):
extra={"model": task.model},
)
def _finalize_usage_on_submit_failure(
self,
candidate_keys: list[dict[str, Any]],
status_code: int | None,
) -> None:
"""
提交失败时结算 pending usage避免遗留 pending 状态)。
从 candidate_keys 中提取最后尝试的 provider 信息,更新 Usage 记录。
"""
from src.services.usage.service import UsageService
# 提取 provider 信息:优先取最后一个有 attempt 的候选
provider_name = "unknown"
provider_id = None
endpoint_id = None
key_id = None
for ck in reversed(candidate_keys):
if ck.get("attempt_status") or ck.get("selected"):
provider_name = ck.get("provider_name") or "unknown"
provider_id = ck.get("provider_id")
endpoint_id = ck.get("endpoint_id")
key_id = ck.get("key_id")
break
try:
# 更新 usage 状态并设置 provider 信息
UsageService.update_usage_status(
self.db,
request_id=self.request_id,
status="failed",
error_message=f"submit_failed (status_code={status_code or 'unknown'})",
provider=provider_name,
provider_id=provider_id,
provider_endpoint_id=endpoint_id,
provider_api_key_id=key_id,
status_code=status_code,
)
except Exception as exc:
logger.warning(
"Failed to finalize usage on submit failure: request_id=%s, error=%s",
self.request_id,
sanitize_error_message(str(exc)),
)
def _build_billing_rule_snapshot(
self, rule_lookup: BillingRuleLookupResult | None
) -> dict[str, Any]:
@@ -290,16 +367,16 @@ class VideoHandlerBase(ABC):
- 无可用候选 / 全部失败:抛 HTTPException(503)
"""
# 延迟导入,避免 handler 基类层引入过多依赖导致循环
from src.services.task.orchestrator import (
from src.services.candidate.service import CandidateService
from src.services.candidate.submit import (
AllCandidatesFailedError,
AsyncTaskOrchestrator,
SubmitOutcome,
UpstreamClientRequestError,
)
orchestrator = AsyncTaskOrchestrator(self.db)
candidate_service = CandidateService(self.db)
try:
return await orchestrator.submit_with_failover(
return await candidate_service.submit_with_failover(
api_format=api_format,
model_name=model_name,
affinity_key=str(self.api_key.id),
@@ -314,8 +391,12 @@ class VideoHandlerBase(ABC):
max_candidates=max_candidates,
)
except UpstreamClientRequestError as exc:
# 将 pending usage 结算为 failed并记录 provider 信息
self._finalize_usage_on_submit_failure(exc.candidate_keys, exc.response.status_code)
return self._build_error_response(exc.response)
except AllCandidatesFailedError as exc:
# 将 pending usage 结算为 failed
self._finalize_usage_on_submit_failure(exc.candidate_keys, exc.last_status_code)
detail = "No available provider for video generation"
if config.billing_require_rule:
detail = "No available provider with billing rule for video generation"

View File

@@ -18,7 +18,6 @@ from src.core.api_format import ApiFamily, get_auth_handler
from src.core.api_format.enums import AuthMethod
from src.core.logger import logger
from src.models.gemini import GeminiRequest
from src.services.gemini_files_mapping import extract_file_names_from_request
from src.services.provider.transport import redact_url_for_log
@@ -63,11 +62,9 @@ class GeminiChatAdapter(ChatAdapterBase):
def detect_capability_requirements(
self,
headers: dict[str, str], # noqa: ARG002 - 预留
request_body: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None, # noqa: ARG002 - 预留
) -> dict[str, bool]:
"""检测是否需要 Gemini Files API 能力"""
if request_body and extract_file_names_from_request(request_body):
return {"gemini_files_api": True}
"""Gemini API 无特殊能力要求"""
return {}
def _merge_path_params(

View File

@@ -40,28 +40,31 @@ class GeminiChatHandler(ChatHandlerBase):
Gemini 文件与上传它的 API Key 绑定,必须使用同一 Key 访问。
此方法从缓存中查找文件→Key 映射,优先使用正确的 Key。
当同一源文件被上传到多个 Key 时,会返回所有可用的 Key ID
让系统能够选择任意可用的 Key。
注意事项:
- 如果映射缺失(缓存过期/重启),会记录警告,请求可能失败
- 如果多个文件属于不同 Key只能使用其中一个其他文件可能无法访问
- 优先返回所有支持该文件的 Key让调度器选择可用的
"""
from src.core.logger import logger
from src.services.gemini_files_mapping import (
extract_file_names_from_request,
get_file_key_mapping,
get_all_key_ids_for_file,
)
file_names = extract_file_names_from_request(request_body or {})
if not file_names:
return None
preferred_key_ids: list[str] = []
unmapped_files: list[str] = [] # 记录找不到映射的文件
all_key_ids: set[str] = set()
unmapped_files: list[str] = []
for file_name in file_names:
key_id = await get_file_key_mapping(file_name)
if key_id:
if key_id not in preferred_key_ids:
preferred_key_ids.append(key_id)
# 获取所有支持该文件的 Key包括通过 source_hash 关联的)
key_ids = await get_all_key_ids_for_file(file_name)
if key_ids:
all_key_ids.update(key_ids)
else:
unmapped_files.append(file_name)
@@ -72,14 +75,10 @@ class GeminiChatHandler(ChatHandlerBase):
"请求可能失败(文件属于其他 Key 或映射已过期)"
)
# 警告:多个文件属于不同 Key
if len(preferred_key_ids) > 1:
logger.warning(
f"[{self.request_id}] 请求使用了多个文件,但它们属于不同的 Key: "
f"{preferred_key_ids},只能使用第一个 Key其他文件可能无法访问"
)
if all_key_ids:
logger.debug(f"[{self.request_id}] 文件引用可用的 Key: {list(all_key_ids)}")
return preferred_key_ids or None
return list(all_key_ids) if all_key_ids else None
def extract_model_from_request(
self,

View File

@@ -4,6 +4,7 @@ Gemini Video Handler - Veo 视频生成实现
from __future__ import annotations
import time
from datetime import datetime, timedelta, timezone
from typing import Any, AsyncIterator
from uuid import uuid4
@@ -34,12 +35,14 @@ from src.core.api_format.conversion.internal_video import (
VideoStatus,
)
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
from src.core.api_format.conversion.registry import format_conversion_registry
from src.core.api_format.headers import HOP_BY_HOP_HEADERS
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.usage.service import UsageService
class GeminiVeoHandler(VideoHandlerBase):
@@ -92,20 +95,93 @@ class GeminiVeoHandler(VideoHandlerBase):
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
# 异步任务:提前创建 pending usage便于前端看到“处理中”
try:
UsageService.create_pending_usage(
db=self.db,
request_id=self.request_id,
user=self.user,
api_key=self.api_key,
model=internal_request.model,
is_stream=False,
request_type="video",
api_format=self.FORMAT_ID,
request_headers=original_headers,
request_body=original_request_body,
)
except Exception as exc:
logger.warning(
"Failed to create pending usage for video request_id=%s: %s",
self.request_id,
sanitize_error_message(str(exc)),
)
# 用于跟踪是否发生了格式转换
format_conversion_info: dict[str, Any] = {
"converted": False,
"provider_format": None,
}
async def _submit(candidate: ProviderCandidate) -> Any:
upstream_key, endpoint, _key, auth_info = await self._resolve_upstream_key(candidate)
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
headers = self._build_upstream_headers(
original_headers, upstream_key, endpoint, auth_info
# 检测目标格式
provider_format = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
client = await HTTPClientPool.get_default_client_async()
return await client.post(upstream_url, headers=headers, json=original_request_body)
needs_conversion = provider_format.upper() != self.FORMAT_ID.upper()
format_conversion_info["provider_format"] = provider_format
format_conversion_info["converted"] = needs_conversion
if needs_conversion and provider_format.upper().startswith("OPENAI:"):
# Gemini -> OpenAI 格式转换
converted_body = format_conversion_registry.convert_video_request(
original_request_body,
self.FORMAT_ID,
provider_format,
)
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string
if "seconds" in converted_body and converted_body["seconds"] is not None:
converted_body["seconds"] = str(converted_body["seconds"])
# 构建 OpenAI 风格的 URL
upstream_url = self._build_openai_upstream_url(endpoint.base_url)
# 构建 OpenAI 风格的请求头
headers = self._build_openai_upstream_headers(
original_headers, upstream_key, endpoint
)
client = await HTTPClientPool.get_default_client_async()
return await client.post(upstream_url, headers=headers, json=converted_body)
else:
# 原始 Gemini 格式
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
headers = self._build_upstream_headers(
original_headers, upstream_key, endpoint, auth_info
)
client = await HTTPClientPool.get_default_client_async()
return await client.post(upstream_url, headers=headers, json=original_request_body)
def _extract_task_id(payload: dict[str, Any]) -> str | None:
value = payload.get("name")
if not value:
return None
return normalize_gemini_operation_id(str(value))
# 根据响应格式提取 task ID
# Gemini: {"name": "operations/..."}
# OpenAI: {"id": "..."}
if "name" in payload:
value = payload.get("name")
logger.debug(
"[GeminiVeoHandler] Upstream response name=%s, keys=%s",
value,
list(payload.keys()) if isinstance(payload, dict) else type(payload),
)
if not value:
return None
return normalize_gemini_operation_id(str(value))
if "id" in payload:
# OpenAI 格式
return str(payload["id"])
return None
outcome_or_response = await self._submit_with_failover(
api_format=self.FORMAT_ID,
@@ -114,7 +190,7 @@ class GeminiVeoHandler(VideoHandlerBase):
submit_func=_submit,
extract_external_task_id=_extract_task_id,
supported_auth_types={"api_key", "vertex_ai"},
allow_format_conversion=False,
allow_format_conversion=True,
max_candidates=10,
)
if isinstance(outcome_or_response, JSONResponse):
@@ -135,35 +211,92 @@ class GeminiVeoHandler(VideoHandlerBase):
external_task_id = outcome.external_task_id
# 如果发生了格式转换,记录转换后的请求体
converted_request_body = original_request_body
if format_conversion_info["converted"]:
try:
converted_request_body = format_conversion_registry.convert_video_request(
original_request_body,
self.FORMAT_ID,
format_conversion_info["provider_format"],
)
except Exception as e:
logger.warning(
"[GeminiVeoHandler] Failed to record converted request: %s",
sanitize_error_message(str(e)),
)
task = self._create_task_record(
external_task_id=external_task_id,
candidate=outcome.candidate,
original_request_body=original_request_body,
converted_request_body=converted_request_body,
internal_request=internal_request,
candidate_keys=outcome.candidate_keys,
original_headers=original_headers,
billing_rule_snapshot=billing_rule_snapshot,
format_converted=format_conversion_info["converted"],
)
try:
self.db.add(task)
self.db.flush() # 先 flush 检测冲突
self.db.commit()
self.db.refresh(task)
logger.info(
f"[GeminiVeoHandler] Task created: id={task.id}, external_task_id={task.external_task_id}, user_id={task.user_id}"
logger.debug(
"[GeminiVeoHandler] Task created: id=%s, external_task_id=%s",
task.id,
task.external_task_id,
)
except IntegrityError:
self.db.rollback()
raise HTTPException(status_code=409, detail="Task already exists")
# 先构建返回给客户端的响应(使用短 ID 对外暴露)
internal_task = InternalVideoTask(
id=task.id,
id=task.short_id,
external_id=external_task_id,
status=VideoStatus.SUBMITTED,
created_at=task.created_at,
original_request=internal_request,
)
response_body = self._normalizer.video_task_from_internal(internal_task)
# 提交成功后立即结算 Usage费用暂时为 0轮询完成后更新
response_time_ms = int((time.time() - self.start_time) * 1000)
try:
# 构建发送给上游的请求头(脱敏)
upstream_request_headers = self._build_upstream_headers(
original_headers,
"", # key 不重要,只是用于记录
outcome.candidate.endpoint,
None, # auth_info
)
UsageService.finalize_submitted(
self.db,
request_id=self.request_id,
provider_name=outcome.candidate.provider.name,
provider_id=outcome.candidate.provider.id,
provider_endpoint_id=outcome.candidate.endpoint.id,
provider_api_key_id=outcome.candidate.key.id,
response_time_ms=response_time_ms,
status_code=outcome.upstream_status_code or 200,
endpoint_api_format=make_signature_key(
str(getattr(outcome.candidate.endpoint, "api_family", "")).strip().lower(),
str(getattr(outcome.candidate.endpoint, "endpoint_kind", "")).strip().lower(),
),
provider_request_headers=upstream_request_headers,
response_headers=outcome.upstream_headers,
response_body=response_body, # 使用我们转换后的响应(包含我们的 ID
)
self.db.commit()
except Exception as exc:
logger.warning(
"Failed to finalize submitted usage for video request_id=%s: %s",
self.request_id,
sanitize_error_message(str(exc)),
)
return JSONResponse(response_body)
async def handle_get_task(
@@ -234,6 +367,28 @@ class GeminiVeoHandler(VideoHandlerBase):
task.status = VideoStatus.CANCELLED.value
task.updated_at = datetime.now(timezone.utc)
# 将 Usage 作废(不收费)
# 尝试 finalize_void处理 pending和 void_settled处理已 settled
try:
voided = UsageService.finalize_void(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
)
if not voided:
# pending 状态未找到,尝试处理已 settled 的记录
UsageService.void_settled(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
)
except Exception as exc:
logger.warning(
"Failed to void usage for cancelled task=%s: %s",
task.id,
sanitize_error_message(str(exc)),
)
self.db.commit()
return JSONResponse({})
@@ -275,12 +430,38 @@ class GeminiVeoHandler(VideoHandlerBase):
if task.video_expires_at < now:
raise HTTPException(status_code=410, detail="Video URL has expired")
# 获取 provider 的认证信息Gemini 下载视频需要带 API Key
endpoint, key = self._get_endpoint_and_key(task)
download_headers: dict[str, str] = {}
if key.api_key:
try:
upstream_key = crypto_service.decrypt(key.api_key)
# Gemini API 使用 x-goog-api-key 头进行认证
download_headers["x-goog-api-key"] = upstream_key
# 如果是 Vertex AI需要使用 OAuth Bearer token
auth_info = await get_provider_auth(endpoint, key)
if auth_info:
download_headers.pop("x-goog-api-key", None)
download_headers[auth_info.auth_header] = auth_info.auth_value
except Exception as exc:
logger.warning(
"[VideoDownload] Failed to get auth for download task=%s: %s",
task.id,
sanitize_error_message(str(exc)),
)
# 继续尝试无认证下载(某些 URL 可能是预签名的)
# 代理下载而非直接重定向,避免暴露上游存储 URL
client = await HTTPClientPool.get_default_client_async()
# 使用 httpx 支持重定向Gemini 视频 URL 会重定向到实际存储位置)
import httpx
try:
request = client.build_request("GET", task.video_url)
# 视频下载可能较大,设置 5 分钟超时
response = await client.send(request, stream=True, timeout=300.0)
# 使用 follow_redirects=True 跟随重定向
async with httpx.AsyncClient(
follow_redirects=True, timeout=httpx.Timeout(300.0)
) as client:
response = await client.get(task.video_url, headers=download_headers)
except Exception as exc:
logger.error(
"[VideoDownload] Upstream fetch failed user=%s task=%s: %s",
@@ -291,21 +472,14 @@ class GeminiVeoHandler(VideoHandlerBase):
raise HTTPException(status_code=502, detail="Failed to fetch video")
if response.status_code >= 400:
await response.aclose()
raise HTTPException(status_code=response.status_code, detail="Upstream error")
async def _iter_bytes() -> AsyncIterator[bytes]:
try:
async for chunk in response.aiter_bytes():
yield chunk
finally:
await response.aclose()
# 返回完整的视频内容(非 streaming因为需要跟随重定向
safe_headers = {
k: v for k, v in response.headers.items() if k.lower() not in HOP_BY_HOP_HEADERS
}
return StreamingResponse(
_iter_bytes(),
return Response(
content=response.content,
status_code=response.status_code,
headers=safe_headers,
media_type=response.headers.get("content-type", "video/mp4"),
@@ -375,6 +549,36 @@ class GeminiVeoHandler(VideoHandlerBase):
"status": error.get("status", "BAD_GATEWAY"),
}
# ------------------------------------------------------------------
# OpenAI format conversion helpers
# ------------------------------------------------------------------
def _build_openai_upstream_url(self, base_url: str | None) -> str:
"""构建 OpenAI Sora API 的上游 URL"""
base = (base_url or "https://api.openai.com").rstrip("/")
if base.endswith("/v1"):
return f"{base}/videos"
return f"{base}/v1/videos"
def _build_openai_upstream_headers(
self,
original_headers: dict[str, str],
upstream_key: str,
endpoint: ProviderEndpoint,
) -> dict[str, str]:
"""构建 OpenAI 格式的请求头"""
extra_headers = get_extra_headers_from_endpoint(endpoint)
endpoint_sig = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
return build_upstream_headers_for_endpoint(
original_headers,
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
)
def _create_task_record(
self,
*,
@@ -385,6 +589,8 @@ class GeminiVeoHandler(VideoHandlerBase):
candidate_keys: list[dict[str, Any]] | None = None,
original_headers: dict[str, str] | None = None,
billing_rule_snapshot: dict[str, Any] | None = None,
converted_request_body: dict[str, Any] | None = None,
format_converted: bool = False,
) -> VideoTask:
now = datetime.now(timezone.utc)
@@ -407,8 +613,14 @@ class GeminiVeoHandler(VideoHandlerBase):
}
request_metadata["request_headers"] = safe_headers
provider_api_format = make_signature_key(
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
)
return VideoTask(
id=str(uuid4()),
request_id=self.request_id,
external_task_id=external_task_id,
user_id=self.user.id,
api_key_id=self.api_key.id,
@@ -416,15 +628,12 @@ class GeminiVeoHandler(VideoHandlerBase):
endpoint_id=candidate.endpoint.id,
key_id=candidate.key.id,
client_api_format=self.FORMAT_ID,
provider_api_format=make_signature_key(
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
),
format_converted=False,
provider_api_format=provider_api_format,
format_converted=format_converted,
model=internal_request.model,
prompt=internal_request.prompt,
original_request_body=original_request_body,
converted_request_body=original_request_body,
converted_request_body=converted_request_body or original_request_body,
duration_seconds=internal_request.duration_seconds,
resolution=internal_request.resolution,
aspect_ratio=internal_request.aspect_ratio,
@@ -438,30 +647,45 @@ class GeminiVeoHandler(VideoHandlerBase):
request_metadata=request_metadata,
)
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
"""按 external_task_id 查找任务Gemini 使用 operations/{id} 格式)"""
normalized_id = normalize_gemini_operation_id(external_id)
logger.info(
f"[GeminiVeoHandler] Looking for task: normalized_id={normalized_id}, user_id={self.user.id}"
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:
"""覆盖父类方法Gemini 使用 short_id 作为对外暴露的 ID"""
try:
status = VideoStatus(task.status)
except ValueError:
status = VideoStatus.PENDING
return InternalVideoTask(
id=task.short_id, # Gemini 使用短 ID
external_id=task.external_task_id,
status=status,
progress_percent=task.progress_percent or 0,
progress_message=task.progress_message,
video_url=task.video_url,
video_urls=task.video_urls or [],
created_at=task.created_at,
completed_at=task.completed_at,
error_code=task.error_code,
error_message=task.error_message,
extra={"model": task.model},
)
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
"""按 short_id 查找任务(我们对外暴露的 operation 格式是 models/{model}/operations/{short_id}"""
from src.api.handlers.base.video_handler_base import extract_short_id_from_operation
short_id = extract_short_id_from_operation(external_id)
# 通过 short_id 查找任务
task = (
self.db.query(VideoTask)
.filter(
VideoTask.external_task_id == normalized_id,
VideoTask.short_id == short_id,
VideoTask.user_id == self.user.id,
)
.first()
)
if not task:
logger.warning(
f"[GeminiVeoHandler] Task not found: normalized_id={normalized_id}, user_id={self.user.id}"
)
logger.debug("[GeminiVeoHandler] Task not found: short_id=%s", short_id)
raise HTTPException(status_code=404, detail="Video task not found")
logger.info(
f"[GeminiVeoHandler] Task found: id={task.id}, external_task_id={task.external_task_id}"
)
return task

View File

@@ -5,6 +5,7 @@ OpenAI Video Handler - Sora 视频生成实现
from __future__ import annotations
import json
import time
from datetime import datetime, timedelta, timezone
from typing import Any, AsyncIterator
from uuid import uuid4
@@ -14,6 +15,7 @@ from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import get_provider_auth
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
from src.clients.http_client import HTTPClientPool
from src.config.settings import config
@@ -30,6 +32,7 @@ from src.core.api_format.conversion.internal_video import (
VideoStatus,
)
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
from src.core.api_format.conversion.registry import format_conversion_registry
from src.core.api_format.headers import HOP_BY_HOP_HEADERS
from src.core.crypto import crypto_service
from src.core.logger import logger
@@ -91,16 +94,91 @@ class OpenAIVideoHandler(VideoHandlerBase):
)
raise HTTPException(status_code=400, detail=str(e))
# 异步任务:提前创建 pending usage便于前端看到“处理中”
try:
UsageService.create_pending_usage(
db=self.db,
request_id=self.request_id,
user=self.user,
api_key=self.api_key,
model=internal_request.model,
is_stream=False,
request_type="video",
api_format=self.FORMAT_ID,
request_headers=original_headers,
request_body=original_request_body,
)
except Exception as exc:
logger.warning(
"Failed to create pending usage for video request_id=%s: %s",
self.request_id,
sanitize_error_message(str(exc)),
)
# 用于跟踪是否发生了格式转换
format_conversion_info: dict[str, Any] = {
"converted": False,
"provider_format": None,
}
async def _submit(candidate: ProviderCandidate) -> Any:
upstream_key, endpoint, _provider_key = await self._resolve_upstream_key(candidate)
upstream_url = self._build_upstream_url(endpoint.base_url)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
return await client.post(upstream_url, headers=headers, json=original_request_body)
# 检测目标格式
provider_format = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
needs_conversion = provider_format.upper() != self.FORMAT_ID.upper()
format_conversion_info["provider_format"] = provider_format
format_conversion_info["converted"] = needs_conversion
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string
request_body = original_request_body.copy()
if "seconds" in request_body and request_body["seconds"] is not None:
request_body["seconds"] = str(request_body["seconds"])
if needs_conversion and provider_format.upper().startswith("GEMINI:"):
# OpenAI -> Gemini 格式转换
converted_body = format_conversion_registry.convert_video_request(
request_body,
self.FORMAT_ID,
provider_format,
)
# 如果 model 不在请求体中,从路径或内部请求中获取
if "model" not in converted_body:
converted_body["model"] = internal_request.model
# 构建 Gemini 风格的 URL
upstream_url = self._build_gemini_upstream_url(
endpoint.base_url, internal_request.model
)
# 构建 Gemini 风格的请求头
auth_info = await get_provider_auth(endpoint, _provider_key)
headers = self._build_gemini_upstream_headers(
original_headers, upstream_key, endpoint, auth_info
)
client = await HTTPClientPool.get_default_client_async()
return await client.post(upstream_url, headers=headers, json=converted_body)
else:
# 原始 OpenAI 格式
upstream_url = self._build_upstream_url(endpoint.base_url)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
return await client.post(upstream_url, headers=headers, json=request_body)
def _extract_task_id(payload: dict[str, Any]) -> str | None:
value = payload.get("id")
return str(value) if value else None
# 根据响应格式提取 task ID
# OpenAI: {"id": "..."}
# Gemini: {"name": "operations/..."}
if "id" in payload:
return str(payload["id"])
if "name" in payload:
# Gemini 格式
return str(payload["name"])
return None
# 捕获提交阶段的所有错误,记录失败任务
try:
@@ -110,8 +188,8 @@ class OpenAIVideoHandler(VideoHandlerBase):
task_type="video",
submit_func=_submit,
extract_external_task_id=_extract_task_id,
supported_auth_types={"api_key"},
allow_format_conversion=False,
supported_auth_types={"api_key", "vertex_ai"},
allow_format_conversion=True,
max_candidates=10,
)
except HTTPException as exc:
@@ -156,14 +234,31 @@ class OpenAIVideoHandler(VideoHandlerBase):
external_task_id = outcome.external_task_id
# 如果发生了格式转换,记录转换后的请求体
converted_request_body = original_request_body
if format_conversion_info["converted"]:
try:
converted_request_body = format_conversion_registry.convert_video_request(
original_request_body,
self.FORMAT_ID,
format_conversion_info["provider_format"],
)
except Exception as e:
logger.warning(
"[OpenAIVideoHandler] Failed to record converted request: %s",
sanitize_error_message(str(e)),
)
task = self._create_task_record(
external_task_id=external_task_id,
candidate=outcome.candidate,
original_request_body=original_request_body,
converted_request_body=converted_request_body,
internal_request=internal_request,
candidate_keys=outcome.candidate_keys,
original_headers=original_headers,
billing_rule_snapshot=billing_rule_snapshot,
format_converted=format_conversion_info["converted"],
)
try:
self.db.add(task)
@@ -174,6 +269,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
self.db.rollback()
raise HTTPException(status_code=409, detail="Task already exists")
# 先构建返回给客户端的响应OpenAI Sora 使用 UUID
internal_task = InternalVideoTask(
id=task.id,
external_id=external_task_id,
@@ -182,6 +278,42 @@ class OpenAIVideoHandler(VideoHandlerBase):
original_request=internal_request,
)
response_body = self._normalizer.video_task_from_internal(internal_task)
# 提交成功后立即结算 Usage费用暂时为 0轮询完成后更新
response_time_ms = int((time.time() - self.start_time) * 1000)
try:
# 构建发送给上游的请求头(脱敏)
upstream_request_headers = self._build_upstream_headers(
original_headers,
"", # key 不重要,只是用于记录
outcome.candidate.endpoint,
)
UsageService.finalize_submitted(
self.db,
request_id=self.request_id,
provider_name=outcome.candidate.provider.name,
provider_id=outcome.candidate.provider.id,
provider_endpoint_id=outcome.candidate.endpoint.id,
provider_api_key_id=outcome.candidate.key.id,
response_time_ms=response_time_ms,
status_code=outcome.upstream_status_code or 200,
endpoint_api_format=make_signature_key(
str(getattr(outcome.candidate.endpoint, "api_family", "")).strip().lower(),
str(getattr(outcome.candidate.endpoint, "endpoint_kind", "")).strip().lower(),
),
provider_request_headers=upstream_request_headers,
response_headers=outcome.upstream_headers,
response_body=response_body, # 使用我们转换后的响应(包含我们的 ID
)
self.db.commit()
except Exception as exc:
logger.warning(
"Failed to finalize submitted usage for video request_id=%s: %s",
self.request_id,
sanitize_error_message(str(exc)),
)
return JSONResponse(response_body)
async def handle_get_task(
@@ -206,17 +338,59 @@ class OpenAIVideoHandler(VideoHandlerBase):
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
tasks = (
self.db.query(VideoTask)
.filter(VideoTask.user_id == self.user.id)
.order_by(VideoTask.created_at.desc())
.limit(100)
.all()
)
params = query_params or {}
# 解析分页参数
after = params.get("after")
try:
limit = min(int(params.get("limit") or 20), 100) # 默认 20最大 100
except (ValueError, TypeError):
limit = 20
order = params.get("order", "desc").lower()
if order not in ("asc", "desc"):
order = "desc"
# 构建查询
query = self.db.query(VideoTask).filter(VideoTask.user_id == self.user.id)
# 处理游标分页after 参数使用 UUID
if after:
after_task = (
self.db.query(VideoTask)
.filter(VideoTask.id == after, VideoTask.user_id == self.user.id)
.first()
)
if after_task and after_task.created_at:
if order == "desc":
query = query.filter(VideoTask.created_at < after_task.created_at)
else:
query = query.filter(VideoTask.created_at > after_task.created_at)
# 排序
if order == "asc":
query = query.order_by(VideoTask.created_at.asc())
else:
query = query.order_by(VideoTask.created_at.desc())
# 获取 limit + 1 条记录以判断是否有更多数据
tasks = query.limit(limit + 1).all()
has_more = len(tasks) > limit
tasks = tasks[:limit]
items = [
self._normalizer.video_task_from_internal(self._task_to_internal(t)) for t in tasks
]
return JSONResponse({"object": "list", "data": items})
response_data: dict[str, Any] = {
"object": "list",
"data": items,
"has_more": has_more,
}
# 如果有更多数据,返回最后一条的 ID 作为下一页游标
if has_more and tasks:
response_data["last_id"] = tasks[-1].id
return JSONResponse(response_data)
async def handle_cancel_task(
self,
@@ -245,9 +419,80 @@ class OpenAIVideoHandler(VideoHandlerBase):
task.status = VideoStatus.CANCELLED.value
task.updated_at = datetime.now(timezone.utc)
# 将 Usage 作废(不收费)
# 尝试 finalize_void处理 pending和 void_settled处理已 settled
try:
voided = UsageService.finalize_void(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
)
if not voided:
# pending 状态未找到,尝试处理已 settled 的记录
UsageService.void_settled(
self.db,
request_id=task.request_id,
reason="cancelled_by_user",
)
except Exception as exc:
logger.warning(
"Failed to void usage for cancelled task=%s: %s",
task.id,
sanitize_error_message(str(exc)),
)
self.db.commit()
return JSONResponse({})
async def handle_delete_task(
self,
*,
task_id: str,
http_request: Request,
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""删除已完成或失败的视频及其存储资源"""
task = self._get_task(task_id)
# 只能删除已完成或失败的视频
if task.status not in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
raise HTTPException(
status_code=400,
detail=f"Can only delete completed or failed videos (current status: {task.status})",
)
# 如果有 external_task_id向上游发送删除请求
if task.external_task_id:
try:
endpoint, key = self._get_endpoint_and_key(task)
if key.api_key:
upstream_key = crypto_service.decrypt(key.api_key)
upstream_url = self._build_upstream_url(
endpoint.base_url, task.external_task_id
)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
response = await client.delete(upstream_url, headers=headers)
if response.status_code >= 400 and response.status_code != 404:
# 404 表示上游已删除,不算错误
return self._build_error_response(response)
except Exception as exc:
logger.warning(
"Failed to delete video from upstream task=%s: %s",
task.id,
sanitize_error_message(str(exc)),
)
# 继续删除本地记录
# 删除本地任务记录
self.db.delete(task)
self.db.commit()
return JSONResponse({"id": task_id, "object": "video", "deleted": True})
async def handle_remix_task(
self,
*,
@@ -280,8 +525,13 @@ class OpenAIVideoHandler(VideoHandlerBase):
)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string
request_body = original_request_body.copy()
if "seconds" in request_body and request_body["seconds"] is not None:
request_body["seconds"] = str(request_body["seconds"])
client = await HTTPClientPool.get_default_client_async()
response = await client.post(upstream_url, headers=headers, json=original_request_body)
response = await client.post(upstream_url, headers=headers, json=request_body)
if response.status_code >= 400:
return self._build_error_response(response)
@@ -350,7 +600,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
raise HTTPException(status_code=409, detail="Task already exists")
internal_task = InternalVideoTask(
id=task.id,
id=task.id, # OpenAI Sora 使用 UUID
external_id=external_task_id,
status=VideoStatus.SUBMITTED,
created_at=task.created_at,
@@ -389,15 +639,25 @@ class OpenAIVideoHandler(VideoHandlerBase):
if task.status == VideoStatus.CANCELLED.value:
raise HTTPException(status_code=404, detail="Video task was cancelled")
# 支持 variant 查询参数: video (默认), thumbnail, spritesheet
variant = (query_params or {}).get("variant", "video")
# 如果 video_url 是完整的 HTTP URL直接代理该 URL适用于不支持 /content 端点的上游如 API易
# 保持流式代理而非重定向,确保客户端行为与官方 OpenAI 一致
if variant == "video" and task.video_url and task.video_url.startswith("http"):
logger.debug(
"[VideoDownload] Proxying direct URL task=%s url=%s",
task_id,
task.video_url,
)
return await self._proxy_direct_url(task.video_url, task_id)
if not task.external_task_id:
raise HTTPException(status_code=500, detail="Task missing external_task_id")
endpoint, key = self._get_endpoint_and_key(task)
if not key.api_key:
raise HTTPException(status_code=500, detail="Provider key not configured")
upstream_key = crypto_service.decrypt(key.api_key)
# 支持 variant 查询参数: video (默认), thumbnail, spritesheet
variant = (query_params or {}).get("variant", "video")
if variant not in {"video", "thumbnail", "spritesheet"}:
raise HTTPException(
status_code=400,
@@ -412,6 +672,12 @@ class OpenAIVideoHandler(VideoHandlerBase):
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
logger.debug(
"[VideoDownload] Requesting upstream url=%s task=%s external_task_id=%s",
upstream_url,
task_id,
task.external_task_id,
)
try:
# 使用 httpx 的 stream 方法并正确管理上下文
# 视频下载可能较大,设置 5 分钟超时
@@ -419,8 +685,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
response = await client.send(request, stream=True, timeout=300.0)
except Exception as exc:
logger.warning(
"[VideoDownload] Upstream connection failed task=%s: %s",
"[VideoDownload] Upstream connection failed task=%s url=%s: %s",
task_id,
upstream_url,
sanitize_error_message(str(exc)),
)
raise HTTPException(status_code=502, detail="Upstream connection failed") from exc
@@ -483,6 +750,46 @@ class OpenAIVideoHandler(VideoHandlerBase):
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
return upstream_key, candidate.endpoint, candidate.key
async def _proxy_direct_url(self, url: str, task_id: str) -> Response | StreamingResponse:
"""代理直接的视频 URL如 CDN URL保持与官方 API 一致的流式返回行为"""
client = await HTTPClientPool.get_default_client_async()
try:
request = client.build_request("GET", url)
response = await client.send(request, stream=True, timeout=300.0)
except Exception as exc:
logger.warning(
"[VideoDownload] Direct URL connection failed task=%s url=%s: %s",
task_id,
url,
sanitize_error_message(str(exc)),
)
raise HTTPException(status_code=502, detail="Video download failed") from exc
if response.status_code >= 400:
await response.aread() # consume body before closing
await response.aclose()
return JSONResponse(
status_code=response.status_code,
content={"error": {"type": "upstream_error", "message": "Video not available"}},
)
async def _iter_bytes() -> AsyncIterator[bytes]:
try:
async for chunk in response.aiter_bytes():
yield chunk
finally:
await response.aclose()
safe_headers = {
k: v for k, v in response.headers.items() if k.lower() not in HOP_BY_HOP_HEADERS
}
return StreamingResponse(
_iter_bytes(),
status_code=response.status_code,
headers=safe_headers,
media_type=response.headers.get("content-type", "video/mp4"),
)
def _build_upstream_url(self, base_url: str | None, suffix: str | None = None) -> str:
base = (base_url or self.DEFAULT_BASE_URL).rstrip("/")
if base.endswith("/v1"):
@@ -508,6 +815,42 @@ class OpenAIVideoHandler(VideoHandlerBase):
endpoint_headers=extra_headers,
)
# ------------------------------------------------------------------
# Gemini format conversion helpers
# ------------------------------------------------------------------
def _build_gemini_upstream_url(self, base_url: str | None, model: str) -> str:
"""构建 Gemini Veo API 的上游 URL"""
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
return f"{base}/v1beta/models/{model}:predictLongRunning"
def _build_gemini_upstream_headers(
self,
original_headers: dict[str, str],
upstream_key: str,
endpoint: ProviderEndpoint,
auth_info: Any | None,
) -> dict[str, str]:
"""构建 Gemini 格式的请求头"""
extra_headers = get_extra_headers_from_endpoint(endpoint)
endpoint_sig = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
headers = build_upstream_headers_for_endpoint(
original_headers,
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
)
if auth_info:
# 覆盖为 OAuth2 BearerVertex AI
headers.pop("x-goog-api-key", None)
headers[auth_info.auth_header] = auth_info.auth_value
return headers
# _build_error_response 继承自基类 VideoHandlerBase
def _create_task_record(
@@ -520,6 +863,8 @@ class OpenAIVideoHandler(VideoHandlerBase):
candidate_keys: list[dict[str, Any]] | None = None,
original_headers: dict[str, str] | None = None,
billing_rule_snapshot: dict[str, Any] | None = None,
converted_request_body: dict[str, Any] | None = None,
format_converted: bool = False,
) -> VideoTask:
now = datetime.now(timezone.utc)
size = internal_request.extra.get("original_size")
@@ -543,8 +888,14 @@ class OpenAIVideoHandler(VideoHandlerBase):
}
request_metadata["request_headers"] = safe_headers
provider_api_format = make_signature_key(
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
)
return VideoTask(
id=str(uuid4()),
request_id=self.request_id,
external_task_id=external_task_id,
user_id=self.user.id,
api_key_id=self.api_key.id,
@@ -552,15 +903,12 @@ class OpenAIVideoHandler(VideoHandlerBase):
endpoint_id=candidate.endpoint.id,
key_id=candidate.key.id,
client_api_format=self.FORMAT_ID,
provider_api_format=make_signature_key(
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
),
format_converted=False,
provider_api_format=provider_api_format,
format_converted=format_converted,
model=internal_request.model,
prompt=internal_request.prompt,
original_request_body=original_request_body,
converted_request_body=original_request_body,
converted_request_body=converted_request_body or original_request_body,
duration_seconds=internal_request.duration_seconds,
resolution=internal_request.resolution,
aspect_ratio=internal_request.aspect_ratio,
@@ -580,8 +928,23 @@ class OpenAIVideoHandler(VideoHandlerBase):
status = VideoStatus(task.status)
except ValueError:
status = VideoStatus.PENDING
# 构建 extra 字段
extra: dict[str, Any] = {
"model": task.model,
"size": task.size,
"seconds": str(task.duration_seconds) if task.duration_seconds else None,
"prompt": task.prompt,
}
# 检查是否是 remix 视频
if task.original_request_body and isinstance(task.original_request_body, dict):
remixed_from = task.original_request_body.get("remix_video_id")
if remixed_from:
extra["remixed_from_video_id"] = remixed_from
return InternalVideoTask(
id=task.id,
id=task.id, # OpenAI Sora 使用 UUID
external_id=task.external_task_id,
status=status,
progress_percent=task.progress_percent or 0,
@@ -596,7 +959,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
expires_at=task.video_expires_at,
error_code=task.error_code,
error_message=task.error_message,
extra={"model": task.model, "size": task.size, "seconds": task.duration_seconds},
extra=extra,
)
async def _record_failed_usage(
@@ -609,8 +972,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
original_headers: dict[str, str],
) -> None:
"""记录失败请求的使用记录(无任务记录)"""
import time
response_time_ms = int((time.time() - self.start_time) * 1000)
safe_headers = {
k: v
@@ -669,8 +1030,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
candidate_keys: list[dict[str, Any]] | None = None,
) -> None:
"""创建失败的任务记录和使用记录"""
import time
now = datetime.now(timezone.utc)
response_time_ms = int((time.time() - self.start_time) * 1000)
@@ -696,6 +1055,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
# 创建失败的任务记录
task = VideoTask(
id=str(uuid4()),
request_id=self.request_id,
external_task_id=None,
user_id=self.user.id,
api_key_id=self.api_key.id,

View File

@@ -14,13 +14,14 @@ from .system_catalog import router as system_catalog_router
from .videos import router as videos_router
router = APIRouter()
# Models API 需要在最前面注册,避免被其他路由的 path 参数捕获
# Video API 路由需要在 Models API 之前注册,因为 Models API 有 /v1beta/models/{path} 通配符路由
# 会错误匹配 /v1beta/models/{model}/operations/{id}/content 等视频路由
router.include_router(videos_router, tags=["Video Generation"])
router.include_router(models_router)
router.include_router(claude_router, tags=["Claude API"])
router.include_router(openai_router)
router.include_router(gemini_router, tags=["Gemini API"])
router.include_router(gemini_files_router, tags=["Gemini Files API"])
router.include_router(videos_router, tags=["Video Generation"])
router.include_router(system_catalog_router, tags=["System Catalog"])
router.include_router(catalog_router)
router.include_router(capabilities_router)

View File

@@ -15,14 +15,17 @@ Gemini Files API 代理端点
参考文档:
https://ai.google.dev/api/files
优化HTTP 代理请求期间不持有数据库连接,避免阻塞其他请求。
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Dict, Optional, Tuple
from urllib.parse import urlencode
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi import APIRouter, HTTPException, Request, Response
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
@@ -30,20 +33,30 @@ from src.clients.http_client import HTTPClientPool
from src.core.api_format import get_auth_handler, get_default_auth_method_for_endpoint
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.database import get_db
from src.database import create_session
from src.models.database import ApiKey, GlobalModel, Model, Provider, ProviderEndpoint, User
from src.services.auth.service import AuthService
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
from src.services.gemini_files_mapping import delete_file_key_mapping, store_file_key_mapping
from src.services.provider.transport import redact_url_for_log
@dataclass
class UpstreamContext:
"""上游请求上下文(不依赖数据库会话)"""
upstream_key: str
base_url: str
file_key_id: str
user_id: str
router = APIRouter(tags=["Gemini Files API"])
# Gemini Files API 基础 URL
GEMINI_FILES_BASE_URL = "https://generativelanguage.googleapis.com"
# Gemini Files API 能力标签
REQUIRED_CAPABILITIES = {"gemini_files_api": True}
# Gemini Files API 能力限制(任何 Gemini key 都可用)
# 需要从客户端请求中移除的头部(这些会由代理重新设置或不应转发)
HEADERS_TO_REMOVE = frozenset(
@@ -184,9 +197,25 @@ async def _select_provider_candidate(
db: Session,
user_api_key: ApiKey,
model_name: str,
require_files_capability: bool = True,
) -> ProviderCandidate | None:
"""选择支持 Files API 的 Provider/Endpoint/Key 组合"""
"""
选择可用的 Provider/Endpoint/Key 组合
Args:
db: 数据库会话
user_api_key: 用户 API Key
model_name: 模型名称
require_files_capability: 是否要求 gemini_files 能力(默认 True
Returns:
匹配的候选,如果没有则返回 None
"""
scheduler = CacheAwareScheduler()
# 要求 gemini_files 能力:只有 Google 官方 API 才支持 Files API
capability_requirements = {"gemini_files": True} if require_files_capability else None
candidates, _global_model_id = await scheduler.list_all_candidates(
db=db,
api_format="gemini:chat",
@@ -194,7 +223,7 @@ async def _select_provider_candidate(
affinity_key=str(user_api_key.id),
user_api_key=user_api_key,
max_candidates=10,
capability_requirements=REQUIRED_CAPABILITIES,
capability_requirements=capability_requirements,
)
for candidate in candidates:
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
@@ -206,12 +235,20 @@ async def _select_provider_candidate(
async def _resolve_upstream_context(
request: Request,
db: Session,
) -> tuple[str, str, str]:
) -> tuple[str, str, str, str]:
"""
解析上游 Key 与 Base URL
解析上游 Key 与 Base URL(需要外部提供 db session
仅允许系统 API Key通过能力标签选择支持 Files API 的 Provider Key。
仅允许系统 API Key选择可用的 Gemini Provider Key(无能力限制)
Args:
request: HTTP 请求
db: 数据库会话
Returns:
(upstream_key, base_url, key_id, user_id)
"""
client_key = _extract_gemini_api_key(request)
if not client_key:
raise HTTPException(
@@ -252,14 +289,19 @@ async def _resolve_upstream_context(
},
)
candidate = await _select_provider_candidate(db, user_api_key, model_name)
# 选择可用的 provider candidate(要求 gemini_files 能力)
candidate = await _select_provider_candidate(
db, user_api_key, model_name, require_files_capability=True
)
if not candidate:
raise HTTPException(
status_code=503,
detail={
"error": {
"code": 503,
"message": "No available key with gemini_files_api capability",
"message": "No available Gemini key with 'gemini_files' capability. "
"Please ensure at least one Provider Key has the 'gemini_files' capability enabled.",
"status": "UNAVAILABLE",
}
},
@@ -268,7 +310,7 @@ async def _resolve_upstream_context(
try:
upstream_key = crypto_service.decrypt(candidate.key.api_key)
except Exception as exc:
logger.error(f"Failed to decrypt provider key for Gemini Files API: {exc}")
logger.error("Failed to decrypt provider key for Gemini Files API: %s", exc)
raise HTTPException(
status_code=500,
detail={
@@ -281,7 +323,29 @@ async def _resolve_upstream_context(
)
base_url = candidate.endpoint.base_url or GEMINI_FILES_BASE_URL
return upstream_key, base_url, str(candidate.key.id)
return upstream_key, base_url, str(candidate.key.id), str(user.id)
async def _resolve_upstream_context_standalone(request: Request) -> UpstreamContext:
"""
解析上游上下文(自管理数据库连接,适用于 HTTP 代理场景)
优化在返回上下文后立即释放数据库连接HTTP 请求期间不持有连接。
Args:
request: HTTP 请求
Returns:
UpstreamContext: 包含所有必要信息的上下文对象
"""
with create_session() as db:
upstream_key, base_url, file_key_id, user_id = await _resolve_upstream_context(request, db)
return UpstreamContext(
upstream_key=upstream_key,
base_url=base_url,
file_key_id=file_key_id,
user_id=user_id,
)
async def _proxy_request(
@@ -291,6 +355,7 @@ async def _proxy_request(
content: bytes | None = None,
json_body: dict[str, Any] | None = None,
file_key_id: str | None = None,
user_id: str | None = None,
) -> Response:
"""
代理请求到上游 Gemini API
@@ -302,6 +367,7 @@ async def _proxy_request(
content: 原始请求体(二进制)
json_body: JSON 请求体
file_key_id: 上游 Provider Key ID用于成功响应时存储 file→key 映射
user_id: 用户 ID用于文件映射的权限验证
Returns:
FastAPI Response 对象
@@ -338,12 +404,28 @@ async def _proxy_request(
try:
payload = response.json()
file_name = None
file_obj = None
if isinstance(payload, dict):
# 单文件上传响应
file_name = payload.get("name")
file_obj = payload
# 嵌套格式:{"file": {...}}
if not file_name and isinstance(payload.get("file"), dict):
file_name = payload["file"].get("name")
if file_name:
await store_file_key_mapping(file_name, file_key_id)
file_obj = payload["file"]
if file_name and file_obj:
display_name = file_obj.get("displayName") or file_obj.get("display_name")
mime_type = file_obj.get("mimeType") or file_obj.get("mime_type")
await store_file_key_mapping(
file_name,
file_key_id,
user_id=user_id,
display_name=display_name,
mime_type=mime_type,
)
logger.debug(
f"Gemini file→key 映射已存储: {file_name} → key_id={file_key_id}"
)
@@ -355,14 +437,26 @@ async def _proxy_request(
mapped_count = 0
for item in files_list:
if isinstance(item, dict) and item.get("name"):
await store_file_key_mapping(item["name"], file_key_id)
item_display_name = item.get("displayName") or item.get(
"display_name"
)
item_mime_type = item.get("mimeType") or item.get("mime_type")
await store_file_key_mapping(
item["name"],
file_key_id,
user_id=user_id,
display_name=item_display_name,
mime_type=item_mime_type,
)
mapped_count += 1
if mapped_count > 0:
logger.debug(
f"Gemini list_files 批量映射已存储: {mapped_count} 个文件 → key_id={file_key_id}"
"Gemini list_files 批量映射已存储: %d 个文件 → key_id=%s",
mapped_count,
file_key_id,
)
except (ValueError, KeyError) as e:
logger.debug(f"Failed to store Gemini file mapping: {e}")
logger.debug("Failed to store Gemini file mapping: %s", e)
return Response(
content=response.content,
@@ -373,7 +467,7 @@ async def _proxy_request(
except Exception as e:
sanitized_error = redact_url_for_log(str(e))
logger.error(f"Gemini Files API proxy error: {sanitized_error}")
logger.error("Gemini Files API proxy error: %s", sanitized_error)
return JSONResponse(
status_code=502,
content={
@@ -394,7 +488,6 @@ async def _proxy_request(
@router.post("/upload/v1beta/files")
async def upload_file(
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
上传文件到 Gemini Files API
@@ -421,26 +514,34 @@ async def upload_file(
}
}
```
"""
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
# 读取请求体
优化HTTP 代理期间不持有数据库连接
"""
# 阶段 1解析上下文短暂持有数据库连接
ctx = await _resolve_upstream_context_standalone(request)
# 阶段 2读取请求体
body = await request.body()
# 构建上游请求
# 阶段 3代理请求不持有数据库连接
upstream_url = _build_upstream_url(
base_url,
ctx.base_url,
"/v1beta/files",
dict(request.query_params),
is_upload=True,
)
headers = _build_upstream_headers(dict(request.headers), upstream_key)
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
logger.debug(f"Gemini Files upload proxy: POST {redact_url_for_log(upstream_url)}")
logger.debug("Gemini Files upload proxy: POST %s", redact_url_for_log(upstream_url))
return await _proxy_request(
"POST", upstream_url, headers, content=body, file_key_id=file_key_id
"POST",
upstream_url,
headers,
content=body,
file_key_id=ctx.file_key_id,
user_id=ctx.user_id,
)
@@ -452,13 +553,14 @@ async def upload_file(
@router.get("/v1beta/files")
async def list_files(
request: Request,
db: Session = Depends(get_db),
pageSize: int | None = None,
pageToken: str | None = None,
) -> Any:
"""
列出已上传的文件
优化HTTP 代理期间不持有数据库连接
**认证方式**:
- `x-goog-api-key` 请求头,或
- `?key=` URL 参数
@@ -488,22 +590,222 @@ async def list_files(
}
```
"""
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
# 阶段 1解析上下文短暂持有数据库连接
ctx = await _resolve_upstream_context_standalone(request)
# 构建查询参数
# 阶段 2代理请求不持有数据库连接
query_params = dict(request.query_params)
if pageSize is not None:
query_params["pageSize"] = pageSize
if pageToken is not None:
query_params["pageToken"] = pageToken
# 构建上游请求
upstream_url = _build_upstream_url(base_url, "/v1beta/files", query_params)
upstream_url = _build_upstream_url(ctx.base_url, "/v1beta/files", query_params)
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
logger.debug("Gemini Files list proxy: GET %s", redact_url_for_log(upstream_url))
return await _proxy_request(
"GET", upstream_url, headers, file_key_id=ctx.file_key_id, user_id=ctx.user_id
)
# ==============================================================================
# 下载文件内容端点(用于视频等媒体文件)
# 注意:必须在 /v1beta/files/{file_name:path} 之前注册,否则会被通配符路由捕获
# ==============================================================================
async def _find_video_task_by_id(
db: Session, short_id: str, user_id: str
) -> tuple[str | None, str | None]:
"""
通过短 ID 查找视频任务,返回其 provider key 和 video_url
Args:
db: 数据库会话
short_id: 视频任务的短 IDVideoTask.short_idGemini 风格)
user_id: 用户 ID用于权限验证
Returns:
(upstream_key, video_url) - 如果找到任务返回 key 和 url否则返回 (None, None)
"""
from src.models.database import ProviderAPIKey, VideoTask
logger.debug(
"[Files Download] Searching video task: short_id=%s, user_id=%s", short_id, user_id
)
# 通过 short_id 查找,同时验证用户权限
task = (
db.query(VideoTask)
.filter(VideoTask.short_id == short_id, VideoTask.user_id == user_id)
.first()
)
if not task:
logger.debug("[Files Download] No video task found: short_id=%s", short_id)
return None, None
if not task.video_url:
logger.debug("[Files Download] Task found but no video_url: short_id=%s", short_id)
return None, None
if not task.key_id:
logger.debug("[Files Download] Task found but no key_id: short_id=%s", short_id)
return None, task.video_url
# 获取 provider key
provider_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == task.key_id).first()
if not provider_key or not provider_key.api_key:
logger.debug("[Files Download] Provider key not found: key_id=%s", task.key_id)
return None, task.video_url
try:
upstream_key = crypto_service.decrypt(provider_key.api_key)
logger.debug("[Files Download] Found key for task: short_id=%s", short_id)
return upstream_key, task.video_url
except Exception as e:
logger.error("[Files Download] Failed to decrypt key: %s", e)
return None, task.video_url
@router.get("/v1beta/files/{file_id}:download")
async def download_file(
file_id: str,
request: Request,
) -> Any:
"""
下载文件(官方 Gemini API 格式)
**认证方式**:
- `x-goog-api-key` 请求头,或
- `?key=` URL 参数
**路径参数**:
- `file_id`: 文件 ID
- 以 `aev_` 开头:视频任务下载(如 `aev_sknuzqlo8sds`Gemini 风格短 ID
- 其他:普通 Gemini 文件下载(透传到上游)
**查询参数**:
- `alt=media`: 可选,保持与官方 API 兼容
**示例**:
```
GET /v1beta/files/aev_{short_id}:download?alt=media # 视频任务
GET /v1beta/files/{gemini_file_id}:download?alt=media # 普通文件
```
优化HTTP 下载期间不持有数据库连接
"""
import httpx
from fastapi import HTTPException
from fastapi.responses import JSONResponse, Response
# ========== 阶段 1数据库操作短暂持有连接==========
client_key = _extract_gemini_api_key(request)
if not client_key:
raise HTTPException(
status_code=401,
detail={
"error": {"code": 401, "message": "API key required", "status": "UNAUTHENTICATED"}
},
)
# 在数据库会话内完成所有查询
with create_session() as db:
auth_result = AuthService.authenticate_api_key(db, client_key)
if not auth_result:
raise HTTPException(
status_code=401,
detail={
"error": {
"code": 401,
"message": "API key not valid",
"status": "UNAUTHENTICATED",
}
},
)
user, _user_api_key = auth_result
# 根据前缀判断处理方式
if file_id.startswith("aev_"):
# 视频任务下载:使用短 ID 查找
short_id = file_id[4:] # 去掉 "aev_" 前缀
logger.debug("[Files Download] Video task: short_id=%s, user_id=%s", short_id, user.id)
upstream_key, video_url = await _find_video_task_by_id(db, short_id, user.id)
if not upstream_key or not video_url:
raise HTTPException(
status_code=404,
detail={
"error": {
"code": 404,
"message": f"Video not found or not ready: {file_id}",
"status": "NOT_FOUND",
}
},
)
upstream_url = video_url
else:
# 普通文件下载:透传到 Gemini
try:
upstream_key, base_url, _file_key_id, _user_id = await _resolve_upstream_context(
request, db
)
except HTTPException:
raise HTTPException(
status_code=404,
detail={
"error": {
"code": 404,
"message": f"File not found: {file_id}",
"status": "NOT_FOUND",
}
},
)
file_name = f"files/{file_id}" if not file_id.startswith("files/") else file_id
upstream_url = _build_upstream_url(
base_url,
f"/v1beta/{file_name}:download",
dict(request.query_params),
)
# ========== 阶段 2HTTP 下载(不持有数据库连接)==========
headers = _build_upstream_headers(dict(request.headers), upstream_key)
logger.debug(f"Gemini Files list proxy: GET {redact_url_for_log(upstream_url)}")
logger.debug("Gemini Files download proxy: GET %s", redact_url_for_log(upstream_url))
return await _proxy_request("GET", upstream_url, headers, file_key_id=file_key_id)
# 使用 follow_redirects=True 跟随重定向Gemini 文件下载会重定向)
try:
async with httpx.AsyncClient(follow_redirects=True, timeout=httpx.Timeout(300.0)) as client:
response = await client.get(upstream_url, headers=headers)
except Exception as exc:
logger.error("Gemini Files download failed: %s", exc)
raise HTTPException(status_code=502, detail="Failed to download file")
if response.status_code >= 400:
content: dict[str, Any]
if response.headers.get("content-type", "").startswith("application/json"):
try:
content = response.json()
except Exception:
content = {"error": response.text}
else:
content = {"error": response.text}
return JSONResponse(content=content, status_code=response.status_code)
# 返回文件内容
return Response(
content=response.content,
status_code=response.status_code,
headers={
k: v
for k, v in response.headers.items()
if k.lower() not in {"transfer-encoding", "connection", "keep-alive"}
},
media_type=response.headers.get("content-type", "application/octet-stream"),
)
# ==============================================================================
@@ -515,7 +817,6 @@ async def list_files(
async def get_file(
file_name: str,
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
获取指定文件的元数据
@@ -542,24 +843,29 @@ async def get_file(
"state": "ACTIVE"
}
```
"""
upstream_key, base_url, file_key_id = await _resolve_upstream_context(request, db)
优化HTTP 代理期间不持有数据库连接
"""
# 阶段 1解析上下文短暂持有数据库连接
ctx = await _resolve_upstream_context_standalone(request)
# 阶段 2代理请求不持有数据库连接
# 规范化文件名(确保以 files/ 开头)
if not file_name.startswith("files/"):
file_name = f"files/{file_name}"
# 构建上游请求
upstream_url = _build_upstream_url(
base_url,
ctx.base_url,
f"/v1beta/{file_name}",
dict(request.query_params),
)
headers = _build_upstream_headers(dict(request.headers), upstream_key)
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
logger.debug(f"Gemini Files get proxy: GET {redact_url_for_log(upstream_url)}")
logger.debug("Gemini Files get proxy: GET %s", redact_url_for_log(upstream_url))
return await _proxy_request("GET", upstream_url, headers, file_key_id=file_key_id)
return await _proxy_request(
"GET", upstream_url, headers, file_key_id=ctx.file_key_id, user_id=ctx.user_id
)
# ==============================================================================
@@ -571,7 +877,6 @@ async def get_file(
async def delete_file(
file_name: str,
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
删除指定文件
@@ -585,29 +890,31 @@ async def delete_file(
**响应格式**:
成功时返回空 JSON 对象:`{}`
"""
upstream_key, base_url, _file_key_id = await _resolve_upstream_context(request, db)
优化HTTP 代理期间不持有数据库连接
"""
# 阶段 1解析上下文短暂持有数据库连接
ctx = await _resolve_upstream_context_standalone(request)
# 阶段 2代理请求不持有数据库连接
# 规范化文件名(确保以 files/ 开头)
if not file_name.startswith("files/"):
file_name = f"files/{file_name}"
# 构建上游请求
upstream_url = _build_upstream_url(
base_url,
ctx.base_url,
f"/v1beta/{file_name}",
dict(request.query_params),
)
headers = _build_upstream_headers(dict(request.headers), upstream_key)
headers = _build_upstream_headers(dict(request.headers), ctx.upstream_key)
logger.debug(f"Gemini Files delete proxy: DELETE {redact_url_for_log(upstream_url)}")
logger.debug("Gemini Files delete proxy: DELETE %s", redact_url_for_log(upstream_url))
del _file_key_id # 显式标记delete 端点不需要存储映射
response = await _proxy_request("DELETE", upstream_url, headers)
if response.status_code < 300:
await delete_file_key_mapping(file_name)
else:
logger.debug(
f"Gemini Files delete failed, skip mapping cleanup: status={response.status_code}"
"Gemini Files delete failed, skip mapping cleanup: status=%s", response.status_code
)
return response

View File

@@ -61,9 +61,10 @@ async def list_video_tasks_sora(http_request: Request, db: Session = Depends(get
@router.delete("/v1/videos/{task_id}")
async def cancel_video_task_sora(
async def delete_video_task_sora(
task_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
"""删除已完成或失败的视频及其存储资源"""
adapter = OpenAIVideoAdapter()
return await pipeline.run(
adapter=adapter,
@@ -71,7 +72,7 @@ async def cancel_video_task_sora(
db=db,
mode=adapter.mode,
api_format_hint=adapter.allowed_api_formats[0],
path_params={"task_id": task_id, "action": "cancel"},
path_params={"task_id": task_id},
)
@@ -121,6 +122,47 @@ async def create_video_veo(model: str, http_request: Request, db: Session = Depe
)
# Gemini Veo operation routes - support both formats:
# 1. models/{model}/operations/{id} (official Gemini Veo format)
# 2. operations/{...} (legacy format for compatibility)
@router.get("/v1beta/models/{model}/operations/{operation_id}")
async def get_video_veo_by_model(
model: str, operation_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
"""Get video task status (Gemini Veo format: models/{model}/operations/{id})"""
adapter = GeminiVeoAdapter()
# Reconstruct full operation name
full_operation_name = f"models/{model}/operations/{operation_id}"
return await pipeline.run(
adapter=adapter,
http_request=http_request,
db=db,
mode=adapter.mode,
api_format_hint=adapter.allowed_api_formats[0],
path_params={"task_id": full_operation_name},
)
@router.post("/v1beta/models/{model}/operations/{operation_id}:cancel")
async def cancel_video_veo_by_model(
model: str, operation_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
"""Cancel video task (Gemini Veo format: models/{model}/operations/{id}:cancel)"""
adapter = GeminiVeoAdapter()
full_operation_name = f"models/{model}/operations/{operation_id}"
return await pipeline.run(
adapter=adapter,
http_request=http_request,
db=db,
mode=adapter.mode,
api_format_hint=adapter.allowed_api_formats[0],
path_params={"task_id": full_operation_name, "action": "cancel"},
)
# Legacy routes for backward compatibility
@router.get("/v1beta/operations/{operation_id:path}")
async def get_video_veo(
operation_id: str, http_request: Request, db: Session = Depends(get_db)
@@ -163,16 +205,4 @@ async def cancel_video_veo(
)
@router.get("/v1beta/operations/{operation_id:path}/content")
async def download_video_content_veo(
operation_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
adapter = GeminiVeoAdapter()
return await pipeline.run(
adapter=adapter,
http_request=http_request,
db=db,
mode=adapter.mode,
api_format_hint=adapter.allowed_api_formats[0],
path_params={"task_id": operation_id},
)
# Video download is now handled by /v1beta/files/{task_id}:download in gemini_files.py

View File

@@ -736,6 +736,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy import or_
from sqlalchemy.orm import load_only
from src.models.database import ProviderEndpoint
from src.utils.database_helpers import escape_like_pattern, safe_truncate_escaped
@@ -877,7 +878,44 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
)
# 计算总数用于分页
total_records = query.count()
# Perf: avoid Query.count() building a subquery selecting many columns
total_records = 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.user_id,
Usage.api_key_id,
Usage.provider_name,
Usage.model,
Usage.target_model,
Usage.input_tokens,
Usage.output_tokens,
Usage.total_tokens,
Usage.total_cost_usd,
Usage.response_time_ms,
Usage.first_byte_time_ms,
Usage.is_stream,
Usage.status,
Usage.created_at,
Usage.cache_creation_input_tokens,
Usage.cache_read_input_tokens,
Usage.status_code,
Usage.error_message,
Usage.api_format,
Usage.endpoint_api_format,
Usage.has_format_conversion,
Usage.input_price_per_1m,
Usage.output_price_per_1m,
Usage.cache_creation_price_per_1m,
Usage.cache_read_price_per_1m,
Usage.actual_total_cost_usd,
Usage.rate_multiplier,
),
load_only(ApiKey.id, ApiKey.name, ApiKey.key_encrypted),
load_only(ProviderEndpoint.id, ProviderEndpoint.api_format),
)
usage_records = (
query.order_by(Usage.created_at.desc()).offset(self.offset).limit(self.limit).all()
)
@@ -1231,7 +1269,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
endpoint_format,
format_acceptance_config,
is_stream=False,
global_conversion_enabled=global_conversion_enabled,
effective_conversion_enabled=global_conversion_enabled,
)
if is_compatible:
provider_to_formats.setdefault(provider_id, set()).add(endpoint_format)

View File

@@ -3,11 +3,20 @@
用于候选筛选时判断端点是否可以处理客户端请求格式。
三层开关优先级(从高到低):
1. 全局开关 ON → 强制允许(跳过后续检查)
2. 全局开关 OFF → 看提供商开关
- 提供商开关 ON → 强制允许(跳过端点检查)
- 提供商开关 OFF → 看端点配置
3. 端点配置format_acceptance_config
- enabled=true + 白名单/黑名单检查 → 允许
- enabled=false 或未配置 → 禁止
转换逻辑:
1. 格式完全匹配 -> 透传(无需转换)
2. 格式不同 -> 需要检查全局开关 + 端点开关
- data_format_id 相同 -> 可透传(无需数据转换),但需全局开关 + 端点开关
- data_format_id 不同 -> 需要转换,检查全局开关 + 端点配置 + 转换器能力
2. 格式不同 -> 需要检查三层开关
- data_format_id 相同 -> 可透传(无需数据转换)
- data_format_id 不同 -> 需要转换,检查转换器能力
"""
from __future__ import annotations
@@ -28,8 +37,10 @@ def is_format_compatible(
endpoint_api_format: str,
endpoint_format_acceptance_config: dict | None,
is_stream: bool,
global_conversion_enabled: bool,
effective_conversion_enabled: bool,
registry: FormatConversionRegistry | None = None,
*,
skip_endpoint_check: bool = False,
) -> tuple[bool, bool, str | None]:
"""
检查端点是否兼容客户端格式
@@ -39,8 +50,9 @@ def is_format_compatible(
endpoint_api_format: 端点的 API 格式
endpoint_format_acceptance_config: 端点的格式接受配置
is_stream: 是否是流式请求
global_conversion_enabled: 全局格式转换开关(来自环境变量 FORMAT_CONVERSION_ENABLED默认 True
effective_conversion_enabled: 有效格式转换开关(全局 OR 提供商
registry: 转换器注册表(可选,默认使用全局单例)
skip_endpoint_check: 是否跳过端点配置检查(当全局或提供商开关为 ON 时设为 True
Returns:
(is_compatible, needs_conversion, skip_reason)
@@ -66,30 +78,36 @@ def is_format_compatible(
if provider_key == client_key:
return True, False, None
# 2. 格式不同 -> 需要检查全局格式转换开关
# 即使 data_format_id 相同(如 claude:chat / claude:cli也需要全局开关启用
if not global_conversion_enabled:
return False, False, "全局格式转换未启用(环境变量 FORMAT_CONVERSION_ENABLED=false"
# 2. 格式不同 -> 需要检查格式转换开关
# 如果有效开关为 False全局 OFF 且提供商 OFF直接拒绝
if not effective_conversion_enabled:
return False, False, "格式转换已禁用(全局和提供商开关均为关闭"
# 3. 格式不同时,统一检查端点配置(核心控制)
if endpoint_format_acceptance_config is None:
return False, False, "端点配置格式接受策略"
# 3. 如果全局或提供商开关为 ON跳过端点配置检查
if not skip_endpoint_check:
# 检查端点配置(第三层开关)
if endpoint_format_acceptance_config is None:
return False, False, "端点未配置格式接受策略"
config = endpoint_format_acceptance_config
if not isinstance(config, dict):
return False, False, "端点格式配置无效"
if not config.get("enabled", False):
return False, False, "端点格式接受未启用"
config = endpoint_format_acceptance_config
if not isinstance(config, dict):
return False, False, "端点格式配置无效"
if not config.get("enabled", False):
return False, False, "端点格式接受未启用"
# 检查 reject_formats优先
reject_formats = config.get("reject_formats", [])
if client_key in [f.upper() for f in reject_formats]:
return False, False, f"端点拒绝 {client_format} 格式"
# 检查 reject_formats优先
reject_formats = config.get("reject_formats", [])
if client_key in [f.upper() for f in reject_formats]:
return False, False, f"端点拒绝 {client_format} 格式"
# 检查 accept_formats
accept_formats = config.get("accept_formats", [])
if accept_formats and client_key not in [f.upper() for f in accept_formats]:
return False, False, f"端点不接受 {client_format} 格式"
# 检查 accept_formats
accept_formats = config.get("accept_formats", [])
if accept_formats and client_key not in [f.upper() for f in accept_formats]:
return False, False, f"端点不接受 {client_format} 格式"
# 检查流式转换
if is_stream and not config.get("stream_conversion", True):
return False, False, "端点不支持流式格式转换"
# 4. 检查是否可以透传data_format_id 相同)
# 例如claude:chat / claude:cli 的 data_format_id 都是 "claude",数据格式相同可透传
@@ -99,11 +117,7 @@ def is_format_compatible(
return True, False, None
# 5. 需要数据转换的情况data_format_id 不同)
# 检查流式转换
if is_stream and not config.get("stream_conversion", True):
return False, False, "端点不支持流式格式转换"
# 6. 检查转换器能力
# 检查转换器能力
if not registry.can_convert_full(
client_key,
provider_key,

View File

@@ -736,10 +736,17 @@ class GeminiNormalizer(FormatNormalizer):
instance = instances[0] if isinstance(instances[0], dict) else {}
params = request.get("parameters") or {}
# 解析 image用于 image-to-video 或第一帧)
# 官方格式: {"image": {"inlineData": {"mimeType": "image/png", "data": "base64..."}}}
image = instance.get("image") if isinstance(instance, dict) else None
image_ref = None
if isinstance(image, dict):
image_ref = image.get("bytesBase64Encoded")
inline_data = image.get("inlineData", {})
if isinstance(inline_data, dict):
image_ref = inline_data.get("data")
# 兼容旧格式
if not image_ref:
image_ref = image.get("bytesBase64Encoded")
prompt = instance.get("prompt") if isinstance(instance, dict) else None
prompt_str = str(prompt).strip() if prompt else ""
@@ -747,7 +754,7 @@ class GeminiNormalizer(FormatNormalizer):
raise ValueError("Video prompt is required")
duration_raw = params.get("durationSeconds")
sample_count_raw = params.get("sampleCount")
sample_count_raw = params.get("sampleCount") or params.get("numberOfVideos")
try:
duration_seconds = int(duration_raw) if duration_raw else 8
@@ -760,6 +767,35 @@ class GeminiNormalizer(FormatNormalizer):
except (ValueError, TypeError):
sample_count = 1
# 构建 extra 字段,保留所有 Veo 特有的参数
extra: dict[str, Any] = {
"personGeneration": params.get("personGeneration"),
"sampleCount": sample_count,
}
# negativePrompt - 负面提示词
if params.get("negativePrompt"):
extra["negativePrompt"] = params["negativePrompt"]
# lastFrame - 最后一帧(用于插值)
last_frame = params.get("lastFrame")
if isinstance(last_frame, dict):
extra["lastFrame"] = last_frame
# referenceImages - 参考图像最多3张仅 Veo 3.1
ref_images = params.get("referenceImages")
if isinstance(ref_images, list) and ref_images:
extra["referenceImages"] = ref_images
# video - 视频扩展输入(用于视频续写)
video_input = instance.get("video") if isinstance(instance, dict) else None
if isinstance(video_input, dict):
extra["video"] = video_input
# seed - 种子值Veo 3
if params.get("seed") is not None:
extra["seed"] = params["seed"]
return InternalVideoRequest(
prompt=prompt_str,
model=str(request.get("model") or "veo-3.1-generate-preview"),
@@ -767,10 +803,7 @@ class GeminiNormalizer(FormatNormalizer):
aspect_ratio=str(params.get("aspectRatio") or "16:9"),
resolution=str(params.get("resolution") or "720p"),
reference_image_url=image_ref,
extra={
"personGeneration": params.get("personGeneration"),
"sampleCount": sample_count,
},
extra=extra,
)
def video_request_from_internal(self, internal: InternalVideoRequest) -> dict[str, Any]:
@@ -783,15 +816,31 @@ class GeminiNormalizer(FormatNormalizer):
"resolution": internal.resolution,
"durationSeconds": internal.duration_seconds,
}
for key in ["personGeneration", "sampleCount"]:
for key in ["personGeneration", "sampleCount", "negativePrompt", "seed"]:
if key in internal.extra:
parameters[key] = internal.extra[key]
return {
# lastFrame 和 referenceImages 需要特殊处理
if internal.extra.get("lastFrame"):
parameters["lastFrame"] = internal.extra["lastFrame"]
if internal.extra.get("referenceImages"):
parameters["referenceImages"] = internal.extra["referenceImages"]
# video 输入(视频续写)
if internal.extra.get("video"):
instance["video"] = internal.extra["video"]
result: dict[str, Any] = {
"instances": [instance],
"parameters": parameters,
}
# 模型信息(用于 URL 构建)
if internal.model:
result["model"] = internal.model
return result
def video_task_to_internal(self, response: dict[str, Any]) -> InternalVideoTask:
operation_name = str(response.get("name") or "")
done = bool(response.get("done"))
@@ -824,20 +873,30 @@ class GeminiNormalizer(FormatNormalizer):
)
def video_task_from_internal(self, internal: InternalVideoTask) -> dict[str, Any]:
# 优先使用 external_id(上游返回的 operation name),否则用内部 id
operation_name = internal.external_id or f"operations/{internal.id}"
if not operation_name.startswith("operations/"):
operation_name = f"operations/{operation_name}"
# external_id 中提取 model 名称,用于构建 operation name
# external_id 格式: models/{model}/operations/{gemini_id}
model_name = "unknown"
if internal.external_id:
parts = internal.external_id.split("/")
if len(parts) >= 2 and parts[0] == "models":
model_name = parts[1]
# 使用我们的内部 task_id 构建 operation name不暴露 Gemini 的 operation ID
# 格式: models/{model}/operations/{our_task_id}
operation_name = f"models/{model_name}/operations/{internal.id}"
if internal.status == VideoStatus.COMPLETED:
urls = internal.video_urls or ([internal.video_url] if internal.video_url else [])
# 使用我们的内部 task_id 构建下载 URL不暴露真实的 Gemini file_id
# 使用 aev_ 前缀标识这是视频任务的下载链接
# 格式:/v1beta/files/aev_{task_id}:download?alt=media
proxy_download_url = f"/v1beta/files/aev_{internal.id}:download?alt=media"
return {
"name": operation_name,
"done": True,
"response": {
"generateVideoResponse": {
"generatedSamples": [
{"video": {"uri": url, "mimeType": "video/mp4"}} for url in urls
{"video": {"uri": proxy_download_url, "mimeType": "video/mp4"}}
]
}
},

View File

@@ -752,10 +752,21 @@ class OpenAINormalizer(FormatNormalizer):
"message": internal.error_message,
}
for key in ["model", "size", "seconds"]:
if key in internal.extra:
# 基本字段
for key in ["model", "size", "prompt"]:
if internal.extra.get(key):
payload[key] = internal.extra[key]
# seconds 必须是字符串类型
seconds = internal.extra.get("seconds")
if seconds is not None:
payload["seconds"] = str(seconds)
# remix 相关字段
remixed_from = internal.extra.get("remixed_from_video_id")
if remixed_from:
payload["remixed_from_video_id"] = remixed_from
return payload
def video_poll_to_internal(self, response: dict[str, Any]) -> InternalVideoPollResult:
@@ -764,14 +775,18 @@ class OpenAINormalizer(FormatNormalizer):
if status == "completed":
expires_at = response.get("expires_at")
# 使用任务 ID 构建内容路径,由调用方拼接完整 URL
# 如果 task_id 不存在,说明上游响应异常
video_url = f"videos/{task_id}/content" if task_id else None
# 优先使用上游返回的直接 URL某些代理如 API易 会返回 CDN URL
# 回退到构建相对路径(标准 OpenAI API 通过 /content 端点下载)
direct_url = (
response.get("video_url") or response.get("url") or response.get("result_url")
)
# 使用直接 URL 或回退到相对路径
video_url = direct_url or (f"videos/{task_id}/content" if task_id else None)
if not video_url:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_task_id",
error_message="Upstream response missing task id",
error_code="missing_video_url",
error_message="Upstream response missing video url",
raw_response=response,
)
return InternalVideoPollResult(

View File

@@ -155,6 +155,110 @@ class FormatConversionRegistry:
except Exception as e:
raise FormatConversionError(source_format, target_format, str(e)) from e
# ==================== 视频格式转换 ====================
def convert_video_request(
self,
request: dict[str, Any],
source_format: str,
target_format: str,
) -> dict[str, Any]:
"""转换视频请求格式OpenAI <-> Gemini
Args:
request: 原始视频请求
source_format: 源格式(如 openai:video, gemini:video
target_format: 目标格式
Returns:
转换后的视频请求
"""
# 统一使用基础格式 ID去掉 :video 后缀)
src_base = self._video_format_to_base(source_format)
tgt_base = self._video_format_to_base(target_format)
if src_base == tgt_base:
return request
src = self._require_normalizer(src_base)
tgt = self._require_normalizer(tgt_base)
with _track_conversion_metrics(
"video_request", str(source_format).upper(), str(target_format).upper()
):
try:
internal = src.video_request_to_internal(request)
return tgt.video_request_from_internal(internal)
except Exception as e:
raise FormatConversionError(source_format, target_format, str(e)) from e
def convert_video_task(
self,
task_response: dict[str, Any],
source_format: str,
target_format: str,
) -> dict[str, Any]:
"""转换视频任务响应格式OpenAI <-> Gemini
Args:
task_response: 原始任务响应
source_format: 源格式
target_format: 目标格式
Returns:
转换后的任务响应
"""
src_base = self._video_format_to_base(source_format)
tgt_base = self._video_format_to_base(target_format)
if src_base == tgt_base:
return task_response
src = self._require_normalizer(src_base)
tgt = self._require_normalizer(tgt_base)
with _track_conversion_metrics(
"video_task", str(source_format).upper(), str(target_format).upper()
):
try:
internal = src.video_task_to_internal(task_response)
return tgt.video_task_from_internal(internal)
except Exception as e:
raise FormatConversionError(source_format, target_format, str(e)) from e
def can_convert_video(self, source_format: str, target_format: str) -> bool:
"""检查是否支持视频格式转换"""
src_base = self._video_format_to_base(source_format)
tgt_base = self._video_format_to_base(target_format)
if src_base == tgt_base:
return True
src = self.get_normalizer(src_base)
tgt = self.get_normalizer(tgt_base)
if src is None or tgt is None:
return False
# 检查是否有视频转换方法
return (
hasattr(src, "video_request_to_internal")
and hasattr(src, "video_task_to_internal")
and hasattr(tgt, "video_request_from_internal")
and hasattr(tgt, "video_task_from_internal")
)
def _video_format_to_base(self, format_id: str) -> str:
"""将视频格式 ID 转换为基础格式 ID
例如: openai:video -> openai:chat, gemini:video -> gemini:chat
"""
upper = str(format_id).upper()
if upper.endswith(":VIDEO"):
base = upper[:-6] # 去掉 :VIDEO
return f"{base}:CHAT"
return upper
# ==================== 流式转换(严格) ====================
def convert_stream_chunk(

View File

@@ -248,10 +248,10 @@ register_capability(
)
register_capability(
name="gemini_files_api",
display_name="Gemini文件上传",
description="支持 Gemini Files API上传、查询、删除),第三方 Key 通常不支持",
match_mode=CapabilityMatchMode.COMPATIBLE, # 需要时选有的,不需要时都可选
config_mode=CapabilityConfigMode.REQUEST_PARAM, # 从请求路径检测
short_name="文件上传",
name="gemini_files",
display_name="Gemini 文件 API",
description="支持 Gemini Files API文件上传/管理),仅 Google 官方 API 支持",
match_mode=CapabilityMatchMode.COMPATIBLE,
config_mode=CapabilityConfigMode.USER_CONFIGURABLE,
short_name="文件API",
)

View File

@@ -236,15 +236,23 @@ class ModuleRegistry:
# 获取启用状态
enabled = self.is_enabled(name, db) if available else False
# 注意:配置验证失败时自动禁用模块
# 自动禁用会在查询方法中产生写操作副作用,违反幂等性原则
# 配置验证状态通过 config_validated/config_error 字段返回,由调用方决定如何处理
# 配置验证失败时自动禁用模块
# 注意:此处故意在 get_status() 中写入,以确保模块状态与配置同步
# 场景:用户删除了模块所依赖的 Provider Key 后,模块应自动关闭
# 权衡:查询方法中的写操作副作用 vs 状态一致性保证
if enabled and not config_validated:
self.set_enabled(name, False, db)
enabled = False
# 计算激活状态available && enabled && config_validated && 依赖模块都激活
is_active = self.is_active(name, db) if available else False
active = is_active and config_validated
return ModuleStatus(
name=name,
available=available,
enabled=enabled,
active=self.is_active(name, db) if available else False,
active=active,
config_validated=config_validated,
config_error=config_error,
display_name=meta.display_name,

View File

@@ -3,7 +3,7 @@
"""
from ..models.database import ApiKey, Base, Usage, User, UserQuota
from .database import create_session, get_db, get_db_url, init_db, log_pool_status
from .database import create_session, get_db, get_db_context, get_db_url, init_db, log_pool_status
__all__ = [
"Base",
@@ -12,6 +12,7 @@ __all__ = [
"Usage",
"UserQuota",
"get_db",
"get_db_context",
"init_db",
"create_session",
"get_db_url",

View File

@@ -4,6 +4,7 @@
import time
from collections.abc import Generator
from contextlib import contextmanager
from typing import Any, cast
from sqlalchemy import create_engine, event
@@ -271,6 +272,31 @@ def create_session() -> Session:
return _SessionLocal()
@contextmanager
def get_db_context() -> Generator[Session, None, None]:
"""
获取数据库会话的上下文管理器
自动管理会话生命周期:创建、提交/回滚、关闭
示例:
with get_db_context() as db:
user = db.query(User).first()
# 事务在 with 块结束时自动提交或回滚
"""
_ensure_engine()
assert _SessionLocal is not None
db = _SessionLocal()
try:
yield db
db.commit()
except Exception:
db.rollback()
raise
finally:
db.close()
def get_db_url() -> str:
"""返回当前配置的数据库连接字符串(供脚本/测试使用)。"""
return config.database_url

View File

@@ -202,14 +202,14 @@ async def lifespan(app: FastAPI) -> Any:
logger.info("启动月卡额度重置调度器...")
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
from src.services.task.task_poller import get_task_poller
from src.services.usage.quota_scheduler import get_quota_scheduler
from src.services.video.task_poller import get_video_task_poller
from src.utils.task_coordinator import StartupTaskCoordinator
quota_scheduler = get_quota_scheduler()
maintenance_scheduler = get_maintenance_scheduler()
model_fetch_scheduler = get_model_fetch_scheduler()
video_task_poller = get_video_task_poller()
task_poller = get_task_poller()
task_coordinator = StartupTaskCoordinator(redis_client)
# 启动额度调度器
@@ -238,14 +238,14 @@ async def lifespan(app: FastAPI) -> Any:
logger.info("检测到其他 worker 已运行模型获取调度器,本实例跳过")
model_fetch_scheduler = None # type: ignore[assignment]
# 启动视频任务轮询服务
video_poller_active = await task_coordinator.acquire("video_task_poller")
if video_poller_active:
logger.info("启动视频任务轮询服务...")
await video_task_poller.start()
# 启动异步任务轮询服务(当前仅视频)
task_poller_active = await task_coordinator.acquire("task_poller:video")
if task_poller_active:
logger.info("启动 TaskPollervideo...")
await task_poller.start()
else:
logger.info("检测到其他 worker 已运行视频任务轮询,本实例跳过")
video_task_poller = None # type: ignore[assignment]
logger.info("检测到其他 worker 已运行 TaskPollervideo,本实例跳过")
task_poller = None # type: ignore[assignment]
# 启动统一的定时任务调度器
from src.services.system.scheduler import get_scheduler
@@ -296,10 +296,10 @@ async def lifespan(app: FastAPI) -> Any:
await model_fetch_scheduler.stop()
await task_coordinator.release("model_fetch_scheduler")
if video_task_poller:
logger.info("停止视频任务轮询...")
await video_task_poller.stop()
await task_coordinator.release("video_task_poller")
if task_poller:
logger.info("停止 TaskPollervideo...")
await task_poller.stop()
await task_coordinator.release("task_poller:video")
# 停止统一的定时任务调度器
logger.info("停止定时任务调度器...")

View File

@@ -272,6 +272,14 @@ class Usage(Base):
"""使用记录模型"""
__tablename__ = "usage"
__table_args__ = (
# Composite indexes for common query patterns (analytics / list pages)
Index("idx_usage_user_created", "user_id", "created_at"),
Index("idx_usage_apikey_created", "api_key_id", "created_at"),
Index("idx_usage_provider_model_created", "provider_name", "model", "created_at"),
Index("idx_usage_provider_created", "provider_name", "created_at"),
Index("idx_usage_model_created", "model", "created_at"),
)
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
user_id = Column(String(36), ForeignKey("users.id", ondelete="SET NULL"), nullable=True)
@@ -347,6 +355,13 @@ class Usage(Base):
# cancelled: 客户端主动断开连接
status = Column(String(20), default="completed", nullable=False, index=True)
# 结算状态(与 status 解耦)
# - pending: 等待结算(任务未完成 / 流式未结束)
# - settled: 已结算cost 已写入,可能 > 0 或 = 0
# - void: 作废(不收费,如任务未开始就取消)
billing_status = Column(String(20), default="settled", nullable=False, index=True)
finalized_at = Column(DateTime(timezone=True), nullable=True) # 结算完成时间(可选)
# 完整请求和响应记录
request_headers = Column(JSON, nullable=True) # 客户端请求头
request_body = Column(JSON, nullable=True) # 请求体7天内未压缩
@@ -654,6 +669,14 @@ class Provider(Base):
# 注意:如果全局配置 KEEP_PRIORITY_ON_CONVERSION=true此字段被忽略所有提供商都保持优先级
keep_priority_on_conversion = Column(Boolean, default=False, nullable=False)
# 是否允许格式转换(默认 True
# - True: 该提供商可以作为格式转换的目标(如 OpenAI 客户端请求可以路由到此 Gemini 提供商)
# - False: 该提供商不接受需要格式转换的请求
# 优先级逻辑:
# - 全局开关 ON → 强制允许所有提供商的格式转换(忽略此字段)
# - 全局开关 OFF → 由此字段决定是否允许该提供商的格式转换
enable_format_conversion = Column(Boolean, default=False, nullable=False)
# 状态
is_active = Column(Boolean, default=True, nullable=False)
@@ -1350,12 +1373,26 @@ class ProviderAPIKey(Base):
provider = relationship("Provider", back_populates="api_keys")
def _generate_short_id(length: int = 12) -> str:
"""生成 Gemini 风格的短 ID小写字母+数字)"""
import secrets
import string
alphabet = string.ascii_lowercase + string.digits
return "".join(secrets.choice(alphabet) for _ in range(length))
class VideoTask(Base):
"""视频生成任务"""
__tablename__ = "video_tasks"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
# Gemini 风格的短 ID用于对外暴露如 operations/xxx
short_id = Column(String(16), unique=True, index=True, default=_generate_short_id)
request_id = Column(
String(100), unique=True, index=True, nullable=False
) # 关联 Usage/RequestCandidate
external_task_id = Column(String(200))
# 关联
@@ -1865,6 +1902,7 @@ class RequestCandidate(Base):
Index("idx_request_candidates_request_id", "request_id"),
Index("idx_request_candidates_status", "status"),
Index("idx_request_candidates_provider_id", "provider_id"),
Index("idx_request_candidates_created_at", "created_at"),
)
# 关系
@@ -2109,5 +2147,60 @@ class StatsUserDaily(Base):
user = relationship("User")
class GeminiFileMapping(Base):
"""
Gemini Files API 文件与 Provider Key 的映射关系
用于持久化存储 file_id → key_id 的绑定关系,
确保后续 generateContent 请求使用上传时的同一 Key。
Gemini 文件有 48 小时有效期,此表中的记录也会在过期后被清理。
"""
__tablename__ = "gemini_file_mappings"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
# 文件名(如 files/abc123xyz
file_name = Column(String(255), nullable=False, unique=True, index=True)
# Provider Key ID关联到 provider_api_keys 表)
key_id = Column(
String(36),
ForeignKey("provider_api_keys.id", ondelete="CASCADE"),
nullable=False,
index=True,
)
# 用户 ID用于权限验证可选
user_id = Column(
String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=True, index=True
)
# 文件元数据(可选,用于调试)
display_name = Column(String(255), nullable=True)
mime_type = Column(String(100), nullable=True)
# 源文件哈希(用于关联相同源文件的不同上传,可选)
# 当同一源文件上传到多个 Key 时,可通过此字段找到所有等效文件
source_hash = Column(String(64), nullable=True, index=True)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
# 过期时间Gemini 文件 48 小时后过期)
expires_at = Column(DateTime(timezone=True), nullable=False, index=True)
# 关系
key = relationship("ProviderAPIKey")
user = relationship("User")
__table_args__ = (
Index("idx_gemini_file_mappings_expires", "expires_at"),
Index("idx_gemini_file_mappings_source_hash", "source_hash"),
)
# 导入扩展的数据库模型
from .database_extensions import ApiKeyProviderMapping, ProviderUsageTracking

View File

@@ -635,6 +635,10 @@ class ProviderUpdateRequest(BaseModel):
None,
description="格式转换时是否保持优先级True=保持原优先级False=需要转换时降级)",
)
enable_format_conversion: bool | None = Field(
None,
description="是否允许格式转换(提供商级别开关)",
)
is_active: bool | None = None
billing_type: str | None = Field(
None, description="计费类型monthly_quota/pay_as_you_go/free_tier"
@@ -667,6 +671,10 @@ class ProviderWithEndpointsSummary(BaseModel):
default=False,
description="格式转换时是否保持优先级True=保持原优先级False=需要转换时降级)",
)
enable_format_conversion: bool = Field(
default=True,
description="是否允许格式转换(提供商级别开关)",
)
is_active: bool
# 计费相关字段

View File

@@ -7,6 +7,7 @@
from src.core.modules.base import ModuleDefinition
# 导入所有模块定义
from src.modules.gemini_files import gemini_files_module
from src.modules.ldap import ldap_module
from src.modules.oauth import oauth_module
@@ -14,6 +15,7 @@ from src.modules.oauth import oauth_module
ALL_MODULES: list[ModuleDefinition] = [
ldap_module,
oauth_module,
gemini_files_module,
]
__all__ = ["ALL_MODULES"]

View File

@@ -0,0 +1,96 @@
"""
Gemini Files 文件管理模块
提供 Gemini Files API 文件上传和管理功能:
- 文件上传到 Google Gemini Files API
- 文件映射管理file_id → key_id
- 文件列表查看和删除
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any
from src.core.modules.base import (
ModuleCategory,
ModuleDefinition,
ModuleHealth,
ModuleMetadata,
)
if TYPE_CHECKING:
from sqlalchemy.orm import Session
def _get_router() -> Any:
"""延迟导入路由"""
from src.api.admin.gemini_files import router
return router
async def _health_check() -> ModuleHealth:
"""健康检查"""
# 检查是否有可用的 gemini_files 能力的 Key
from src.database import create_session
from src.models.database import ProviderAPIKey
db = create_session()
try:
# 查找有 gemini_files 能力的 Key
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.is_active.is_(True)).all()
has_capable_key = any(
key.capabilities and key.capabilities.get("gemini_files", False) for key in keys
)
if has_capable_key:
return ModuleHealth.HEALTHY
return ModuleHealth.DEGRADED
except Exception:
return ModuleHealth.UNKNOWN
finally:
db.close()
def _validate_config(db: Session) -> tuple[bool, str]:
"""
验证 Gemini Files 模块配置
检查项:
1. 至少有一个有 gemini_files 能力的 Provider Key
"""
from src.models.database import ProviderAPIKey
# 查找有 gemini_files 能力的 Key
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)
]
if not capable_keys:
return (
False,
"至少启用一个具有「Gemini 文件 API」能力的 Key",
)
return True, ""
gemini_files_module = ModuleDefinition(
metadata=ModuleMetadata(
name="gemini_files",
display_name="文件缓存",
description="管理 Gemini Files API 上传的文件,支持文件上传、查看和删除",
category=ModuleCategory.INTEGRATION,
env_key="GEMINI_FILES_AVAILABLE",
default_available=True,
required_packages=[],
api_prefix="/api/admin/gemini-files",
admin_route="/admin/gemini-files",
admin_menu_icon="FileUp",
admin_menu_group="system",
admin_menu_order=60,
),
router_factory=_get_router,
health_check=_health_check,
validate_config=_validate_config,
)

View File

@@ -804,8 +804,33 @@ class OAuthService:
cfg = OAuthService._get_provider_config(db, provider_type)
# Read all required fields first, then release DB connection before any awaits.
# This prevents holding a pooled connection while doing network I/O.
auth_url = provider.get_effective_authorization_url(cfg)
token_url = provider.get_effective_token_url(cfg)
redirect_uri = cfg.redirect_uri
client_id = cfg.client_id
has_secret = bool(cfg.client_secret_encrypted)
client_secret = cfg.get_client_secret() if has_secret else None
# Release DB connection (safe only when session has no pending changes).
try:
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
except Exception:
has_pending_changes = False
if not has_pending_changes:
original_expire_on_commit = getattr(db, "expire_on_commit", True)
db.expire_on_commit = False
try:
if db.in_transaction():
db.commit()
except Exception:
try:
db.rollback()
except Exception:
pass
finally:
db.expire_on_commit = original_expire_on_commit
async def _reachable(url: str) -> bool:
try:
@@ -823,7 +848,7 @@ class OAuthService:
secret_status = "unknown"
details = ""
if cfg.client_secret_encrypted:
if has_secret and client_secret:
# 使用无效 code 做一次 token 请求(仅做粗略判定)
try:
async with httpx.AsyncClient(
@@ -834,9 +859,9 @@ class OAuthService:
data={
"grant_type": "authorization_code",
"code": "invalid",
"redirect_uri": cfg.redirect_uri,
"client_id": cfg.client_id,
"client_secret": cfg.get_client_secret(),
"redirect_uri": redirect_uri,
"client_id": client_id,
"client_secret": client_secret,
},
)
try:

View File

@@ -29,6 +29,8 @@ from src.services.billing.models import (
CostBreakdown,
StandardizedUsage,
)
from src.services.billing.schema import BillingSnapshot, CostResult
from src.services.billing.service import BillingService
from src.services.billing.templates import BILLING_TEMPLATE_REGISTRY, BillingTemplates
from src.services.billing.usage_mapper import UsageMapper, map_usage, map_usage_from_response
@@ -44,6 +46,10 @@ __all__ = [
# 计算器
"BillingCalculator",
"calculate_request_cost",
# 统一入口Phase2
"BillingService",
"BillingSnapshot",
"CostResult",
# 映射器
"UsageMapper",
"map_usage",

View File

@@ -0,0 +1,64 @@
"""
Billing schema (stable contracts)
These dataclasses are meant to be stored in `Usage.request_metadata` / `Task.request_metadata`
for auditability. They are internal-only and MUST NOT be exposed to end users without sanitizing.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Literal
BILLING_SNAPSHOT_SCHEMA_VERSION = "1.0"
BillingSnapshotStatus = Literal["complete", "incomplete", "no_rule", "legacy"]
@dataclass(frozen=True)
class BillingSnapshot:
"""Stable billing snapshot for audit."""
schema_version: str = BILLING_SNAPSHOT_SCHEMA_VERSION
# Rule info (optional for legacy/no_rule)
rule_id: str | None = None
rule_name: str | None = None
scope: str | None = None
# Rule expression (internal, do not expose to clients)
expression: str | None = None
# Dimensions
dimensions_used: dict[str, Any] = field(default_factory=dict)
missing_required: list[str] = field(default_factory=list)
# Result
cost: float = 0.0
status: BillingSnapshotStatus = "no_rule"
# Audit
calculated_at: str = "" # ISO 8601
def to_dict(self) -> dict[str, Any]:
return {
"schema_version": self.schema_version,
"rule_id": self.rule_id,
"rule_name": self.rule_name,
"scope": self.scope,
"expression": self.expression,
"dimensions_used": self.dimensions_used,
"missing_required": self.missing_required,
"cost": self.cost,
"status": self.status,
"calculated_at": self.calculated_at,
}
@dataclass(frozen=True)
class CostResult:
"""Billing calculation output."""
cost: float
status: BillingSnapshotStatus
snapshot: BillingSnapshot

View File

@@ -0,0 +1,145 @@
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.logger import logger
from src.services.billing.dimension_collector_service import DimensionCollectorService
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
from src.services.billing.rule_service import BillingRuleService
from src.services.model.cost import ModelCostService
from .schema import BILLING_SNAPSHOT_SCHEMA_VERSION, BillingSnapshot, CostResult
class BillingService:
"""
BillingService (pure-ish application helper for billing domain).
Notes:
- This service **does not** write Usage rows.
- It may read billing rules & collectors from DB.
"""
def __init__(self, db: Session):
self.db = db
self._formula_engine = FormulaEngine()
self._dimension_collector = DimensionCollectorService(db)
def collect_dimensions(
self,
*,
api_format: str | None,
task_type: str | None,
request: dict[str, Any] | None = None,
response: dict[str, Any] | None = None,
metadata: dict[str, Any] | None = None,
base_dimensions: dict[str, Any] | None = None,
) -> dict[str, Any]:
return self._dimension_collector.collect_dimensions(
api_format=api_format,
task_type=task_type,
request=request,
response=response,
metadata=metadata,
base_dimensions=base_dimensions,
)
def calculate(
self,
*,
task_type: str,
model: str,
provider_id: str,
dimensions: dict[str, Any],
strict_mode: bool | None = None,
) -> CostResult:
"""
Calculate cost for a task.
Returns:
CostResult (includes BillingSnapshot)
Raises:
BillingIncompleteError: when strict_mode=True and required dims missing.
"""
strict = config.billing_strict_mode if strict_mode is None else bool(strict_mode)
lookup = BillingRuleService.find_rule(
self.db,
provider_id=provider_id,
model_name=model,
task_type=task_type,
)
if lookup and lookup.rule and lookup.rule.expression:
rule = lookup.rule
result = self._formula_engine.evaluate(
expression=rule.expression,
variables=rule.variables or {},
dimensions=dimensions,
dimension_mappings=rule.dimension_mappings or {},
strict_mode=strict,
)
cost = float(result.cost) if result.status == "complete" else 0.0
snapshot = BillingSnapshot(
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
rule_id=str(rule.id),
rule_name=str(rule.name),
scope=str(getattr(lookup, "scope", None) or ""),
expression=str(rule.expression),
dimensions_used=dimensions,
missing_required=result.missing_required,
cost=cost,
status=result.status,
calculated_at=datetime.now(timezone.utc).isoformat(),
)
return CostResult(cost=cost, status=result.status, snapshot=snapshot)
# No rule fallback
if task_type in ("chat", "cli"):
input_tokens = int(dimensions.get("input_tokens") or 0)
output_tokens = int(dimensions.get("output_tokens") or 0)
cost = float(
ModelCostService.calculate_cost(
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
)
)
snapshot = BillingSnapshot(
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
rule_id=None,
rule_name=None,
scope=None,
expression=None,
dimensions_used=dimensions,
missing_required=[],
cost=cost,
status="legacy",
calculated_at=datetime.now(timezone.utc).isoformat(),
)
return CostResult(cost=cost, status="legacy", snapshot=snapshot)
logger.warning(
"No billing rule for task (task_type=%s, model=%s, provider_id=%s)",
task_type,
model,
provider_id,
)
snapshot = BillingSnapshot(
schema_version=BILLING_SNAPSHOT_SCHEMA_VERSION,
rule_id=None,
rule_name=None,
scope=None,
expression=None,
dimensions_used=dimensions,
missing_required=[],
cost=0.0,
status="no_rule",
calculated_at=datetime.now(timezone.utc).isoformat(),
)
return CostResult(cost=0.0, status="no_rule", snapshot=snapshot)

View File

@@ -182,6 +182,43 @@ class CacheAwareScheduler:
"last_reservation_result": None,
}
@staticmethod
def _release_db_connection_before_await(db: Session) -> None:
"""
Best-effort: end a read-only transaction before awaiting async I/O.
This scheduler does a lot of async work (cache/Redis) mixed with sync SQLAlchemy reads.
If a SELECT has already started a transaction, the pooled connection can remain checked
out while we await, causing pool pressure under concurrency.
Safety:
- Only commits when the Session has no ORM pending changes.
- Temporarily disables expire_on_commit to keep already-loaded ORM objects usable.
"""
try:
if db is None:
return
has_pending_changes = bool(db.new) or bool(db.dirty) or bool(db.deleted)
if has_pending_changes:
return
if not db.in_transaction():
return
original_expire_on_commit = getattr(db, "expire_on_commit", True)
db.expire_on_commit = False
try:
db.commit()
except Exception:
try:
db.rollback()
except Exception:
pass
finally:
db.expire_on_commit = original_expire_on_commit
except Exception:
# Never let this optimization break scheduling
return
async def _ensure_initialized(self) -> None:
"""确保所有异步组件已初始化"""
if self._affinity_manager is None:
@@ -577,6 +614,8 @@ class CacheAwareScheduler:
Returns:
(候选列表, global_model_id) - global_model_id 用于缓存亲和性
"""
# If the caller already touched the DB, release the connection before we do async work.
self._release_db_connection_before_await(db)
await self._ensure_initialized()
target_format = normalize_endpoint_signature(api_format)
@@ -648,6 +687,9 @@ class CacheAwareScheduler:
provider_limit=provider_limit,
)
# Provider query starts a transaction; release connection before entering async candidate build.
self._release_db_connection_before_await(db)
logger.debug(
"[Scheduler] Found %d active providers",
len(providers),
@@ -680,8 +722,13 @@ class CacheAwareScheduler:
# 2. 构建候选列表(传入 is_stream 和 capability_requirements 用于过滤)
from src.config.settings import config
from src.services.system.config import SystemConfigService
global_conversion_enabled = config.format_conversion_enabled
# 全局格式转换开关:优先使用数据库配置,回退到环境变量
global_conversion_enabled = SystemConfigService.is_format_conversion_enabled(db)
# 如果环境变量明确禁用,则禁用(环境变量可作为强制禁用开关)
if not config.format_conversion_enabled:
global_conversion_enabled = False
candidates = await self._build_candidates(
db=db,
providers=providers,
@@ -801,6 +848,9 @@ class CacheAwareScheduler:
- supported_capabilities: 模型支持的能力列表
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
"""
# Avoid holding a DB connection while awaiting cache/Redis inside ModelCacheService.
self._release_db_connection_before_await(db)
# 使用 ModelCacheService 解析模型名称(支持映射名)
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(
db, model_name
@@ -1120,18 +1170,34 @@ class CacheAwareScheduler:
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
# 计算格式转换的有效开关状态(三层优先级)
# 全局 ON → 强制允许(跳过端点检查)
# 全局 OFF → 提供商 ON → 强制允许(跳过端点检查)
# 全局 OFF → 提供商 OFF → 看端点配置
provider_allows_conversion = getattr(provider, "enable_format_conversion", True)
effective_conversion_enabled = (
global_conversion_enabled or provider_allows_conversion
)
# 如果全局或提供商开关为 ON跳过端点配置检查
skip_endpoint_check = global_conversion_enabled or provider_allows_conversion
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
client_format_str,
endpoint_format_str,
getattr(endpoint, "format_acceptance_config", None),
is_stream,
global_conversion_enabled,
effective_conversion_enabled,
skip_endpoint_check=skip_endpoint_check,
)
logger.debug(
"[Scheduler] Format compatibility: client=%s, endpoint=%s, compatible=%s, reason=%s",
"[Scheduler] Format compatibility: client=%s, endpoint=%s, compatible=%s, "
"global=%s, provider=%s, skip_endpoint=%s, reason=%s",
client_format_str,
endpoint_format_str,
is_compatible,
global_conversion_enabled,
provider_allows_conversion,
skip_endpoint_check,
_compat_reason,
)
if not is_compatible:

View File

@@ -0,0 +1,28 @@
"""
Candidate domain (Phase2)
This package centralizes:
- candidate resolving (Provider/Endpoint/Key combinations)
- request_candidates recording & audit
- failover execution policies
"""
from src.services.candidate.policy import RetryMode, RetryPolicy, SkipPolicy
from src.services.candidate.schema import (
CANDIDATE_KEY_SCHEMA_VERSION,
CandidateKey,
CandidateResult,
)
from src.services.candidate.service import CandidateService
__all__ = [
"CandidateService",
# schema
"CANDIDATE_KEY_SCHEMA_VERSION",
"CandidateKey",
"CandidateResult",
# policies
"RetryMode",
"RetryPolicy",
"SkipPolicy",
]

View File

@@ -0,0 +1,57 @@
from __future__ import annotations
from typing import Any, Protocol
from sqlalchemy.orm import Session
from src.services.cache.aware_scheduler import ProviderCandidate
from .policy import RetryPolicy, SkipPolicy
from .schema import CandidateKey, CandidateResult
class AttemptFunc(Protocol):
async def __call__(self, candidate: ProviderCandidate) -> Any: ...
class FailoverEngine:
"""
FailoverEngine executes candidate attempts under policies.
Phase2 scaffolding: implementation will gradually replace legacy orchestrators.
"""
def __init__(self, db: Session):
self.db = db
async def execute(
self,
*,
candidates: list[ProviderCandidate],
attempt_func: AttemptFunc,
retry_policy: RetryPolicy,
skip_policy: SkipPolicy,
request_id: str | None = None,
max_candidates: int | None = None,
) -> CandidateResult:
# NOTE: intentionally minimal for now; legacy orchestrators still in use.
# This will be implemented when migrating video/chat flows to CandidateService.
_ = (retry_policy, skip_policy, request_id, max_candidates)
candidate_keys: list[CandidateKey] = []
for idx, cand in enumerate(candidates):
candidate_keys.append(
CandidateKey(
candidate_index=idx,
provider_id=str(cand.provider.id),
provider_name=str(cand.provider.name),
endpoint_id=str(cand.endpoint.id),
key_id=str(cand.key.id),
key_name=str(getattr(cand.key, "name", "") or ""),
auth_type=str(getattr(cand.key, "auth_type", "") or ""),
priority=int(getattr(cand.key, "priority", 0) or 0),
is_cached=bool(getattr(cand, "is_cached", False)),
status="available",
)
)
raise NotImplementedError("FailoverEngine.execute is not implemented yet")

View File

@@ -0,0 +1,41 @@
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
class RetryMode(str, Enum):
"""Retry mode for candidate attempts."""
PRE_EXPAND = "pre_expand" # pre-create retry slots (sync)
ON_DEMAND = "on_demand" # create retry record only when retry happens
DISABLED = "disabled" # no retry
@dataclass(frozen=True)
class RetryPolicy:
"""Unified retry policy."""
mode: RetryMode = RetryMode.DISABLED
max_retries: int = 1
retry_on_cached_only: bool = True
@classmethod
def for_sync_task(cls) -> "RetryPolicy":
return cls(mode=RetryMode.PRE_EXPAND, max_retries=2)
@classmethod
def for_async_task(cls) -> "RetryPolicy":
return cls(mode=RetryMode.DISABLED, max_retries=1)
@classmethod
def for_async_submit_with_retry(cls) -> "RetryPolicy":
return cls(mode=RetryMode.ON_DEMAND, max_retries=2)
@dataclass(frozen=True)
class SkipPolicy:
"""Rules for skipping unsupported candidates."""
allow_format_conversion: bool = True
supported_auth_types: set[str] | None = None

View File

@@ -0,0 +1,59 @@
from __future__ import annotations
from sqlalchemy.orm import Session
from src.models.database import RequestCandidate
from .schema import CandidateKey
class CandidateRecorder:
"""Read helpers for RequestCandidate audit data."""
def __init__(self, db: Session):
self.db = db
def get_candidate_keys(self, request_id: str) -> list[CandidateKey]:
rows: list[RequestCandidate] = (
self.db.query(RequestCandidate)
.filter(RequestCandidate.request_id == request_id)
.order_by(RequestCandidate.candidate_index.asc(), RequestCandidate.retry_index.asc())
.all()
)
result: list[CandidateKey] = []
for row in rows:
provider_name = None
if getattr(row, "provider", None) is not None:
provider_name = getattr(row.provider, "name", None)
key_name = None
auth_type = None
priority = None
if getattr(row, "key", None) is not None:
key_name = getattr(row.key, "name", None)
auth_type = getattr(row.key, "auth_type", None)
priority = getattr(row.key, "priority", None)
result.append(
CandidateKey(
candidate_index=int(row.candidate_index or 0),
retry_index=int(row.retry_index or 0),
provider_id=str(row.provider_id) if row.provider_id else None,
provider_name=str(provider_name) if provider_name else None,
endpoint_id=str(row.endpoint_id) if row.endpoint_id else None,
key_id=str(row.key_id) if row.key_id else None,
key_name=str(key_name) if key_name else None,
auth_type=str(auth_type) if auth_type else None,
priority=int(priority) if priority is not None else None,
is_cached=bool(getattr(row, "is_cached", False)),
status=str(getattr(row, "status", "") or "pending"),
skip_reason=getattr(row, "skip_reason", None),
error_type=getattr(row, "error_type", None),
error_message=getattr(row, "error_message", None),
status_code=getattr(row, "status_code", None),
latency_ms=getattr(row, "latency_ms", None),
)
)
return result

View File

@@ -0,0 +1,10 @@
"""
CandidateResolver facade import.
Phase2 keeps the implementation in `services/orchestration/` for compatibility,
and gradually migrates it into this package.
"""
from src.services.orchestration.candidate_resolver import CandidateResolver
__all__ = ["CandidateResolver"]

View File

@@ -0,0 +1,71 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from src.services.cache.aware_scheduler import ProviderCandidate
CANDIDATE_KEY_SCHEMA_VERSION = "1.0"
@dataclass(frozen=True)
class CandidateKey:
"""Stable candidate key snapshot for audit."""
schema_version: str = CANDIDATE_KEY_SCHEMA_VERSION
candidate_index: int = 0
retry_index: int = 0
provider_id: str | None = None
provider_name: str | None = None
endpoint_id: str | None = None
key_id: str | None = None
key_name: str | None = None
auth_type: str | None = None
priority: int | None = None
is_cached: bool = False
status: str = "pending" # pending/success/failed/skipped/available/...
skip_reason: str | None = None
error_type: str | None = None
error_message: str | None = None
status_code: int | None = None
latency_ms: int | None = None
def to_dict(self) -> dict[str, Any]:
data: dict[str, Any] = {
"schema_version": self.schema_version,
"candidate_index": self.candidate_index,
"retry_index": self.retry_index,
"provider_id": self.provider_id,
"provider_name": self.provider_name,
"endpoint_id": self.endpoint_id,
"key_id": self.key_id,
"key_name": self.key_name,
"auth_type": self.auth_type,
"priority": self.priority,
"is_cached": self.is_cached,
"status": self.status,
"skip_reason": self.skip_reason,
"error_type": self.error_type,
"error_message": self.error_message,
"status_code": self.status_code,
"latency_ms": self.latency_ms,
}
# drop Nones for compact audit payload
return {k: v for k, v in data.items() if v is not None}
@dataclass(slots=True)
class CandidateResult:
"""Failover execution result."""
success: bool
selected: ProviderCandidate | None
selected_index: int | None
candidate_keys: list[CandidateKey]
external_task_id: str | None = None
error: Exception | None = None
last_status_code: int | None = None

View File

@@ -0,0 +1,515 @@
from __future__ import annotations
import re
from datetime import datetime, timezone
from typing import Any
import httpx
from sqlalchemy import update
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
from src.services.candidate.submit import (
AllCandidatesFailedError,
SubmitOutcome,
UpstreamClientRequestError,
)
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.system.config import SystemConfigService
from .recorder import CandidateRecorder
from .resolver import CandidateResolver
from .schema import CandidateKey
_SENSITIVE_PATTERN = re.compile(
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
re.IGNORECASE,
)
def _sanitize(message: str, max_length: int = 200) -> str:
if not message:
return "request_failed"
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
class CandidateService:
"""
CandidateService (Facade).
Phase2 note: this is introduced as a new domain entrypoint. Legacy orchestrators
still exist and will be migrated gradually to use this service.
"""
def __init__(self, db: Session, redis_client: Any | None = None) -> None:
self.db = db
self.redis = redis_client
self._cache_scheduler = None
self._resolver: CandidateResolver | None = None
self._error_classifier: ErrorClassifier | None = None
self._recorder = CandidateRecorder(db)
async def _ensure_initialized(self) -> None:
if self._cache_scheduler is not None:
return
priority_mode = SystemConfigService.get_config(
self.db,
"provider_priority_mode",
"provider",
)
scheduling_mode = SystemConfigService.get_config(
self.db,
"scheduling_mode",
"cache_affinity",
)
self._cache_scheduler = await get_cache_aware_scheduler(
self.redis,
priority_mode=priority_mode,
scheduling_mode=scheduling_mode,
)
self._resolver = CandidateResolver(db=self.db, cache_scheduler=self._cache_scheduler)
self._error_classifier = ErrorClassifier(db=self.db, cache_scheduler=self._cache_scheduler)
async def resolve(
self,
*,
api_format: str,
model_name: str,
affinity_key: str,
user_api_key: ApiKey | None = None,
request_id: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
preferred_key_ids: list[str] | None = None,
) -> tuple[list[ProviderCandidate], str]:
await self._ensure_initialized()
assert self._resolver is not None
return await self._resolver.fetch_candidates(
api_format=api_format,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
request_id=request_id,
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
)
def _should_stop_on_http_error(self, *, status_code: int, error_text: str) -> bool:
"""
Decide whether an upstream HTTP error is a client error (no failover).
Rules:
- 401/403/429 are usually key/permission/ratelimit issues -> allow failover
- other 4xx: stop only if ErrorClassifier says it's a client error
"""
if status_code in (401, 403, 429):
return False
if 400 <= status_code < 500:
assert self._error_classifier is not None
return self._error_classifier.is_client_error(error_text)
return False
async def submit_with_failover(
self,
*,
api_format: str,
model_name: str,
affinity_key: str,
user_api_key: ApiKey,
request_id: str | None,
task_type: str,
submit_func: Any,
extract_external_task_id: Any,
supported_auth_types: set[str] | None = None,
allow_format_conversion: bool = False,
capability_requirements: dict[str, bool] | None = None,
max_candidates: int | None = None,
) -> SubmitOutcome:
"""
Submit async task with failover, returning the selected candidate + external_task_id.
Phase2 submit entrypoint (replaces legacy submit orchestrator).
"""
# IMPORTANT:
# This method awaits upstream HTTP calls. If we have an open DB transaction before awaiting,
# the connection can be held for a long time (pool exhaustion under concurrency).
#
# Also note SQLAlchemy's default expire_on_commit=True would expire ORM objects and may
# trigger unexpected lazy DB loads after we commit (potentially during the await).
# We disable it temporarily to keep candidate/provider/key objects in-memory.
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
self.db.expire_on_commit = False
await self._ensure_initialized()
assert self._resolver is not None
try:
candidates, _global_model_id = await self._resolver.fetch_candidates(
api_format=api_format,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
request_id=request_id,
is_stream=False,
capability_requirements=capability_requirements,
)
if not candidates:
raise ProviderNotAvailableException("No candidates available")
if max_candidates is not None and max_candidates > 0:
candidates = candidates[:max_candidates]
# Pre-create RequestCandidate records (no retry expand for async submit stage)
record_map: dict[tuple[int, int], str] = {}
if request_id:
try:
record_map = self.create_candidate_records(
candidates=candidates,
request_id=request_id,
user_api_key=user_api_key,
required_capabilities=capability_requirements,
expand_retries=False,
)
except Exception as exc:
logger.warning(
"[CandidateService] Failed to create candidate records: %s",
_sanitize(str(exc)),
)
record_map = {}
candidate_keys: list[dict[str, Any]] = []
eligible_count = 0
last_status_code: int | None = None
for idx, cand in enumerate(candidates):
now = datetime.now(timezone.utc)
auth_type = getattr(cand.key, "auth_type", "api_key") or "api_key"
candidate_info: dict[str, Any] = {
"index": idx,
"provider_id": cand.provider.id,
"provider_name": cand.provider.name,
"endpoint_id": cand.endpoint.id,
"key_id": cand.key.id,
"key_name": getattr(cand.key, "name", None),
"auth_type": auth_type,
"priority": getattr(cand.key, "priority", 0) or 0,
"is_cached": bool(getattr(cand, "is_cached", False)),
}
candidate_keys.append(candidate_info)
record_id = record_map.get((idx, 0))
# Scheduler marked skip
if getattr(cand, "is_skipped", False):
skip_reason = getattr(cand, "skip_reason", None) or "skipped"
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
if record_id:
# record is usually already skipped, but keep it consistent
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(status="skipped", skip_reason=skip_reason)
)
continue
# Format conversion checks
# 优先级:全局开关 ON 强制允许,全局开关 OFF 看提供商开关
needs_conversion = bool(getattr(cand, "needs_conversion", False))
if needs_conversion:
# 1. Check handler-level switch (handler 不支持则直接跳过)
if not allow_format_conversion:
skip_reason = "format_conversion_not_supported"
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(status="skipped", skip_reason=skip_reason)
)
continue
# 2. Check global + provider switches
# 全局 ON → 允许;全局 OFF → 看提供商
from src.services.system.config import SystemConfigService
global_enabled = SystemConfigService.is_format_conversion_enabled(self.db)
provider_enabled = getattr(cand.provider, "enable_format_conversion", True)
effective_enabled = global_enabled or provider_enabled
if not effective_enabled:
skip_reason = "format_conversion_disabled"
candidate_info.update(
{
"skipped": True,
"skip_reason": skip_reason,
"global_conversion_enabled": global_enabled,
"provider_conversion_enabled": provider_enabled,
}
)
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(status="skipped", skip_reason=skip_reason)
)
continue
# auth_type filter
if supported_auth_types is not None and auth_type not in supported_auth_types:
skip_reason = f"unsupported_auth_type:{auth_type}"
candidate_info.update({"skipped": True, "skip_reason": skip_reason})
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(status="skipped", skip_reason=skip_reason)
)
continue
# billing rule filter
rule_lookup: BillingRuleLookupResult | None = None
has_billing_rule = True
if config.billing_require_rule:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=cand.provider.id,
model_name=model_name,
task_type=task_type,
)
has_billing_rule = rule_lookup is not None
if not has_billing_rule:
skip_reason = "billing_rule_missing"
candidate_info.update(
{"has_billing_rule": False, "skipped": True, "skip_reason": skip_reason}
)
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(status="skipped", skip_reason=skip_reason)
)
continue
candidate_info["has_billing_rule"] = has_billing_rule
eligible_count += 1
# Mark pending
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(status="pending", started_at=now)
)
# Flush/commit BEFORE awaiting upstream submit to avoid holding DB connections
# during potentially slow network operations.
if self.db.in_transaction():
try:
self.db.commit()
except Exception:
self.db.rollback()
raise
# Attempt submit (upstream HTTP)
try:
response: httpx.Response = await submit_func(cand)
except Exception as exc:
finished_at = datetime.now(timezone.utc)
error_msg = _sanitize(str(exc))
candidate_info.update(
{
"attempt_status": "exception",
"error_type": type(exc).__name__,
"error_message": error_msg,
}
)
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(
status="failed",
error_type=type(exc).__name__,
error_message=error_msg,
finished_at=finished_at,
)
)
continue
last_status_code = int(getattr(response, "status_code", 0) or 0)
if response.status_code >= 400:
finished_at = datetime.now(timezone.utc)
try:
error_text = response.text or ""
except Exception:
error_text = ""
error_msg = _sanitize(error_text)
candidate_info.update(
{
"attempt_status": "http_error",
"status_code": response.status_code,
"error_message": error_msg,
}
)
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(
status="failed",
status_code=response.status_code,
error_type="http_error",
error_message=error_msg,
finished_at=finished_at,
)
)
if self._should_stop_on_http_error(
status_code=response.status_code, error_text=error_text
):
try:
self.db.commit()
except Exception:
self.db.rollback()
raise UpstreamClientRequestError(
response=response,
candidate_keys=candidate_keys,
)
continue
# Parse JSON
payload: dict[str, Any] | None = None
try:
data = response.json()
if isinstance(data, dict):
payload = data
except Exception as exc:
finished_at = datetime.now(timezone.utc)
error_msg = _sanitize(str(exc))
candidate_info.update(
{
"attempt_status": "invalid_json",
"error_type": type(exc).__name__,
"error_message": error_msg,
}
)
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(
status="failed",
status_code=response.status_code,
error_type="invalid_json",
error_message=error_msg,
finished_at=finished_at,
)
)
continue
external_task_id = extract_external_task_id(payload or {})
if not external_task_id:
finished_at = datetime.now(timezone.utc)
candidate_info.update(
{
"attempt_status": "empty_task_id",
"error_message": "Upstream returned empty task id",
}
)
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(
status="failed",
status_code=response.status_code,
error_type="empty_task_id",
error_message="Upstream returned empty task id",
finished_at=finished_at,
)
)
continue
# Success
finished_at = datetime.now(timezone.utc)
candidate_info.update({"attempt_status": "success", "selected": True})
if record_id:
self.db.execute(
update(RequestCandidate)
.where(RequestCandidate.id == record_id)
.values(
status="success",
status_code=response.status_code,
finished_at=finished_at,
)
)
try:
self.db.commit()
except Exception:
self.db.rollback()
return SubmitOutcome(
candidate=cand,
candidate_keys=candidate_keys,
external_task_id=str(external_task_id),
rule_lookup=rule_lookup,
upstream_payload=payload,
upstream_headers=dict(response.headers),
upstream_status_code=response.status_code,
)
# Persist candidate records before raising
try:
self.db.commit()
except Exception:
self.db.rollback()
if eligible_count == 0:
reason = "no_eligible_candidates"
if config.billing_require_rule:
reason = "no_candidate_with_billing_rule"
raise AllCandidatesFailedError(
reason=reason,
candidate_keys=candidate_keys,
last_status_code=last_status_code,
)
raise AllCandidatesFailedError(
reason="all_candidates_failed",
candidate_keys=candidate_keys,
last_status_code=last_status_code,
)
finally:
# Restore Session behavior for the rest of the request lifecycle.
self.db.expire_on_commit = original_expire_on_commit
def create_candidate_records(
self,
*,
candidates: list[ProviderCandidate],
request_id: str,
user_api_key: ApiKey,
required_capabilities: dict[str, bool] | None = None,
expand_retries: bool = True,
) -> dict[tuple[int, int], str]:
# CandidateResolver.create_candidate_records is currently the canonical implementation
assert self._resolver is not None, "Call resolve() once before create_candidate_records()"
return self._resolver.create_candidate_records(
all_candidates=candidates,
request_id=request_id,
user_id=str(user_api_key.user_id),
user_api_key=user_api_key,
required_capabilities=required_capabilities,
expand_retries=expand_retries,
)
def get_candidate_keys(self, request_id: str) -> list["CandidateKey"]:
return self._recorder.get_candidate_keys(request_id)

View File

@@ -0,0 +1,66 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Protocol, runtime_checkable
import httpx
from src.services.billing.rule_service import BillingRuleLookupResult
from src.services.cache.aware_scheduler import ProviderCandidate
@runtime_checkable
class SubmitFunc(Protocol):
async def __call__(self, candidate: ProviderCandidate) -> httpx.Response: ...
@runtime_checkable
class ExtractExternalTaskIdFunc(Protocol):
def __call__(self, payload: dict[str, Any]) -> str | None: ...
class UpstreamClientRequestError(RuntimeError):
"""可判定为客户端请求问题(不应 failover的上游错误。"""
def __init__(
self,
*,
response: httpx.Response,
candidate_keys: list[dict[str, Any]],
) -> None:
self.response = response
self.candidate_keys = candidate_keys
super().__init__(f"Upstream client error: HTTP {response.status_code}")
class AllCandidatesFailedError(RuntimeError):
def __init__(
self,
*,
reason: str,
candidate_keys: list[dict[str, Any]],
last_status_code: int | None = None,
) -> None:
self.reason = reason
self.candidate_keys = candidate_keys
self.last_status_code = last_status_code
super().__init__(f"All candidates failed: {reason}")
class CandidateUnsupportedError(RuntimeError):
"""候选不被当前任务支持(如 auth_type/格式转换需求不支持)。"""
class CandidateSubmissionError(RuntimeError):
"""候选提交异常(网络/解密/解析等)。"""
@dataclass(slots=True)
class SubmitOutcome:
candidate: ProviderCandidate
candidate_keys: list[dict[str, Any]]
external_task_id: str
rule_lookup: BillingRuleLookupResult | None
upstream_payload: dict[str, Any] | None = None
upstream_headers: dict[str, str] | None = None
upstream_status_code: int | None = None

View File

@@ -1,21 +1,41 @@
"""
Gemini Files API - 文件与 Key 绑定缓存
Gemini Files API - 文件与 Key 绑定映射服务
用于在上传文件后记录 file_id -> provider_key_id
并在后续 generateContent 请求中优先使用同一 Key。
存储策略:
- 数据库(持久化):主存储,支持服务重启后恢复
- Redis缓存加速读取TTL=48小时
读取策略:
1. 先查 Redis 缓存
2. 缓存未命中时回查数据库
3. 从数据库读取后回填缓存
清理策略:
- 数据库中 expires_at 过期的记录由定时任务清理
- Redis 缓存由 TTL 自动过期
"""
from __future__ import annotations
from typing import Any, Dict, Optional, Set
import uuid
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import delete
from sqlalchemy.orm import Session
from src.core.cache_service import CacheService
from src.core.logger import logger
FILE_MAPPING_TTL_SECONDS = 60 * 60 * 48 # 48小时
FILE_MAPPING_CACHE_PREFIX = "gemini_files:key"
def _normalize_file_name(file_name: str) -> str:
"""规范化文件名,确保以 files/ 开头"""
name = (file_name or "").strip()
if not name:
return ""
@@ -23,31 +43,294 @@ def _normalize_file_name(file_name: str) -> str:
def build_file_mapping_key(file_name: str) -> str:
"""构建 Redis 缓存键"""
normalized = _normalize_file_name(file_name)
return f"{FILE_MAPPING_CACHE_PREFIX}:{normalized}" if normalized else ""
async def store_file_key_mapping(file_name: str, key_id: str) -> None:
cache_key = build_file_mapping_key(file_name)
if not cache_key or not key_id:
# =============================================================================
# 异步接口(用于请求处理流程)
# =============================================================================
async def store_file_key_mapping(
file_name: str,
key_id: str,
user_id: str | None = None,
display_name: str | None = None,
mime_type: str | None = None,
source_hash: str | None = None,
) -> None:
"""
存储文件→Key 映射(同时写入 Redis 和数据库)
Args:
file_name: 文件名(如 files/abc123
key_id: Provider Key ID
user_id: 用户 ID可选用于权限验证
display_name: 文件显示名(可选)
mime_type: 文件 MIME 类型(可选)
source_hash: 源文件哈希(可选,用于关联相同源文件的不同上传)
"""
normalized_name = _normalize_file_name(file_name)
if not normalized_name or not key_id:
return
# 1. 写入 Redis 缓存
cache_key = build_file_mapping_key(normalized_name)
await CacheService.set(cache_key, str(key_id), ttl_seconds=FILE_MAPPING_TTL_SECONDS)
# 2. 写入数据库(异步执行,不阻塞主流程)
try:
await _store_to_database(
file_name=normalized_name,
key_id=key_id,
user_id=user_id,
display_name=display_name,
mime_type=mime_type,
source_hash=source_hash,
)
except Exception as e:
# 数据库写入失败只记录警告,不影响主流程
logger.warning(f"Failed to persist Gemini file mapping to database: {e}")
async def _store_to_database(
file_name: str,
key_id: str,
user_id: str | None = None,
display_name: str | None = None,
mime_type: str | None = None,
source_hash: str | None = None,
) -> None:
"""将映射写入数据库"""
from src.database import get_db_context
from src.models.database import GeminiFileMapping
now = datetime.now(timezone.utc)
expires_at = now + timedelta(hours=48)
with get_db_context() as db:
# 使用 upsert 逻辑:存在则更新,不存在则插入
existing = (
db.query(GeminiFileMapping).filter(GeminiFileMapping.file_name == file_name).first()
)
if existing:
# 更新现有记录
existing.key_id = key_id
existing.user_id = user_id
existing.display_name = display_name
existing.mime_type = mime_type
existing.source_hash = source_hash
existing.expires_at = expires_at
else:
# 插入新记录
mapping = GeminiFileMapping(
id=str(uuid.uuid4()),
file_name=file_name,
key_id=key_id,
user_id=user_id,
display_name=display_name,
mime_type=mime_type,
source_hash=source_hash,
created_at=now,
expires_at=expires_at,
)
db.add(mapping)
db.commit()
async def get_file_key_mapping(file_name: str) -> str | None:
cache_key = build_file_mapping_key(file_name)
if not cache_key:
"""
获取文件→Key 映射
读取策略:
1. 先查 Redis 缓存
2. 缓存未命中时回查数据库
3. 从数据库读取后回填缓存
Args:
file_name: 文件名(如 files/abc123
Returns:
Provider Key ID如果不存在或已过期则返回 None
"""
normalized_name = _normalize_file_name(file_name)
if not normalized_name:
return None
value = await CacheService.get(cache_key)
if value:
return str(value)
cache_key = build_file_mapping_key(normalized_name)
# 1. 先查 Redis 缓存
cached_value = await CacheService.get(cache_key)
if cached_value:
return str(cached_value)
# 2. 缓存未命中,回查数据库
key_id = await _get_from_database(normalized_name)
if key_id:
# 3. 回填缓存(使用剩余有效期或默认 TTL
await CacheService.set(cache_key, key_id, ttl_seconds=FILE_MAPPING_TTL_SECONDS)
logger.debug(f"Gemini file mapping cache refilled from database: {normalized_name}")
return key_id
async def _get_from_database(file_name: str) -> str | None:
"""从数据库查询映射"""
from src.database import get_db_context
from src.models.database import GeminiFileMapping
now = datetime.now(timezone.utc)
try:
with get_db_context() as db:
mapping = (
db.query(GeminiFileMapping)
.filter(
GeminiFileMapping.file_name == file_name,
GeminiFileMapping.expires_at > now, # 只返回未过期的
)
.first()
)
if mapping:
return str(mapping.key_id)
except Exception as e:
logger.warning(f"Failed to query Gemini file mapping from database: {e}")
return None
async def get_all_key_ids_for_file(file_name: str) -> list[str]:
"""
获取支持指定文件的所有 Key ID 列表
当同一个源文件被上传到多个 Key 时,返回所有可用的 Key ID。
这允许系统在首选 Key 不可用时选择其他 Key。
Args:
file_name: 文件名(如 files/abc123
Returns:
所有支持该文件的 Key ID 列表(包括原始映射和具有相同 source_hash 的映射)
"""
from src.database import get_db_context
from src.models.database import GeminiFileMapping
normalized_name = _normalize_file_name(file_name)
if not normalized_name:
return []
now = datetime.now(timezone.utc)
try:
with get_db_context() as db:
# 首先获取原始映射
original_mapping = (
db.query(GeminiFileMapping)
.filter(
GeminiFileMapping.file_name == normalized_name,
GeminiFileMapping.expires_at > now,
)
.first()
)
if not original_mapping:
return []
key_ids = [str(original_mapping.key_id)]
# 如果有 source_hash查找所有具有相同 source_hash 的映射
if original_mapping.source_hash:
related_mappings = (
db.query(GeminiFileMapping)
.filter(
GeminiFileMapping.source_hash == original_mapping.source_hash,
GeminiFileMapping.expires_at > now,
GeminiFileMapping.file_name != normalized_name, # 排除原始映射
)
.all()
)
for mapping in related_mappings:
kid = str(mapping.key_id)
if kid not in key_ids:
key_ids.append(kid)
return key_ids
except Exception as e:
logger.warning(f"Failed to query related Gemini file mappings: {e}")
return []
async def delete_file_key_mapping(file_name: str) -> None:
cache_key = build_file_mapping_key(file_name)
if cache_key:
await CacheService.delete(cache_key)
"""
删除文件→Key 映射(同时从 Redis 和数据库删除)
Args:
file_name: 文件名(如 files/abc123
"""
normalized_name = _normalize_file_name(file_name)
if not normalized_name:
return
# 1. 从 Redis 删除
cache_key = build_file_mapping_key(normalized_name)
await CacheService.delete(cache_key)
# 2. 从数据库删除
try:
await _delete_from_database(normalized_name)
except Exception as e:
logger.warning(f"Failed to delete Gemini file mapping from database: {e}")
async def _delete_from_database(file_name: str) -> None:
"""从数据库删除映射"""
from src.database import get_db_context
from src.models.database import GeminiFileMapping
with get_db_context() as db:
db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.file_name == file_name))
db.commit()
# =============================================================================
# 同步接口(用于定时任务等场景)
# =============================================================================
def cleanup_expired_mappings(db: Session) -> int:
"""
清理过期的文件映射记录(同步方法,供定时任务调用)
Args:
db: 数据库会话
Returns:
删除的记录数
"""
from src.models.database import GeminiFileMapping
now = datetime.now(timezone.utc)
result = db.execute(delete(GeminiFileMapping).where(GeminiFileMapping.expires_at <= now))
db.commit()
deleted_count = result.rowcount
if deleted_count > 0:
logger.info(f"Cleaned up {deleted_count} expired Gemini file mappings")
return deleted_count
# =============================================================================
# 请求解析工具函数
# =============================================================================
def _extract_file_name_from_uri(file_uri: str) -> str | None:

View File

@@ -39,6 +39,10 @@ MAX_CONCURRENT_REQUESTS = 5
# 单个 Key 处理的超时时间(秒)
KEY_FETCH_TIMEOUT_SECONDS = 120
# 模型获取 HTTP 请求超时时间(秒)
# 使用较短的超时10秒避免不支持 /models 端点的提供商长时间阻塞
MODEL_FETCH_HTTP_TIMEOUT = 10.0
# 上游模型缓存 TTL与定时任务间隔保持一致
UPSTREAM_MODELS_CACHE_TTL_SECONDS = MODEL_FETCH_INTERVAL_MINUTES * 60
@@ -260,7 +264,48 @@ class ModelFetchScheduler:
logger.exception(f"更新 Key {key_id} 错误信息失败")
async def _fetch_models_for_key_by_id(self, key_id: str) -> str:
"""根据 Key ID 获取模型并更新,返回结果状态"""
"""
根据 Key ID 获取模型并更新,返回结果状态
优化分两个阶段处理HTTP 请求期间不持有数据库连接,避免阻塞其他请求
"""
# ========== 阶段 1准备数据短暂持有连接==========
fetch_context = self._prepare_fetch_context(key_id)
if fetch_context is None:
return "skip"
if isinstance(fetch_context, str):
return fetch_context # "error" or "skip"
key_id, provider_id, provider_name, api_key_value, endpoint_configs = fetch_context
# ========== 阶段 2HTTP 请求(不持有数据库连接)==========
# 使用较短的超时时间10秒避免长时间阻塞
all_models, errors, has_success = await fetch_models_from_endpoints(
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
)
# ========== 阶段 3更新数据库获取新连接==========
return await self._update_key_after_fetch(
key_id=key_id,
provider_id=provider_id,
provider_name=provider_name,
all_models=all_models,
errors=errors,
has_success=has_success,
)
def _prepare_fetch_context(
self, key_id: str
) -> tuple[str, str, str, str, list[dict]] | str | None:
"""
准备获取模型所需的上下文数据
Returns:
- tuple: (key_id, provider_id, provider_name, api_key_value, endpoint_configs)
- "skip": 跳过该 Key
- "error": 出错
- None: Key 不存在
"""
with create_session() as db:
key = (
db.query(ProviderAPIKey)
@@ -271,142 +316,154 @@ class ModelFetchScheduler:
if not key:
logger.warning(f"Key {key_id} 不存在,跳过")
return "skip"
return None
if not key.is_active or not key.auto_fetch_models:
logger.debug(f"Key {key_id} 已禁用或关闭自动获取,跳过")
return "skip"
try:
result = await self._fetch_models_for_key(db, key)
now = datetime.now(timezone.utc)
provider_id = key.provider_id
# 获取 Provider 和 Endpoints
provider = (
db.query(Provider)
.options(joinedload(Provider.endpoints))
.filter(Provider.id == provider_id)
.first()
)
if not provider:
logger.warning(f"Provider {provider_id} 不存在,跳过 Key {key.id}")
key.last_models_fetch_error = "Provider not found"
key.last_models_fetch_at = now
db.commit()
return result
return "error"
# Vertex AI 类型不支持自动获取模型
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
if auth_type == "vertex_ai":
key.last_models_fetch_error = "auto_fetch_models 暂不支持 Vertex AI 类型的 Key"
key.last_models_fetch_at = now
db.commit()
logger.info(f"Key {key.id} 为 Vertex AI 类型,跳过自动获取模型")
return "skip"
# 解密 API Key
if not key.api_key:
logger.warning(f"Key {key.id} 没有 API Key跳过")
key.last_models_fetch_error = "No API key configured"
key.last_models_fetch_at = now
db.commit()
return "error"
try:
api_key_value = crypto_service.decrypt(key.api_key)
except Exception:
db.rollback()
raise
logger.error(f"解密 Key {key.id} 失败")
key.last_models_fetch_error = "Decrypt error"
key.last_models_fetch_at = now
db.commit()
return "error"
async def _fetch_models_for_key(
# 构建 api_format -> endpoint 映射
format_to_endpoint: dict[str, object] = {}
for endpoint in provider.endpoints: # type: ignore[attr-defined]
if endpoint.is_active:
format_to_endpoint[endpoint.api_format] = endpoint
if not format_to_endpoint:
logger.warning(f"Provider {provider.name} 没有活跃的端点,跳过 Key {key.id}")
key.last_models_fetch_error = "No active endpoints"
key.last_models_fetch_at = now
db.commit()
return "error"
# 使用公共函数构建所有格式的端点配置
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint) # type: ignore[arg-type]
return (key_id, provider_id, provider.name, api_key_value, endpoint_configs)
async def _update_key_after_fetch(
self,
db: Session,
key: ProviderAPIKey,
key_id: str,
provider_id: str,
provider_name: str,
all_models: list[dict],
errors: list[str],
has_success: bool,
) -> str:
"""为单个 Key 获取模型并更新 allowed_models返回结果状态"""
"""
HTTP 请求完成后更新数据库
使用新的数据库连接来更新 Key 的 allowed_models
"""
now = datetime.now(timezone.utc)
provider_id = key.provider_id
# 获取 Provider 和 Endpoints
provider = (
db.query(Provider)
.options(joinedload(Provider.endpoints))
.filter(Provider.id == provider_id)
.first()
)
with create_session() as db:
# 重新获取 Key因为之前的连接已关闭
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
logger.warning(f"Key {key_id} 在更新时不存在")
return "error"
if not provider:
logger.warning(f"Provider {provider_id} 不存在,跳过 Key {key.id}")
key.last_models_fetch_error = "Provider not found"
# 记录获取时间
key.last_models_fetch_at = now
return "error"
# Vertex AI 类型不支持自动获取模型(需要使用 Service Account 认证
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
if auth_type == "vertex_ai":
key.last_models_fetch_error = "auto_fetch_models 暂不支持 Vertex AI 类型的 Key"
key.last_models_fetch_at = now
logger.info(f"Key {key.id} 为 Vertex AI 类型,跳过自动获取模型")
return "skip"
# 如果没有任何成功的响应,不更新 allowed_models保留旧数据
if not has_success:
error_msg = "; ".join(errors) if errors else "All endpoints failed"
key.last_models_fetch_error = error_msg
logger.warning(
f"Provider {provider_name} Key {key.id} 所有端点获取失败,保留现有模型列表"
)
db.commit()
return "error"
# 解密 API Key
if not key.api_key:
logger.warning(f"Key {key.id} 没有 API Key跳过")
key.last_models_fetch_error = "No API key configured"
key.last_models_fetch_at = now
return "error"
# 有成功的响应,清除错误状态
key.last_models_fetch_error = None
try:
api_key_value = crypto_service.decrypt(key.api_key)
except Exception:
# 不记录异常详情,避免泄露密钥信息
logger.error(f"解密 Key {key.id} 失败")
key.last_models_fetch_error = "Decrypt error"
key.last_models_fetch_at = now
return "error"
# 去重获取模型 ID 列表
fetched_model_ids: set[str] = set()
for model in all_models:
model_id = model.get("id")
if model_id:
fetched_model_ids.add(model_id)
# 构建 api_format -> endpoint 映射
format_to_endpoint: dict[str, object] = {}
for endpoint in provider.endpoints: # type: ignore[attr-defined]
if endpoint.is_active:
format_to_endpoint[endpoint.api_format] = endpoint
if not format_to_endpoint:
logger.warning(f"Provider {provider.name} 没有活跃的端点,跳过 Key {key.id}")
key.last_models_fetch_error = "No active endpoints"
key.last_models_fetch_at = now
return "error"
# 使用公共函数构建所有格式的端点配置
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint) # type: ignore[arg-type]
# 并发获取模型
all_models, errors, has_success = await fetch_models_from_endpoints(endpoint_configs)
# 记录获取时间
key.last_models_fetch_at = now
# 如果没有任何成功的响应,不更新 allowed_models保留旧数据
if not has_success:
# 所有端点都失败时,记录错误
error_msg = "; ".join(errors) if errors else "All endpoints failed"
key.last_models_fetch_error = error_msg
logger.warning(
f"Provider {provider.name} Key {key.id} 所有端点获取失败,保留现有模型列表"
)
return "error"
# 有成功的响应,清除错误状态(部分失败不算失败)
key.last_models_fetch_error = None
# 去重获取模型 ID 列表
fetched_model_ids: set[str] = set()
for model in all_models:
model_id = model.get("id")
if model_id:
fetched_model_ids.add(model_id)
logger.info(
f"Provider {provider.name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型"
)
# 写入上游模型缓存(按 model id + api_format 去重后的完整模型信息)
seen_keys: set[str] = set()
unique_models: list[dict] = []
for model in all_models:
model_id = model.get("id")
api_format = model.get("api_format", "")
unique_key = f"{model_id}:{api_format}"
if model_id and unique_key not in seen_keys:
seen_keys.add(unique_key)
unique_models.append(model)
await set_upstream_models_to_cache(
provider_id, # type: ignore[arg-type]
key.id, # type: ignore[arg-type]
unique_models,
)
# 更新 allowed_models保留 locked_models
has_changed = self._update_key_allowed_models(key, fetched_model_ids)
# 如果白名单有变化,触发缓存失效和自动关联检查
if has_changed and provider_id:
from src.services.model.global_model import on_key_allowed_models_changed
await on_key_allowed_models_changed(
db=db,
provider_id=provider_id,
allowed_models=list(key.allowed_models or []),
logger.info(
f"Provider {provider_name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型"
)
return "success"
# 写入上游模型缓存(按 model id + api_format 去重后的完整模型信息)
seen_keys: set[str] = set()
unique_models: list[dict] = []
for model in all_models:
model_id = model.get("id")
api_format = model.get("api_format", "")
unique_key = f"{model_id}:{api_format}"
if model_id and unique_key not in seen_keys:
seen_keys.add(unique_key)
unique_models.append(model)
await set_upstream_models_to_cache(provider_id, key.id, unique_models)
# 更新 allowed_models保留 locked_models
has_changed = self._update_key_allowed_models(key, fetched_model_ids)
db.commit()
# 如果白名单有变化,触发缓存失效和自动关联检查
if has_changed and provider_id:
from src.services.model.global_model import on_key_allowed_models_changed
# 使用新会话处理后续操作
with create_session() as db2:
await on_key_allowed_models_changed(
db=db2,
provider_id=provider_id,
allowed_models=list(key.allowed_models or []),
)
return "success"
def _update_key_allowed_models(self, key: ProviderAPIKey, fetched_model_ids: set[str]) -> bool:
"""

View File

@@ -156,6 +156,8 @@ class CandidateResolver:
user_id: str,
user_api_key: ApiKey,
required_capabilities: dict[str, bool] | None = None,
*,
expand_retries: bool = True,
) -> dict[tuple[int, int], str]:
"""
为所有候选预先创建 available 状态记录(批量插入优化)
@@ -211,9 +213,12 @@ class CandidateResolver:
candidate_record_map[(candidate_index, 0)] = record_id
else:
# max_retries 已从 Endpoint 迁移到 ProviderEndpoint 仍可能保留旧字段用于兼容)
max_retries_for_candidate = (
int(provider.max_retries or 2) if candidate.is_cached else 1
)
if not expand_retries:
max_retries_for_candidate = 1
else:
max_retries_for_candidate = (
int(provider.max_retries or 2) if candidate.is_cached else 1
)
for retry_index in range(max_retries_for_candidate):
record_id = str(uuid.uuid4())

View File

@@ -38,6 +38,19 @@ BALANCE_CACHE_TTL = 86400
# 认证失败缓存 TTL60 秒,避免频繁重试但允许用户修正后快速重试)
AUTH_FAILED_CACHE_TTL = 60
# 后台余额刷新并发限制(避免启动时耗尽连接池)
# 使用较小的值3确保不会对连接池造成过大压力
_balance_refresh_semaphore: asyncio.Semaphore | None = None
def _get_balance_refresh_semaphore() -> asyncio.Semaphore:
"""获取余额刷新信号量(延迟初始化)"""
global _balance_refresh_semaphore
if _balance_refresh_semaphore is None:
# 限制为 3 个并发,确保后台任务不会占用太多连接
_balance_refresh_semaphore = asyncio.Semaphore(3)
return _balance_refresh_semaphore
def _get_batch_balance_concurrency() -> int:
"""
@@ -98,6 +111,44 @@ class ProviderOpsService:
# 连接器缓存 {provider_id: ProviderConnector}
self._connectors: dict[str, ProviderConnector] = {}
def _release_db_connection_before_await(self) -> None:
"""
Release pooled DB connection before long awaits (network/Redis).
SQLAlchemy Session will keep a connection checked out while a transaction is open,
even for read-only queries. In async code, this can exhaust the pool if we `await`
network I/O while holding that transaction.
Safety:
- Only commits when the session has no pending changes (new/dirty/deleted).
- Temporarily disables expire_on_commit to avoid unexpected lazy reloads.
"""
try:
has_pending_changes = bool(self.db.new) or bool(self.db.dirty) or bool(self.db.deleted)
except Exception:
has_pending_changes = False
if has_pending_changes:
return
try:
if not self.db.in_transaction():
return
except Exception:
return
original_expire_on_commit = getattr(self.db, "expire_on_commit", True)
self.db.expire_on_commit = False
try:
self.db.commit()
except Exception:
try:
self.db.rollback()
except Exception:
pass
finally:
self.db.expire_on_commit = original_expire_on_commit
# ==================== 配置管理 ====================
def get_config(self, provider_id: str) -> ProviderOpsConfig | None:
@@ -251,6 +302,9 @@ class ProviderOpsService:
if not actual_credentials:
return False, "未提供凭据"
# Avoid holding a DB connection while awaiting network I/O.
self._release_db_connection_before_await()
# 建立连接
logger.info(
f"尝试连接: provider_id={provider_id}, "
@@ -333,6 +387,8 @@ class ProviderOpsService:
)
connector = self._connectors.get(provider_id)
# Avoid holding a DB connection while awaiting authentication checks.
self._release_db_connection_before_await()
if not connector or not await connector.is_authenticated():
return ActionResult(
status=ActionStatus.AUTH_EXPIRED,
@@ -373,6 +429,9 @@ class ProviderOpsService:
# 创建操作实例
action = architecture.get_action(action_type, merged_config)
# Avoid holding a DB connection while awaiting the upstream action.
self._release_db_connection_before_await()
# 执行操作
async with connector.get_client() as client:
result = await action.execute(client)
@@ -422,6 +481,9 @@ class ProviderOpsService:
Returns:
操作结果(可能是缓存的)
"""
# Avoid holding a DB connection while awaiting Redis/cache I/O.
self._release_db_connection_before_await()
# 尝试从缓存获取
cached = await self._get_cached_balance(provider_id)
@@ -453,7 +515,20 @@ class ProviderOpsService:
注意:这是一个后台任务,使用独立的短生命周期 session
避免长时间占用连接池资源。
使用信号量限制并发数,避免启动时多个刷新任务同时运行导致连接池耗尽。
"""
semaphore = _get_balance_refresh_semaphore()
# 尝试获取信号量,如果无法立即获取则跳过本次刷新
# 这样可以避免在连接池紧张时阻塞
try:
# 使用 wait_for 设置超时,避免无限等待
await asyncio.wait_for(semaphore.acquire(), timeout=5.0)
except asyncio.TimeoutError:
logger.debug(f"异步刷新余额跳过(并发限制): provider_id={provider_id}")
return
db = None
try:
# 后台任务需要创建独立的 session因为原请求的 session 可能已关闭
@@ -469,6 +544,8 @@ class ProviderOpsService:
db.close()
except Exception:
pass
# 释放信号量
semaphore.release()
async def _clear_balance_cache(self, provider_id: str) -> None:
"""清除余额缓存"""
@@ -762,6 +839,9 @@ class ProviderOpsService:
if not provider_ids:
return {}
# Release the DB connection before awaiting many async cache refreshes.
self._release_db_connection_before_await()
# 使用信号量限制并发数,避免同时发起过多请求耗尽连接池
concurrency = _get_batch_balance_concurrency()
semaphore = asyncio.Semaphore(concurrency)
@@ -822,6 +902,9 @@ class ProviderOpsService:
from src.utils.ssl_utils import get_ssl_context
# Avoid holding a DB connection while awaiting verify pre-processing / network.
self._release_db_connection_before_await()
# 移除 base_url 末尾的斜杠
base_url = base_url.rstrip("/")

View File

@@ -136,6 +136,11 @@ class SystemConfigService:
"value": [],
"description": "邮箱后缀列表,配合 email_suffix_mode 使用",
},
# 格式转换开关
"enable_format_conversion": {
"value": False,
"description": "全局格式转换开关:开启时强制允许所有提供商的格式转换;关闭时由各提供商自行决定",
},
"audit_log_retention_days": {
"value": 30,
"description": "审计日志保留天数,超过此天数的审计日志将被自动清理",
@@ -358,6 +363,11 @@ class SystemConfigService:
"""获取敏感请求头列表"""
return cls.get_config(db, "sensitive_headers", [])
@classmethod
def is_format_conversion_enabled(cls, db: Session) -> bool:
"""检查全局格式转换是否启用"""
return bool(cls.get_config(db, "enable_format_conversion", True))
@classmethod
def mask_sensitive_headers(cls, db: Session, headers: dict[str, Any]) -> dict[str, Any]:
"""脱敏敏感请求头"""

View File

@@ -8,6 +8,7 @@
- 审计日志清理:定期清理过期的审计日志
- 连接池监控:定期检查数据库连接池状态
- Pending 状态清理:清理异常的 Pending 状态记录
- Gemini 文件映射清理:清理过期的 Gemini 文件→Key 映射
使用 APScheduler 进行任务调度,支持时区配置。
"""
@@ -103,6 +104,14 @@ class MaintenanceScheduler:
name="审计日志清理",
)
# Gemini 文件映射清理 - 每小时执行
scheduler.add_interval_job(
self._scheduled_gemini_file_mapping_cleanup,
hours=1,
job_id="gemini_file_mapping_cleanup",
name="Gemini文件映射清理",
)
# Provider 签到任务 - 凌晨 1:05 执行
scheduler.add_cron_job(
self._scheduled_provider_checkin,
@@ -117,8 +126,9 @@ class MaintenanceScheduler:
async def _run_startup_tasks(self) -> None:
"""启动时执行的初始化任务"""
# 延迟一点执行,确保系统完全启动
await asyncio.sleep(2)
# 延迟执行,等待系统完全启动Redis 连接、其他后台任务稳定)
# 增加延迟时间避免与 UsageQueueConsumer 等后台任务竞争数据库连接
await asyncio.sleep(10)
try:
logger.info("启动时执行首次清理任务...")
@@ -170,6 +180,10 @@ class MaintenanceScheduler:
"""审计日志清理任务(定时调用)"""
await self._perform_audit_cleanup()
async def _scheduled_gemini_file_mapping_cleanup(self) -> None:
"""Gemini 文件映射清理任务(定时调用)"""
await self._perform_gemini_file_mapping_cleanup()
async def _scheduled_provider_checkin(self) -> None:
"""Provider 签到任务(定时调用)"""
await self._perform_provider_checkin()
@@ -483,6 +497,26 @@ class MaintenanceScheduler:
finally:
db.close()
async def _perform_gemini_file_mapping_cleanup(self) -> None:
"""清理过期的 Gemini 文件映射记录"""
db = create_session()
try:
from src.services.gemini_files_mapping import cleanup_expired_mappings
deleted_count = cleanup_expired_mappings(db)
if deleted_count > 0:
logger.info(f"清理了 {deleted_count} 条过期的 Gemini 文件映射")
except Exception as e:
logger.exception(f"Gemini 文件映射清理失败: {e}")
try:
db.rollback()
except Exception:
pass
finally:
db.close()
async def _perform_provider_checkin(self) -> None:
"""执行 Provider 签到任务
@@ -508,8 +542,21 @@ class MaintenanceScheduler:
logger.info(f"开始执行 Provider 签到,共 {len(provider_ids)} 个...")
# 创建 ProviderOpsService 并执行批量余额查询(会触发签到)
service = ProviderOpsService(db)
# 释放主 session 的连接,避免在整个签到期间占用连接池
# (后续每个 provider 将使用独立短生命周期 session
try:
if db.in_transaction():
db.commit()
except Exception:
try:
db.rollback()
except Exception:
pass
try:
db.close()
except Exception:
pass
db = None
# 使用信号量限制并发,避免同时发起过多请求
concurrency = 3 # 签到任务并发数
@@ -518,7 +565,9 @@ class MaintenanceScheduler:
async def _checkin_provider(provider_id: str) -> tuple[str, bool, str]:
"""执行单个 Provider 的签到"""
async with semaphore:
task_db = create_session()
try:
service = ProviderOpsService(task_db)
# 触发余额查询(会先执行签到)
result = await service.query_balance(provider_id)
# 检查签到结果
@@ -537,6 +586,11 @@ class MaintenanceScheduler:
except Exception as e:
logger.warning(f"Provider {provider_id} 签到失败: {e}")
return provider_id, False, str(e)
finally:
try:
task_db.close()
except Exception:
pass
# 并行执行签到
tasks = [_checkin_provider(pid) for pid in provider_ids]
@@ -556,7 +610,8 @@ class MaintenanceScheduler:
except Exception as e:
logger.exception(f"Provider 签到任务执行失败: {e}")
finally:
db.close()
if db is not None:
db.close()
async def _perform_cleanup(self) -> None:
"""执行清理任务"""

View File

@@ -1,25 +1,13 @@
"""
异步任务服务层
任务服务层Phase2
提供视频/图片/音频等异步任务的
- 提交阶段故障转移AsyncTaskOrchestrator
- 终态计费与 Usage 写入VideoTelemetry 等)
统一任务框架相关的应用层入口
- 候选提交阶段`services.candidate.CandidateService`
- 终态结算:`services.task.application.TaskApplicationService`
"""
from .orchestrator import (
AllCandidatesFailedError,
AsyncTaskOrchestrator,
CandidateSubmissionError,
CandidateUnsupportedError,
SubmitOutcome,
UpstreamClientRequestError,
)
from .application import TaskApplicationService
__all__ = [
"AsyncTaskOrchestrator",
"SubmitOutcome",
"AllCandidatesFailedError",
"UpstreamClientRequestError",
"CandidateUnsupportedError",
"CandidateSubmissionError",
"TaskApplicationService",
]

View File

@@ -0,0 +1,312 @@
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.logger import logger
from src.models.database import ApiKey, Provider, Usage, User, VideoTask
from src.services.billing.dimension_collector_service import DimensionCollectorService
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
from src.services.billing.rule_service import BillingRuleService
from src.services.usage.service import UsageService
class TaskApplicationService:
"""
TaskApplicationService (Phase2)
当前仅先收敛"终态结算"入口,用于替代旧版 VideoTelemetry 直写 Usage 的流程。
后续将扩展 submit/cancel 并迁移候选编排逻辑。
"""
def __init__(self, db: Session, *, redis_client: Any | None = None) -> None:
self.db = db
self.redis = redis_client
async def finalize_video_task(self, task: VideoTask) -> bool:
"""
更新视频任务的计费信息(轮询完成后调用)。
异步任务的计费流程:
1. 提交成功时Usage 已结算billing_status='settled',费用=0
2. 轮询完成时更新实际费用成功则计费失败则保持0
返回 True 表示成功更新False 表示无需更新(如已是最终状态)
"""
request_id = getattr(task, "request_id", None) or task.id
existing = self.db.query(Usage).filter(Usage.request_id == request_id).first()
if not existing:
# Usage 不存在,尝试创建并结算(兜底逻辑)
logger.warning(
"Usage not found for video task, creating fallback: task_id=%s request_id=%s",
task.id,
request_id,
)
return await self._create_fallback_usage(task, request_id)
# 检查是否已有计费更新标记(避免重复计费)
metadata = existing.request_metadata or {}
if metadata.get("billing_updated_at"):
logger.debug(
"Video task billing already updated: task_id=%s request_id=%s",
task.id,
request_id,
)
return False
# 计算异步任务总耗时ms
response_time_ms: int | None = None
if task.submitted_at and task.completed_at:
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
# === 收集计费维度 ===
base_dimensions: dict[str, Any] = {
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size or "",
"retry_count": task.retry_count,
}
collector_metadata: dict[str, Any] = {
"task": {
"id": task.id,
"external_task_id": task.external_task_id,
"model": task.model,
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size,
"retry_count": task.retry_count,
"video_size_bytes": task.video_size_bytes,
},
"result": {
"video_url": task.video_url,
"video_urls": task.video_urls or [],
},
}
dims = DimensionCollectorService(self.db).collect_dimensions(
api_format=task.provider_api_format,
task_type="video",
request=task.original_request_body or {},
response=(
(task.request_metadata or {}).get("poll_raw_response")
if isinstance(task.request_metadata, dict)
else None
),
metadata=collector_metadata,
base_dimensions=base_dimensions,
)
# === 计算成本(优先使用冻结的 billing_rule_snapshot===
rule_snapshot = None
if isinstance(task.request_metadata, dict):
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
expression = None
variables: dict[str, Any] | None = None
dimension_mappings: dict[str, dict[str, Any]] | None = None
rule_id = None
rule_name = None
rule_scope = None
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
rule_id = rule_snapshot.get("rule_id")
rule_name = rule_snapshot.get("rule_name")
rule_scope = rule_snapshot.get("scope")
expression = rule_snapshot.get("expression")
variables = rule_snapshot.get("variables") or {}
dimension_mappings = rule_snapshot.get("dimension_mappings") or {}
else:
lookup = BillingRuleService.find_rule(
self.db,
provider_id=task.provider_id,
model_name=task.model,
task_type="video",
)
if lookup:
rule = lookup.rule
rule_id = rule.id
rule_name = rule.name
rule_scope = getattr(lookup, "scope", None)
expression = rule.expression
variables = rule.variables or {}
dimension_mappings = rule.dimension_mappings or {}
billing_snapshot: dict[str, Any] = {
"schema_version": "1.0",
"rule_id": str(rule_id) if rule_id else None,
"rule_name": str(rule_name) if rule_name else None,
"scope": str(rule_scope) if rule_scope else None,
"expression": str(expression) if expression else None,
"dimensions_used": dims,
"missing_required": [],
"cost": 0.0,
"status": "no_rule",
"calculated_at": datetime.now(timezone.utc).isoformat(),
}
cost = 0.0
# 只有任务成功时才计费
if task.status == "completed" and expression:
engine = FormulaEngine()
try:
result = engine.evaluate(
expression=str(expression),
variables=variables,
dimensions=dims,
dimension_mappings=dimension_mappings,
strict_mode=config.billing_strict_mode,
)
billing_snapshot["status"] = result.status
billing_snapshot["missing_required"] = result.missing_required
if result.status == "complete":
cost = float(result.cost)
billing_snapshot["cost"] = cost
except BillingIncompleteError as exc:
# strict_mode=true标记任务失败并隐藏产物避免"免费放行"
task.status = "failed"
task.error_code = "billing_incomplete"
task.error_message = f"Missing required dimensions: {exc.missing_required}"
task.video_url = None
task.video_urls = None
billing_snapshot["status"] = "incomplete"
billing_snapshot["missing_required"] = exc.missing_required
billing_snapshot["cost"] = 0.0
except Exception as exc:
billing_snapshot["status"] = "incomplete"
billing_snapshot["error"] = str(exc)
billing_snapshot["cost"] = 0.0
# 回写到 task.request_metadata 便于审计/重算
# 重新赋值整个字典,确保 SQLAlchemy 检测到变更
metadata = dict(task.request_metadata) if task.request_metadata else {}
metadata["billing_snapshot"] = billing_snapshot
task.request_metadata = metadata
# === 更新已结算的 Usage 计费信息 ===
updated = UsageService.update_settled_billing(
self.db,
request_id=request_id,
total_cost_usd=cost,
request_cost_usd=cost,
status="completed" if task.status == "completed" else "failed",
status_code=200 if task.status == "completed" else 500,
error_message=(
None
if task.status == "completed"
else (task.error_message or task.error_code or "video_task_failed")
),
response_time_ms=response_time_ms,
billing_snapshot=billing_snapshot,
extra_metadata={
"dimensions": dims,
"raw_response_ref": {
"video_task_id": task.id,
"field": "video_tasks.request_metadata.poll_raw_response",
},
},
)
if updated:
logger.debug(
"Updated video task billing: task_id=%s request_id=%s cost=%.6f",
task.id,
request_id,
cost,
)
else:
logger.warning(
"Failed to update video task billing (may already be updated): "
"task_id=%s request_id=%s",
task.id,
request_id,
)
return updated
async def _create_fallback_usage(self, task: VideoTask, request_id: str) -> bool:
"""
兜底逻辑:当 Usage 不存在时创建完整记录。
这种情况理论上不应发生submit 阶段已创建),但保留以防万一。
"""
user_obj = self.db.query(User).filter(User.id == task.user_id).first()
api_key_obj = (
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
if task.api_key_id
else None
)
provider_obj = (
self.db.query(Provider).filter(Provider.id == task.provider_id).first()
if task.provider_id
else None
)
provider_name = provider_obj.name if provider_obj else "unknown"
# 计算响应时间
response_time_ms: int | None = None
if task.submitted_at and task.completed_at:
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
try:
await UsageService.record_usage_with_custom_cost(
db=self.db,
user=user_obj,
api_key=api_key_obj,
provider=provider_name,
model=task.model,
request_type="video",
total_cost_usd=0.0, # 兜底记录不计费
request_cost_usd=0.0,
input_tokens=0,
output_tokens=0,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
api_format=task.client_api_format,
endpoint_api_format=task.provider_api_format,
has_format_conversion=bool(task.format_converted),
is_stream=False,
response_time_ms=response_time_ms,
first_byte_time_ms=None,
status_code=200 if task.status == "completed" else 500,
error_message=(
None
if task.status == "completed"
else (task.error_message or task.error_code or "video_task_failed")
),
metadata={
"fallback_created": True,
"video_task_id": task.id,
},
request_headers=(
(task.request_metadata or {}).get("request_headers")
if isinstance(task.request_metadata, dict)
else None
),
request_body=task.original_request_body,
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=request_id,
provider_id=task.provider_id,
provider_endpoint_id=task.endpoint_id,
provider_api_key_id=task.key_id,
status="completed" if task.status == "completed" else "failed",
target_model=None,
)
return True
except Exception as exc:
logger.exception(
"Failed to create fallback usage for video task=%s: %s",
task.id,
str(exc),
)
return False

View File

@@ -0,0 +1,36 @@
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
class TaskMode(str, Enum):
SYNC = "sync"
ASYNC = "async"
@dataclass(slots=True)
class TaskContext:
"""
TaskContext (pure DTO)
- Only primitive types / IDs
- Serializable & safe to pass across processes
"""
request_id: str
task_type: str # chat/cli/video/image/audio
task_mode: TaskMode
user_id: str
api_key_id: str
client_ip: str = ""
user_agent: str = ""
start_time: float = 0.0
api_format: str | None = None
model: str | None = None
mapped_model: str | None = None
capability_requirements: dict[str, bool] = field(default_factory=dict)

View File

@@ -1,3 +1,7 @@
"""Task telemetry implementations for concrete task types (video/image/audio)."""
"""Per-task-type implementations (Phase2).
Currently includes:
- video: polling adapter
"""
__all__ = []

View File

@@ -0,0 +1,626 @@
"""
Video task poller adapter.
Implements the video-specific poll/normalize/update logic used by TaskPollerService.
优化HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import ProviderAuthInfo, get_provider_auth
from src.api.handlers.base.video_handler_base import (
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.clients.http_client import HTTPClientPool
from src.config.settings import config
from src.core.api_format import (
build_upstream_headers_for_endpoint,
get_extra_headers_from_endpoint,
make_signature_key,
)
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.database import create_session
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.services.task.application import TaskApplicationService
@dataclass(slots=True)
class VideoPollContext:
"""视频轮询上下文,保存 HTTP 请求所需的数据(不依赖数据库会话)"""
task_id: str
external_task_id: str
provider_api_format: str
base_url: str
upstream_key: str
headers: dict[str, str]
# 用于更新任务的原始数据
poll_count: int
retry_count: int
poll_interval_seconds: int
max_poll_count: int
current_status: str
# 永久性错误指示词(用于降级判断,不应重试)
_PERMANENT_ERROR_INDICATORS = frozenset(
{
"not found",
"404",
"unauthorized",
"401",
"forbidden",
"403",
"invalid request",
"invalid api key",
"does not exist",
}
)
class PollHTTPError(RuntimeError):
"""HTTP 轮询错误,携带状态码便于区分临时/永久错误"""
def __init__(self, status_code: int, message: str):
# 确保错误信息包含状态码
full_message = f"HTTP {status_code}: {message}" if message else f"HTTP {status_code}"
super().__init__(full_message)
self.status_code = status_code
self.original_message = message
class VideoTaskPollerAdapter:
task_type = "video"
# scheduler
job_id = "task_poller:video"
job_name = "视频任务轮询"
interval_seconds = config.video_poll_interval_seconds
# distributed lock
lock_key = "task_poller:video:lock"
lock_ttl = 60
# execution
batch_size = config.video_poll_batch_size
concurrency = config.video_poll_concurrency
consecutive_failure_alert_threshold = 5
max_backoff_seconds = 300
def __init__(self) -> None:
self._openai_normalizer = OpenAINormalizer()
self._gemini_normalizer = GeminiNormalizer()
def sanitize_error_message(self, message: str) -> str:
return sanitize_error_message(message)
def list_due_task_ids(self, db: Session, *, now: datetime, limit: int) -> list[str]:
tasks = (
db.query(VideoTask)
.filter(
VideoTask.status.in_(
[
VideoStatus.SUBMITTED.value,
VideoStatus.QUEUED.value,
VideoStatus.PROCESSING.value,
]
),
VideoTask.next_poll_at <= now,
VideoTask.poll_count < VideoTask.max_poll_count,
)
.order_by(VideoTask.next_poll_at.asc())
.limit(limit)
.all()
)
return [t.id for t in tasks]
def get_task(self, db: Session, task_id: str) -> VideoTask | None:
# SQLAlchemy 1.4+ API
return db.get(VideoTask, task_id)
# ==================== 分阶段处理方法(优化数据库连接占用)====================
async def prepare_poll_context(
self, db: Session, task: VideoTask
) -> VideoPollContext | InternalVideoPollResult:
"""
阶段 1准备轮询上下文短暂持有数据库连接
Returns:
VideoPollContext: 成功时返回上下文
InternalVideoPollResult: 失败时返回错误结果
"""
if not task.endpoint_id or not task.key_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_provider_info",
error_message="Task missing endpoint_id or key_id",
)
endpoint = self._get_endpoint(db, task.endpoint_id)
key = self._get_key(db, task.key_id)
if not key.api_key:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="provider_config_error",
error_message="Provider key not properly configured",
)
try:
upstream_key = crypto_service.decrypt(key.api_key)
except Exception:
logger.warning("Failed to decrypt provider key for task %s", task.id)
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="decryption_error",
error_message="Failed to decrypt provider key",
)
provider_format = (task.provider_api_format or "").strip().lower()
if not provider_format:
provider_format = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
# 构建请求头
if provider_format.startswith("gemini:"):
auth_info = await get_provider_auth(endpoint, key)
else:
auth_info = None
headers = self._build_headers(provider_format, upstream_key, endpoint, auth_info)
return VideoPollContext(
task_id=task.id,
external_task_id=task.external_task_id or "",
provider_api_format=provider_format,
base_url=endpoint.base_url or "",
upstream_key=upstream_key,
headers=headers,
poll_count=task.poll_count,
retry_count=task.retry_count,
poll_interval_seconds=task.poll_interval_seconds,
max_poll_count=task.max_poll_count,
current_status=task.status,
)
async def poll_task_http(self, ctx: VideoPollContext) -> InternalVideoPollResult:
"""
阶段 2执行 HTTP 请求(不持有数据库连接)
Args:
ctx: 轮询上下文
Returns:
InternalVideoPollResult: 轮询结果
"""
if not ctx.external_task_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_external_task_id",
error_message="Task missing external_task_id",
)
if ctx.provider_api_format.startswith("gemini:"):
return await self._poll_gemini_with_context(ctx)
return await self._poll_openai_with_context(ctx)
async def update_task_after_poll(
self,
task_id: str,
result: InternalVideoPollResult,
ctx: VideoPollContext | None,
redis_client: Any | None,
error_exception: Exception | None = None,
) -> None:
"""
阶段 3更新数据库获取新的数据库连接
Args:
task_id: 任务 ID
result: 轮询结果
ctx: 轮询上下文(准备阶段就失败时为 None
redis_client: Redis 客户端
error_exception: 如果 HTTP 请求失败,传入异常对象
"""
with create_session() as db:
task = db.get(VideoTask, task_id)
if not task:
logger.warning("Task %s disappeared during poll update", task_id)
return
if error_exception is not None and ctx is not None:
# HTTP 请求失败(需要 ctx 来计算 backoff
self._handle_poll_error(task, error_exception, ctx)
elif result.status == VideoStatus.COMPLETED:
task.status = VideoStatus.COMPLETED.value
task.video_url = result.video_url
task.video_expires_at = result.expires_at
task.completed_at = datetime.now(timezone.utc)
task.progress_percent = 100
if result.video_urls:
task.video_urls = result.video_urls
self._attach_poll_raw_response(task, result)
elif result.status == VideoStatus.FAILED:
task.status = VideoStatus.FAILED.value
task.error_code = result.error_code
task.error_message = result.error_message
task.completed_at = datetime.now(timezone.utc)
self._attach_poll_raw_response(task, result)
else:
task.poll_count += 1
task.progress_percent = result.progress_percent
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
seconds=task.poll_interval_seconds
)
# 超时检查
task.updated_at = datetime.now(timezone.utc)
if task.poll_count >= task.max_poll_count and task.status not in [
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
]:
task.status = VideoStatus.FAILED.value
task.error_code = "poll_timeout"
task.error_message = f"Task timed out after {task.poll_count} polls"
task.completed_at = datetime.now(timezone.utc)
# 终态结算
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try:
await TaskApplicationService(db, redis_client=redis_client).finalize_video_task(
task
)
except Exception as exc:
logger.exception(
"Failed to record video usage for task=%s: %s",
task.id,
sanitize_error_message(str(exc)),
)
db.commit()
def _handle_poll_error(self, task: VideoTask, exc: Exception, ctx: VideoPollContext) -> None:
"""处理轮询错误"""
task.poll_count += 1
error_msg = sanitize_error_message(str(exc))
logger.warning("Poll error for task %s: %s", task.id, error_msg)
task.progress_message = f"Poll error: {error_msg}"
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
is_permanent = self._is_permanent_error(exc, status_code=status_code)
if is_permanent:
task.status = VideoStatus.FAILED.value
task.error_code = "poll_permanent_error"
task.error_message = error_msg
task.completed_at = datetime.now(timezone.utc)
else:
backoff = min(
ctx.poll_interval_seconds * (2 ** min(ctx.retry_count, 5)),
self.max_backoff_seconds,
)
task.retry_count += 1
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
async def _poll_openai_with_context(self, ctx: VideoPollContext) -> InternalVideoPollResult:
"""使用上下文进行 OpenAI 轮询(不需要数据库)"""
url = self._build_openai_url(ctx.base_url, ctx.external_task_id)
client = await HTTPClientPool.get_default_client_async()
response = await client.get(url, headers=ctx.headers)
if response.status_code >= 400:
error_message = self._extract_error_message(response.text, response.status_code)
raise PollHTTPError(response.status_code, error_message)
payload = response.json()
return self._openai_normalizer.video_poll_to_internal(payload)
async def _poll_gemini_with_context(self, ctx: VideoPollContext) -> InternalVideoPollResult:
"""使用上下文进行 Gemini 轮询(不需要数据库)"""
operation_name = normalize_gemini_operation_id(ctx.external_task_id)
url = self._build_gemini_url(ctx.base_url, operation_name)
logger.debug(
"[VideoPoller] Gemini poll: task=%s external_id=%s url=%s",
ctx.task_id,
ctx.external_task_id,
url,
)
client = await HTTPClientPool.get_default_client_async()
response = await client.get(url, headers=ctx.headers)
if response.status_code >= 400:
logger.warning(
"[VideoPoller] Gemini poll failed: task=%s status=%s response=%s",
ctx.task_id,
response.status_code,
response.text[:500] if response.text else "(empty)",
)
error_message = self._extract_error_message(response.text, response.status_code)
raise PollHTTPError(response.status_code, error_message)
payload = response.json()
return self._gemini_normalizer.video_poll_to_internal(payload)
# ==================== 旧版方法(保留兼容性)====================
async def poll_single_task(
self, db: Session, task: VideoTask, *, redis_client: Any | None
) -> None:
"""
旧版单任务轮询方法(保留向后兼容)
注意:此方法在 HTTP 请求期间持有数据库连接,建议使用分阶段方法。
"""
try:
result = await self._poll_task_status(db, task)
if result.status == VideoStatus.COMPLETED:
task.status = VideoStatus.COMPLETED.value
task.video_url = result.video_url
task.video_expires_at = result.expires_at
task.completed_at = datetime.now(timezone.utc)
task.progress_percent = 100
if result.video_urls:
task.video_urls = result.video_urls
self._attach_poll_raw_response(task, result)
elif result.status == VideoStatus.FAILED:
task.status = VideoStatus.FAILED.value
task.error_code = result.error_code
task.error_message = result.error_message
task.completed_at = datetime.now(timezone.utc)
self._attach_poll_raw_response(task, result)
else:
task.poll_count += 1
task.progress_percent = result.progress_percent
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
seconds=task.poll_interval_seconds
)
except Exception as exc:
task.poll_count += 1
error_msg = sanitize_error_message(str(exc))
logger.warning("Poll error for task %s: %s", task.id, error_msg)
task.progress_message = f"Poll error: {error_msg}"
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
is_permanent = self._is_permanent_error(exc, status_code=status_code)
if is_permanent:
task.status = VideoStatus.FAILED.value
task.error_code = "poll_permanent_error"
task.error_message = error_msg
task.completed_at = datetime.now(timezone.utc)
else:
backoff = min(
task.poll_interval_seconds * (2 ** min(task.retry_count, 5)),
self.max_backoff_seconds,
)
task.retry_count += 1
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
# 超时:超过最大轮询次数且未进入终态
task.updated_at = datetime.now(timezone.utc)
if task.poll_count >= task.max_poll_count and task.status not in [
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
]:
task.status = VideoStatus.FAILED.value
task.error_code = "poll_timeout"
task.error_message = f"Task timed out after {task.poll_count} polls"
task.completed_at = datetime.now(timezone.utc)
# 终态结算
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try:
await TaskApplicationService(db, redis_client=redis_client).finalize_video_task(
task
)
except Exception as exc:
logger.exception(
"Failed to record video usage for task=%s: %s",
task.id,
sanitize_error_message(str(exc)),
)
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
if not result.raw_response:
return
# 重新赋值整个字典,确保 SQLAlchemy 检测到变更
# (直接修改 JSON 字段内部不会自动标记为 dirty
metadata = dict(task.request_metadata) if task.request_metadata else {}
metadata["poll_raw_response"] = result.raw_response
task.request_metadata = metadata
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
if status_code is not None:
return 400 <= status_code < 500 and status_code != 429
error_msg = str(exc).lower()
return any(indicator in error_msg for indicator in _PERMANENT_ERROR_INDICATORS)
async def _poll_task_status(self, db: Session, task: VideoTask) -> InternalVideoPollResult:
if not task.endpoint_id or not task.key_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_provider_info",
error_message="Task missing endpoint_id or key_id",
)
endpoint = self._get_endpoint(db, task.endpoint_id)
key = self._get_key(db, task.key_id)
if not key.api_key:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="provider_config_error",
error_message="Provider key not properly configured",
)
try:
upstream_key = crypto_service.decrypt(key.api_key)
except Exception:
logger.warning("Failed to decrypt provider key for task %s", task.id)
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="decryption_error",
error_message="Failed to decrypt provider key",
)
provider_format = (task.provider_api_format or "").strip().lower()
if not provider_format:
provider_format = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
if provider_format.startswith("gemini:"):
auth_info = await get_provider_auth(endpoint, key)
return await self._poll_gemini(task, endpoint, upstream_key, auth_info)
return await self._poll_openai(task, endpoint, upstream_key)
async def _poll_openai(
self,
task: VideoTask,
endpoint: ProviderEndpoint,
upstream_key: str,
) -> InternalVideoPollResult:
if not task.external_task_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_external_task_id",
error_message="Task missing external_task_id",
)
url = self._build_openai_url(endpoint.base_url, task.external_task_id)
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
headers = self._build_headers(endpoint_sig, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
response = await client.get(url, headers=headers)
if response.status_code >= 400:
error_message = self._extract_error_message(response.text, response.status_code)
raise PollHTTPError(response.status_code, error_message)
payload = response.json()
return self._openai_normalizer.video_poll_to_internal(payload)
async def _poll_gemini(
self,
task: VideoTask,
endpoint: ProviderEndpoint,
upstream_key: str,
auth_info: ProviderAuthInfo | None,
) -> InternalVideoPollResult:
if not task.external_task_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_external_task_id",
error_message="Task missing external_task_id",
)
operation_name = normalize_gemini_operation_id(task.external_task_id)
url = self._build_gemini_url(endpoint.base_url, operation_name)
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
headers = self._build_headers(endpoint_sig, upstream_key, endpoint, auth_info)
logger.debug(
"[VideoPoller] Gemini poll: task=%s external_id=%s url=%s",
task.id,
task.external_task_id,
url,
)
client = await HTTPClientPool.get_default_client_async()
response = await client.get(url, headers=headers)
if response.status_code >= 400:
logger.warning(
"[VideoPoller] Gemini poll failed: task=%s status=%s response=%s",
task.id,
response.status_code,
response.text[:500] if response.text else "(empty)",
)
error_message = self._extract_error_message(response.text, response.status_code)
raise PollHTTPError(response.status_code, error_message)
payload = response.json()
return self._gemini_normalizer.video_poll_to_internal(payload)
def _build_openai_url(self, base_url: str | None, task_id: str) -> str:
base = (base_url or "https://api.openai.com").rstrip("/")
if base.endswith("/v1"):
return f"{base}/videos/{task_id}"
return f"{base}/v1/videos/{task_id}"
def _build_gemini_url(self, base_url: str | None, operation_name: str) -> str:
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
return f"{base}/v1beta/{operation_name}"
def _build_headers(
self,
endpoint_sig: str,
upstream_key: str,
endpoint: ProviderEndpoint,
auth_info: ProviderAuthInfo | None = None,
) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint)
headers = build_upstream_headers_for_endpoint(
{},
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
)
if auth_info:
headers.pop("x-goog-api-key", None)
headers[auth_info.auth_header] = auth_info.auth_value
return headers
def _get_endpoint(self, db: Session, endpoint_id: str) -> ProviderEndpoint:
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
if not endpoint:
raise RuntimeError("Provider endpoint not found")
return endpoint
def _get_key(self, db: Session, key_id: str) -> ProviderAPIKey:
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise RuntimeError("Provider key not found")
return key
def _extract_error_message(self, response_text: str | None, status_code: int) -> str:
"""从响应中提取有意义的错误信息"""
if not response_text:
return f"Request failed with status {status_code}"
# 尝试解析 JSON 格式的错误
try:
data = json.loads(response_text)
# OpenAI 格式: {"error": {"message": "..."}}
if isinstance(data.get("error"), dict):
error_obj = data["error"]
message = error_obj.get("message") or error_obj.get("detail") or str(error_obj)
return sanitize_error_message(message)
# Gemini 格式: {"error": {"message": "...", "code": 404}}
if "message" in data:
return sanitize_error_message(data["message"])
except (json.JSONDecodeError, TypeError, KeyError):
pass
# 回退到原始文本(截断)
return sanitize_error_message(response_text[:500])

View File

@@ -1,320 +0,0 @@
"""
VideoTelemetryPhase3
将 Video 异步任务的“终态计费 + Usage 写入 + required 缺失告警”从 poller 中抽离出来,
便于未来 Image/Audio 复用相同框架。
"""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.api.handlers.base.video_handler_base import sanitize_error_message
from src.config.settings import config
from src.core.api_format.conversion.internal_video import VideoStatus
from src.core.logger import logger
from src.models.database import ApiKey, Provider, User, VideoTask
from src.services.billing.dimension_collector_service import DimensionCollectorService
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
from src.services.billing.rule_service import BillingRuleService
from src.services.usage.service import UsageService
class VideoTelemetry:
def __init__(self, db: Session, *, redis_client: Any | None = None) -> None:
self.db = db
self.redis = redis_client
self._formula_engine = FormulaEngine()
async def record_terminal_usage(self, task: VideoTask) -> None:
"""
为视频任务终态写入 Usage
- COMPLETED: 使用 FormulaEngine 计算 cost或 no_rule / incomplete -> cost=0
- FAILED: cost=0
该方法可能会在 strict_mode 缺失 required 维度时将任务降级为 FAILED 并隐藏产物。
"""
request_id = None
if isinstance(task.request_metadata, dict):
request_id = task.request_metadata.get("request_id")
request_id = request_id or task.id
# 计算异步任务总耗时ms
response_time_ms = None
if task.submitted_at and task.completed_at:
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
# 基础维度(无需 collectors 也可计费)
base_dimensions: dict[str, Any] = {
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size or "",
"retry_count": task.retry_count,
}
# collectors 可用的 metadata结构稳定便于配置 path
collector_metadata: dict[str, Any] = {
"task": {
"id": task.id,
"external_task_id": task.external_task_id,
"model": task.model,
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size,
"retry_count": task.retry_count,
"video_size_bytes": task.video_size_bytes,
},
"result": {
"video_url": task.video_url,
"video_urls": task.video_urls or [],
},
}
# 维度采集base + collectors 覆盖/补全
dims = DimensionCollectorService(self.db).collect_dimensions(
api_format=task.provider_api_format,
task_type="video",
request=task.original_request_body or {},
response=(
(task.request_metadata or {}).get("poll_raw_response")
if isinstance(task.request_metadata, dict)
else None
),
metadata=collector_metadata,
base_dimensions=base_dimensions,
)
# 取冻结的 rule_snapshot若缺失则回退 DB 查找(兼容旧任务)
rule_snapshot = None
if isinstance(task.request_metadata, dict):
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
billing_snapshot: dict[str, Any] = {
"status": "complete",
"missing_required": [],
"strict_mode": config.billing_strict_mode,
}
cost = 0.0
if task.status == VideoStatus.FAILED.value:
billing_snapshot["billed_reason"] = "task_failed"
else:
# COMPLETED计算成本
expression = None
variables = None
dimension_mappings = None
rule_id = None
rule_name = None
rule_scope = None
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
rule_id = rule_snapshot.get("rule_id")
rule_name = rule_snapshot.get("rule_name")
rule_scope = rule_snapshot.get("scope")
expression = rule_snapshot.get("expression")
variables = rule_snapshot.get("variables")
dimension_mappings = rule_snapshot.get("dimension_mappings")
else:
lookup = BillingRuleService.find_rule(
self.db,
provider_id=task.provider_id,
model_name=task.model,
task_type="video",
)
if lookup:
rule = lookup.rule
rule_id = rule.id
rule_name = rule.name
rule_scope = lookup.scope
expression = rule.expression
variables = rule.variables
dimension_mappings = rule.dimension_mappings
if not expression:
billing_snapshot["status"] = "no_rule"
billing_snapshot["cost_breakdown"] = {"total": 0.0}
logger.warning(
"No billing rule for video task (request_id=%s, model=%s, provider_id=%s)",
request_id,
task.model,
task.provider_id,
)
else:
billing_snapshot.update(
{
"rule_id": rule_id,
"rule_name": rule_name,
"rule_scope": rule_scope,
"expression": expression,
"variables": variables or {},
}
)
try:
result = self._formula_engine.evaluate(
expression=expression,
variables=variables or {},
dimensions=dims,
dimension_mappings=dimension_mappings or {},
strict_mode=config.billing_strict_mode,
)
billing_snapshot["status"] = result.status
billing_snapshot["missing_required"] = result.missing_required
billing_snapshot["resolved_values"] = result.resolved_values
if result.status == "complete":
cost = result.cost
else:
logger.error(
"Billing incomplete due to missing required dimensions "
"(request_id=%s, model=%s, missing=%s)",
request_id,
task.model,
result.missing_required,
)
cost = 0.0
await self._maybe_alert_missing_required(
model=task.model,
missing_required=result.missing_required,
)
if result.error:
billing_snapshot["error"] = result.error
except BillingIncompleteError as exc:
logger.error(
"Billing strict mode triggered (request_id=%s, model=%s, missing=%s)",
request_id,
task.model,
exc.missing_required,
)
billing_snapshot["status"] = "incomplete"
billing_snapshot["missing_required"] = exc.missing_required
billing_snapshot["resolved_values"] = {}
billing_snapshot["error"] = "strict_mode_missing_required"
cost = 0.0
# strict_mode=true标记任务失败并隐藏产物避免"免费放行"
task.status = VideoStatus.FAILED.value
task.error_code = "billing_incomplete"
task.error_message = f"Missing required dimensions: {exc.missing_required}"
task.video_url = None
task.video_urls = None
await self._maybe_alert_missing_required(
model=task.model,
missing_required=exc.missing_required,
)
billing_snapshot["cost_breakdown"] = {"total": cost}
# 将 billing_snapshot 回写到 task.request_metadata 便于对账(不会影响 usage 的单独存档)
if task.request_metadata is None:
task.request_metadata = {}
if isinstance(task.request_metadata, dict):
task.request_metadata["billing_snapshot"] = billing_snapshot
usage_metadata: dict[str, Any] = {
"billing_snapshot": billing_snapshot,
"dimensions": dims,
"raw_response_ref": {
"video_task_id": task.id,
"field": "video_tasks.request_metadata.poll_raw_response",
},
}
# 查询关联对象(用于写入 usage.user_id/api_key_id 等)
user_obj = self.db.query(User).filter(User.id == task.user_id).first()
api_key_obj = (
self.db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
if task.api_key_id
else None
)
provider_obj = (
self.db.query(Provider).filter(Provider.id == task.provider_id).first()
if task.provider_id
else None
)
provider_name = provider_obj.name if provider_obj else "unknown"
await UsageService.record_usage_with_custom_cost(
db=self.db,
user=user_obj,
api_key=api_key_obj,
provider=provider_name,
model=task.model,
request_type="video",
total_cost_usd=cost,
request_cost_usd=cost,
input_tokens=0,
output_tokens=0,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
api_format=task.client_api_format,
endpoint_api_format=task.provider_api_format,
has_format_conversion=bool(task.format_converted),
is_stream=False,
response_time_ms=response_time_ms,
first_byte_time_ms=None,
status_code=200 if task.status == VideoStatus.COMPLETED.value else 500,
error_message=(
None
if task.status == VideoStatus.COMPLETED.value
else (task.error_message or task.error_code or "video_task_failed")
),
metadata=usage_metadata,
request_headers=(
(task.request_metadata or {}).get("request_headers")
if isinstance(task.request_metadata, dict)
else None
),
request_body=task.original_request_body,
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=request_id,
provider_id=task.provider_id,
provider_endpoint_id=task.endpoint_id,
provider_api_key_id=task.key_id,
status="completed" if task.status == VideoStatus.COMPLETED.value else "failed",
target_model=None,
)
async def _maybe_alert_missing_required(
self, *, model: str, missing_required: list[str]
) -> None:
"""required 维度缺失告警:同一 (model, dimension) 1 小时内 >= 10 次触发升级告警。"""
if not missing_required:
return
if not self.redis:
logger.error(
"Missing required billing dimensions (model=%s): %s", model, missing_required
)
return
# 按小时 bucket 聚合
now = datetime.now(timezone.utc)
hour_bucket = now.strftime("%Y%m%d%H")
for dim in missing_required:
key = f"billing:missing_required:{model}:{dim}:{hour_bucket}"
try:
count = await self.redis.incr(key)
if count == 1:
await self.redis.expire(key, 3700)
if count >= 10:
logger.warning(
"Billing required dimension missing frequently (model=%s, dim=%s, count=%s/hour)",
model,
dim,
count,
)
except Exception as exc:
logger.warning(
"Failed to record billing alert counter: %s", sanitize_error_message(str(exc))
)
__all__ = ["VideoTelemetry"]

View File

@@ -0,0 +1,27 @@
from __future__ import annotations
from enum import Enum
class TaskStatus(str, Enum):
"""Generic task status (progress)."""
PENDING = "pending"
STREAMING = "streaming"
SUBMITTED = "submitted"
QUEUED = "queued"
PROCESSING = "processing"
COMPLETED = "completed"
FAILED = "failed"
CANCELLED = "cancelled"
EXPIRED = "expired"
class BillingStatus(str, Enum):
"""Billing settlement status (Usage.billing_status)."""
PENDING = "pending"
SETTLED = "settled"
VOID = "void"

View File

@@ -1,635 +0,0 @@
"""
AsyncTaskOrchestrator
提交阶段故障转移(多候选尝试):
- 目标:拿到 external_task_id 后锁定 provider/endpoint/key后续轮询不再切换。
- 仅覆盖“提交阶段”;轮询阶段由各 task poller 使用已锁定的信息执行。
"""
from __future__ import annotations
import re
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any, Protocol, runtime_checkable
import httpx
from redis.asyncio import Redis
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.models.database import ApiKey, RequestCandidate
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import ProviderCandidate, get_cache_aware_scheduler
from src.services.orchestration.candidate_resolver import CandidateResolver
from src.services.orchestration.error_classifier import ErrorClassifier
from src.services.system.config import SystemConfigService
_SENSITIVE_PATTERN = re.compile(
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
re.IGNORECASE,
)
def _sanitize(message: str, max_length: int = 200) -> str:
if not message:
return "request_failed"
return _SENSITIVE_PATTERN.sub("[REDACTED]", message)[:max_length]
@runtime_checkable
class SubmitFunc(Protocol):
async def __call__(self, candidate: ProviderCandidate) -> httpx.Response: ...
@runtime_checkable
class ExtractExternalTaskIdFunc(Protocol):
def __call__(self, payload: dict[str, Any]) -> str | None: ...
class UpstreamClientRequestError(RuntimeError):
"""可判定为客户端请求问题(不应 failover的上游错误。"""
def __init__(
self,
*,
response: httpx.Response,
candidate_keys: list[dict[str, Any]],
) -> None:
self.response = response
self.candidate_keys = candidate_keys
super().__init__(f"Upstream client error: HTTP {response.status_code}")
class AllCandidatesFailedError(RuntimeError):
def __init__(
self,
*,
reason: str,
candidate_keys: list[dict[str, Any]],
last_status_code: int | None = None,
) -> None:
self.reason = reason
self.candidate_keys = candidate_keys
self.last_status_code = last_status_code
super().__init__(f"All candidates failed: {reason}")
class CandidateUnsupportedError(RuntimeError):
"""候选不被当前任务支持(如 auth_type/格式转换需求不支持)。"""
class CandidateSubmissionError(RuntimeError):
"""候选提交异常(网络/解密/解析等)。"""
@dataclass(slots=True)
class SubmitOutcome:
candidate: ProviderCandidate
candidate_keys: list[dict[str, Any]]
external_task_id: str
rule_lookup: BillingRuleLookupResult | None
upstream_payload: dict[str, Any] | None = None
class AsyncTaskOrchestrator:
"""
异步任务编排器:只负责提交阶段的候选遍历与错误处理策略。
"""
def __init__(self, db: Session, *, redis_client: Redis | None = None) -> None:
self.db = db
self.redis = redis_client
self._candidate_resolver: CandidateResolver | None = None
self._error_classifier: ErrorClassifier | None = None
self._cache_scheduler = None
# 候选记录映射:{candidate_index: RequestCandidate}
self._candidate_records: dict[int, RequestCandidate] = {}
def _create_candidate_records(
self,
candidates: list[ProviderCandidate],
request_id: str | None,
user_api_key: ApiKey,
) -> dict[int, RequestCandidate]:
"""
为所有候选预创建 RequestCandidate 记录。
Args:
candidates: 候选列表
request_id: 请求 ID
user_api_key: 用户 API Key
Returns:
{candidate_index: RequestCandidate} 映射
"""
if not request_id:
return {}
now = datetime.now(timezone.utc)
records: dict[int, RequestCandidate] = {}
for idx, cand in enumerate(candidates):
record = RequestCandidate(
id=str(uuid.uuid4()),
request_id=request_id,
candidate_index=idx,
retry_index=0,
user_id=user_api_key.user_id if user_api_key else None,
api_key_id=user_api_key.id if user_api_key else None,
provider_id=cand.provider.id,
endpoint_id=cand.endpoint.id,
key_id=cand.key.id,
status="available",
is_cached=bool(getattr(cand, "is_cached", False)),
created_at=now,
)
self.db.add(record)
records[idx] = record
try:
self.db.flush()
except Exception as exc:
logger.warning(
"[AsyncTaskOrchestrator] Failed to create candidate records: %s",
str(exc),
)
self.db.rollback()
return {}
return records
def _update_candidate_record(
self,
idx: int,
*,
status: str,
skip_reason: str | None = None,
status_code: int | None = None,
error_type: str | None = None,
error_message: str | None = None,
started_at: datetime | None = None,
finished_at: datetime | None = None,
) -> None:
"""更新候选记录状态。"""
record = self._candidate_records.get(idx)
if not record:
return
record.status = status
if skip_reason is not None:
record.skip_reason = skip_reason
if status_code is not None:
record.status_code = status_code
if error_type is not None:
record.error_type = error_type
if error_message is not None:
record.error_message = error_message
if started_at is not None:
record.started_at = started_at
if finished_at is not None:
record.finished_at = finished_at
try:
self.db.flush()
except Exception as exc:
logger.warning(
"[AsyncTaskOrchestrator] Failed to update candidate record %d: %s",
idx,
str(exc),
)
def _commit_candidate_records(self) -> None:
"""提交候选记录到数据库。"""
if not self._candidate_records:
return
try:
self.db.commit()
except Exception as exc:
logger.warning(
"[AsyncTaskOrchestrator] Failed to commit candidate records: %s",
str(exc),
)
self.db.rollback()
async def _ensure_initialized(self) -> None:
if self._cache_scheduler is not None:
return
# 使用 SystemConfigService 读取运行时调度策略(与 Chat/CLI 一致)
priority_mode = SystemConfigService.get_config(
self.db,
"provider_priority_mode",
"provider",
)
scheduling_mode = SystemConfigService.get_config(
self.db,
"scheduling_mode",
"cache_affinity",
)
self._cache_scheduler = await get_cache_aware_scheduler(
self.redis,
priority_mode=priority_mode,
scheduling_mode=scheduling_mode,
)
self._candidate_resolver = CandidateResolver(
db=self.db,
cache_scheduler=self._cache_scheduler,
)
self._error_classifier = ErrorClassifier(db=self.db, cache_scheduler=self._cache_scheduler)
def _should_stop_on_http_error(self, *, status_code: int, error_text: str) -> bool:
"""
判断某个上游 HTTP 错误是否为“客户端错误”(不应 failover
规则:
- 401/403/429一般是 key/权限/限流问题,优先 failover
- 其他 4xx若 ErrorClassifier 判断为客户端请求错误,则停止
"""
if status_code in (401, 403, 429):
return False
if 400 <= status_code < 500:
assert self._error_classifier is not None
return self._error_classifier.is_client_error(error_text)
return False
async def submit_with_failover(
self,
*,
api_format: str,
model_name: str,
affinity_key: str,
user_api_key: ApiKey,
request_id: str | None,
task_type: str,
submit_func: SubmitFunc,
extract_external_task_id: ExtractExternalTaskIdFunc,
supported_auth_types: set[str] | None = None,
allow_format_conversion: bool = False,
capability_requirements: dict[str, bool] | None = None,
max_candidates: int | None = None,
) -> SubmitOutcome:
"""
提交异步任务并在失败时自动尝试下一个候选,直到拿到 external_task_id。
Returns:
SubmitOutcome包含选中的候选 + external_task_id + candidate_keys + billing rule lookup
Raises:
UpstreamClientRequestError: 判定为客户端请求错误(不应 failover
ProviderNotAvailableException: 没有可用候选(调度器层面)
AllCandidatesFailedError: 有候选但全部提交失败
"""
await self._ensure_initialized()
assert self._candidate_resolver is not None
logger.info(
"[AsyncTaskOrchestrator] submit_with_failover: "
"api_format=%s, model=%s, task_type=%s, request_id=%s",
api_format,
model_name,
task_type,
request_id,
)
candidates, _global_model_id = await self._candidate_resolver.fetch_candidates(
api_format=api_format,
model_name=model_name,
affinity_key=affinity_key,
user_api_key=user_api_key,
request_id=request_id,
is_stream=False,
capability_requirements=capability_requirements,
)
logger.info(
"[AsyncTaskOrchestrator] fetch_candidates returned %d candidates for model=%s",
len(candidates),
model_name,
)
# 如果没有候选,直接抛出异常
if not candidates:
logger.error(
"[AsyncTaskOrchestrator] No candidates returned from fetch_candidates for model=%s",
model_name,
)
raise ProviderNotAvailableException("No candidates available")
if max_candidates is not None and max_candidates > 0:
candidates = candidates[:max_candidates]
# 创建候选记录(用于链路追踪)
self._candidate_records = self._create_candidate_records(
candidates=candidates,
request_id=request_id,
user_api_key=user_api_key,
)
candidate_keys: list[dict[str, Any]] = []
eligible_count = 0
last_status_code: int | None = None
for idx, cand in enumerate(candidates):
submit_started_at = datetime.now(timezone.utc)
auth_type = getattr(cand.key, "auth_type", "api_key") or "api_key"
candidate_info: dict[str, Any] = {
"index": idx,
"provider_id": cand.provider.id,
"provider_name": cand.provider.name,
"endpoint_id": cand.endpoint.id,
"key_id": cand.key.id,
"key_name": cand.key.name,
"auth_type": auth_type,
"priority": getattr(cand.key, "priority", 0) or 0,
"is_cached": bool(getattr(cand, "is_cached", False)),
}
candidate_keys.append(candidate_info)
logger.info(
"[AsyncTaskOrchestrator] Checking candidate %d: provider=%s, is_skipped=%s, skip_reason=%s, needs_conversion=%s, auth_type=%s",
idx,
cand.provider.name,
getattr(cand, "is_skipped", False),
getattr(cand, "skip_reason", None),
getattr(cand, "needs_conversion", False),
auth_type,
)
# 调度器层面标记为跳过(健康/熔断/并发等)
if getattr(cand, "is_skipped", False):
skip_reason = getattr(cand, "skip_reason", None) or "skipped"
candidate_info.update(
{
"skipped": True,
"skip_reason": skip_reason,
}
)
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
logger.info(
"[AsyncTaskOrchestrator] Candidate %d skipped: is_skipped=True, reason=%s",
idx,
cand.skip_reason,
)
continue
# 视频/图片等直连 upstream 的 handler 目前不支持跨格式转换
if not allow_format_conversion and bool(getattr(cand, "needs_conversion", False)):
candidate_info.update(
{"skipped": True, "skip_reason": "format_conversion_not_supported"}
)
self._update_candidate_record(
idx, status="skipped", skip_reason="format_conversion_not_supported"
)
logger.info("[AsyncTaskOrchestrator] Candidate %d skipped: needs_conversion", idx)
continue
# auth_type 过滤
if supported_auth_types is not None and auth_type not in supported_auth_types:
skip_reason = f"unsupported_auth_type:{auth_type}"
candidate_info.update(
{
"skipped": True,
"skip_reason": skip_reason,
}
)
self._update_candidate_record(idx, status="skipped", skip_reason=skip_reason)
logger.info(
"[AsyncTaskOrchestrator] Candidate %d skipped: unsupported_auth_type=%s",
idx,
auth_type,
)
continue
# billing rule 过滤(可选)
rule_lookup: BillingRuleLookupResult | None = None
has_billing_rule = True
if config.billing_require_rule:
logger.info(
"[AsyncTaskOrchestrator] Checking billing rule for candidate %d (billing_require_rule=True)",
idx,
)
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=cand.provider.id,
model_name=model_name,
task_type=task_type,
)
has_billing_rule = rule_lookup is not None
logger.info(
"[AsyncTaskOrchestrator] Billing rule lookup result: has_rule=%s",
has_billing_rule,
)
if not has_billing_rule:
candidate_info.update(
{
"has_billing_rule": False,
"skipped": True,
"skip_reason": "billing_rule_missing",
}
)
self._update_candidate_record(
idx, status="skipped", skip_reason="billing_rule_missing"
)
logger.info(
"[AsyncTaskOrchestrator] Candidate %d skipped: billing_rule_missing", idx
)
continue
candidate_info["has_billing_rule"] = has_billing_rule
logger.info("[AsyncTaskOrchestrator] Candidate %d eligible, attempting submit", idx)
eligible_count += 1
# 更新记录为 pending 状态(开始尝试)
self._update_candidate_record(idx, status="pending", started_at=submit_started_at)
# 尝试提交
try:
response = await submit_func(cand)
except Exception as exc:
finished_at = datetime.now(timezone.utc)
logger.error(
"[AsyncTaskOrchestrator] Candidate %d submit exception: %s: %s",
idx,
type(exc).__name__,
str(exc),
)
error_msg = _sanitize(str(exc))
candidate_info.update(
{
"attempt_status": "exception",
"error_type": type(exc).__name__,
"error_message": error_msg,
}
)
self._update_candidate_record(
idx,
status="failed",
error_type=type(exc).__name__,
error_message=error_msg,
finished_at=finished_at,
)
continue
logger.info(
"[AsyncTaskOrchestrator] Candidate %d submit response: status_code=%d",
idx,
response.status_code,
)
last_status_code = int(getattr(response, "status_code", 0) or 0)
# 上游错误:决定是否停止
if response.status_code >= 400:
finished_at = datetime.now(timezone.utc)
error_text = ""
try:
error_text = response.text or ""
except Exception:
error_text = ""
error_msg = _sanitize(error_text)
candidate_info.update(
{
"attempt_status": "http_error",
"status_code": response.status_code,
"error_message": error_msg,
}
)
self._update_candidate_record(
idx,
status="failed",
status_code=response.status_code,
error_type="http_error",
error_message=error_msg,
finished_at=finished_at,
)
if self._should_stop_on_http_error(
status_code=response.status_code, error_text=error_text
):
self._commit_candidate_records()
raise UpstreamClientRequestError(
response=response,
candidate_keys=candidate_keys,
)
continue
# 解析任务 ID200 但缺字段也视为失败并 failover
payload: dict[str, Any] | None = None
try:
data = response.json()
if isinstance(data, dict):
payload = data
logger.info(
"[AsyncTaskOrchestrator] Candidate %d response payload: %s",
idx,
str(payload)[:500] if payload else "None",
)
except Exception as exc:
finished_at = datetime.now(timezone.utc)
logger.error(
"[AsyncTaskOrchestrator] Candidate %d invalid JSON: %s",
idx,
str(exc),
)
error_msg = _sanitize(str(exc))
candidate_info.update(
{
"attempt_status": "invalid_json",
"error_type": type(exc).__name__,
"error_message": error_msg,
}
)
self._update_candidate_record(
idx,
status="failed",
status_code=response.status_code,
error_type="invalid_json",
error_message=error_msg,
finished_at=finished_at,
)
continue
external_task_id = extract_external_task_id(payload or {})
logger.info(
"[AsyncTaskOrchestrator] Candidate %d extracted task_id: %s",
idx,
external_task_id,
)
if not external_task_id:
finished_at = datetime.now(timezone.utc)
candidate_info.update(
{
"attempt_status": "empty_task_id",
"error_message": "Upstream returned empty task id",
}
)
self._update_candidate_record(
idx,
status="failed",
status_code=response.status_code,
error_type="empty_task_id",
error_message="Upstream returned empty task id",
finished_at=finished_at,
)
logger.warning(
"[AsyncTaskOrchestrator] Candidate %d: empty task_id, payload keys: %s",
idx,
list(payload.keys()) if payload else [],
)
continue
# 成功
finished_at = datetime.now(timezone.utc)
candidate_info.update({"attempt_status": "success", "selected": True})
self._update_candidate_record(
idx,
status="success",
status_code=response.status_code,
finished_at=finished_at,
)
self._commit_candidate_records()
return SubmitOutcome(
candidate=cand,
candidate_keys=candidate_keys,
external_task_id=str(external_task_id),
rule_lookup=rule_lookup,
upstream_payload=payload,
)
# 没有任何候选可尝试
if not candidates:
raise ProviderNotAvailableException("No candidates available")
# 提交所有候选记录
self._commit_candidate_records()
if eligible_count == 0:
reason = "no_eligible_candidates"
if config.billing_require_rule:
reason = "no_candidate_with_billing_rule"
raise AllCandidatesFailedError(
reason=reason,
candidate_keys=candidate_keys,
last_status_code=last_status_code,
)
raise AllCandidatesFailedError(
reason="all_candidates_failed",
candidate_keys=candidate_keys,
last_status_code=last_status_code,
)
__all__ = [
"AsyncTaskOrchestrator",
"SubmitOutcome",
"AllCandidatesFailedError",
"UpstreamClientRequestError",
"CandidateUnsupportedError",
"CandidateSubmissionError",
]

View File

@@ -0,0 +1,240 @@
"""
Task poller (Phase2)
Provides a generic polling skeleton for async tasks.
Currently wired with a video poller adapter.
优化HTTP 请求期间不持有数据库连接,避免阻塞其他请求。
采用三阶段处理:准备数据 -> HTTP 请求 -> 更新数据库。
"""
from __future__ import annotations
import asyncio
from datetime import datetime, timezone
from typing import Any, Protocol, runtime_checkable
from uuid import uuid4
from sqlalchemy.orm import Session
from src.core.api_format.conversion.internal_video import InternalVideoPollResult
from src.core.logger import logger
from src.database import create_session
from src.services.system.scheduler import get_scheduler
from src.services.task.impl.video_poller import VideoPollContext, VideoTaskPollerAdapter
@runtime_checkable
class TaskPollerAdapter(Protocol):
task_type: str
# scheduler
job_id: str
job_name: str
interval_seconds: int
# distributed lock (optional, best-effort)
lock_key: str
lock_ttl: int
# execution
batch_size: int
concurrency: int
consecutive_failure_alert_threshold: int
def list_due_task_ids(self, db: Session, *, now: datetime, limit: int) -> list[str]: ...
def get_task(self, db: Session, task_id: str) -> Any | None: ...
# 分阶段处理方法(推荐使用)
async def prepare_poll_context(
self, db: Session, task: Any
) -> Any: ... # Returns context or error result
async def poll_task_http(self, ctx: Any) -> Any: ... # Returns poll result
async def update_task_after_poll(
self,
task_id: str,
result: Any,
ctx: Any,
redis_client: Any | None,
error_exception: Exception | None = None,
) -> None: ...
# 旧版方法(保留兼容性)
async def poll_single_task(
self, db: Session, task: Any, *, redis_client: Any | None
) -> None: ...
def sanitize_error_message(self, message: str) -> str: ...
class TaskPollerService:
"""Generic background poller for async tasks."""
def __init__(self, adapter: TaskPollerAdapter) -> None:
self.adapter = adapter
self._lock = asyncio.Lock()
self.redis: Any | None = None
self._semaphore: asyncio.Semaphore | None = None
self._consecutive_failures = 0
async def start(self) -> None:
if self._semaphore is None:
self._semaphore = asyncio.Semaphore(self.adapter.concurrency)
# lazy import to avoid redis hard dependency in local runs
from src.clients.redis_client import get_redis_client
if self.redis is None:
self.redis = await get_redis_client(require_redis=False)
scheduler = get_scheduler()
scheduler.add_interval_job(
self.poll_pending_tasks,
seconds=self.adapter.interval_seconds,
job_id=self.adapter.job_id,
name=self.adapter.job_name,
)
async def stop(self) -> None:
scheduler = get_scheduler()
scheduler.remove_job(self.adapter.job_id)
async def poll_pending_tasks(self) -> None:
async with self._lock:
token = await self._acquire_redis_lock()
if token is None:
return
try:
with create_session() as db:
now = datetime.now(timezone.utc)
task_ids = self.adapter.list_due_task_ids(
db, now=now, limit=self.adapter.batch_size
)
if not task_ids:
self._consecutive_failures = 0
return
poll_results: list[bool] = []
if self._semaphore is None:
self._semaphore = asyncio.Semaphore(self.adapter.concurrency)
semaphore = self._semaphore
async def poll_with_semaphore(task_id: str) -> None:
async with semaphore:
try:
# ========== 阶段 1准备数据短暂持有连接==========
with create_session() as task_db:
task_obj = self.adapter.get_task(task_db, task_id)
if not task_obj:
logger.warning(
"[%s] Task %s disappeared during poll",
self.adapter.task_type,
task_id,
)
poll_results.append(True)
return
ctx_or_result = await self.adapter.prepare_poll_context(
task_db, task_obj
)
# 检查是否是错误结果(而非上下文)
if isinstance(ctx_or_result, InternalVideoPollResult):
# 准备阶段就失败了,直接更新任务状态
await self.adapter.update_task_after_poll(
task_id=task_id,
result=ctx_or_result,
ctx=None, # type: ignore[arg-type]
redis_client=self.redis,
)
poll_results.append(True)
return
ctx: VideoPollContext = ctx_or_result
# ========== 阶段 2HTTP 请求(不持有数据库连接)==========
error_exception: Exception | None = None
try:
result = await self.adapter.poll_task_http(ctx)
except Exception as http_exc:
# HTTP 请求失败,记录异常以便后续处理
error_exception = http_exc
result = InternalVideoPollResult(
status=None, # type: ignore[arg-type]
error_message=str(http_exc),
)
# ========== 阶段 3更新数据库获取新连接==========
await self.adapter.update_task_after_poll(
task_id=task_id,
result=result,
ctx=ctx,
redis_client=self.redis,
error_exception=error_exception,
)
poll_results.append(True)
except Exception as exc:
logger.exception(
"[%s] Unexpected error polling task %s: %s",
self.adapter.task_type,
task_id,
self.adapter.sanitize_error_message(str(exc)),
)
poll_results.append(False)
async with asyncio.TaskGroup() as tg:
for tid in task_ids:
tg.create_task(poll_with_semaphore(tid))
batch_failures = sum(1 for r in poll_results if r is False)
if batch_failures == len(task_ids):
self._consecutive_failures += 1
if (
self._consecutive_failures
>= self.adapter.consecutive_failure_alert_threshold
):
logger.error(
"[ALERT] %s poller: %d consecutive batches failed.",
self.adapter.task_type,
self._consecutive_failures,
)
else:
self._consecutive_failures = 0
finally:
await self._release_redis_lock(token)
async def _acquire_redis_lock(self) -> str | None:
if not self.redis:
return "no_redis"
token = str(uuid4())
acquired = await self.redis.set(
self.adapter.lock_key, token, nx=True, ex=self.adapter.lock_ttl
)
return token if acquired else None
async def _release_redis_lock(self, token: str) -> None:
if not self.redis or token == "no_redis":
return
script = """
if redis.call('GET', KEYS[1]) == ARGV[1] then
return redis.call('DEL', KEYS[1])
end
return 0
"""
await self.redis.eval(script, 1, self.adapter.lock_key, token)
_task_poller: TaskPollerService | None = None
def get_task_poller() -> TaskPollerService:
global _task_poller
if _task_poller is None:
_task_poller = TaskPollerService(VideoTaskPollerAdapter())
return _task_poller

View File

@@ -576,6 +576,12 @@ class UsageService:
"""更新已存在的 Usage 记录(内部方法)"""
# 更新关键字段
existing_usage.provider_name = usage_params["provider_name"]
existing_usage.model = usage_params["model"]
existing_usage.request_type = usage_params["request_type"]
existing_usage.api_format = usage_params["api_format"]
existing_usage.endpoint_api_format = usage_params["endpoint_api_format"]
existing_usage.has_format_conversion = usage_params["has_format_conversion"]
existing_usage.is_stream = usage_params["is_stream"]
existing_usage.status = usage_params["status"]
existing_usage.status_code = usage_params["status_code"]
existing_usage.error_message = usage_params["error_message"]
@@ -621,6 +627,10 @@ class UsageService:
existing_usage.provider_endpoint_id = usage_params["provider_endpoint_id"]
existing_usage.provider_api_key_id = usage_params["provider_api_key_id"]
# 更新元数据(如 billing_snapshot/dimensions 等)
if usage_params.get("request_metadata") is not None:
existing_usage.request_metadata = usage_params["request_metadata"]
# 更新模型映射信息
if target_model is not None:
existing_usage.target_model = target_model
@@ -1000,6 +1010,11 @@ class UsageService:
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
)
# 结算标记record_usage_async 写入的 Usage 通常为终态记录
if status not in ("pending", "streaming"):
usage.billing_status = "settled"
usage.finalized_at = datetime.now(timezone.utc)
db.commit() # 立即提交事务,释放数据库锁
return usage
@@ -1172,6 +1187,11 @@ class UsageService:
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
)
# 结算标记:终态请求写入 settled + finalized_at
if status not in ("pending", "streaming"):
usage.billing_status = "settled"
usage.finalized_at = datetime.now(timezone.utc)
# 提交事务
try:
db.commit()
@@ -1297,17 +1317,39 @@ class UsageService:
is_free_tier=is_free_tier,
)
# Upsert与 record_usage 保持一致
# Upsert并发幂等:优先用 billing_status 作为结算闸门
from sqlalchemy import update
existing_usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if existing_usage:
# 避免重复记账:若已是终态记录,直接返回(批量接口也采用该策略
if existing_usage.status not in ("pending", "streaming"):
# 避免重复记账:若已结算/作废,直接返回(防止并发重复加计数
if getattr(existing_usage, "billing_status", None) in ("settled", "void"):
logger.debug(
"record_usage_with_custom_cost: request_id=%s already finalized (status=%s), skip",
"record_usage_with_custom_cost: request_id=%s already finalized (billing_status=%s), skip",
request_id,
existing_usage.status,
getattr(existing_usage, "billing_status", None),
)
return existing_usage
# 并发闸门:只有 billing_status='pending' 的那一次调用可以继续
now = datetime.now(timezone.utc)
claim = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "pending",
)
.values(billing_status="settled", finalized_at=now)
)
if claim.rowcount != 1:
# 已被其他 worker 抢先处理(或被 VOID
latest = db.query(Usage).filter(Usage.request_id == request_id).first()
return latest or existing_usage
# 同步 ORM 对象(避免后续代码读到旧值)
existing_usage.billing_status = "settled"
existing_usage.finalized_at = now
cls._update_existing_usage(existing_usage, usage_params, target_model)
usage = existing_usage
else:
@@ -1382,6 +1424,11 @@ class UsageService:
.values(monthly_used_usd=Provider.monthly_used_usd + actual_total_cost)
)
# 结算标记record_usage_with_custom_cost 写入/更新的 Usage 通常为终态记录
if status not in ("pending", "streaming"):
usage.billing_status = "settled"
usage.finalized_at = datetime.now(timezone.utc)
try:
db.commit()
except Exception as e:
@@ -2147,35 +2194,31 @@ class UsageService:
# ========== 请求状态追踪方法 ==========
@classmethod
def create_pending_usage(
def begin_pending_usage(
cls,
db: Session,
request_id: str,
user: User | None,
api_key: ApiKey | None,
model: str,
*,
is_stream: bool = False,
request_type: str = "chat",
api_format: str | None = None,
request_headers: dict[str, Any] | None = None,
request_body: Any | None = None,
) -> Usage:
"""
创建 pending 状态的使用记录(在请求开始时调用)
创建(或返回已有)pending Usage 记录,但**不提交事务**。
Args:
db: 数据库会话
request_id: 请求ID
user: 用户对象
api_key: API Key 对象
model: 模型名称
is_stream: 是否流式请求
api_format: API 格式
request_headers: 请求头
request_body: 请求体
Returns:
创建的 Usage 记录
适用场景:
- ApplicationService 在同一事务内创建 pending usage + task + candidates
- submit 幂等:重复调用同一 request_id 时返回已有记录
"""
existing = db.query(Usage).filter(Usage.request_id == request_id).first()
if existing:
return existing
# 根据配置决定是否记录请求详情
should_log_headers = SystemConfigService.should_log_headers(db)
should_log_body = SystemConfigService.should_log_body(db)
@@ -2204,21 +2247,359 @@ class UsageService:
output_tokens=0,
total_tokens=0,
total_cost_usd=0.0,
request_type="chat",
request_type=request_type,
api_format=api_format,
is_stream=is_stream,
status="pending",
billing_status="pending",
request_headers=processed_request_headers,
request_body=processed_request_body,
)
db.add(usage)
db.flush()
return usage
@classmethod
def create_pending_usage(
cls,
db: Session,
request_id: str,
user: User | None,
api_key: ApiKey | None,
model: str,
is_stream: bool = False,
request_type: str = "chat",
api_format: str | None = None,
request_headers: dict[str, Any] | None = None,
request_body: Any | None = None,
) -> Usage:
"""
创建 pending 状态的使用记录(在请求开始时调用)
Args:
db: 数据库会话
request_id: 请求ID
user: 用户对象
api_key: API Key 对象
model: 模型名称
is_stream: 是否流式请求
api_format: API 格式
request_headers: 请求头
request_body: 请求体
Returns:
创建的 Usage 记录
"""
usage = cls.begin_pending_usage(
db,
request_id=request_id,
user=user,
api_key=api_key,
model=model,
is_stream=is_stream,
request_type=request_type,
api_format=api_format,
request_headers=request_headers,
request_body=request_body,
)
db.commit()
logger.debug(f"创建 pending 使用记录: request_id={request_id}, model={model}")
return usage
# ========== billing_status 并发幂等 finalize ==========
@classmethod
def finalize_settled(
cls,
db: Session,
request_id: str,
*,
total_cost_usd: float,
request_cost_usd: float | None = None,
status: str = "completed",
status_code: int = 200,
error_message: str | None = None,
response_time_ms: int | None = None,
billing_snapshot: dict[str, Any] | None = None,
extra_metadata: dict[str, Any] | None = None,
) -> bool:
"""
并发安全的幂等 finalizesettled
约定:
- 仅当 billing_status='pending' 时才会生效rowcount==1
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
cost = float(total_cost_usd)
request_cost = float(request_cost_usd) if request_cost_usd is not None else cost
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "pending",
)
.values(
billing_status="settled",
finalized_at=now,
total_cost_usd=cost,
request_cost_usd=request_cost,
status=status,
status_code=status_code,
error_message=error_message,
response_time_ms=response_time_ms,
)
)
if result.rowcount != 1:
return False
# 写入审计快照(只在本次 finalize 生效时执行)
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if usage:
metadata = usage.request_metadata or {}
if billing_snapshot is not None:
metadata["billing_snapshot"] = billing_snapshot
if extra_metadata:
metadata.update(extra_metadata)
usage.request_metadata = metadata
return True
@classmethod
def finalize_void(
cls,
db: Session,
request_id: str,
*,
reason: str | None = None,
status_code: int = 499,
) -> bool:
"""
并发安全的幂等 finalizevoid不收费
约定:
- 仅当 billing_status='pending' 时才会生效rowcount==1
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "pending",
)
.values(
billing_status="void",
finalized_at=now,
total_cost_usd=0.0,
request_cost_usd=0.0,
status="cancelled",
status_code=status_code,
error_message=reason,
response_time_ms=None,
)
)
return result.rowcount == 1
@classmethod
def finalize_submitted(
cls,
db: Session,
request_id: str,
*,
provider_name: str,
provider_id: str | None = None,
provider_endpoint_id: str | None = None,
provider_api_key_id: str | None = None,
response_time_ms: int | None = None,
status_code: int = 200,
endpoint_api_format: str | None = None,
provider_request_headers: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
response_body: Any | None = None,
) -> bool:
"""
异步任务提交成功时的幂等结算。
将 pending 使用记录标记为 settled费用暂时为 0。
后续轮询完成后通过 update_settled_billing 更新实际费用。
约定:
- 仅当 billing_status='pending' 时才会生效rowcount==1
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
# 处理响应头和响应体
should_log_headers = SystemConfigService.should_log_headers(db)
should_log_body = SystemConfigService.should_log_body(db)
processed_provider_headers = None
if should_log_headers and provider_request_headers:
processed_provider_headers = SystemConfigService.mask_sensitive_headers(
db, provider_request_headers
)
processed_response_headers = None
if should_log_headers and response_headers:
processed_response_headers = dict(response_headers)
processed_response_body = None
if should_log_body and response_body:
processed_response_body = SystemConfigService.truncate_body(
db, response_body, is_request=False
)
values: dict[str, Any] = {
"billing_status": "settled",
"finalized_at": now,
"total_cost_usd": 0.0,
"request_cost_usd": 0.0,
"status": "completed",
"status_code": status_code,
"response_time_ms": response_time_ms,
"provider_name": provider_name,
"provider_id": provider_id,
"provider_endpoint_id": provider_endpoint_id,
"provider_api_key_id": provider_api_key_id,
"endpoint_api_format": endpoint_api_format,
}
if processed_provider_headers is not None:
values["provider_request_headers"] = processed_provider_headers
if processed_response_headers is not None:
values["response_headers"] = processed_response_headers
if processed_response_body is not None:
values["response_body"] = processed_response_body
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "pending",
)
.values(**values)
)
return result.rowcount == 1
@classmethod
def update_settled_billing(
cls,
db: Session,
request_id: str,
*,
total_cost_usd: float,
request_cost_usd: float | None = None,
status: str = "completed",
status_code: int = 200,
error_message: str | None = None,
response_time_ms: int | None = None,
billing_snapshot: dict[str, Any] | None = None,
extra_metadata: dict[str, Any] | None = None,
) -> bool:
"""
更新已结算记录的计费信息(用于异步任务轮询完成后)。
与 finalize_settled 不同:
- finalize_settled: pending -> settled首次结算
- update_settled_billing: settled -> settled更新费用
约定:
- 仅当 billing_status='settled' 时才会生效
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
cost = float(total_cost_usd)
request_cost = float(request_cost_usd) if request_cost_usd is not None else cost
values: dict[str, Any] = {
"total_cost_usd": cost,
"request_cost_usd": request_cost,
"status": status,
"status_code": status_code,
}
if error_message is not None:
values["error_message"] = error_message
if response_time_ms is not None:
values["response_time_ms"] = response_time_ms
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "settled",
)
.values(**values)
)
if result.rowcount != 1:
return False
# 写入审计快照
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if usage:
metadata = usage.request_metadata or {}
if billing_snapshot is not None:
metadata["billing_snapshot"] = billing_snapshot
if extra_metadata:
metadata.update(extra_metadata)
metadata["billing_updated_at"] = now.isoformat()
usage.request_metadata = metadata
return True
@classmethod
def void_settled(
cls,
db: Session,
request_id: str,
*,
reason: str | None = None,
status_code: int = 499,
) -> bool:
"""
将已结算的记录作废(用于异步任务取消)。
与 finalize_void 不同:
- finalize_void: pending -> void未结算时作废
- void_settled: settled -> void已结算后取消费用归零
约定:
- 仅当 billing_status='settled' 时才会生效
- 不在本方法内 commit由调用方决定事务提交时机
"""
from sqlalchemy import update
now = datetime.now(timezone.utc)
result = db.execute(
update(Usage)
.where(
Usage.request_id == request_id,
Usage.billing_status == "settled",
)
.values(
billing_status="void",
finalized_at=now,
total_cost_usd=0.0,
request_cost_usd=0.0,
status="cancelled",
status_code=status_code,
error_message=reason,
)
)
return result.rowcount == 1
@classmethod
def update_usage_status(
cls,
@@ -2235,6 +2616,7 @@ class UsageService:
api_format: str | None = None,
endpoint_api_format: str | None = None,
has_format_conversion: bool | None = None,
status_code: int | None = None,
) -> Usage | None:
"""
快速更新使用记录状态
@@ -2253,6 +2635,7 @@ class UsageService:
api_format: API 格式(可选,用于获取按格式配置的倍率)
endpoint_api_format: 端点原生 API 格式(可选)
has_format_conversion: 是否发生了格式转换(可选)
status_code: HTTP 状态码(可选)
Returns:
更新后的 Usage 记录,如果未找到则返回 None
@@ -2295,6 +2678,16 @@ class UsageService:
usage.endpoint_api_format = endpoint_api_format
if has_format_conversion is not None:
usage.has_format_conversion = has_format_conversion
if status_code is not None:
usage.status_code = status_code
# 结算状态:当请求进入终态时,将 billing_status 标记为 settled
# 注意:取消是否应 VOID/部分结算由更高层策略决定;这里默认终态均视为已结算。
if status in ("completed", "failed", "cancelled"):
if getattr(usage, "billing_status", None) == "pending":
usage.billing_status = "settled"
if getattr(usage, "finalized_at", None) is None:
usage.finalized_at = datetime.now(timezone.utc)
db.commit()

View File

@@ -1,10 +0,0 @@
"""
视频相关服务
"""
from src.services.video.task_poller import VideoTaskPollerService, get_video_task_poller
__all__ = [
"VideoTaskPollerService",
"get_video_task_poller",
]

View File

@@ -1,446 +0,0 @@
"""
视频任务后台轮询服务
"""
from __future__ import annotations
import asyncio
from datetime import datetime, timedelta, timezone
from typing import Any
from uuid import uuid4
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import ProviderAuthInfo, get_provider_auth
from src.api.handlers.base.video_handler_base import (
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.clients.http_client import HTTPClientPool
from src.clients.redis_client import get_redis_client
from src.config.settings import config
from src.core.api_format import (
build_upstream_headers_for_endpoint,
get_extra_headers_from_endpoint,
make_signature_key,
)
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.database import create_session
from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.services.system.scheduler import get_scheduler
from src.services.task.impl.video_telemetry import VideoTelemetry
# 永久性错误指示词(用于降级判断,不应重试)
_PERMANENT_ERROR_INDICATORS = frozenset(
{
"not found",
"404",
"unauthorized",
"401",
"forbidden",
"403",
"invalid request",
"invalid api key",
"does not exist",
}
)
class PollHTTPError(RuntimeError):
"""HTTP 轮询错误,携带状态码便于区分临时/永久错误"""
def __init__(self, status_code: int, message: str):
super().__init__(message)
self.status_code = status_code
class VideoTaskPollerService:
"""后台轮询视频生成任务状态"""
LOCK_KEY = "video_task_poller:lock"
LOCK_TTL = 60
MAX_BACKOFF_SECONDS = 300
# 连续失败告警阈值
CONSECUTIVE_FAILURE_ALERT_THRESHOLD = 5
def __init__(self) -> None:
self._lock = asyncio.Lock()
self.redis = None
self._openai_normalizer = OpenAINormalizer()
self._gemini_normalizer = GeminiNormalizer()
# 追踪连续失败次数(用于告警)
self._consecutive_failures = 0
# 从配置读取参数
self._batch_size = config.video_poll_batch_size
self._concurrency = config.video_poll_concurrency
# Semaphore 延迟初始化,避免在事件循环外创建
self._semaphore: asyncio.Semaphore | None = None
async def start(self) -> None:
# 在事件循环内初始化 Semaphore
if self._semaphore is None:
self._semaphore = asyncio.Semaphore(self._concurrency)
if self.redis is None:
self.redis = await get_redis_client(require_redis=False)
scheduler = get_scheduler()
scheduler.add_interval_job(
self.poll_pending_tasks,
seconds=config.video_poll_interval_seconds,
job_id="video_task_poller",
name="视频任务轮询",
)
async def stop(self) -> None:
"""停止轮询服务"""
scheduler = get_scheduler()
scheduler.remove_job("video_task_poller")
async def poll_pending_tasks(self) -> None:
async with self._lock:
token = await self._acquire_redis_lock()
if token is None:
return
try:
with create_session() as db:
now = datetime.now(timezone.utc)
tasks = (
db.query(VideoTask)
.filter(
VideoTask.status.in_(
[
VideoStatus.SUBMITTED.value,
VideoStatus.QUEUED.value,
VideoStatus.PROCESSING.value,
]
),
VideoTask.next_poll_at <= now,
VideoTask.poll_count < VideoTask.max_poll_count,
)
.order_by(VideoTask.next_poll_at.asc())
.limit(self._batch_size)
.all()
)
if not tasks:
# 无任务时重置连续失败计数
self._consecutive_failures = 0
return
# 提取任务 ID 列表,释放查询 session 后逐个轮询
task_ids = [t.id for t in tasks]
# 并发轮询:每个任务使用独立 session避免共享 session 的并发风险
poll_results: list[bool] = []
# 确保 semaphore 已初始化(在 start 中初始化,此处防御性检查)
if self._semaphore is None:
self._semaphore = asyncio.Semaphore(self._concurrency)
semaphore = self._semaphore
async def poll_with_semaphore(task_id: str) -> None:
"""带信号量的轮询,结果写入 poll_results"""
async with semaphore:
try:
with create_session() as task_db:
task_obj = task_db.query(VideoTask).get(task_id)
if not task_obj:
logger.warning("Task %s disappeared during poll", task_id)
poll_results.append(True)
return
await self._poll_single_task(task_db, task_obj)
task_db.commit()
poll_results.append(True)
except Exception as exc:
logger.exception(
"Unexpected error polling task %s: %s",
task_id,
sanitize_error_message(str(exc)),
)
poll_results.append(False)
async with asyncio.TaskGroup() as tg:
for tid in task_ids:
tg.create_task(poll_with_semaphore(tid))
batch_failures = sum(1 for r in poll_results if r is False)
# 更新连续失败计数并检查告警阈值
if batch_failures == len(task_ids):
self._consecutive_failures += 1
if self._consecutive_failures >= self.CONSECUTIVE_FAILURE_ALERT_THRESHOLD:
logger.error(
"[ALERT] Video task poller: %d consecutive batches failed. "
"Provider connectivity or configuration issue suspected.",
self._consecutive_failures,
)
else:
self._consecutive_failures = 0
finally:
await self._release_redis_lock(token)
async def _poll_single_task(self, db: Session, task: VideoTask) -> None:
try:
result = await self._poll_task_status(db, task)
if result.status == VideoStatus.COMPLETED:
task.status = VideoStatus.COMPLETED.value
task.video_url = result.video_url
task.video_expires_at = result.expires_at
task.completed_at = datetime.now(timezone.utc)
task.progress_percent = 100
# 存储多视频 URLGemini sampleCount > 1 时)
if result.video_urls:
task.video_urls = result.video_urls
# 保存上游原始响应(用于审计/重算)
self._attach_poll_raw_response(task, result)
elif result.status == VideoStatus.FAILED:
task.status = VideoStatus.FAILED.value
task.error_code = result.error_code
task.error_message = result.error_message
task.completed_at = datetime.now(timezone.utc)
self._attach_poll_raw_response(task, result)
else:
task.poll_count += 1
task.progress_percent = result.progress_percent
task.next_poll_at = datetime.now(timezone.utc) + timedelta(
seconds=task.poll_interval_seconds
)
except Exception as exc:
task.poll_count += 1
error_msg = sanitize_error_message(str(exc))
logger.warning("Poll error for task %s: %s", task.id, error_msg)
task.progress_message = f"Poll error: {error_msg}"
# 区分临时性错误和永久性错误
status_code = exc.status_code if isinstance(exc, PollHTTPError) else None
is_permanent = self._is_permanent_error(exc, status_code=status_code)
if is_permanent:
task.status = VideoStatus.FAILED.value
task.error_code = "poll_permanent_error"
task.error_message = error_msg
task.completed_at = datetime.now(timezone.utc)
else:
# 临时性错误:指数退避重试
backoff = min(
task.poll_interval_seconds * (2 ** min(task.retry_count, 5)),
self.MAX_BACKOFF_SECONDS,
)
task.retry_count += 1
task.next_poll_at = datetime.now(timezone.utc) + timedelta(seconds=backoff)
# 检查是否超过最大轮询次数(超时)
task.updated_at = datetime.now(timezone.utc)
if task.poll_count >= task.max_poll_count and task.status not in [
VideoStatus.COMPLETED.value,
VideoStatus.FAILED.value,
VideoStatus.CANCELLED.value,
]:
task.status = VideoStatus.FAILED.value
task.error_code = "poll_timeout"
task.error_message = f"Task timed out after {task.poll_count} polls"
task.completed_at = datetime.now(timezone.utc)
# 终态写入 Usage复用外层 per-task session
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try:
await VideoTelemetry(db, redis_client=self.redis).record_terminal_usage(task)
except Exception as exc:
logger.exception(
"Failed to record video usage for task=%s: %s",
task.id,
sanitize_error_message(str(exc)),
)
def _attach_poll_raw_response(self, task: VideoTask, result: InternalVideoPollResult) -> None:
if not result.raw_response:
return
if task.request_metadata is None:
task.request_metadata = {}
# 仅在终态写一次,避免污染 request_metadata
task.request_metadata["poll_raw_response"] = result.raw_response
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
"""判断是否为永久性错误(不应重试)"""
# 优先使用 HTTP 状态码判断
if status_code is not None:
# 4xx 客户端错误(除 429 限流)通常是永久性错误
return 400 <= status_code < 500 and status_code != 429
# 降级到字符串匹配
error_msg = str(exc).lower()
return any(indicator in error_msg for indicator in _PERMANENT_ERROR_INDICATORS)
async def _poll_task_status(self, db: Session, task: VideoTask) -> InternalVideoPollResult:
if not task.endpoint_id or not task.key_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_provider_info",
error_message="Task missing endpoint_id or key_id",
)
endpoint = self._get_endpoint(db, task.endpoint_id)
key = self._get_key(db, task.key_id)
if not key.api_key:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="provider_config_error",
error_message="Provider key not properly configured",
)
try:
upstream_key = crypto_service.decrypt(key.api_key)
except Exception:
logger.warning("Failed to decrypt provider key for task %s", task.id)
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="decryption_error",
error_message="Failed to decrypt provider key",
)
provider_format = (task.provider_api_format or "").strip().lower()
if not provider_format:
provider_format = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
if provider_format.startswith("gemini:"):
auth_info = await get_provider_auth(endpoint, key)
return await self._poll_gemini(task, endpoint, upstream_key, auth_info)
return await self._poll_openai(task, endpoint, upstream_key)
async def _poll_openai(
self,
task: VideoTask,
endpoint: ProviderEndpoint,
upstream_key: str,
) -> InternalVideoPollResult:
if not task.external_task_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_external_task_id",
error_message="Task missing external_task_id",
)
url = self._build_openai_url(endpoint.base_url, task.external_task_id)
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
headers = self._build_headers(endpoint_sig, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
response = await client.get(url, headers=headers)
if response.status_code >= 400:
raise PollHTTPError(
response.status_code,
sanitize_error_message(response.text or "Poll error"),
)
payload = response.json()
return self._openai_normalizer.video_poll_to_internal(payload)
async def _poll_gemini(
self,
task: VideoTask,
endpoint: ProviderEndpoint,
upstream_key: str,
auth_info: ProviderAuthInfo | None,
) -> InternalVideoPollResult:
if not task.external_task_id:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_external_task_id",
error_message="Task missing external_task_id",
)
operation_name = normalize_gemini_operation_id(task.external_task_id)
url = self._build_gemini_url(endpoint.base_url, operation_name)
endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
headers = self._build_headers(endpoint_sig, upstream_key, endpoint, auth_info)
client = await HTTPClientPool.get_default_client_async()
response = await client.get(url, headers=headers)
if response.status_code >= 400:
raise PollHTTPError(
response.status_code,
sanitize_error_message(response.text or "Poll error"),
)
payload = response.json()
return self._gemini_normalizer.video_poll_to_internal(payload)
def _build_openai_url(self, base_url: str | None, task_id: str) -> str:
base = (base_url or "https://api.openai.com").rstrip("/")
if base.endswith("/v1"):
return f"{base}/videos/{task_id}"
return f"{base}/v1/videos/{task_id}"
def _build_gemini_url(self, base_url: str | None, operation_name: str) -> str:
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
return f"{base}/v1beta/{operation_name}"
def _build_headers(
self,
endpoint_sig: str,
upstream_key: str,
endpoint: ProviderEndpoint,
auth_info: ProviderAuthInfo | None = None,
) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint)
headers = build_upstream_headers_for_endpoint(
{},
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
)
if auth_info:
headers.pop("x-goog-api-key", None)
headers[auth_info.auth_header] = auth_info.auth_value
return headers
def _get_endpoint(self, db: Session, endpoint_id: str) -> ProviderEndpoint:
endpoint = db.query(ProviderEndpoint).filter(ProviderEndpoint.id == endpoint_id).first()
if not endpoint:
raise RuntimeError("Provider endpoint not found")
return endpoint
def _get_key(self, db: Session, key_id: str) -> ProviderAPIKey:
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise RuntimeError("Provider key not found")
return key
async def _acquire_redis_lock(self) -> str | None:
if not self.redis:
return "no_redis"
token = str(uuid4())
acquired = await self.redis.set(self.LOCK_KEY, token, nx=True, ex=self.LOCK_TTL)
return token if acquired else None
async def _release_redis_lock(self, token: str) -> None:
if not self.redis or token == "no_redis":
return
script = """
if redis.call('GET', KEYS[1]) == ARGV[1] then
return redis.call('DEL', KEYS[1])
end
return 0
"""
await self.redis.eval(script, 1, self.LOCK_KEY, token)
_video_task_poller: VideoTaskPollerService | None = None
def get_video_task_poller() -> VideoTaskPollerService:
global _video_task_poller
if _video_task_poller is None:
_video_task_poller = VideoTaskPollerService()
return _video_task_poller