mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 09:50: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:
@@ -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