mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +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},
|
||||
)
|
||||
Reference in New Issue
Block a user