feat: 视频计费增强与影子计费系统

This commit is contained in:
fawney19
2026-02-03 18:48:39 +08:00
parent ed68aebfb0
commit 00442c41fa
84 changed files with 6507 additions and 1678 deletions

View File

@@ -23,6 +23,7 @@ from src.core.exceptions import InvalidRequestException, NotFoundException
from src.database import get_db
from src.models.database import BillingRule, DimensionCollector
from src.services.billing.formula_engine import SafeExpressionEvaluator, UnsafeExpressionError
from src.services.billing.presets import BillingPresetService, PresetApplyMode, list_preset_packs
router = APIRouter(prefix="/api/admin/billing", tags=["Admin - Billing"])
pipeline = ApiRequestPipeline()
@@ -127,6 +128,30 @@ class DimensionCollectorResponse(BaseModel):
)
class BillingPresetInfoResponse(BaseModel):
name: str
version: str
description: str
collector_count: int
class ApplyBillingPresetRequest(BaseModel):
preset: str = Field(..., min_length=1, max_length=100)
mode: PresetApplyMode = "merge"
@router.get("/presets")
async def list_billing_presets(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = BillingPresetListAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/presets/apply")
async def apply_billing_preset(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = BillingPresetApplyAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/rules")
async def list_billing_rules(
request: Request,
@@ -525,3 +550,38 @@ def _validate_dimension_collector_request(
raise InvalidRequestException(
"default_value already exists for this (api_format, task_type, dimension_name)"
)
@dataclass
class BillingPresetListAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
items = []
for p in list_preset_packs():
items.append(
BillingPresetInfoResponse(
name=p.name,
version=p.version,
description=p.description,
collector_count=len(p.collectors or []),
).model_dump()
)
return {"items": items}
class BillingPresetApplyAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
payload = context.ensure_json_body()
try:
req = ApplyBillingPresetRequest.model_validate(payload)
except Exception as exc:
raise InvalidRequestException(f"Invalid request body: {exc}")
result = BillingPresetService.apply_preset(
context.db,
preset_name=req.preset,
mode=req.mode,
)
if result.errors:
# still return counts; caller can display partial results
return {"ok": False, **result.to_dict()}
return {"ok": True, **result.to_dict()}

View File

@@ -726,9 +726,9 @@ class AdminGetApiFormatsAdapter(AdminApiAdapter):
def _label_for(sig: str) -> str:
fam, kind = (sig.split(":", 1) + [""])[:2]
fam_title = {"claude": "Claude", "openai": "OpenAI", "gemini": "Gemini"}.get(fam, fam)
if kind == "chat":
return fam_title
kind_title = {"cli": "CLI", "video": "Video", "image": "Image"}.get(kind, kind)
kind_title = {"chat": "Chat", "cli": "CLI", "video": "Video", "image": "Image"}.get(
kind, kind
)
return f"{fam_title} {kind_title}".strip()
endpoint_defs = list_endpoint_definitions()

View File

@@ -4,7 +4,7 @@ from __future__ import annotations
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime
from datetime import datetime, timedelta, timezone
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request
@@ -14,6 +14,8 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.config.constants import CacheTTL
from src.config.settings import config
from src.database import get_db
from src.models.database import (
ApiKey,
@@ -25,11 +27,31 @@ from src.models.database import (
User,
)
from src.services.usage.service import UsageService
from src.utils.cache_decorator import cache_result
router = APIRouter(prefix="/api/admin/usage", tags=["Admin - Usage"])
pipeline = ApiRequestPipeline()
def _apply_admin_default_range(
start_date: datetime | None, end_date: datetime | None
) -> tuple[datetime | None, datetime | None]:
"""
Apply a default time range for admin usage endpoints to protect DB from unbounded scans.
Enabled by setting ADMIN_USAGE_DEFAULT_DAYS>0.
"""
if start_date is not None or end_date is not None:
return start_date, end_date
days = int(getattr(config, "admin_usage_default_days", 0) or 0)
if days <= 0:
return start_date, end_date
now = datetime.now(timezone.utc)
return now - timedelta(days=days), now
# ==================== RESTful Routes ====================
@@ -267,9 +289,14 @@ async def get_usage_detail(
class AdminUsageStatsAdapter(AdminApiAdapter):
def __init__(self, start_date: datetime | None, end_date: datetime | None):
self.start_date = start_date
self.end_date = end_date
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
@cache_result(
key_prefix="admin:usage:stats",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=["start_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
@@ -347,10 +374,15 @@ class AdminActivityHeatmapAdapter(AdminApiAdapter):
class AdminUsageByModelAdapter(AdminApiAdapter):
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
self.start_date = start_date
self.end_date = end_date
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
self.limit = limit
@cache_result(
key_prefix="admin:usage:agg:model",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=["start_date", "end_date", "limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = db.query(
@@ -394,10 +426,15 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
class AdminUsageByUserAdapter(AdminApiAdapter):
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
self.start_date = start_date
self.end_date = end_date
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
self.limit = limit
@cache_result(
key_prefix="admin:usage:agg:user",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=["start_date", "end_date", "limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = (
@@ -444,10 +481,15 @@ class AdminUsageByUserAdapter(AdminApiAdapter):
class AdminUsageByProviderAdapter(AdminApiAdapter):
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
self.start_date = start_date
self.end_date = end_date
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
self.limit = limit
@cache_result(
key_prefix="admin:usage:agg:provider",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=["start_date", "end_date", "limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
@@ -558,10 +600,15 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
class AdminUsageByApiFormatAdapter(AdminApiAdapter):
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
self.start_date = start_date
self.end_date = end_date
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
self.limit = limit
@cache_result(
key_prefix="admin:usage:agg:api_format",
ttl=CacheTTL.ADMIN_USAGE_AGGREGATION,
user_specific=False,
vary_by=["start_date", "end_date", "limit"],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = db.query(
@@ -624,8 +671,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
limit: int,
offset: int,
):
self.start_date = start_date
self.end_date = end_date
self.start_date, self.end_date = _apply_admin_default_range(start_date, end_date)
self.search = search
self.user_id = user_id
self.username = username
@@ -635,6 +681,23 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
self.limit = limit
self.offset = offset
@cache_result(
key_prefix="admin:usage:records",
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
user_specific=False,
vary_by=[
"start_date",
"end_date",
"search",
"user_id",
"username",
"model",
"provider",
"status",
"limit",
"offset",
],
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy import or_
from sqlalchemy.orm import load_only
@@ -955,7 +1018,11 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 先通过主键 id 查找,如果找不到再尝试通过 request_id 查找
usage_record = db.query(Usage).filter(Usage.id == self.usage_id).first()
if not usage_record:
# 兼容通过 request_id 查找(用于异步任务等场景)
usage_record = db.query(Usage).filter(Usage.request_id == self.usage_id).first()
if not usage_record:
raise HTTPException(status_code=404, detail="Usage record not found")
@@ -970,6 +1037,9 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
usage_id=self.usage_id,
)
# 提取视频/图像/音频计费信息
video_billing_info = self._extract_video_billing_info(usage_record)
return {
"id": usage_record.id,
"request_id": usage_record.request_id,
@@ -1022,6 +1092,7 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
"response_body": usage_record.get_response_body(),
"metadata": usage_record.request_metadata,
"tiered_pricing": tiered_pricing_info,
"video_billing": video_billing_info,
}
async def _get_tiered_pricing_info(self, db: Session, usage_record: Any) -> dict | None:
@@ -1077,6 +1148,75 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
"source": pricing_source, # 定价来源: 'provider' 或 'global'
}
def _extract_video_billing_info(self, usage_record: Any) -> dict | None:
"""
从 request_metadata.billing_snapshot 和 dimensions 中提取视频/图像/音频计费信息。
返回结构:
{
"task_type": "video" | "image" | "audio",
"duration_seconds": 10.5, # 视频时长(秒)
"resolution": "1080p", # 分辨率
"video_price_per_second": 0.1, # 每秒单价
"video_cost": 1.05, # 视频费用
"rule_name": "...", # 计费规则名称
"expression": "...", # 计费公式
"status": "complete", # 计费状态
}
"""
request_type = getattr(usage_record, "request_type", None)
if request_type not in {"video", "image", "audio"}:
return None
metadata = getattr(usage_record, "request_metadata", None)
if not metadata:
return None
billing_snapshot = metadata.get("billing_snapshot") if isinstance(metadata, dict) else None
dimensions = metadata.get("dimensions") if isinstance(metadata, dict) else None
result: dict = {
"task_type": request_type,
}
# 从 billing_snapshot 中提取计费规则信息
if billing_snapshot and isinstance(billing_snapshot, dict):
result["rule_name"] = billing_snapshot.get("rule_name")
result["expression"] = billing_snapshot.get("expression")
result["status"] = billing_snapshot.get("status")
result["cost"] = billing_snapshot.get("cost")
# 从 dimensions_used 中提取维度
dims_used = billing_snapshot.get("dimensions_used")
if dims_used and isinstance(dims_used, dict):
if "duration_seconds" in dims_used:
result["duration_seconds"] = dims_used["duration_seconds"]
if "video_resolution_key" in dims_used:
result["resolution"] = dims_used["video_resolution_key"]
if "video_price_per_second" in dims_used:
result["video_price_per_second"] = dims_used["video_price_per_second"]
if "video_cost" in dims_used:
result["video_cost"] = dims_used["video_cost"]
# 补充从 dimensions 中提取(备用)
if dimensions and isinstance(dimensions, dict):
if "duration_seconds" not in result and "duration_seconds" in dimensions:
result["duration_seconds"] = dimensions["duration_seconds"]
if "resolution" not in result and "video_resolution_key" in dimensions:
result["resolution"] = dimensions["video_resolution_key"]
# 如果没有有意义的视频计费信息,返回 None
has_video_info = (
result.get("duration_seconds")
or result.get("resolution")
or result.get("video_cost")
or result.get("cost")
)
if not has_video_info:
return None
return result
# ==================== 缓存亲和性分析 ====================

View File

@@ -7,18 +7,22 @@ from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from typing import Any, AsyncIterator
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from fastapi.responses import StreamingResponse
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.api.dashboard.routes import DashboardAdapter
from src.clients.http_client import HTTPClientPool
from src.core.crypto import crypto_service
from src.core.enums import UserRole
from src.core.logger import logger
from src.database import get_db
from src.models.database import Provider, ProviderEndpoint, User, VideoTask
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, User, VideoTask
router = APIRouter(prefix="/api/admin/video-tasks", tags=["Admin - Video Tasks"])
pipeline = ApiRequestPipeline()
@@ -119,6 +123,117 @@ async def cancel_video_task(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/{task_id}/video")
async def proxy_video_stream(
task_id: str,
request: Request,
token: str | None = Query(None, description="JWT access token"),
db: Session = Depends(get_db),
) -> StreamingResponse:
"""
代理视频流(用于需要认证的视频链接)
**路径参数**:
- `task_id`: 任务 ID
**查询参数**:
- `token`: JWT access token用于 video 标签请求)
**返回**:
- 视频流
"""
from src.services.auth.service import AuthService
# 尝试从多个来源获取 tokenquery param > cookie > header
auth_token = token
if not auth_token:
auth_token = request.cookies.get("access_token")
if not auth_token:
auth_header = request.headers.get("Authorization")
if auth_header and auth_header.startswith("Bearer "):
auth_token = auth_header[7:]
if not auth_token:
raise HTTPException(status_code=401, detail="Unauthorized")
try:
# 验证 token 并获取 payload
payload = await AuthService.verify_token(auth_token, token_type="access")
user_id = payload.get("user_id") or payload.get("sub")
if not user_id:
raise HTTPException(status_code=401, detail="Unauthorized")
# 查询用户
user = db.query(User).filter(User.id == user_id).first()
if not user:
raise HTTPException(status_code=401, detail="Unauthorized")
except HTTPException:
raise
except Exception:
raise HTTPException(status_code=401, detail="Unauthorized")
# 查询任务
query = db.query(VideoTask).filter(VideoTask.id == task_id)
if user.role != UserRole.ADMIN:
query = query.filter(VideoTask.user_id == user.id)
task = query.first()
if not task:
raise HTTPException(status_code=404, detail="Video task not found")
if not task.video_url:
raise HTTPException(status_code=404, detail="Video not available")
# 检查是否需要代理Google API 链接需要认证)
video_url = task.video_url
needs_proxy = "generativelanguage.googleapis.com" in video_url
if not needs_proxy:
# 不需要代理,重定向到原始 URL
from fastapi.responses import RedirectResponse
return RedirectResponse(url=video_url)
# 需要代理:获取 provider key 进行认证
if not task.key_id:
raise HTTPException(status_code=500, detail="Missing provider key")
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == task.key_id).first()
if not key or not key.api_key:
raise HTTPException(status_code=500, detail="Provider key not found")
try:
api_key = crypto_service.decrypt(key.api_key)
except Exception:
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
# 构建认证头
headers = {"x-goog-api-key": api_key}
async def stream_video() -> AsyncIterator[bytes]:
"""流式下载并返回视频"""
try:
client = await HTTPClientPool.get_default_client_async()
async with client.stream("GET", video_url, headers=headers) as response:
if response.status_code >= 400:
logger.warning(
"Video proxy failed: task={} status={}", task_id, response.status_code
)
return
async for chunk in response.aiter_bytes(chunk_size=65536):
yield chunk
except Exception as e:
logger.exception("Video proxy error: task={} error={}", task_id, str(e))
return StreamingResponse(
stream_video(),
media_type="video/mp4",
headers={
"Content-Disposition": f'inline; filename="video_{task_id}.mp4"',
"Cache-Control": "private, max-age=3600",
},
)
# ==================== Adapters ====================
@@ -359,6 +474,7 @@ class VideoTaskDetailAdapter(DashboardAdapter):
"video_urls": task.video_urls,
"thumbnail_url": task.thumbnail_url,
"video_size_bytes": task.video_size_bytes,
"video_duration_seconds": task.video_duration_seconds,
"video_expires_at": (
task.video_expires_at.isoformat() if task.video_expires_at else None
),

View File

@@ -211,8 +211,10 @@ class AdminDashboardStatsAdapter(AdminApiAdapter):
month_start = month_start_local.astimezone(timezone.utc)
# ==================== 使用预聚合数据 ====================
# 今日实时数据只查询一次,避免重复扫描 Usage 表
today_stats = StatsAggregatorService.get_today_realtime_stats(db)
# 从 stats_summary + 今日实时数据获取全局统计
combined_stats = StatsAggregatorService.get_combined_stats(db)
combined_stats = StatsAggregatorService.get_combined_stats(db, today_stats=today_stats)
all_time_requests = combined_stats["total_requests"]
all_time_success_requests = combined_stats["success_requests"]
@@ -237,7 +239,6 @@ class AdminDashboardStatsAdapter(AdminApiAdapter):
)
# ==================== 今日实时统计 ====================
today_stats = StatsAggregatorService.get_today_realtime_stats(db)
requests_today = today_stats["total_requests"]
cost_today = today_stats["total_cost"]
actual_cost_today = today_stats["actual_total_cost"]
@@ -951,26 +952,9 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
today_stats = StatsAggregatorService.get_today_realtime_stats(db)
today_str = today_local.date().isoformat()
if today_stats["total_requests"] > 0:
# 今日平均响应时间需要单独查询
today_avg_rt = (
db.query(func.avg(Usage.response_time_ms))
.filter(Usage.created_at >= today, Usage.response_time_ms.isnot(None))
.scalar()
or 0
)
# 今日 unique_models 和 unique_providers
today_unique_models = (
db.query(func.count(func.distinct(Usage.model)))
.filter(Usage.created_at >= today)
.scalar()
or 0
)
today_unique_providers = (
db.query(func.count(func.distinct(Usage.provider_name)))
.filter(Usage.created_at >= today)
.scalar()
or 0
)
today_avg_rt_ms = float(today_stats.get("avg_response_time_ms") or 0.0)
today_unique_models = int(today_stats.get("unique_models") or 0)
today_unique_providers = int(today_stats.get("unique_providers") or 0)
# 今日 fallback_count
today_fallback_count = (
db.query(func.count())
@@ -996,7 +980,7 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
+ today_stats["cache_read_tokens"]
),
"cost": today_stats["total_cost"],
"avg_response_time": float(today_avg_rt) / 1000.0 if today_avg_rt else 0,
"avg_response_time": today_avg_rt_ms / 1000.0 if today_avg_rt_ms else 0,
"unique_models": today_unique_models,
"unique_providers": today_unique_providers,
"fallback_count": today_fallback_count,

View File

@@ -100,10 +100,7 @@ class StreamTelemetryRecorder:
return
actual_request_body = ctx.provider_request_body or original_request_body
response_body = None
if (
not isinstance(writer, QueueTelemetryWriter)
or config.usage_queue_include_bodies
):
if not isinstance(writer, QueueTelemetryWriter) or writer.include_bodies:
response_body = ctx.build_response_body(response_time_ms)
try:
@@ -403,10 +400,25 @@ class StreamTelemetryRecorder:
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
) -> TelemetryWriter | None:
if config.usage_queue_enabled and self.user_id and self.api_key_id:
from src.services.system.config import SystemConfigService
# Queue payload detail follows system config request_record_level.
log_level = SystemConfigService.get_request_record_level(bg_db).value
sensitive_headers = SystemConfigService.get_sensitive_headers(bg_db) or []
max_request_body_size = int(
SystemConfigService.get_config(bg_db, "max_request_body_size", 5242880) or 0
)
max_response_body_size = int(
SystemConfigService.get_config(bg_db, "max_response_body_size", 5242880) or 0
)
return QueueTelemetryWriter(
request_id=self.request_id,
user_id=self.user_id,
api_key_id=self.api_key_id,
log_level=log_level,
sensitive_headers=sensitive_headers,
max_request_body_size=max_request_body_size,
max_response_body_size=max_response_body_size,
)
db_writer = self._build_db_writer(bg_db)
if db_writer is None:

View File

@@ -75,6 +75,17 @@ class GeminiVeoHandler(VideoHandlerBase):
)
self._normalizer = GeminiNormalizer()
@staticmethod
def _get_request_base_url(http_request: Request) -> str:
"""从 HTTP 请求中获取基础 URL协议 + 主机)"""
# 优先使用 X-Forwarded-Proto 和 X-Forwarded-Host代理场景
proto = http_request.headers.get("x-forwarded-proto") or http_request.url.scheme
host = http_request.headers.get("x-forwarded-host") or http_request.headers.get("host")
if host:
return f"{proto}://{host}"
# 回退到 request.url
return f"{http_request.url.scheme}://{http_request.url.netloc}"
async def handle_create_task(
self,
*,
@@ -259,7 +270,8 @@ class GeminiVeoHandler(VideoHandlerBase):
created_at=task.created_at,
original_request=internal_request,
)
response_body = self._normalizer.video_task_from_internal(internal_task)
base_url = self._get_request_base_url(http_request)
response_body = self._normalizer.video_task_from_internal(internal_task, base_url=base_url)
# 提交成功后立即结算 Usage费用暂时为 0轮询完成后更新
response_time_ms = int((time.time() - self.start_time) * 1000)
@@ -313,7 +325,8 @@ class GeminiVeoHandler(VideoHandlerBase):
# 直接从数据库返回任务状态(后台轮询服务会持续更新状态)
internal_task = self._task_to_internal(task)
response_body = self._normalizer.video_task_from_internal(internal_task)
base_url = self._get_request_base_url(http_request)
response_body = self._normalizer.video_task_from_internal(internal_task, base_url=base_url)
return JSONResponse(response_body)
async def handle_list_tasks(
@@ -331,8 +344,10 @@ class GeminiVeoHandler(VideoHandlerBase):
.limit(100)
.all()
)
base_url = self._get_request_base_url(http_request)
items = [
self._normalizer.video_task_from_internal(self._task_to_internal(t)) for t in tasks
self._normalizer.video_task_from_internal(self._task_to_internal(t), base_url=base_url)
for t in tasks
]
return JSONResponse({"operations": items})

View File

@@ -10,6 +10,7 @@ from datetime import datetime, timedelta, timezone
from typing import Any, AsyncIterator
from uuid import uuid4
import httpx
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.exc import IntegrityError
@@ -650,8 +651,10 @@ class OpenAIVideoHandler(VideoHandlerBase):
try:
# 使用 httpx 的 stream 方法并正确管理上下文
# 视频下载可能较大,设置 5 分钟超时
request = client.build_request("GET", upstream_url, headers=headers)
response = await client.send(request, stream=True, timeout=300.0)
request = client.build_request(
"GET", upstream_url, headers=headers, timeout=httpx.Timeout(300.0)
)
response = await client.send(request, stream=True)
except Exception as exc:
logger.warning(
"[VideoDownload] Upstream connection failed task={} url={}: {}",
@@ -723,8 +726,8 @@ class OpenAIVideoHandler(VideoHandlerBase):
"""代理直接的视频 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)
request = client.build_request("GET", url, timeout=httpx.Timeout(300.0))
response = await client.send(request, stream=True)
except Exception as exc:
logger.warning(
"[VideoDownload] Direct URL connection failed task={} url={}: {}",

View File

@@ -73,11 +73,11 @@ def _get_formats_for_api(api_format: str) -> list[str]:
return _OPENAI_FORMATS
def _is_format_conversion_enabled() -> bool:
"""检查全局格式转换开关(从环境变量读取,默认开启)"""
from src.config.settings import config
def _is_format_conversion_enabled(db: Session) -> bool:
"""检查全局格式转换开关(从数据库配置读取,默认开启)"""
from src.services.system.config import SystemConfigService
return config.format_conversion_enabled
return SystemConfigService.is_format_conversion_enabled(db)
def _get_convertible_formats(client_format: str, global_conversion_enabled: bool) -> list[str]:
@@ -500,7 +500,7 @@ async def list_models(
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled()
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats, empty_response = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
@@ -604,7 +604,7 @@ async def retrieve_model(
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled()
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats, _ = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
@@ -687,7 +687,7 @@ async def list_models_gemini(
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled()
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats, empty_response = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format
@@ -766,7 +766,7 @@ async def get_model_gemini(
restrictions = AccessRestrictions.from_api_key_and_user(key_record, user)
# 获取可用格式(包括可转换的格式)
global_conversion_enabled = _is_format_conversion_enabled()
global_conversion_enabled = _is_format_conversion_enabled(db)
candidate_formats = _get_convertible_formats(api_format, global_conversion_enabled)
candidate_formats, _ = _filter_formats_by_restrictions(
candidate_formats, restrictions, api_format

View File

@@ -1110,7 +1110,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
from sqlalchemy import or_
from src.api.base.models_service import AccessRestrictions
from src.config.settings import config as app_config
from src.services.system.config import SystemConfigService
db = context.db
user = context.user
@@ -1118,8 +1118,8 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
# 使用 AccessRestrictions 类来处理限制(与 /v1/models 逻辑一致)
restrictions = AccessRestrictions.from_api_key_and_user(api_key=None, user=user)
# 检查全局格式转换开关
global_conversion_enabled = app_config.format_conversion_enabled
# 检查全局格式转换开关(从数据库配置读取)
global_conversion_enabled = SystemConfigService.is_format_conversion_enabled(db)
# 获取所有可用的 Provider ID考虑格式转换
available_provider_ids = self._get_all_available_provider_ids(db, global_conversion_enabled)