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

@@ -0,0 +1,112 @@
"""Add video_tasks table
Revision ID: b6f1a2c5d8e9
Revises: 7f6f8065f517
Create Date: 2026-01-30 18:00:00.000000
"""
from typing import Sequence, Union
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "b6f1a2c5d8e9"
down_revision: Union[str, None] = "7f6f8065f517"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def upgrade() -> None:
if table_exists("video_tasks"):
return
op.create_table(
"video_tasks",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("external_task_id", sa.String(200), nullable=True, index=False),
sa.Column("user_id", sa.String(36), sa.ForeignKey("users.id"), nullable=False),
sa.Column("api_key_id", sa.String(36), sa.ForeignKey("api_keys.id"), nullable=True),
sa.Column("provider_id", sa.String(36), sa.ForeignKey("providers.id"), nullable=True),
sa.Column(
"endpoint_id", sa.String(36), sa.ForeignKey("provider_endpoints.id"), nullable=True
),
sa.Column("key_id", sa.String(36), sa.ForeignKey("provider_api_keys.id"), nullable=True),
sa.Column("client_api_format", sa.String(50), nullable=False),
sa.Column("provider_api_format", sa.String(50), nullable=False),
sa.Column("format_converted", sa.Boolean(), server_default=sa.false()),
sa.Column("model", sa.String(100), nullable=False),
sa.Column("prompt", sa.Text(), nullable=False),
sa.Column("original_request_body", sa.JSON(), nullable=True),
sa.Column("converted_request_body", sa.JSON(), nullable=True),
sa.Column("duration_seconds", sa.Integer(), server_default=sa.text("4")),
sa.Column("resolution", sa.String(20), server_default=sa.text("'720p'")),
sa.Column("aspect_ratio", sa.String(10), server_default=sa.text("'16:9'")),
sa.Column("size", sa.String(20), nullable=True),
sa.Column("status", sa.String(20), server_default=sa.text("'pending'")),
sa.Column("progress_percent", sa.Integer(), server_default=sa.text("0")),
sa.Column("progress_message", sa.String(500), nullable=True),
sa.Column("video_url", sa.String(2000), nullable=True),
sa.Column("video_urls", sa.JSON(), nullable=True),
sa.Column("thumbnail_url", sa.String(2000), nullable=True),
sa.Column("video_size_bytes", sa.BigInteger(), nullable=True),
sa.Column("video_expires_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("stored_video_path", sa.String(500), nullable=True),
sa.Column("storage_provider", sa.String(50), nullable=True),
sa.Column("error_code", sa.String(50), nullable=True),
sa.Column("error_message", sa.Text(), nullable=True),
sa.Column("retry_count", sa.Integer(), server_default=sa.text("0")),
sa.Column("max_retries", sa.Integer(), server_default=sa.text("3")),
sa.Column("poll_interval_seconds", sa.Integer(), server_default=sa.text("10")),
sa.Column("next_poll_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("poll_count", sa.Integer(), server_default=sa.text("0")),
sa.Column("max_poll_count", sa.Integer(), server_default=sa.text("360")),
sa.Column(
"remixed_from_task_id",
sa.String(36),
sa.ForeignKey("video_tasks.id", ondelete="SET NULL"),
nullable=True,
),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
server_default=sa.text("CURRENT_TIMESTAMP"),
),
sa.Column("submitted_at", sa.DateTime(timezone=True), nullable=True),
sa.Column("completed_at", sa.DateTime(timezone=True), nullable=True),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
server_default=sa.text("CURRENT_TIMESTAMP"),
),
)
op.create_index("idx_video_tasks_user_status", "video_tasks", ["user_id", "status"])
op.create_index("idx_video_tasks_next_poll", "video_tasks", ["next_poll_at"])
op.create_index("idx_video_tasks_external_id", "video_tasks", ["external_task_id"])
# 唯一约束:同一用户不能有重复的 external_task_id
op.create_unique_constraint(
"uq_video_tasks_user_external_id",
"video_tasks",
["user_id", "external_task_id"],
)
def downgrade() -> None:
if not table_exists("video_tasks"):
return
op.drop_constraint("uq_video_tasks_user_external_id", "video_tasks", type_="unique")
op.drop_index("idx_video_tasks_external_id", table_name="video_tasks")
op.drop_index("idx_video_tasks_next_poll", table_name="video_tasks")
op.drop_index("idx_video_tasks_user_status", table_name="video_tasks")
op.drop_table("video_tasks")

View File

@@ -81,7 +81,7 @@ dev-dependencies = [
[tool.black]
line-length = 100
target-version = ['py314']
target-version = ['py313']
[tool.isort]
profile = "black"

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"]

View File

@@ -11,6 +11,7 @@ from .models import router as models_router
from .modules import router as modules_router
from .openai import router as openai_router
from .system_catalog import router as system_catalog_router
from .videos import router as videos_router
router = APIRouter()
# Models API 需要在最前面注册,避免被其他路由的 path 参数捕获
@@ -19,6 +20,7 @@ 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

@@ -27,7 +27,7 @@ from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
from src.clients.http_client import HTTPClientPool
from src.core.api_format import APIFormat, extract_client_api_key_with_query
from src.core.api_format import APIFormat, get_auth_handler, get_default_auth_method
from src.core.api_format.metadata import get_api_format_definition
from src.core.crypto import crypto_service
from src.core.logger import logger
@@ -50,14 +50,16 @@ GEMINI_FILES_BASE_URL = "https://generativelanguage.googleapis.com"
REQUIRED_CAPABILITIES = {"gemini_files_api": True}
# 需要从客户端请求中移除的头部(这些会由代理重新设置或不应转发)
HEADERS_TO_REMOVE = frozenset({
"host",
"content-length",
"transfer-encoding",
"connection",
"x-goog-api-key",
"authorization",
})
HEADERS_TO_REMOVE = frozenset(
{
"host",
"content-length",
"transfer-encoding",
"connection",
"x-goog-api-key",
"authorization",
}
)
def _extract_gemini_api_key(request: Request) -> str | None:
@@ -68,11 +70,9 @@ def _extract_gemini_api_key(request: Request) -> str | None:
1. URL 参数 ?key=
2. x-goog-api-key 请求头
"""
return extract_client_api_key_with_query(
dict(request.headers),
dict(request.query_params),
APIFormat.GEMINI,
)
auth_method = get_default_auth_method(APIFormat.GEMINI)
handler = get_auth_handler(auth_method)
return handler.extract_credentials(request)
def _build_upstream_headers(
@@ -127,7 +127,7 @@ def _build_upstream_url(
# 处理 base_url 可能包含 /v1beta 的情况,避免重复
normalized_base_url = base_url.rstrip("/")
if normalized_base_url.endswith("/v1beta"):
normalized_base_url = normalized_base_url[:-len("/v1beta")]
normalized_base_url = normalized_base_url[: -len("/v1beta")]
# 上传端点使用不同的路径前缀
if is_upload:
@@ -319,13 +319,9 @@ async def _proxy_request(
response = await client.delete(upstream_url, headers=headers)
elif method.upper() == "POST":
if content is not None:
response = await client.post(
upstream_url, headers=headers, content=content
)
response = await client.post(upstream_url, headers=headers, content=content)
elif json_body is not None:
response = await client.post(
upstream_url, headers=headers, json=json_body
)
response = await client.post(upstream_url, headers=headers, json=json_body)
else:
response = await client.post(upstream_url, headers=headers)
else:
@@ -619,5 +615,3 @@ async def delete_file(
f"Gemini Files delete failed, skip mapping cleanup: status={response.status_code}"
)
return response

View File

@@ -15,15 +15,17 @@ from src.api.base.models_service import (
AccessRestrictions,
ModelInfo,
find_model_by_id,
get_compatible_provider_formats,
get_available_provider_ids,
get_compatible_provider_formats,
list_available_models,
)
from src.core.api_format import (
API_FORMAT_DEFINITIONS,
APIFormat,
ApiFormatDefinition,
detect_format_and_key_from_starlette,
detect_request_context,
get_auth_handler,
get_default_auth_method,
)
from src.core.api_format.conversion import (
format_conversion_registry,
@@ -43,34 +45,20 @@ _GEMINI_FORMATS = [APIFormat.GEMINI.value, APIFormat.GEMINI_CLI.value]
# 所有格式(用于格式转换时的查询)
_ALL_CHAT_FORMATS = [
APIFormat.CLAUDE.value, APIFormat.CLAUDE_CLI.value,
APIFormat.OPENAI.value, APIFormat.OPENAI_CLI.value,
APIFormat.GEMINI.value, APIFormat.GEMINI_CLI.value,
APIFormat.CLAUDE.value,
APIFormat.CLAUDE_CLI.value,
APIFormat.OPENAI.value,
APIFormat.OPENAI_CLI.value,
APIFormat.GEMINI.value,
APIFormat.GEMINI_CLI.value,
]
def _extract_api_key_from_request(
request: Request, definition: ApiFormatDefinition
) -> str | None:
def _extract_api_key_from_request(request: Request, definition: ApiFormatDefinition) -> str | None:
"""根据格式定义从请求中提取 API Key"""
auth_header = definition.auth_header.lower()
auth_type = definition.auth_type
header_value = request.headers.get(auth_header)
if not header_value:
# Gemini 还支持 ?key= 参数
if definition.api_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
return request.query_params.get("key")
return None
if auth_type == "bearer":
# Bearer token: "Bearer xxx"
if header_value.lower().startswith("bearer "):
return header_value[7:].strip()
return None
else:
# header 类型: 直接使用值
return header_value
auth_method = get_default_auth_method(definition.api_format)
handler = get_auth_handler(auth_method)
return handler.extract_credentials(request)
def _detect_api_format_and_key(request: Request) -> tuple[str, str | None]:
@@ -85,8 +73,8 @@ def _detect_api_format_and_key(request: Request) -> tuple[str, str | None]:
Returns:
(api_format, api_key) 元组
"""
format_name, api_key, _auth_method = detect_format_and_key_from_starlette(request)
return format_name, api_key
context = detect_request_context(request)
return context.data_format.value.lower(), context.credentials
def _get_formats_for_api(api_format: str) -> list[str]:
@@ -102,6 +90,7 @@ def _get_formats_for_api(api_format: str) -> list[str]:
def _is_format_conversion_enabled() -> bool:
"""检查全局格式转换开关(从环境变量读取,默认开启)"""
from src.config.settings import config
return config.format_conversion_enabled
@@ -373,8 +362,12 @@ def _build_gemini_model_response(model_info: ModelInfo) -> dict:
"version": "001",
"displayName": model_info.display_name,
"description": model_info.description or f"Model {model_info.id}",
"inputTokenLimit": model_info.context_limit if model_info.context_limit is not None else 128000,
"outputTokenLimit": model_info.output_limit if model_info.output_limit is not None else 8192,
"inputTokenLimit": (
model_info.context_limit if model_info.context_limit is not None else 128000
),
"outputTokenLimit": (
model_info.output_limit if model_info.output_limit is not None else 8192
),
"supportedGenerationMethods": ["generateContent", "countTokens"],
"temperature": 1.0,
"maxTemperature": 2.0,

163
src/api/public/videos.py Normal file
View File

@@ -0,0 +1,163 @@
"""
Video Generation API 路由
"""
from __future__ import annotations
from typing import Any
from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session
from src.api.base.pipeline import ApiRequestPipeline
from src.api.handlers.gemini.video_adapter import GeminiVeoAdapter
from src.api.handlers.openai.video_adapter import OpenAIVideoAdapter
from src.database import get_db
router = APIRouter(tags=["Video Generation"])
pipeline = ApiRequestPipeline()
# -------------------- OpenAI Sora compatible --------------------
@router.post("/v1/videos")
async def create_video_sora(http_request: Request, db: Session = Depends(get_db)) -> Any:
adapter = OpenAIVideoAdapter()
return await pipeline.run(
adapter=adapter,
http_request=http_request,
db=db,
mode=adapter.mode,
api_format_hint=adapter.allowed_api_formats[0],
)
@router.get("/v1/videos/{task_id}")
async def get_video_task_sora(
task_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
adapter = OpenAIVideoAdapter()
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": task_id},
)
@router.get("/v1/videos")
async def list_video_tasks_sora(http_request: Request, db: Session = Depends(get_db)) -> Any:
adapter = OpenAIVideoAdapter()
return await pipeline.run(
adapter=adapter,
http_request=http_request,
db=db,
mode=adapter.mode,
api_format_hint=adapter.allowed_api_formats[0],
)
@router.delete("/v1/videos/{task_id}")
async def cancel_video_task_sora(
task_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
adapter = OpenAIVideoAdapter()
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": task_id, "action": "cancel"},
)
@router.get("/v1/videos/{task_id}/content")
async def download_video_content_sora(
task_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
adapter = OpenAIVideoAdapter()
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": task_id},
)
# -------------------- Gemini Veo compatible --------------------
@router.post("/v1beta/models/{model}:predictLongRunning")
async def create_video_veo(model: 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={"model": model},
)
@router.get("/v1beta/operations/{operation_id}")
async def get_video_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},
)
@router.get("/v1beta/operations")
async def list_video_tasks_veo(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],
)
@router.post("/v1beta/operations/{operation_id}:cancel")
async def cancel_video_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, "action": "cancel"},
)
@router.get("/v1beta/operations/{operation_id}/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},
)

View File

@@ -10,13 +10,26 @@ API 格式核心模块
- utils.py: 工具函数is_cli_format, get_base_format 等)
- detection.py: 格式检测(从请求头、响应内容检测格式)
"""
from src.core.api_format.auth import (
ApiKeyAuthHandler,
AuthHandler,
BearerAuthHandler,
GoogApiKeyAuthHandler,
OAuth2AuthHandler,
QueryKeyAuthHandler,
get_auth_handler,
get_default_auth_method,
)
from src.core.api_format.detection import (
RequestContext,
detect_cli_format_from_path,
detect_format_and_key_from_starlette,
detect_format_from_request,
detect_format_from_response,
detect_request_context,
)
from src.core.api_format.enums import APIFormat
from src.core.api_format.enums import APIFormat, AuthMethod, EndpointType
from src.core.api_format.headers import (
CORE_REDACT_HEADERS,
HOP_BY_HOP_HEADERS,
@@ -65,6 +78,8 @@ from src.core.api_format.utils import (
__all__ = [
# Enums
"APIFormat",
"AuthMethod",
"EndpointType",
# Metadata
"ApiFormatDefinition",
"API_FORMAT_DEFINITIONS",
@@ -111,4 +126,15 @@ __all__ = [
"detect_format_and_key_from_starlette",
"detect_format_from_response",
"detect_cli_format_from_path",
"detect_request_context",
"RequestContext",
# Auth
"AuthHandler",
"BearerAuthHandler",
"ApiKeyAuthHandler",
"GoogApiKeyAuthHandler",
"OAuth2AuthHandler",
"QueryKeyAuthHandler",
"get_auth_handler",
"get_default_auth_method",
]

129
src/core/api_format/auth.py Normal file
View File

@@ -0,0 +1,129 @@
"""
认证处理器
将认证逻辑从 API 格式中解耦,支持多种认证方式。
"""
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING
from src.core.api_format.enums import APIFormat, AuthMethod
if TYPE_CHECKING:
from starlette.requests import Request
class AuthHandler(ABC):
"""认证处理器基类"""
@abstractmethod
def extract_credentials(self, request: Request) -> str | None:
"""从请求中提取凭证"""
@abstractmethod
def build_headers(self, credentials: str) -> dict[str, str]:
"""构造上游请求的认证 Header"""
class BearerAuthHandler(AuthHandler):
"""Authorization: Bearer <token>"""
def extract_credentials(self, request: Request) -> str | None:
auth = request.headers.get("authorization", "")
if auth.lower().startswith("bearer "):
return auth[7:].strip()
return None
def build_headers(self, credentials: str) -> dict[str, str]:
return {"Authorization": f"Bearer {credentials}"}
class ApiKeyAuthHandler(AuthHandler):
"""x-api-key: <key>"""
def extract_credentials(self, request: Request) -> str | None:
return request.headers.get("x-api-key")
def build_headers(self, credentials: str) -> dict[str, str]:
return {"x-api-key": credentials}
class GoogApiKeyAuthHandler(AuthHandler):
"""x-goog-api-key: <key> (支持 ?key=)"""
def extract_credentials(self, request: Request) -> str | None:
return request.query_params.get("key") or request.headers.get("x-goog-api-key")
def build_headers(self, credentials: str) -> dict[str, str]:
return {"x-goog-api-key": credentials}
class QueryKeyAuthHandler(AuthHandler):
"""?key= 参数认证(仅提取,通常用于 Gemini"""
def extract_credentials(self, request: Request) -> str | None:
return request.query_params.get("key")
def build_headers(self, credentials: str) -> dict[str, str]:
return {"x-goog-api-key": credentials}
class OAuth2AuthHandler(AuthHandler):
"""
Google OAuth2 / Service Account 认证
目前使用 Authorization: Bearer 透传 access token。
"""
def extract_credentials(self, request: Request) -> str | None:
auth = request.headers.get("authorization", "")
if auth.lower().startswith("bearer "):
return auth[7:].strip()
return None
def build_headers(self, credentials: str) -> dict[str, str]:
return {"Authorization": f"Bearer {credentials}"}
_AUTH_HANDLERS: dict[AuthMethod, AuthHandler] = {
AuthMethod.BEARER: BearerAuthHandler(),
AuthMethod.API_KEY: ApiKeyAuthHandler(),
AuthMethod.GOOG_API_KEY: GoogApiKeyAuthHandler(),
AuthMethod.OAUTH2: OAuth2AuthHandler(),
AuthMethod.QUERY_KEY: QueryKeyAuthHandler(),
}
def get_auth_handler(auth_method: AuthMethod) -> AuthHandler:
"""获取认证处理器实例"""
handler = _AUTH_HANDLERS.get(auth_method)
if not handler:
raise ValueError(f"Unsupported auth method: {auth_method}")
return handler
def get_default_auth_method(api_format: APIFormat) -> AuthMethod:
"""从 APIFormat 推断默认 AuthMethod兼容旧逻辑"""
mapping = {
APIFormat.OPENAI: AuthMethod.BEARER,
APIFormat.OPENAI_CLI: AuthMethod.BEARER,
APIFormat.CLAUDE: AuthMethod.API_KEY,
APIFormat.CLAUDE_CLI: AuthMethod.BEARER,
APIFormat.GEMINI: AuthMethod.GOOG_API_KEY,
APIFormat.GEMINI_CLI: AuthMethod.GOOG_API_KEY,
}
return mapping.get(api_format, AuthMethod.BEARER)
__all__ = [
"AuthHandler",
"BearerAuthHandler",
"ApiKeyAuthHandler",
"GoogApiKeyAuthHandler",
"OAuth2AuthHandler",
"QueryKeyAuthHandler",
"get_auth_handler",
"get_default_auth_method",
]

View File

@@ -0,0 +1,84 @@
"""
视频格式转换内部表示Internal Video Format
用于 Video API 的 Hub-and-Spoke 统一中间表示。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import datetime
from enum import Enum
from typing import Any
class VideoStatus(str, Enum):
PENDING = "pending"
SUBMITTED = "submitted"
QUEUED = "queued"
PROCESSING = "processing"
COMPLETED = "completed"
FAILED = "failed"
CANCELLED = "cancelled"
EXPIRED = "expired"
@dataclass
class InternalVideoRequest:
"""统一的视频生成请求格式"""
prompt: str
model: str = "sora-2"
duration_seconds: int = 4
aspect_ratio: str = "16:9"
resolution: str = "720p"
reference_image_url: str | None = None # base64 或 URL
character_ids: list[str] = field(default_factory=list)
extra: dict[str, Any] = field(default_factory=dict)
preferred_provider: str | None = None
preferred_format: str | None = None
@dataclass
class InternalVideoTask:
"""统一的视频任务状态"""
id: str
external_id: str | None = None
status: VideoStatus = VideoStatus.PENDING
progress_percent: int = 0
progress_message: str | None = None
video_url: str | None = None
video_urls: list[str] = field(default_factory=list)
thumbnail_url: str | None = None
video_duration_seconds: int | None = None
video_size_bytes: int | None = None
created_at: datetime | None = None
completed_at: datetime | None = None
expires_at: datetime | None = None
error_code: str | None = None
error_message: str | None = None
original_request: InternalVideoRequest | None = None
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
class InternalVideoPollResult:
"""轮询结果"""
status: VideoStatus
progress_percent: int = 0
video_url: str | None = None
video_urls: list[str] = field(default_factory=list)
expires_at: datetime | None = None
error_code: str | None = None
error_message: str | None = None
raw_response: dict[str, Any] | None = None
__all__ = [
"VideoStatus",
"InternalVideoRequest",
"InternalVideoTask",
"InternalVideoPollResult",
]

View File

@@ -5,11 +5,11 @@
再从 internal 输出到目标格式。
"""
from abc import ABC, abstractmethod
from typing import Any
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
from .internal_video import InternalVideoPollResult, InternalVideoRequest, InternalVideoTask
from .stream_events import InternalStreamEvent
from .stream_state import StreamState
@@ -88,8 +88,29 @@ class FormatNormalizer(ABC):
"""将内部错误表示转换为格式特定错误"""
raise NotImplementedError
# ============ 视频转换(可选) ============
def video_request_to_internal(self, request: dict[str, Any]) -> InternalVideoRequest:
"""将视频请求转换为内部表示"""
raise NotImplementedError(f"{self.__class__.__name__} does not support video conversion")
def video_request_from_internal(self, internal: InternalVideoRequest) -> dict[str, Any]:
"""将内部视频请求转换为格式特定请求"""
raise NotImplementedError(f"{self.__class__.__name__} does not support video conversion")
def video_task_to_internal(self, response: dict[str, Any]) -> InternalVideoTask:
"""将视频任务响应转换为内部表示"""
raise NotImplementedError(f"{self.__class__.__name__} does not support video conversion")
def video_task_from_internal(self, internal: InternalVideoTask) -> dict[str, Any]:
"""将内部视频任务转换为格式特定响应"""
raise NotImplementedError(f"{self.__class__.__name__} does not support video conversion")
def video_poll_to_internal(self, response: dict[str, Any]) -> InternalVideoPollResult:
"""将视频轮询响应转换为内部表示"""
raise NotImplementedError(f"{self.__class__.__name__} does not support video conversion")
__all__ = [
"FormatNormalizer",
]

View File

@@ -11,7 +11,6 @@ Gemini (GenerateContent / streamGenerateContent) Normalizer
- 响应/流式通常为 camelCasecandidates/finishReason/usageMetadata/modelVersion
"""
import json
from typing import Any
@@ -43,6 +42,12 @@ from src.core.api_format.conversion.internal import (
UnknownBlock,
UsageInfo,
)
from src.core.api_format.conversion.internal_video import (
InternalVideoPollResult,
InternalVideoRequest,
InternalVideoTask,
VideoStatus,
)
from src.core.api_format.conversion.normalizer import FormatNormalizer
from src.core.api_format.conversion.stream_events import (
ContentBlockStartEvent,
@@ -111,7 +116,9 @@ class GeminiNormalizer(FormatNormalizer):
if isinstance(contents, list):
for content in contents:
if not isinstance(content, dict):
dropped["gemini_content_non_dict"] = dropped.get("gemini_content_non_dict", 0) + 1
dropped["gemini_content_non_dict"] = (
dropped.get("gemini_content_non_dict", 0) + 1
)
continue
imsg, md = self._content_to_internal_message(content)
self._merge_dropped(dropped, md)
@@ -128,9 +135,7 @@ class GeminiNormalizer(FormatNormalizer):
else None
)
temperature = self._optional_float(
generation_config.get("temperature")
if isinstance(generation_config, dict)
else None
generation_config.get("temperature") if isinstance(generation_config, dict) else None
)
top_p = self._optional_float(
generation_config.get("top_p") if isinstance(generation_config, dict) else None
@@ -144,9 +149,7 @@ class GeminiNormalizer(FormatNormalizer):
tools = self._gemini_tools_to_internal(request.get("tools"))
tool_choice = self._gemini_tool_config_to_tool_choice(
request.get("tool_config")
if "tool_config" in request
else request.get("toolConfig")
request.get("tool_config") if "tool_config" in request else request.get("toolConfig")
)
# 构建 extra保留原始 gemini 字段
@@ -253,9 +256,15 @@ class GeminiNormalizer(FormatNormalizer):
orig_gc = gemini_extra.get("generation_config") or gemini_extra.get("generationConfig")
if isinstance(orig_gc, dict):
# responseModalities
if "responseModalities" in orig_gc and "responseModalities" not in generation_config:
if (
"responseModalities" in orig_gc
and "responseModalities" not in generation_config
):
generation_config["responseModalities"] = orig_gc["responseModalities"]
if "response_modalities" in orig_gc and "responseModalities" not in generation_config:
if (
"response_modalities" in orig_gc
and "responseModalities" not in generation_config
):
generation_config["responseModalities"] = orig_gc["response_modalities"]
# thinkingConfig
if "thinkingConfig" in orig_gc and "thinkingConfig" not in generation_config:
@@ -370,14 +379,19 @@ class GeminiNormalizer(FormatNormalizer):
finish_reason = None
if internal.stop_reason is not None:
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(internal.stop_reason.value, "OTHER")
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(
internal.stop_reason.value, "OTHER"
)
usage_metadata: dict[str, Any] = {}
if internal.usage:
usage_metadata = {
"promptTokenCount": int(internal.usage.input_tokens),
"candidatesTokenCount": int(internal.usage.output_tokens),
"totalTokenCount": int(internal.usage.total_tokens or (internal.usage.input_tokens + internal.usage.output_tokens)),
"totalTokenCount": int(
internal.usage.total_tokens
or (internal.usage.input_tokens + internal.usage.output_tokens)
),
}
if internal.usage.cache_read_tokens:
usage_metadata["cachedContentTokenCount"] = int(internal.usage.cache_read_tokens)
@@ -410,7 +424,9 @@ class GeminiNormalizer(FormatNormalizer):
# Streaming
# =========================
def stream_chunk_to_internal(self, chunk: dict[str, Any], state: StreamState) -> list[InternalStreamEvent]:
def stream_chunk_to_internal(
self, chunk: dict[str, Any], state: StreamState
) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID)
events: list[InternalStreamEvent] = []
@@ -457,7 +473,9 @@ class GeminiNormalizer(FormatNormalizer):
if delta:
if not ss.get("text_block_started"):
ss["text_block_started"] = True
events.append(ContentBlockStartEvent(block_index=0, block_type=ContentType.TEXT))
events.append(
ContentBlockStartEvent(block_index=0, block_type=ContentType.TEXT)
)
events.append(ContentDeltaEvent(block_index=0, text_delta=delta))
continue
@@ -500,7 +518,9 @@ class GeminiNormalizer(FormatNormalizer):
inline_data = part.get("inline_data")
if isinstance(inline_data, dict):
mime_type = str(inline_data.get("mimeType") or inline_data.get("mime_type") or "").strip()
mime_type = str(
inline_data.get("mimeType") or inline_data.get("mime_type") or ""
).strip()
data = str(inline_data.get("data") or "").strip()
# 确保 mime_type 和 data 都非空
@@ -635,7 +655,9 @@ class GeminiNormalizer(FormatNormalizer):
if isinstance(event, MessageStopEvent):
finish_reason = None
if event.stop_reason is not None:
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(event.stop_reason.value, "OTHER")
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(
event.stop_reason.value, "OTHER"
)
chunk: dict[str, Any] = base_chunk([])
if finish_reason is not None:
@@ -645,10 +667,15 @@ class GeminiNormalizer(FormatNormalizer):
chunk["usageMetadata"] = {
"promptTokenCount": int(event.usage.input_tokens),
"candidatesTokenCount": int(event.usage.output_tokens),
"totalTokenCount": int(event.usage.total_tokens or (event.usage.input_tokens + event.usage.output_tokens)),
"totalTokenCount": int(
event.usage.total_tokens
or (event.usage.input_tokens + event.usage.output_tokens)
),
}
if event.usage.cache_read_tokens:
chunk["usageMetadata"]["cachedContentTokenCount"] = int(event.usage.cache_read_tokens)
chunk["usageMetadata"]["cachedContentTokenCount"] = int(
event.usage.cache_read_tokens
)
out.append(chunk)
return out
@@ -698,11 +725,171 @@ class GeminiNormalizer(FormatNormalizer):
}
return {"error": payload}
# =========================
# Video conversion
# =========================
def video_request_to_internal(self, request: dict[str, Any]) -> InternalVideoRequest:
instances = request.get("instances")
if not instances or not isinstance(instances, list) or len(instances) == 0:
raise ValueError("Video request requires at least one instance")
instance = instances[0] if isinstance(instances[0], dict) else {}
params = request.get("parameters") or {}
image = instance.get("image") if isinstance(instance, dict) else None
image_ref = None
if isinstance(image, dict):
image_ref = image.get("bytesBase64Encoded")
prompt = instance.get("prompt") if isinstance(instance, dict) else None
prompt_str = str(prompt).strip() if prompt else ""
if not prompt_str:
raise ValueError("Video prompt is required")
duration_raw = params.get("durationSeconds")
sample_count_raw = params.get("sampleCount")
try:
duration_seconds = int(duration_raw) if duration_raw else 8
except (ValueError, TypeError):
duration_seconds = 8
# 安全解析 sampleCount
try:
sample_count = int(sample_count_raw) if sample_count_raw else 1
except (ValueError, TypeError):
sample_count = 1
return InternalVideoRequest(
prompt=prompt_str,
model=str(request.get("model") or "veo-3.1-generate-preview"),
duration_seconds=duration_seconds,
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,
},
)
def video_request_from_internal(self, internal: InternalVideoRequest) -> dict[str, Any]:
instance: dict[str, Any] = {"prompt": internal.prompt}
if internal.reference_image_url:
instance["image"] = {"bytesBase64Encoded": internal.reference_image_url}
parameters: dict[str, Any] = {
"aspectRatio": internal.aspect_ratio,
"resolution": internal.resolution,
"durationSeconds": internal.duration_seconds,
}
for key in ["personGeneration", "sampleCount"]:
if key in internal.extra:
parameters[key] = internal.extra[key]
return {
"instances": [instance],
"parameters": parameters,
}
def video_task_to_internal(self, response: dict[str, Any]) -> InternalVideoTask:
operation_name = str(response.get("name") or "")
done = bool(response.get("done"))
if done:
video_response = response.get("response", {}).get("generateVideoResponse", {})
samples = video_response.get("generatedSamples", [])
video_urls = [
s.get("video", {}).get("uri")
for s in samples
if isinstance(s, dict) and s.get("video", {}).get("uri")
]
return InternalVideoTask(
id=operation_name.replace("operations/", ""),
external_id=operation_name,
status=VideoStatus.COMPLETED,
progress_percent=100,
video_url=video_urls[0] if video_urls else None,
video_urls=video_urls,
extra={"raw_response": video_response},
)
metadata = response.get("metadata", {})
return InternalVideoTask(
id=operation_name.replace("operations/", ""),
external_id=operation_name,
status=VideoStatus.PROCESSING,
progress_percent=50,
extra={"metadata": metadata},
)
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}"
if internal.status == VideoStatus.COMPLETED:
urls = internal.video_urls or ([internal.video_url] if internal.video_url else [])
return {
"name": operation_name,
"done": True,
"response": {
"generateVideoResponse": {
"generatedSamples": [
{"video": {"uri": url, "mimeType": "video/mp4"}} for url in urls
]
}
},
}
return {
"name": operation_name,
"done": False,
"metadata": internal.extra.get("metadata", {}),
}
def video_poll_to_internal(self, response: dict[str, Any]) -> InternalVideoPollResult:
done = bool(response.get("done"))
if done:
error = response.get("error")
if error:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code=str(error.get("code", "unknown")),
error_message=error.get("message", "Unknown error"),
raw_response=response,
)
video_response = response.get("response", {}).get("generateVideoResponse", {})
samples = video_response.get("generatedSamples", [])
video_urls = [
s.get("video", {}).get("uri")
for s in samples
if isinstance(s, dict) and s.get("video", {}).get("uri")
]
video_url = video_urls[0] if video_urls else None
return InternalVideoPollResult(
status=VideoStatus.COMPLETED,
progress_percent=100,
video_url=video_url,
video_urls=video_urls,
raw_response=response,
)
return InternalVideoPollResult(
status=VideoStatus.PROCESSING,
progress_percent=50,
raw_response=response,
)
# =========================
# Helpers
# =========================
def _content_to_internal_message(self, content: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
def _content_to_internal_message(
self, content: dict[str, Any]
) -> tuple[InternalMessage | None, dict[str, int]]:
dropped: dict[str, int] = {}
role_raw = str(content.get("role") or "user")
@@ -749,12 +936,16 @@ class GeminiNormalizer(FormatNormalizer):
if inline is None:
inline = part.get("inlineData")
if isinstance(inline, dict):
mime_type = inline.get("mime_type") if "mime_type" in inline else inline.get("mimeType")
mime_type = (
inline.get("mime_type") if "mime_type" in inline else inline.get("mimeType")
)
data = inline.get("data")
if isinstance(mime_type, str) and mime_type and isinstance(data, str) and data:
blocks.append(ImageBlock(data=data, media_type=mime_type))
else:
dropped["gemini_inline_data_invalid"] = dropped.get("gemini_inline_data_invalid", 0) + 1
dropped["gemini_inline_data_invalid"] = (
dropped.get("gemini_inline_data_invalid", 0) + 1
)
blocks.append(UnknownBlock(raw_type="inline_data", payload=part))
continue
@@ -859,7 +1050,9 @@ class GeminiNormalizer(FormatNormalizer):
return {"role": role, "parts": parts}
def _collapse_system_instruction(self, system_instruction: Any) -> tuple[str | None, dict[str, int]]:
def _collapse_system_instruction(
self, system_instruction: Any
) -> tuple[str | None, dict[str, int]]:
dropped: dict[str, int] = {}
if system_instruction is None:
return None, dropped
@@ -875,12 +1068,18 @@ class GeminiNormalizer(FormatNormalizer):
joined = "".join(texts)
return (joined or None), dropped
dropped["gemini_system_instruction_unsupported"] = dropped.get("gemini_system_instruction_unsupported", 0) + 1
dropped["gemini_system_instruction_unsupported"] = (
dropped.get("gemini_system_instruction_unsupported", 0) + 1
)
return None, dropped
def _get_generation_config(self, request: dict[str, Any]) -> dict[str, Any]:
# 兼容 snake_case 与 camelCase
gc = request.get("generation_config") if "generation_config" in request else request.get("generationConfig")
gc = (
request.get("generation_config")
if "generation_config" in request
else request.get("generationConfig")
)
if not isinstance(gc, dict):
return {}
@@ -936,8 +1135,16 @@ class GeminiNormalizer(FormatNormalizer):
ToolDefinition(
name=name,
description=decl.get("description"),
parameters=decl.get("parameters") if isinstance(decl.get("parameters"), dict) else None,
extra={"gemini_function_declaration": self._extract_extra(decl, {"name", "description", "parameters"})},
parameters=(
decl.get("parameters")
if isinstance(decl.get("parameters"), dict)
else None
),
extra={
"gemini_function_declaration": self._extract_extra(
decl, {"name", "description", "parameters"}
)
},
)
)
@@ -966,7 +1173,11 @@ class GeminiNormalizer(FormatNormalizer):
if mode in ("ANY", "REQUIRED"):
return ToolChoice(type=ToolChoiceType.REQUIRED, extra={"gemini": tool_config})
if isinstance(allowed, list) and len(allowed) == 1:
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=str(allowed[0] or ""), extra={"gemini": tool_config})
return ToolChoice(
type=ToolChoiceType.TOOL,
tool_name=str(allowed[0] or ""),
extra={"gemini": tool_config},
)
return ToolChoice(type=ToolChoiceType.AUTO, extra={"gemini": tool_config})
@@ -1010,7 +1221,9 @@ class GeminiNormalizer(FormatNormalizer):
pass
if "total_tokens" not in fields:
fields["total_tokens"] = int(fields.get("input_tokens", 0) + fields.get("output_tokens", 0))
fields["total_tokens"] = int(
fields.get("input_tokens", 0) + fields.get("output_tokens", 0)
)
return UsageInfo(
input_tokens=int(fields.get("input_tokens", 0)),

View File

@@ -7,12 +7,11 @@ OpenAI Chat Completions Normalizer
- 可选OpenAI error <-> InternalError
"""
import json
import time
from datetime import datetime, timezone
from typing import Any
from src.core.logger import logger
from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS,
RETRYABLE_ERROR_TYPES,
@@ -39,6 +38,12 @@ from src.core.api_format.conversion.internal import (
UnknownBlock,
UsageInfo,
)
from src.core.api_format.conversion.internal_video import (
InternalVideoPollResult,
InternalVideoRequest,
InternalVideoTask,
VideoStatus,
)
from src.core.api_format.conversion.normalizer import FormatNormalizer
from src.core.api_format.conversion.stream_events import (
ContentBlockStartEvent,
@@ -51,6 +56,7 @@ from src.core.api_format.conversion.stream_events import (
ToolCallDeltaEvent,
)
from src.core.api_format.conversion.stream_state import StreamState
from src.core.logger import logger
class OpenAINormalizer(FormatNormalizer):
@@ -81,6 +87,28 @@ class OpenAINormalizer(FormatNormalizer):
StopReason.UNKNOWN: "stop",
}
# 视频尺寸映射: (resolution, aspect_ratio) -> size
_VIDEO_SIZE_MAP: dict[tuple[str, str], str] = {
("480p", "16:9"): "854x480",
("480p", "9:16"): "480x854",
("480p", "1:1"): "480x480",
("720p", "16:9"): "1280x720",
("720p", "9:16"): "720x1280",
("720p", "1:1"): "720x720",
("1080p", "16:9"): "1920x1080",
("1080p", "9:16"): "1080x1920",
("1080p", "1:1"): "1080x1080",
}
# OpenAI Sora 特定尺寸(非标准分辨率,需单独处理)
_SORA_SIZE_REVERSE: dict[str, tuple[str, str]] = {
"1792x1024": ("1080p", "16:9"),
"1024x1792": ("1080p", "9:16"),
}
_VIDEO_SIZE_REVERSE: dict[str, tuple[str, str]] = {
**{value: key for key, value in _VIDEO_SIZE_MAP.items()},
**_SORA_SIZE_REVERSE,
}
# InternalError.type -> OpenAI error.type最佳努力
_ERROR_TYPE_TO_OPENAI: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "invalid_request_error",
@@ -139,9 +167,7 @@ class OpenAINormalizer(FormatNormalizer):
# 兼容新旧参数名:优先使用 max_completion_tokens回退到 max_tokens
mct = request.get("max_completion_tokens")
max_tokens_value = self._optional_int(
mct if mct is not None else request.get("max_tokens")
)
max_tokens_value = self._optional_int(mct if mct is not None else request.get("max_tokens"))
# 构建 extra保留未识别字段
extra: dict[str, Any] = {"openai": self._extract_extra(request, {"messages"})}
@@ -304,7 +330,9 @@ class OpenAINormalizer(FormatNormalizer):
message["content"] = content_value
if tool_blocks:
message["tool_calls"] = [self._tool_use_block_to_openai_call(b, idx) for idx, b in enumerate(tool_blocks)]
message["tool_calls"] = [
self._tool_use_block_to_openai_call(b, idx) for idx, b in enumerate(tool_blocks)
]
finish_reason = None
if internal.stop_reason is not None:
@@ -334,7 +362,9 @@ class OpenAINormalizer(FormatNormalizer):
# Streaming
# =========================
def stream_chunk_to_internal(self, chunk: dict[str, Any], state: StreamState) -> list[InternalStreamEvent]:
def stream_chunk_to_internal(
self, chunk: dict[str, Any], state: StreamState
) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID)
events: list[InternalStreamEvent] = []
@@ -393,7 +423,9 @@ class OpenAINormalizer(FormatNormalizer):
tc_name = str(fn.get("name") or "")
tc_args = fn.get("arguments")
block_index = self._ensure_tool_block_index(ss, tc_id or str(tool_call.get("index") or ""))
block_index = self._ensure_tool_block_index(
ss, tc_id or str(tool_call.get("index") or "")
)
# tool start只在首次见到该 tool_id 时发)
started_key = f"tool_started:{block_index}"
@@ -493,7 +525,12 @@ class OpenAINormalizer(FormatNormalizer):
image_data = event.extra.get("image_data")
image_media_type = event.extra.get("image_media_type")
# 确保图片数据有效base64 数据至少有一定长度)
if image_data and image_media_type and isinstance(image_data, str) and len(image_data) > 10:
if (
image_data
and image_media_type
and isinstance(image_data, str)
and len(image_data) > 10
):
# 构造 data URL 格式的图片
data_url = f"data:{image_media_type};base64,{image_data}"
# 存储图片数据,在 ContentBlockStopEvent 时输出
@@ -602,11 +639,171 @@ class OpenAINormalizer(FormatNormalizer):
payload["param"] = internal.param
return {"error": payload}
# =========================
# Video conversion
# =========================
def video_request_to_internal(self, request: dict[str, Any]) -> InternalVideoRequest:
prompt = str(request.get("prompt") or "").strip()
if not prompt:
raise ValueError("Video prompt is required")
size = str(request.get("size") or "720x1280")
resolution, aspect_ratio = self._VIDEO_SIZE_REVERSE.get(size, ("720p", "9:16"))
input_reference = request.get("input_reference")
reference_url = str(input_reference) if input_reference else None
# 安全解析 seconds 字段
seconds_raw = request.get("seconds")
try:
duration_seconds = int(seconds_raw) if seconds_raw else 4
except (ValueError, TypeError):
duration_seconds = 4
return InternalVideoRequest(
prompt=prompt,
model=str(request.get("model") or "sora-2"),
duration_seconds=duration_seconds,
resolution=resolution,
aspect_ratio=aspect_ratio,
character_ids=request.get("character_ids") or [],
reference_image_url=reference_url,
extra={"original_size": size},
)
def video_request_from_internal(self, internal: InternalVideoRequest) -> dict[str, Any]:
size = self._VIDEO_SIZE_MAP.get((internal.resolution, internal.aspect_ratio), "720x1280")
payload: dict[str, Any] = {
"prompt": internal.prompt,
"model": internal.model,
"seconds": internal.duration_seconds,
"size": size,
"character_ids": internal.character_ids,
}
if internal.reference_image_url:
payload["input_reference"] = internal.reference_image_url
return payload
def video_task_to_internal(self, response: dict[str, Any]) -> InternalVideoTask:
status_map = {
"queued": VideoStatus.QUEUED,
"processing": VideoStatus.PROCESSING,
"completed": VideoStatus.COMPLETED,
"failed": VideoStatus.FAILED,
}
status = status_map.get(str(response.get("status") or ""), VideoStatus.PENDING)
error = response.get("error") or {}
error_code = error.get("code") if isinstance(error, dict) else None
error_message = error.get("message") if isinstance(error, dict) else None
created_at = response.get("created_at")
completed_at = response.get("completed_at")
expires_at = response.get("expires_at")
return InternalVideoTask(
id=str(response.get("id") or ""),
status=status,
progress_percent=int(response.get("progress") or 0),
created_at=datetime.fromtimestamp(created_at, tz=timezone.utc) if created_at else None,
completed_at=(
datetime.fromtimestamp(completed_at, tz=timezone.utc) if completed_at else None
),
expires_at=datetime.fromtimestamp(expires_at, tz=timezone.utc) if expires_at else None,
error_code=error_code,
error_message=error_message,
extra={
"object": response.get("object"),
"model": response.get("model"),
"size": response.get("size"),
"seconds": response.get("seconds"),
},
)
def video_task_from_internal(self, internal: InternalVideoTask) -> dict[str, Any]:
status_map = {
VideoStatus.PENDING: "queued",
VideoStatus.SUBMITTED: "queued",
VideoStatus.QUEUED: "queued",
VideoStatus.PROCESSING: "processing",
VideoStatus.COMPLETED: "completed",
VideoStatus.FAILED: "failed",
VideoStatus.CANCELLED: "failed",
VideoStatus.EXPIRED: "failed",
}
payload: dict[str, Any] = {
"id": internal.id,
"object": "video",
"status": status_map.get(internal.status, "queued"),
"progress": internal.progress_percent,
}
if internal.created_at:
payload["created_at"] = int(internal.created_at.timestamp())
if internal.completed_at:
payload["completed_at"] = int(internal.completed_at.timestamp())
if internal.expires_at:
payload["expires_at"] = int(internal.expires_at.timestamp())
if internal.error_code:
payload["error"] = {
"code": internal.error_code,
"message": internal.error_message,
}
for key in ["model", "size", "seconds"]:
if key in internal.extra:
payload[key] = internal.extra[key]
return payload
def video_poll_to_internal(self, response: dict[str, Any]) -> InternalVideoPollResult:
status = str(response.get("status") or "")
task_id = response.get("id")
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
if not video_url:
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code="missing_task_id",
error_message="Upstream response missing task id",
raw_response=response,
)
return InternalVideoPollResult(
status=VideoStatus.COMPLETED,
progress_percent=100,
video_url=video_url,
expires_at=(
datetime.fromtimestamp(expires_at, tz=timezone.utc) if expires_at else None
),
raw_response=response,
)
if status == "failed":
error = response.get("error") or {}
return InternalVideoPollResult(
status=VideoStatus.FAILED,
error_code=error.get("code") if isinstance(error, dict) else None,
error_message=error.get("message") if isinstance(error, dict) else None,
raw_response=response,
)
return InternalVideoPollResult(
status=VideoStatus.PROCESSING,
progress_percent=int(response.get("progress") or 0),
raw_response=response,
)
# =========================
# Helpers
# =========================
def _openai_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
def _openai_message_to_internal(
self, msg: dict[str, Any]
) -> tuple[InternalMessage | None, dict[str, int]]:
dropped: dict[str, int] = {}
role_raw = str(msg.get("role") or "unknown")
@@ -620,7 +817,11 @@ class OpenAINormalizer(FormatNormalizer):
if tr_block is None:
return None, dropped
return (
InternalMessage(role=Role.USER, content=[tr_block], extra=self._extract_extra(msg, {"role", "content"})),
InternalMessage(
role=Role.USER,
content=[tr_block],
extra=self._extract_extra(msg, {"role", "content"}),
),
dropped,
)
@@ -668,24 +869,34 @@ class OpenAINormalizer(FormatNormalizer):
blocks: list[ContentBlock] = []
for part in content:
if not isinstance(part, dict):
dropped["openai_content_part_non_dict"] = dropped.get("openai_content_part_non_dict", 0) + 1
dropped["openai_content_part_non_dict"] = (
dropped.get("openai_content_part_non_dict", 0) + 1
)
continue
ptype = str(part.get("type") or "unknown")
if ptype == "text":
text = str(part.get("text") or "")
if text:
blocks.append(TextBlock(text=text, extra=self._extract_extra(part, {"type", "text"})))
blocks.append(
TextBlock(text=text, extra=self._extract_extra(part, {"type", "text"}))
)
continue
if ptype == "image_url":
url = (part.get("image_url") or {}).get("url") if isinstance(part.get("image_url"), dict) else None
url = (
(part.get("image_url") or {}).get("url")
if isinstance(part.get("image_url"), dict)
else None
)
if isinstance(url, str) and url:
img = self._image_url_to_block(url)
img.extra.update(self._extract_extra(part, {"type", "image_url"}))
blocks.append(img)
else:
dropped["openai_image_url_missing"] = dropped.get("openai_image_url_missing", 0) + 1
dropped["openai_image_url_missing"] = (
dropped.get("openai_image_url_missing", 0) + 1
)
blocks.append(UnknownBlock(raw_type="image_url", payload=part))
continue
@@ -731,7 +942,9 @@ class OpenAINormalizer(FormatNormalizer):
parameters=params_raw if isinstance(params_raw, dict) else None,
extra={
"openai_tool": self._extract_extra(tool, {"type", "function"}),
"openai_function": self._extract_extra(function, {"name", "description", "parameters"}),
"openai_function": self._extract_extra(
function, {"name", "description", "parameters"}
),
},
)
)
@@ -756,7 +969,9 @@ class OpenAINormalizer(FormatNormalizer):
fn_raw = tool_choice.get("function")
fn: dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
name = str(fn.get("name") or "")
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, extra={"openai": tool_choice})
return ToolChoice(
type=ToolChoiceType.TOOL, tool_name=name, extra={"openai": tool_choice}
)
return ToolChoice(type=ToolChoiceType.AUTO, extra={"openai": tool_choice})
@@ -771,7 +986,9 @@ class OpenAINormalizer(FormatNormalizer):
return {"type": "function", "function": {"name": tool_choice.tool_name or ""}}
return "auto"
def _openai_tool_call_to_block(self, tool_call: Any) -> tuple[ToolUseBlock | None, dict[str, int]]:
def _openai_tool_call_to_block(
self, tool_call: Any
) -> tuple[ToolUseBlock | None, dict[str, int]]:
dropped: dict[str, int] = {}
if not isinstance(tool_call, dict):
dropped["openai_tool_call_non_dict"] = dropped.get("openai_tool_call_non_dict", 0) + 1
@@ -808,12 +1025,16 @@ class OpenAINormalizer(FormatNormalizer):
dropped,
)
def _legacy_function_call_to_block(self, func_call: dict[str, Any]) -> tuple[ToolUseBlock | None, dict[str, int]]:
def _legacy_function_call_to_block(
self, func_call: dict[str, Any]
) -> tuple[ToolUseBlock | None, dict[str, int]]:
dropped: dict[str, int] = {}
name = str(func_call.get("name") or "")
args_str = str(func_call.get("arguments") or "")
if not name:
dropped["openai_function_call_missing_name"] = dropped.get("openai_function_call_missing_name", 0) + 1
dropped["openai_function_call_missing_name"] = (
dropped.get("openai_function_call_missing_name", 0) + 1
)
return None, dropped
tool_input: dict[str, Any]
@@ -844,7 +1065,12 @@ class OpenAINormalizer(FormatNormalizer):
dropped: dict[str, int] = {}
content = msg.get("content")
if content is None:
return ToolResultBlock(tool_use_id=tool_call_id, output=None, content_text=None, extra={"openai": msg}), dropped
return (
ToolResultBlock(
tool_use_id=tool_call_id, output=None, content_text=None, extra={"openai": msg}
),
dropped,
)
if isinstance(content, str):
parsed: Any = None
@@ -906,7 +1132,9 @@ class OpenAINormalizer(FormatNormalizer):
extra={"openai": extra} if extra else {},
)
def _blocks_to_openai_content(self, blocks: list[ContentBlock]) -> str | list[dict[str, Any]] | None:
def _blocks_to_openai_content(
self, blocks: list[ContentBlock]
) -> str | list[dict[str, Any]] | None:
parts: list[dict[str, Any]] = []
text_parts: list[str] = []
@@ -946,7 +1174,9 @@ class OpenAINormalizer(FormatNormalizer):
# OpenAI content 可以是空字符串;但作为响应 message.content 通常允许为 ""/None。
return ""
def _split_blocks(self, blocks: list[ContentBlock]) -> tuple[list[ContentBlock], list[ToolUseBlock]]:
def _split_blocks(
self, blocks: list[ContentBlock]
) -> tuple[list[ContentBlock], list[ToolUseBlock]]:
content_blocks: list[ContentBlock] = []
tool_blocks: list[ToolUseBlock] = []
for b in blocks:
@@ -979,7 +1209,13 @@ class OpenAINormalizer(FormatNormalizer):
def flush_user() -> None:
nonlocal pending
# 丢弃 Unknown/Tool blockstool_result 在 flush 时不会出现)
content = self._blocks_to_openai_content([b for b in pending if not isinstance(b, (UnknownBlock, ToolUseBlock, ToolResultBlock))])
content = self._blocks_to_openai_content(
[
b
for b in pending
if not isinstance(b, (UnknownBlock, ToolUseBlock, ToolResultBlock))
]
)
if content is None:
content = ""
out.append({"role": "user", "content": content})
@@ -1025,7 +1261,9 @@ class OpenAINormalizer(FormatNormalizer):
out["content"] = content_value if content_value is not None else ""
if tool_blocks:
out["tool_calls"] = [self._tool_use_block_to_openai_call(b, idx) for idx, b in enumerate(tool_blocks)]
out["tool_calls"] = [
self._tool_use_block_to_openai_call(b, idx) for idx, b in enumerate(tool_blocks)
]
return out

View File

@@ -6,13 +6,13 @@ API 格式检测
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from starlette.requests import Request
from src.core.api_format.enums import APIFormat
from src.core.api_format.enums import APIFormat, AuthMethod, EndpointType
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS, ApiFormatDefinition
@@ -64,6 +64,78 @@ def _extract_api_key_by_definition(
return header_value, "header"
@dataclass(frozen=True)
class RequestContext:
"""请求上下文 - 三维度信息"""
data_format: APIFormat
endpoint_type: EndpointType
auth_method: AuthMethod
credentials: str | None
def _detect_endpoint_type(path: str) -> EndpointType:
normalized = path.lower()
if normalized.startswith("/upload/v1beta/files") or normalized.startswith("/v1beta/files"):
return EndpointType.FILES
if normalized.startswith("/v1/videos") or (
normalized.startswith("/v1beta/") and "predictlongrunning" in normalized
):
return EndpointType.VIDEO
# Gemini operations (视频轮询) 也归类为 VIDEO
if normalized.startswith("/v1beta/operations"):
return EndpointType.VIDEO
if normalized.startswith("/v1/models"):
return EndpointType.MODELS
if "/embeddings" in normalized:
return EndpointType.EMBEDDING
if "/images" in normalized:
return EndpointType.IMAGE
if "/audio" in normalized:
return EndpointType.AUDIO
return EndpointType.CHAT
def _detect_data_format(
path: str, headers: dict[str, str], query_params: dict[str, str] | None
) -> APIFormat:
normalized = path.lower()
if normalized.startswith("/v1/messages"):
return APIFormat.CLAUDE
if normalized.startswith("/v1beta/") or normalized.startswith("/upload/v1beta/"):
return APIFormat.GEMINI
if normalized.startswith("/v1/chat/completions") or normalized.startswith("/v1/videos"):
return APIFormat.OPENAI
api_format, _api_key, _auth_method = detect_format_from_request(headers, query_params)
return api_format
def _detect_auth_method(
headers: dict[str, str], query_params: dict[str, str] | None
) -> tuple[AuthMethod, str | None]:
# Query key (Gemini) has highest priority
query_key = query_params.get("key") if query_params else None
if query_key:
return AuthMethod.QUERY_KEY, query_key
x_goog_key = headers.get("x-goog-api-key")
if x_goog_key:
return AuthMethod.GOOG_API_KEY, x_goog_key
x_api_key = headers.get("x-api-key")
if x_api_key:
return AuthMethod.API_KEY, x_api_key
auth_header = headers.get("authorization", "")
if auth_header.lower().startswith("bearer "):
return AuthMethod.BEARER, auth_header[7:].strip()
return AuthMethod.BEARER, None
def detect_format_from_request(
headers: dict[str, str],
query_params: dict[str, str] | None = None,
@@ -86,20 +158,26 @@ def detect_format_from_request(
"""
# Claude: x-api-key + anthropic-version (必须同时存在)
claude_def = API_FORMAT_DEFINITIONS[APIFormat.CLAUDE]
claude_key, claude_auth_method = _extract_api_key_by_definition(headers, query_params, claude_def)
claude_key, claude_auth_method = _extract_api_key_by_definition(
headers, query_params, claude_def
)
if claude_key and headers.get("anthropic-version"):
return APIFormat.CLAUDE, claude_key, claude_auth_method
# Gemini: x-goog-api-key (header 类型) 或 ?key=
gemini_def = API_FORMAT_DEFINITIONS[APIFormat.GEMINI]
gemini_key, gemini_auth_method = _extract_api_key_by_definition(headers, query_params, gemini_def)
gemini_key, gemini_auth_method = _extract_api_key_by_definition(
headers, query_params, gemini_def
)
if gemini_key:
return APIFormat.GEMINI, gemini_key, gemini_auth_method
# OpenAI: Authorization: Bearer (默认)
# 注意: 如果只有 x-api-key 但没有 anthropic-version也走 OpenAI 格式
openai_def = API_FORMAT_DEFINITIONS[APIFormat.OPENAI]
openai_key, openai_auth_method = _extract_api_key_by_definition(headers, query_params, openai_def)
openai_key, openai_auth_method = _extract_api_key_by_definition(
headers, query_params, openai_def
)
# 如果 OpenAI 格式没有 key但有 x-api-key也用它兼容
if not openai_key and claude_key:
openai_key = claude_key
@@ -134,6 +212,28 @@ def detect_format_and_key_from_starlette(
return format_name, api_key, auth_method
def detect_request_context(request: Request) -> RequestContext:
"""
从 Request 中检测三维度信息
Returns:
RequestContext(data_format, endpoint_type, auth_method, credentials)
"""
headers = {k.lower(): v for k, v in request.headers.items()}
query_params = dict(request.query_params)
endpoint_type = _detect_endpoint_type(request.url.path)
data_format = _detect_data_format(request.url.path, headers, query_params)
auth_method, credentials = _detect_auth_method(headers, query_params)
return RequestContext(
data_format=data_format,
endpoint_type=endpoint_type,
auth_method=auth_method,
credentials=credentials,
)
def detect_format_from_response(
response_data: dict,
) -> APIFormat | None:
@@ -197,4 +297,6 @@ __all__ = [
"detect_format_and_key_from_starlette",
"detect_format_from_response",
"detect_cli_format_from_path",
"detect_request_context",
"RequestContext",
]

View File

@@ -18,4 +18,26 @@ class APIFormat(Enum):
GEMINI_CLI = "GEMINI_CLI" # Gemini CLI API 格式
__all__ = ["APIFormat"]
class AuthMethod(str, Enum):
"""认证方式 - 决定如何构造认证 Header"""
BEARER = "bearer" # Authorization: Bearer {token}
API_KEY = "api_key" # x-api-key: {key}
GOOG_API_KEY = "goog_key" # x-goog-api-key: {key}
OAUTH2 = "oauth2" # Google OAuth2 / Service Account
QUERY_KEY = "query_key" # ?key={key} (Gemini 备用)
class EndpointType(str, Enum):
"""端点类型 - 决定 API 功能类别"""
CHAT = "chat" # Chat/Completion API
VIDEO = "video" # Video Generation API
FILES = "files" # Files API
IMAGE = "image" # Image Generation API
AUDIO = "audio" # Audio API
EMBEDDING = "embedding" # Embedding API
MODELS = "models" # Models API
__all__ = ["APIFormat", "AuthMethod", "EndpointType"]

View File

@@ -5,9 +5,8 @@
from __future__ import annotations
from typing import Any
from contextlib import asynccontextmanager
from typing import Any
import uvicorn
from fastapi import FastAPI, HTTPException
@@ -30,12 +29,10 @@ from src.core.exceptions import ExceptionHandlers, ProxyException
from src.core.logger import logger
from src.core.modules import get_module_registry
from src.database import init_db
from src.middleware.plugin_middleware import PluginMiddleware
from src.plugins.manager import get_plugin_manager
async def initialize_providers() -> None:
"""从数据库初始化提供商(仅用于日志记录)"""
from sqlalchemy.orm import Session, selectinload
@@ -131,7 +128,9 @@ async def lifespan(app: FastAPI) -> Any:
if redis_client:
logger.info("[OK] Redis客户端初始化成功缓存亲和性功能已启用")
else:
logger.warning("[WARN] Redis未启用或连接失败将使用内存缓存亲和性仅适用于单实例/开发环境)")
logger.warning(
"[WARN] Redis未启用或连接失败将使用内存缓存亲和性仅适用于单实例/开发环境)"
)
except RuntimeError as e:
if config.require_redis:
logger.exception("[ERROR] Redis连接失败应用启动中止")
@@ -201,14 +200,16 @@ async def lifespan(app: FastAPI) -> Any:
# 启动月卡额度重置调度器(仅一个 worker 执行)
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.usage.quota_scheduler import get_quota_scheduler
from src.services.model.fetch_scheduler import get_model_fetch_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_coordinator = StartupTaskCoordinator(redis_client)
# 启动额度调度器
@@ -237,6 +238,15 @@ 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()
else:
logger.info("检测到其他 worker 已运行视频任务轮询,本实例跳过")
video_task_poller = None # type: ignore[assignment]
# 启动统一的定时任务调度器
from src.services.system.scheduler import get_scheduler
@@ -286,6 +296,11 @@ 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")
# 停止统一的定时任务调度器
logger.info("停止定时任务调度器...")
task_scheduler.stop()
@@ -411,7 +426,7 @@ app = FastAPI(
docs_url="/docs" if config.docs_enabled else None,
redoc_url="/redoc" if config.docs_enabled else None,
openapi_url="/openapi.json" if config.docs_enabled else None,
openapi_tags=openapi_tags
openapi_tags=openapi_tags,
)
# 注册全局异常处理器
@@ -458,7 +473,6 @@ app.include_router(public_router) # 公开API端点用户可查看提供商
app.include_router(monitoring_router) # 监控端点
def main() -> Any:
# 初始化新日志系统
debug_mode = config.environment == "development"

View File

@@ -4,12 +4,12 @@
from __future__ import annotations
from typing import Any
import hashlib
import secrets
import uuid
from datetime import datetime, timezone
from enum import Enum as PyEnum
from typing import Any
import bcrypt
from sqlalchemy import (
@@ -1225,6 +1225,97 @@ class ProviderAPIKey(Base):
provider = relationship("Provider", back_populates="api_keys")
class VideoTask(Base):
"""视频生成任务"""
__tablename__ = "video_tasks"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
external_task_id = Column(String(200))
# 关联
user_id = Column(String(36), ForeignKey("users.id"), nullable=False)
api_key_id = Column(String(36), ForeignKey("api_keys.id"))
provider_id = Column(String(36), ForeignKey("providers.id"))
endpoint_id = Column(String(36), ForeignKey("provider_endpoints.id"))
key_id = Column(String(36), ForeignKey("provider_api_keys.id"))
# 格式转换追踪
client_api_format = Column(String(50), nullable=False)
provider_api_format = Column(String(50), nullable=False)
format_converted = Column(Boolean, default=False)
# 任务配置
model = Column(String(100), nullable=False)
prompt = Column(Text, nullable=False)
original_request_body = Column(JSON)
converted_request_body = Column(JSON)
# 视频参数 (统一内部格式)
duration_seconds = Column(Integer, default=4)
resolution = Column(String(20), default="720p")
aspect_ratio = Column(String(10), default="16:9")
size = Column(String(20))
# 状态
status = Column(String(20), default="pending")
progress_percent = Column(Integer, default=0)
progress_message = Column(String(500))
# 结果
video_url = Column(String(2000))
video_urls = Column(JSON)
thumbnail_url = Column(String(2000))
video_size_bytes = Column(BigInteger)
video_expires_at = Column(DateTime(timezone=True))
# 存储 (可选)
stored_video_path = Column(String(500))
storage_provider = Column(String(50))
# 错误
error_code = Column(String(50))
error_message = Column(Text)
retry_count = Column(Integer, default=0)
max_retries = Column(Integer, default=3)
# 轮询配置
poll_interval_seconds = Column(Integer, default=10)
next_poll_at = Column(DateTime(timezone=True)) # 索引在 __table_args__ 中定义
poll_count = Column(Integer, default=0)
max_poll_count = Column(Integer, default=360)
# Remix 支持
remixed_from_task_id = Column(
String(36), ForeignKey("video_tasks.id", ondelete="SET NULL"), nullable=True
)
# 时间戳
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
submitted_at = Column(DateTime(timezone=True))
completed_at = Column(DateTime(timezone=True))
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
# 关系
user = relationship("User", backref="video_tasks")
remixed_from = relationship("VideoTask", remote_side=[id], backref="remixes")
# 复合索引和唯一约束
__table_args__ = (
Index("idx_video_tasks_user_status", "user_id", "status"),
Index("idx_video_tasks_next_poll", "next_poll_at"),
Index("idx_video_tasks_external_id", "external_task_id"),
UniqueConstraint("user_id", "external_task_id", name="uq_video_tasks_user_external_id"),
)
class UserPreference(Base):
"""用户偏好设置表"""

View File

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

View File

@@ -0,0 +1,371 @@
"""
视频任务后台轮询服务
"""
from __future__ import annotations
import asyncio
from datetime import datetime, timedelta, timezone
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 sanitize_error_message
from src.clients.http_client import HTTPClientPool
from src.clients.redis_client import get_redis_client
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint
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
# 永久性错误指示词(用于降级判断,不应重试)
_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
BATCH_SIZE = 50
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
async def start(self) -> None:
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=10,
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
batch_failures = 0
for task in tasks:
try:
await self._poll_single_task(db, task)
except Exception as exc:
batch_failures += 1
# 单个任务失败不影响其他任务处理
logger.exception(
"Unexpected error polling task %s: %s",
task.id,
sanitize_error_message(str(exc)),
)
# 更新连续失败计数并检查告警阈值
if batch_failures == len(tasks):
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
db.commit()
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
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)
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)
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",
)
if (task.provider_api_format or "").upper() == "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)
headers = self._build_headers(APIFormat.OPENAI, 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 = task.external_task_id
if not operation_name.startswith("operations/"):
operation_name = f"operations/{operation_name}"
url = self._build_gemini_url(endpoint.base_url, operation_name)
headers = self._build_headers(APIFormat.GEMINI, 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,
api_format: APIFormat,
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(
{},
api_format,
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