mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -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 实例
|
||||
# =========================================================================
|
||||
|
||||
@@ -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 实例
|
||||
# =========================================================================
|
||||
|
||||
137
src/api/handlers/base/video_adapter_base.py
Normal file
137
src/api/handlers/base/video_adapter_base.py
Normal 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"]
|
||||
208
src/api/handlers/base/video_handler_base.py
Normal file
208
src/api/handlers/base/video_handler_base.py
Normal 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"]
|
||||
@@ -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()
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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,
|
||||
|
||||
22
src/api/handlers/gemini/video_adapter.py
Normal file
22
src/api/handlers/gemini/video_adapter.py
Normal 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"]
|
||||
402
src/api/handlers/gemini/video_handler.py
Normal file
402
src/api/handlers/gemini/video_handler.py
Normal 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 Bearer(Vertex 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"]
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
22
src/api/handlers/openai/video_adapter.py
Normal file
22
src/api/handlers/openai/video_adapter.py
Normal 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"]
|
||||
401
src/api/handlers/openai/video_handler.py
Normal file
401
src/api/handlers/openai/video_handler.py
Normal 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"]
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
163
src/api/public/videos.py
Normal 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},
|
||||
)
|
||||
@@ -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
129
src/core/api_format/auth.py
Normal 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",
|
||||
]
|
||||
84
src/core/api_format/conversion/internal_video.py
Normal file
84
src/core/api_format/conversion/internal_video.py
Normal 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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""
|
||||
"""
|
||||
Gemini (GenerateContent / streamGenerateContent) Normalizer
|
||||
|
||||
负责:
|
||||
@@ -11,7 +11,6 @@ Gemini (GenerateContent / streamGenerateContent) Normalizer
|
||||
- 响应/流式通常为 camelCase(candidates/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)),
|
||||
|
||||
@@ -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 blocks(tool_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
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
30
src/main.py
30
src/main.py
@@ -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"
|
||||
|
||||
@@ -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):
|
||||
"""用户偏好设置表"""
|
||||
|
||||
|
||||
10
src/services/video/__init__.py
Normal file
10
src/services/video/__init__.py
Normal file
@@ -0,0 +1,10 @@
|
||||
"""
|
||||
视频相关服务
|
||||
"""
|
||||
|
||||
from src.services.video.task_poller import VideoTaskPollerService, get_video_task_poller
|
||||
|
||||
__all__ = [
|
||||
"VideoTaskPollerService",
|
||||
"get_video_task_poller",
|
||||
]
|
||||
371
src/services/video/task_poller.py
Normal file
371
src/services/video/task_poller.py
Normal 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
|
||||
# 存储多视频 URL(Gemini 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
|
||||
Reference in New Issue
Block a user