feat: 添加视频生成 API 支持和认证抽象层重构

- 新增 Video Generation API 路由和处理器(支持 Gemini Veo 和 OpenAI Sora 兼容格式)
- 新增 AuthHandler 策略模式,统一 API key 提取逻辑(Bearer/ApiKey/GoogApiKey/OAuth2/QueryKey)
- 新增 RequestContext 三维度检测(数据格式/端点类型/认证方式)
- 新增 EndpointType 和 AuthMethod 枚举
- Gemini/OpenAI normalizer 添加视频格式转换(InternalVideoRequest/Task/PollResult)
- 新增视频任务轮询服务和数据库迁移(video_tasks 表)
- 代码格式化:修复 black 行宽限制,调整 import 排序,target-version 降级至 py313
This commit is contained in:
fawney19
2026-01-30 22:41:42 +08:00
parent 16fb06ff4c
commit 772cb90f64
31 changed files with 3004 additions and 203 deletions

View File

@@ -18,7 +18,6 @@ Chat Adapter 通用基类
from __future__ import annotations
from sqlalchemy.orm import Session
import time
import traceback
from abc import abstractmethod
@@ -27,11 +26,19 @@ from typing import Any
import httpx
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import APIFormat
from src.core.api_format import (
APIFormat,
build_adapter_base_headers,
build_adapter_headers,
get_adapter_protected_keys,
get_auth_handler,
get_default_auth_method,
)
from src.core.exceptions import (
InvalidRequestException,
ModelNotSupportedException,
@@ -43,19 +50,12 @@ from src.core.exceptions import (
QuotaExceededException,
UpstreamClientException,
)
from src.core.api_format import (
build_adapter_base_headers,
build_adapter_headers,
extract_client_api_key,
get_adapter_protected_keys,
)
from src.core.logger import logger
from src.services.billing import calculate_request_cost as _calculate_request_cost
from src.services.request.result import RequestResult
from src.services.usage.recorder import UsageRecorder
class ChatAdapterBase(ApiAdapter):
"""
Chat Adapter 通用基类
@@ -124,8 +124,10 @@ class ChatAdapterBase(ApiAdapter):
return build_test_request_body(cls.FORMAT_ID, request_data)
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""
return extract_client_api_key(dict(request.headers), self._get_api_format())
"""从请求中提取 API 密钥,使用 AuthHandler 新流程"""
auth_method = get_default_auth_method(self._get_api_format())
handler = get_auth_handler(auth_method)
return handler.extract_credentials(request)
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
@@ -170,8 +172,10 @@ class ChatAdapterBase(ApiAdapter):
)
# 请求开始日志
logger.info(f"[REQ] {request_id[:8]} | {self.FORMAT_ID} | {getattr(api_key, 'name', 'unknown')} | "
f"{model} | {'stream' if stream else 'sync'} | quota:{quota_display}")
logger.info(
f"[REQ] {request_id[:8]} | {self.FORMAT_ID} | {getattr(api_key, 'name', 'unknown')} | "
f"{model} | {'stream' if stream else 'sync'} | quota:{quota_display}"
)
try:
# 检查客户端连接
@@ -218,9 +222,7 @@ class ChatAdapterBase(ApiAdapter):
logger.info(f"客户端请求错误: {e.error_type}")
return self._error_response(
status_code=e.status_code,
error_type=(
"invalid_request_error" if e.status_code == 400 else "quota_exceeded"
),
error_type=("invalid_request_error" if e.status_code == 400 else "quota_exceeded"),
message=e.message,
)
@@ -306,7 +308,9 @@ class ChatAdapterBase(ApiAdapter):
return merged
@abstractmethod
def _validate_request_body(self, original_request_body: dict, path_params: dict | None = None) -> None:
def _validate_request_body(
self, original_request_body: dict, path_params: dict | None = None
) -> None:
"""
验证请求体 - 子类必须实现
@@ -383,9 +387,7 @@ class ChatAdapterBase(ApiAdapter):
# 确定错误消息
if isinstance(e, ProviderAuthException):
error_message = (
"上游服务认证失败"
if result.metadata.provider != "unknown"
else "服务暂时不可用"
"上游服务认证失败" if result.metadata.provider != "unknown" else "服务暂时不可用"
)
result.error_message = error_message
@@ -438,7 +440,8 @@ class ChatAdapterBase(ApiAdapter):
if isinstance(e, ProxyException):
logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}")
else:
logger.error(f"{self.FORMAT_ID} 请求处理意外异常",
logger.error(
f"{self.FORMAT_ID} 请求处理意外异常",
exception=e,
extra_data={
"exception_class": e.__class__.__name__,
@@ -480,9 +483,8 @@ class ChatAdapterBase(ApiAdapter):
logger.error(f"记录失败请求时出错: {record_error}")
return self._error_response(
status_code=500,
error_type="internal_server_error",
message="处理请求时发生内部错误")
status_code=500, error_type="internal_server_error", message="处理请求时发生内部错误"
)
def _error_response(self, status_code: int, error_type: str, message: str) -> JSONResponse:
"""生成错误响应 - 子类可覆盖以自定义格式"""
@@ -681,6 +683,7 @@ class ChatAdapterBase(ApiAdapter):
model_name=model_name or request_data.get("model"),
)
# =========================================================================
# Adapter 注册表 - 用于根据 API format 获取 Adapter 实例
# =========================================================================

View File

@@ -23,13 +23,20 @@ from typing import Any
import httpx
from fastapi import HTTPException, Request
from sqlalchemy.orm import Session
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.core.api_format import APIFormat
from src.core.api_format import (
APIFormat,
build_adapter_base_headers,
build_adapter_headers,
get_adapter_protected_keys,
get_auth_handler,
get_default_auth_method,
)
from src.core.exceptions import (
InvalidRequestException,
ModelNotSupportedException,
@@ -41,19 +48,12 @@ from src.core.exceptions import (
QuotaExceededException,
UpstreamClientException,
)
from src.core.api_format import (
build_adapter_base_headers,
build_adapter_headers,
extract_client_api_key,
get_adapter_protected_keys,
)
from src.core.logger import logger
from src.services.billing import calculate_request_cost as _calculate_request_cost
from src.services.request.result import RequestResult
from src.services.usage.recorder import UsageRecorder
class CliAdapterBase(ApiAdapter):
"""
CLI Adapter 通用基类
@@ -94,9 +94,11 @@ class CliAdapterBase(ApiAdapter):
"""
从请求中提取 API 密钥
使用统一的头部处理函数,根据 API 格式自动识别认证头
使用 AuthHandler 新流程,根据 API 格式选择认证方式
"""
return extract_client_api_key(dict(request.headers), self._get_api_format())
auth_method = get_default_auth_method(self._get_api_format())
handler = get_auth_handler(auth_method)
return handler.extract_credentials(request)
@classmethod
def build_base_headers(cls, api_key: str) -> dict[str, str]:
@@ -171,8 +173,10 @@ class CliAdapterBase(ApiAdapter):
)
# 请求开始日志
logger.info(f"[REQ] {request_id[:8]} | {self.FORMAT_ID} | {getattr(api_key, 'name', 'unknown')} | "
f"{model} | {'stream' if stream else 'sync'} | quota:{quota_display}")
logger.info(
f"[REQ] {request_id[:8]} | {self.FORMAT_ID} | {getattr(api_key, 'name', 'unknown')} | "
f"{model} | {'stream' if stream else 'sync'} | quota:{quota_display}"
)
try:
# 检查客户端连接
@@ -220,9 +224,7 @@ class CliAdapterBase(ApiAdapter):
logger.debug(f"客户端请求错误: {e.error_type}")
return self._error_response(
status_code=e.status_code,
error_type=(
"invalid_request_error" if e.status_code == 400 else "quota_exceeded"
),
error_type=("invalid_request_error" if e.status_code == 400 else "quota_exceeded"),
message=e.message,
)
@@ -366,9 +368,7 @@ class CliAdapterBase(ApiAdapter):
# 确定错误消息
if isinstance(e, ProviderAuthException):
error_message = (
"上游服务认证失败"
if result.metadata.provider != "unknown"
else "服务暂时不可用"
"上游服务认证失败" if result.metadata.provider != "unknown" else "服务暂时不可用"
)
result.error_message = error_message
@@ -421,7 +421,8 @@ class CliAdapterBase(ApiAdapter):
if isinstance(e, ProxyException):
logger.error(f"{self.FORMAT_ID} 请求处理业务异常: {type(e).__name__}")
else:
logger.error(f"{self.FORMAT_ID} 请求处理意外异常",
logger.error(
f"{self.FORMAT_ID} 请求处理意外异常",
exception=e,
extra_data={
"exception_class": e.__class__.__name__,
@@ -460,9 +461,8 @@ class CliAdapterBase(ApiAdapter):
await recorder.record_failure(result, original_headers, original_request_body)
return self._error_response(
status_code=500,
error_type="internal_server_error",
message="处理请求时发生内部错误")
status_code=500, error_type="internal_server_error", message="处理请求时发生内部错误"
)
def _error_response(self, status_code: int, error_type: str, message: str) -> JSONResponse:
"""生成错误响应"""
@@ -672,7 +672,9 @@ class CliAdapterBase(ApiAdapter):
# =========================================================================
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
def build_endpoint_url(
cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None
) -> str:
"""
构建CLI API端点URL - 子类应覆盖
@@ -727,6 +729,7 @@ class CliAdapterBase(ApiAdapter):
headers["User-Agent"] = cli_user_agent
return headers
# =========================================================================
# CLI Adapter 注册表 - 用于根据 API format 获取 CLI Adapter 实例
# =========================================================================

View File

@@ -0,0 +1,137 @@
"""
Video Adapter 通用基类
负责请求分发、Handler 创建、认证头提取等通用逻辑。
"""
from __future__ import annotations
from typing import Any
from fastapi import HTTPException, Request
from fastapi.responses import Response
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.video_handler_base import VideoHandlerBase
from src.core.api_format import APIFormat, get_auth_handler, get_default_auth_method
from src.core.logger import logger
class VideoAdapterBase(ApiAdapter):
"""视频生成适配器基类"""
FORMAT_ID: str = "UNKNOWN"
HANDLER_CLASS: type[VideoHandlerBase]
name: str = "video.base"
mode = ApiMode.STANDARD
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
@classmethod
def _get_api_format(cls) -> APIFormat:
try:
return APIFormat[cls.FORMAT_ID]
except KeyError:
return APIFormat.OPENAI
def extract_api_key(self, request: Request) -> str | None:
auth_method = get_default_auth_method(self._get_api_format())
handler = get_auth_handler(auth_method)
return handler.extract_credentials(request)
async def handle(self, context: ApiRequestContext) -> Response:
http_request = context.request
path_params = context.path_params or {}
if context.api_key is None or context.user is None:
raise HTTPException(status_code=401, detail="Unauthorized")
handler = self._create_handler(context)
method = http_request.method.upper()
path = http_request.url.path.lower()
task_id = path_params.get("task_id")
if method in {"POST", "PUT", "PATCH"}:
original_request_body = context.ensure_json_body()
else:
original_request_body = {}
logger.debug(
"[VideoAdapter] dispatch method=%s path=%s task_id=%s",
method,
path,
task_id,
)
# Download content
if method == "GET" and path.endswith("/content") and task_id:
return await handler.handle_download_content(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# Cancel task
if method in {"DELETE", "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"
)
return await handler.handle_cancel_task(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# Get task
if method == "GET" and task_id:
return await handler.handle_get_task(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# List tasks
if method == "GET" and not task_id:
return await handler.handle_list_tasks(
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# Create task (default)
return await handler.handle_create_task(
http_request=http_request,
original_headers=context.original_headers,
original_request_body=original_request_body,
query_params=context.query_params,
path_params=path_params,
)
def _create_handler(self, context: ApiRequestContext) -> VideoHandlerBase:
return self.HANDLER_CLASS(
db=context.db,
user=context.user,
api_key=context.api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
allowed_api_formats=self.allowed_api_formats,
)
__all__ = ["VideoAdapterBase"]

View File

@@ -0,0 +1,208 @@
"""
Video Handler 基类
定义视频生成相关操作的统一接口。
"""
from __future__ import annotations
import re
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.orm import Session
from src.core.api_format.conversion.internal_video import InternalVideoTask, VideoStatus
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
if TYPE_CHECKING:
import httpx
# 敏感信息匹配正则(预编译提升性能)
_SENSITIVE_PATTERN = re.compile(
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
re.IGNORECASE,
)
def sanitize_error_message(message: str, max_length: int = 200) -> str:
"""
移除错误消息中可能包含的敏感信息
Args:
message: 原始错误消息
max_length: 最大长度,默认 200
Returns:
脱敏后的消息
"""
if not message:
return "Request failed"
# 先脱敏再截断,确保敏感信息不会因截断位置而泄露
sanitized = _SENSITIVE_PATTERN.sub("[REDACTED]", message)
return sanitized[:max_length]
class VideoHandlerBase(ABC):
"""视频处理器基类"""
FORMAT_ID: str = ""
def __init__(
self,
db: Session,
user: User,
api_key: ApiKey,
request_id: str,
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: list[str] | None = None,
):
self.db = db
self.user = user
self.api_key = api_key
self.request_id = request_id
self.client_ip = client_ip
self.user_agent = user_agent
self.start_time = start_time
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
@abstractmethod
async def handle_create_task(
self,
*,
http_request: Request,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""创建视频任务"""
@abstractmethod
async def handle_get_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:
"""获取视频任务状态"""
@abstractmethod
async def handle_list_tasks(
self,
*,
http_request: Request,
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""列出任务"""
@abstractmethod
async def handle_cancel_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:
"""取消任务"""
@abstractmethod
async def handle_download_content(
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,
) -> Response | StreamingResponse:
"""下载视频内容"""
def _build_error_response(self, response: "httpx.Response") -> JSONResponse:
"""
构建脱敏后的错误响应
子类可重写 _format_error_payload 自定义错误格式。
"""
content_type = response.headers.get("content-type", "")
if "application/json" in content_type:
try:
error_data = response.json()
if isinstance(error_data, dict) and "error" in error_data:
payload = self._format_error_payload(error_data["error"], response.status_code)
return JSONResponse(
status_code=response.status_code,
content={"error": payload},
)
except ValueError, KeyError, TypeError:
pass
message = sanitize_error_message(response.text or "Upstream error")
fallback_payload = self._format_error_payload({"message": message}, response.status_code)
return JSONResponse(
status_code=response.status_code,
content={"error": fallback_payload},
)
def _format_error_payload(self, error: dict[str, Any], status_code: int) -> dict[str, Any]:
"""
格式化错误负载,子类可重写以匹配特定 API 格式
默认返回 OpenAI 风格格式。
"""
return {
"type": error.get("type", "upstream_error"),
"message": sanitize_error_message(error.get("message", "Request failed")),
}
def _get_task(self, task_id: str) -> VideoTask:
task = (
self.db.query(VideoTask)
.filter(VideoTask.id == task_id, VideoTask.user_id == self.user.id)
.first()
)
if not task:
raise HTTPException(status_code=404, detail="Video task not found")
return task
def _get_endpoint_and_key(self, task: VideoTask) -> tuple[ProviderEndpoint, ProviderAPIKey]:
endpoint = (
self.db.query(ProviderEndpoint).filter(ProviderEndpoint.id == task.endpoint_id).first()
)
key = self.db.query(ProviderAPIKey).filter(ProviderAPIKey.id == task.key_id).first()
if not endpoint or not key:
raise HTTPException(status_code=500, detail="Provider endpoint or key not found")
return endpoint, key
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:
try:
status = VideoStatus(task.status)
except ValueError:
status = VideoStatus.PENDING
return InternalVideoTask(
id=task.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},
)
__all__ = ["VideoHandlerBase", "sanitize_error_message"]

View File

@@ -98,7 +98,9 @@ class ClaudeChatAdapter(ChatAdapterBase):
"""
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
def _validate_request_body(self, original_request_body: dict, path_params: dict | None = None) -> None:
def _validate_request_body(
self, original_request_body: dict, path_params: dict | None = None
) -> None:
"""验证请求体"""
try:
if not isinstance(original_request_body, dict):
@@ -220,15 +222,16 @@ class ClaudeTokenCountAdapter(ApiAdapter):
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
# 优先检查 x-api-key
api_key = request.headers.get("x-api-key")
from src.core.api_format import get_auth_handler
from src.core.api_format.enums import AuthMethod
handler = get_auth_handler(AuthMethod.API_KEY)
api_key = handler.extract_credentials(request)
if api_key:
return api_key
# 降级到 Authorization: Bearer
authorization = request.headers.get("authorization")
if authorization and authorization.startswith("Bearer "):
return authorization.replace("Bearer ", "")
return None
bearer_handler = get_auth_handler(AuthMethod.BEARER)
return bearer_handler.extract_credentials(request)
async def handle(self, context: ApiRequestContext) -> Any:
payload = context.ensure_json_body()

View File

@@ -7,10 +7,14 @@ Gemini API Handler 模块
from src.api.handlers.gemini.adapter import GeminiChatAdapter, build_gemini_adapter
from src.api.handlers.gemini.handler import GeminiChatHandler
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
from src.api.handlers.gemini.video_adapter import GeminiVeoAdapter
from src.api.handlers.gemini.video_handler import GeminiVeoHandler
__all__ = [
"GeminiChatAdapter",
"GeminiChatHandler",
"GeminiStreamParser",
"build_gemini_adapter",
"GeminiVeoAdapter",
"GeminiVeoHandler",
]

View File

@@ -14,7 +14,8 @@ from fastapi.responses import JSONResponse
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import extract_client_api_key_with_query
from src.core.api_format import 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
@@ -43,7 +44,9 @@ class GeminiChatAdapter(ChatAdapterBase):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["GEMINI"])
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
logger.info(
f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}"
)
def extract_api_key(self, request: Request) -> str | None:
"""
@@ -53,11 +56,8 @@ class GeminiChatAdapter(ChatAdapterBase):
1. URL 参数 ?key=
2. x-goog-api-key 请求头
"""
return extract_client_api_key_with_query(
dict(request.headers),
dict(request.query_params),
self._get_api_format(),
)
handler = get_auth_handler(AuthMethod.GOOG_API_KEY)
return handler.extract_credentials(request)
def detect_capability_requirements(
self,
@@ -91,7 +91,9 @@ class GeminiChatAdapter(ChatAdapterBase):
"""
return original_request_body.copy()
def _validate_request_body(self, original_request_body: dict, path_params: dict | None = None) -> None:
def _validate_request_body(
self, original_request_body: dict, path_params: dict | None = None
) -> None:
"""验证请求体"""
path_params = path_params or {}
is_stream = path_params.get("stream", False)
@@ -200,7 +202,9 @@ class GeminiChatAdapter(ChatAdapterBase):
try:
response = await client.get(models_url, headers=headers)
logger.debug(f"Gemini models request to {redact_url_for_log(models_url)}: status={response.status_code}")
logger.debug(
f"Gemini models request to {redact_url_for_log(models_url)}: status={response.status_code}"
)
if response.status_code == 200:
data = response.json()
if "models" in data:
@@ -218,13 +222,17 @@ class GeminiChatAdapter(ChatAdapterBase):
else:
error_body = response.text[:500] if response.text else "(empty)"
error_msg = f"HTTP {response.status_code}: {error_body}"
logger.warning(f"Gemini models request to {redact_url_for_log(models_url)} failed: {error_msg}")
logger.warning(
f"Gemini models request to {redact_url_for_log(models_url)} failed: {error_msg}"
)
return [], error_msg
except Exception as e:
# 异常信息可能包含带 key 参数的 URL需要脱敏
sanitized_error = redact_url_for_log(str(e))
error_msg = f"Request error: {sanitized_error}"
logger.warning(f"Failed to fetch Gemini models from {redact_url_for_log(models_url)}: {sanitized_error}")
logger.warning(
f"Failed to fetch Gemini models from {redact_url_for_log(models_url)}: {sanitized_error}"
)
return [], error_msg
@classmethod
@@ -273,6 +281,7 @@ class GeminiChatAdapter(ChatAdapterBase):
# 使用基类的通用endpoint checker
from src.api.handlers.base.endpoint_checker import run_endpoint_check
return await run_endpoint_check(
client=client,
url=url,

View File

@@ -0,0 +1,22 @@
"""
Gemini Video Adapter - 基于 VideoAdapterBase 的 Veo 适配器
"""
from __future__ import annotations
from src.api.handlers.base.video_adapter_base import VideoAdapterBase
from src.api.handlers.base.video_handler_base import VideoHandlerBase
class GeminiVeoAdapter(VideoAdapterBase):
FORMAT_ID = "GEMINI"
name = "gemini.video"
@property
def HANDLER_CLASS(self) -> type[VideoHandlerBase]:
from src.api.handlers.gemini.video_handler import GeminiVeoHandler
return GeminiVeoHandler
__all__ = ["GeminiVeoAdapter"]

View File

@@ -0,0 +1,402 @@
"""
Gemini Video Handler - Veo 视频生成实现
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any, AsyncIterator
from uuid import uuid4
from fastapi import HTTPException, Request
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.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint
from src.core.api_format.conversion.internal_video import (
InternalVideoRequest,
InternalVideoTask,
VideoStatus,
)
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
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.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
class GeminiVeoHandler(VideoHandlerBase):
FORMAT_ID = "GEMINI"
DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com"
POLL_INTERVAL_SECONDS = 10
MAX_POLL_COUNT = 360
def __init__(
self,
db: Session,
user: User,
api_key: ApiKey,
request_id: str,
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: list[str] | None = None,
):
super().__init__(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip=client_ip,
user_agent=user_agent,
start_time=start_time,
allowed_api_formats=allowed_api_formats,
)
self._normalizer = GeminiNormalizer()
async def handle_create_task(
self,
*,
http_request: Request,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
# 将路径中的 model 合并到请求体再解析
model = path_params.get("model") if path_params else None
request_with_model = {**original_request_body}
if model:
request_with_model["model"] = str(model)
try:
internal_request = self._normalizer.video_request_to_internal(request_with_model)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
candidate = await self._select_candidate(internal_request.model)
if not candidate:
raise HTTPException(
status_code=503, detail="No available provider for video generation"
)
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)
client = await HTTPClientPool.get_default_client_async()
response = await client.post(upstream_url, headers=headers, json=original_request_body)
if response.status_code >= 400:
return self._build_error_response(response)
payload = response.json()
external_task_id = str(payload.get("name") or "")
if not external_task_id:
raise HTTPException(status_code=502, detail="Upstream returned empty task id")
task = self._create_task_record(
external_task_id=external_task_id,
candidate=candidate,
original_request_body=original_request_body,
internal_request=internal_request,
)
try:
self.db.add(task)
self.db.flush() # 先 flush 检测冲突
self.db.commit()
self.db.refresh(task)
except IntegrityError:
self.db.rollback()
raise HTTPException(status_code=409, detail="Task already exists")
internal_task = InternalVideoTask(
id=task.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)
return JSONResponse(response_body)
async def handle_get_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:
# Gemini 使用 operations/{id} 格式,需要按 external_task_id 查找
task = self._get_task_by_external_id(task_id)
internal_task = self._task_to_internal(task)
response_body = self._normalizer.video_task_from_internal(internal_task)
return JSONResponse(response_body)
async def handle_list_tasks(
self,
*,
http_request: Request,
original_headers: dict[str, str],
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()
)
items = [
self._normalizer.video_task_from_internal(self._task_to_internal(t)) for t in tasks
]
return JSONResponse({"operations": items})
async def handle_cancel_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_by_external_id(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)
operation_name = task.external_task_id
if not operation_name.startswith("operations/"):
operation_name = f"operations/{operation_name}"
upstream_url = self._build_cancel_url(endpoint.base_url, operation_name)
auth_info = await get_provider_auth(endpoint, key)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint, auth_info)
client = await HTTPClientPool.get_default_client_async()
response = await client.post(upstream_url, headers=headers, json={})
if response.status_code >= 400:
return self._build_error_response(response)
task.status = VideoStatus.CANCELLED.value
task.updated_at = datetime.now(timezone.utc)
self.db.commit()
return JSONResponse({})
async def handle_download_content(
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,
) -> Response | StreamingResponse:
task = self._get_task_by_external_id(task_id)
# 根据任务状态返回不同的错误码
if not task.video_url:
if task.status in (
VideoStatus.PENDING.value,
VideoStatus.SUBMITTED.value,
VideoStatus.QUEUED.value,
VideoStatus.PROCESSING.value,
):
# 任务仍在处理中,返回 202 Accepted
raise HTTPException(
status_code=202,
detail=f"Video is still processing (status: {task.status})",
)
if task.status == VideoStatus.FAILED.value:
raise HTTPException(
status_code=422,
detail=f"Video generation failed: {task.error_message or 'Unknown error'}",
)
# 其他状态(如 CANCELLED
raise HTTPException(status_code=404, detail="Video not available")
# 检查视频是否已过期
if task.video_expires_at:
now = datetime.now(timezone.utc)
if task.video_expires_at < now:
raise HTTPException(status_code=410, detail="Video URL has expired")
# 代理下载而非直接重定向,避免暴露上游存储 URL
client = await HTTPClientPool.get_default_client_async()
try:
request = client.build_request("GET", task.video_url)
# 视频下载可能较大,设置 5 分钟超时
response = await client.send(request, stream=True, timeout=300.0)
except Exception as exc:
logger.error(
"[VideoDownload] Upstream fetch failed user=%s task=%s: %s",
self.user.id,
task.id,
sanitize_error_message(str(exc)),
)
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()
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"),
)
# ------------------------------------------------------------------
# Helpers
# ------------------------------------------------------------------
async def _select_candidate(self, model_name: str) -> ProviderCandidate | None:
scheduler = CacheAwareScheduler()
candidates, _ = await scheduler.list_all_candidates(
db=self.db,
api_format=APIFormat.GEMINI,
model_name=model_name,
affinity_key=str(self.api_key.id),
user_api_key=self.api_key,
max_candidates=10,
)
for candidate in candidates:
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
if auth_type in {"api_key", "vertex_ai"}:
return candidate
return None
async def _resolve_upstream_key(
self, candidate: ProviderCandidate
) -> tuple[str, ProviderEndpoint, ProviderAPIKey, Any | None]:
try:
upstream_key = crypto_service.decrypt(candidate.key.api_key)
except Exception as exc:
logger.error(
"Failed to decrypt provider key id=%s: %s",
candidate.key.id,
sanitize_error_message(str(exc)),
)
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
auth_info = await get_provider_auth(candidate.endpoint, candidate.key)
return upstream_key, candidate.endpoint, candidate.key, auth_info
def _build_upstream_url(self, base_url: str | None, model: str) -> str:
base = (base_url or self.DEFAULT_BASE_URL).rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
return f"{base}/v1beta/models/{model}:predictLongRunning"
def _build_cancel_url(self, base_url: str | None, operation_name: str) -> str:
base = (base_url or self.DEFAULT_BASE_URL).rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
return f"{base}/v1beta/{operation_name}:cancel"
def _build_upstream_headers(
self,
original_headers: dict[str, str],
upstream_key: str,
endpoint: ProviderEndpoint,
auth_info: Any | None,
) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint)
headers = build_upstream_headers(
original_headers,
APIFormat.GEMINI,
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
def _format_error_payload(self, error: dict[str, Any], status_code: int) -> dict[str, Any]:
"""Gemini 风格错误格式"""
return {
"code": error.get("code", status_code),
"message": sanitize_error_message(error.get("message", "Request failed")),
"status": error.get("status", "BAD_GATEWAY"),
}
def _create_task_record(
self,
*,
external_task_id: str,
candidate: ProviderCandidate,
original_request_body: dict[str, Any],
internal_request: Any,
) -> VideoTask:
now = datetime.now(timezone.utc)
return VideoTask(
id=str(uuid4()),
external_task_id=external_task_id,
user_id=self.user.id,
api_key_id=self.api_key.id,
provider_id=candidate.provider.id,
endpoint_id=candidate.endpoint.id,
key_id=candidate.key.id,
client_api_format="GEMINI",
provider_api_format=str(candidate.endpoint.api_format),
format_converted=False,
model=internal_request.model,
prompt=internal_request.prompt,
original_request_body=original_request_body,
converted_request_body=original_request_body,
duration_seconds=internal_request.duration_seconds,
resolution=internal_request.resolution,
aspect_ratio=internal_request.aspect_ratio,
status=VideoStatus.SUBMITTED.value,
progress_percent=0,
poll_interval_seconds=self.POLL_INTERVAL_SECONDS,
next_poll_at=now + timedelta(seconds=self.POLL_INTERVAL_SECONDS),
poll_count=0,
max_poll_count=self.MAX_POLL_COUNT,
submitted_at=now,
)
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
"""按 external_task_id 查找任务Gemini 使用 operations/{id} 格式)"""
normalized_id = external_id
if not normalized_id.startswith("operations/"):
normalized_id = f"operations/{normalized_id}"
task = (
self.db.query(VideoTask)
.filter(
VideoTask.external_task_id == normalized_id,
VideoTask.user_id == self.user.id,
)
.first()
)
if not task:
raise HTTPException(status_code=404, detail="Video task not found")
return task
__all__ = ["GeminiVeoHandler"]

View File

@@ -15,7 +15,8 @@ from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.api.handlers.gemini.adapter import GeminiChatAdapter
from src.config.settings import config
from src.core.api_format import extract_client_api_key_with_query
from src.core.api_format import get_auth_handler
from src.core.api_format.enums import AuthMethod
@register_cli_adapter
@@ -48,11 +49,8 @@ class GeminiCliAdapter(CliAdapterBase):
1. URL 参数 ?key=
2. x-goog-api-key 请求头
"""
return extract_client_api_key_with_query(
dict(request.headers),
dict(request.query_params),
self._get_api_format(),
)
handler = get_auth_handler(AuthMethod.GOOG_API_KEY)
return handler.extract_credentials(request)
def _merge_path_params(
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
@@ -129,16 +127,16 @@ class GeminiCliAdapter(CliAdapterBase):
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
if extra_headers:
cli_headers.update(extra_headers)
models, error = await GeminiChatAdapter.fetch_models(
client, base_url, api_key, cli_headers
)
models, error = await GeminiChatAdapter.fetch_models(client, base_url, api_key, cli_headers)
# 更新 api_format 为 CLI 格式
for m in models:
m["api_format"] = cls.FORMAT_ID
return models, error
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
def build_endpoint_url(
cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None
) -> str:
"""构建Gemini CLI API端点URL"""
effective_model_name = model_name or request_data.get("model", "")
if not effective_model_name:

View File

@@ -4,8 +4,12 @@ OpenAI Chat API 处理器
from src.api.handlers.openai.adapter import OpenAIChatAdapter
from src.api.handlers.openai.handler import OpenAIChatHandler
from src.api.handlers.openai.video_adapter import OpenAIVideoAdapter
from src.api.handlers.openai.video_handler import OpenAIVideoHandler
__all__ = [
"OpenAIChatAdapter",
"OpenAIChatHandler",
"OpenAIVideoAdapter",
"OpenAIVideoHandler",
]

View File

@@ -0,0 +1,22 @@
"""
OpenAI Video Adapter - 基于 VideoAdapterBase 的 Sora 适配器
"""
from __future__ import annotations
from src.api.handlers.base.video_adapter_base import VideoAdapterBase
from src.api.handlers.base.video_handler_base import VideoHandlerBase
class OpenAIVideoAdapter(VideoAdapterBase):
FORMAT_ID = "OPENAI"
name = "openai.video"
@property
def HANDLER_CLASS(self) -> type[VideoHandlerBase]:
from src.api.handlers.openai.video_handler import OpenAIVideoHandler
return OpenAIVideoHandler
__all__ = ["OpenAIVideoAdapter"]

View File

@@ -0,0 +1,401 @@
"""
OpenAI Video Handler - Sora 视频生成实现
"""
from __future__ import annotations
import json
from datetime import datetime, timedelta, timezone
from typing import Any, AsyncIterator
from uuid import uuid4
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
from src.clients.http_client import HTTPClientPool
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint
from src.core.api_format.conversion.internal_video import (
InternalVideoRequest,
InternalVideoTask,
VideoStatus,
)
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
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.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
class OpenAIVideoHandler(VideoHandlerBase):
FORMAT_ID = "OPENAI"
DEFAULT_BASE_URL = "https://api.openai.com"
POLL_INTERVAL_SECONDS = 10
MAX_POLL_COUNT = 360
def __init__(
self,
db: Session,
user: User,
api_key: ApiKey,
request_id: str,
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: list[str] | None = None,
):
super().__init__(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip=client_ip,
user_agent=user_agent,
start_time=start_time,
allowed_api_formats=allowed_api_formats,
)
self._normalizer = OpenAINormalizer()
async def handle_create_task(
self,
*,
http_request: Request,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
try:
internal_request = self._normalizer.video_request_to_internal(original_request_body)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
candidate = await self._select_candidate(internal_request.model)
if not candidate:
raise HTTPException(
status_code=503, detail="No available provider for video generation"
)
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()
response = await client.post(upstream_url, headers=headers, json=original_request_body)
if response.status_code >= 400:
return self._build_error_response(response)
payload = response.json()
external_task_id = str(payload.get("id") or "")
if not external_task_id:
raise HTTPException(status_code=502, detail="Upstream returned empty task id")
task = self._create_task_record(
external_task_id=external_task_id,
candidate=candidate,
original_request_body=original_request_body,
internal_request=internal_request,
)
try:
self.db.add(task)
self.db.flush() # 先 flush 检测冲突
self.db.commit()
self.db.refresh(task)
except IntegrityError:
self.db.rollback()
raise HTTPException(status_code=409, detail="Task already exists")
internal_task = InternalVideoTask(
id=task.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)
return JSONResponse(response_body)
async def handle_get_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)
internal_task = self._task_to_internal(task)
response_body = self._normalizer.video_task_from_internal(internal_task)
return JSONResponse(response_body)
async def handle_list_tasks(
self,
*,
http_request: Request,
original_headers: dict[str, str],
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()
)
items = [
self._normalizer.video_task_from_internal(self._task_to_internal(t)) for t in tasks
]
return JSONResponse({"object": "list", "data": items})
async def handle_cancel_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 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)
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:
return self._build_error_response(response)
task.status = VideoStatus.CANCELLED.value
task.updated_at = datetime.now(timezone.utc)
self.db.commit()
return JSONResponse({})
async def handle_download_content(
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,
) -> Response | StreamingResponse:
task = self._get_task(task_id)
# 根据本地任务状态提前返回适当错误(避免不必要的上游请求)
if task.status in (
VideoStatus.PENDING.value,
VideoStatus.SUBMITTED.value,
VideoStatus.QUEUED.value,
VideoStatus.PROCESSING.value,
):
raise HTTPException(
status_code=202,
detail=f"Video is still processing (status: {task.status})",
)
if task.status == VideoStatus.FAILED.value:
raise HTTPException(
status_code=422,
detail=f"Video generation failed: {task.error_message or 'Unknown error'}",
)
if task.status == VideoStatus.CANCELLED.value:
raise HTTPException(status_code=404, detail="Video task was cancelled")
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)
upstream_url = self._build_upstream_url(
endpoint.base_url, f"{task.external_task_id}/content"
)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
try:
# 使用 httpx 的 stream 方法并正确管理上下文
# 视频下载可能较大,设置 5 分钟超时
request = client.build_request("GET", upstream_url, headers=headers)
response = await client.send(request, stream=True, timeout=300.0)
except Exception as exc:
logger.warning(
"[VideoDownload] Upstream connection failed task=%s: %s",
task_id,
sanitize_error_message(str(exc)),
)
raise HTTPException(status_code=502, detail="Upstream connection failed") from exc
if response.status_code >= 400:
error_body = await response.aread()
await response.aclose()
content_type = response.headers.get("content-type", "")
if "application/json" in content_type:
try:
data = json.loads(error_body)
# 脱敏:移除可能的敏感信息
if isinstance(data, dict) and isinstance(data.get("error"), dict):
if "message" in data["error"]:
data["error"]["message"] = sanitize_error_message(
str(data["error"]["message"])
)
return JSONResponse(status_code=response.status_code, content=data)
except json.JSONDecodeError:
pass
message = sanitize_error_message(error_body.decode(errors="ignore"))
return JSONResponse(
status_code=response.status_code,
content={"error": {"type": "upstream_error", "message": message}},
)
async def _iter_bytes() -> AsyncIterator[bytes]:
try:
async for chunk in response.aiter_bytes():
yield chunk
finally:
await response.aclose()
# 过滤 hop-by-hop 和系统管理头部
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", "application/octet-stream"),
)
# ------------------------------------------------------------------
# Helpers
# ------------------------------------------------------------------
async def _select_candidate(self, model_name: str) -> ProviderCandidate | None:
scheduler = CacheAwareScheduler()
candidates, _ = await scheduler.list_all_candidates(
db=self.db,
api_format=APIFormat.OPENAI,
model_name=model_name,
affinity_key=str(self.api_key.id),
user_api_key=self.api_key,
max_candidates=10,
)
for candidate in candidates:
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
if auth_type == "api_key":
return candidate
return None
async def _resolve_upstream_key(
self, candidate: ProviderCandidate
) -> tuple[str, ProviderEndpoint, ProviderAPIKey]:
try:
upstream_key = crypto_service.decrypt(candidate.key.api_key)
except Exception as exc:
logger.error(
"Failed to decrypt provider key id=%s: %s",
candidate.key.id,
sanitize_error_message(str(exc)),
)
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
return upstream_key, candidate.endpoint, candidate.key
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"):
url = f"{base}/videos"
else:
url = f"{base}/v1/videos"
if suffix:
return f"{url}/{suffix}"
return url
def _build_upstream_headers(
self, original_headers: dict[str, str], upstream_key: str, endpoint: ProviderEndpoint
) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint)
return build_upstream_headers(
original_headers,
APIFormat.OPENAI,
upstream_key,
endpoint_headers=extra_headers,
)
# _build_error_response 继承自基类 VideoHandlerBase
def _create_task_record(
self,
*,
external_task_id: str,
candidate: ProviderCandidate,
original_request_body: dict[str, Any],
internal_request: InternalVideoRequest,
) -> VideoTask:
now = datetime.now(timezone.utc)
size = internal_request.extra.get("original_size")
return VideoTask(
id=str(uuid4()),
external_task_id=external_task_id,
user_id=self.user.id,
api_key_id=self.api_key.id,
provider_id=candidate.provider.id,
endpoint_id=candidate.endpoint.id,
key_id=candidate.key.id,
client_api_format="OPENAI",
provider_api_format=str(candidate.endpoint.api_format),
format_converted=False,
model=internal_request.model,
prompt=internal_request.prompt,
original_request_body=original_request_body,
converted_request_body=original_request_body,
duration_seconds=internal_request.duration_seconds,
resolution=internal_request.resolution,
aspect_ratio=internal_request.aspect_ratio,
size=size,
status=VideoStatus.SUBMITTED.value,
progress_percent=0,
poll_interval_seconds=self.POLL_INTERVAL_SECONDS,
next_poll_at=now + timedelta(seconds=self.POLL_INTERVAL_SECONDS),
poll_count=0,
max_poll_count=self.MAX_POLL_COUNT,
submitted_at=now,
)
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:
try:
status = VideoStatus(task.status)
except ValueError:
status = VideoStatus.PENDING
return InternalVideoTask(
id=task.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 [],
thumbnail_url=task.thumbnail_url,
video_duration_seconds=task.duration_seconds,
video_size_bytes=task.video_size_bytes,
created_at=task.created_at,
completed_at=task.completed_at,
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},
)
__all__ = ["OpenAIVideoHandler"]