mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/ - 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层 - 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构 - 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations) - 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image - 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
28
_deprecated_py_src/api/handlers/gemini/__init__.py
Normal file
28
_deprecated_py_src/api/handlers/gemini/__init__.py
Normal file
@@ -0,0 +1,28 @@
|
||||
"""Gemini handler package (lazy exports)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
|
||||
"GeminiChatAdapter": (".adapter", "GeminiChatAdapter"),
|
||||
"build_gemini_adapter": (".adapter", "build_gemini_adapter"),
|
||||
"GeminiChatHandler": (".handler", "GeminiChatHandler"),
|
||||
"GeminiStreamParser": (".stream_parser", "GeminiStreamParser"),
|
||||
"GeminiVeoAdapter": (".video_adapter", "GeminiVeoAdapter"),
|
||||
"GeminiVeoHandler": (".video_handler", "GeminiVeoHandler"),
|
||||
}
|
||||
|
||||
__all__ = list(_LAZY_EXPORTS.keys())
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
if name not in _LAZY_EXPORTS:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
module_name, attr_name = _LAZY_EXPORTS[name]
|
||||
module = import_module(module_name, __name__)
|
||||
value = getattr(module, attr_name)
|
||||
globals()[name] = value
|
||||
return value
|
||||
438
_deprecated_py_src/api/handlers/gemini/adapter.py
Normal file
438
_deprecated_py_src/api/handlers/gemini/adapter.py
Normal file
@@ -0,0 +1,438 @@
|
||||
"""
|
||||
Gemini Chat Adapter
|
||||
|
||||
处理 Gemini API 格式的请求适配
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
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 ApiFamily, get_auth_handler, resolve_header_name_case
|
||||
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
|
||||
|
||||
|
||||
class GeminiCapabilityDetector:
|
||||
"""Gemini API 能力检测器"""
|
||||
|
||||
@staticmethod
|
||||
def detect_from_request(
|
||||
headers: dict[str, str], # noqa: ARG004 - 预留
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""
|
||||
从请求体检测 Gemini 能力需求
|
||||
|
||||
检测规则:
|
||||
- fileData.fileUri -> gemini_files: True
|
||||
"""
|
||||
requirements: dict[str, bool] = {}
|
||||
if request_body and extract_file_names_from_request(request_body):
|
||||
requirements["gemini_files"] = True
|
||||
return requirements
|
||||
|
||||
|
||||
@register_adapter
|
||||
class GeminiChatAdapter(ChatAdapterBase):
|
||||
"""
|
||||
Gemini Chat API 适配器
|
||||
|
||||
处理 Gemini Chat 格式的请求
|
||||
端点: /v1beta/models/{model}:generateContent
|
||||
"""
|
||||
|
||||
FORMAT_ID = "gemini:chat"
|
||||
API_FAMILY = ApiFamily.GEMINI
|
||||
name = "gemini.chat"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.gemini.handler import GeminiChatHandler
|
||||
|
||||
return GeminiChatHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats)
|
||||
logger.info(
|
||||
f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}"
|
||||
)
|
||||
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
|
||||
|
||||
优先级(与 Google SDK 行为一致):
|
||||
1. URL 参数 ?key=
|
||||
2. x-goog-api-key 请求头
|
||||
"""
|
||||
handler = get_auth_handler(AuthMethod.GOOG_API_KEY)
|
||||
return handler.extract_credentials(request)
|
||||
|
||||
def detect_capability_requirements(
|
||||
self,
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""从请求体检测 Gemini 能力需求(fileData.fileUri -> gemini_files)"""
|
||||
return GeminiCapabilityDetector.detect_from_request(headers, request_body)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - Gemini 特化版本
|
||||
|
||||
Gemini API 特点:
|
||||
- model 不合并到请求体(通过 extract_model_from_request 从 path_params 获取)
|
||||
- stream 不合并到请求体(Gemini API 通过 URL 端点区分流式/非流式)
|
||||
|
||||
Handler 层的 extract_model_from_request 会从 path_params 获取 model,
|
||||
prepare_provider_request_body 会确保发送给 Gemini API 的请求体不含 model。
|
||||
|
||||
Args:
|
||||
original_request_body: 原始请求体字典
|
||||
path_params: URL 路径参数字典(不使用)
|
||||
|
||||
Returns:
|
||||
原始请求体(不合并任何 path_params)
|
||||
"""
|
||||
return original_request_body.copy()
|
||||
|
||||
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)
|
||||
model = path_params.get("model", "unknown")
|
||||
|
||||
try:
|
||||
if not isinstance(original_request_body, dict):
|
||||
raise ValueError("Request body must be a JSON object")
|
||||
|
||||
# Gemini 必需字段: contents
|
||||
if "contents" not in original_request_body:
|
||||
raise ValueError("Missing required field: contents")
|
||||
|
||||
request = GeminiRequest.model_validate(
|
||||
original_request_body,
|
||||
strict=False,
|
||||
)
|
||||
except ValueError as e:
|
||||
logger.error(f"请求体基本验证失败: {str(e)}")
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
logger.warning(f"Pydantic验证警告(将继续处理): {str(e)}")
|
||||
request = GeminiRequest.model_construct(
|
||||
contents=original_request_body.get("contents", []),
|
||||
)
|
||||
|
||||
# 设置 model(从 path_params 获取,用于日志和审计)
|
||||
request.model = model
|
||||
# 设置 stream 属性(用于 ChatAdapterBase 判断流式模式)
|
||||
request.stream = is_stream
|
||||
return request
|
||||
|
||||
def _extract_message_count(self, payload: dict[str, Any], request_obj: Any) -> int:
|
||||
"""提取消息数量"""
|
||||
contents = payload.get("contents", [])
|
||||
if hasattr(request_obj, "contents"):
|
||||
contents = request_obj.contents
|
||||
return len(contents) if isinstance(contents, list) else 0
|
||||
|
||||
def _build_audit_metadata(self, payload: dict[str, Any], request_obj: Any) -> dict[str, Any]:
|
||||
"""构建 Gemini Chat 特定的审计元数据"""
|
||||
role_counts: dict[str, int] = {}
|
||||
|
||||
contents = getattr(request_obj, "contents", []) or []
|
||||
for content in contents:
|
||||
if isinstance(content, dict):
|
||||
role = content.get("role", "unknown")
|
||||
else:
|
||||
role = getattr(content, "role", None) or "unknown"
|
||||
role_counts[role] = role_counts.get(role, 0) + 1
|
||||
|
||||
generation_config = getattr(request_obj, "generation_config", None) or {}
|
||||
if hasattr(generation_config, "dict"):
|
||||
generation_config = generation_config.dict()
|
||||
elif not isinstance(generation_config, dict):
|
||||
generation_config = {}
|
||||
|
||||
# 判断流式模式
|
||||
stream = getattr(request_obj, "stream", False)
|
||||
|
||||
return {
|
||||
"action": "gemini_generate_content",
|
||||
"model": getattr(request_obj, "model", payload.get("model", "unknown")),
|
||||
"stream": bool(stream),
|
||||
"max_output_tokens": generation_config.get("max_output_tokens"),
|
||||
"temperature": generation_config.get("temperature"),
|
||||
"top_p": generation_config.get("top_p"),
|
||||
"top_k": generation_config.get("top_k"),
|
||||
"contents_count": len(contents),
|
||||
"content_roles": role_counts,
|
||||
"tools_count": len(getattr(request_obj, "tools", None) or []),
|
||||
"system_instruction_present": bool(getattr(request_obj, "system_instruction", None)),
|
||||
"safety_settings_count": len(getattr(request_obj, "safety_settings", None) or []),
|
||||
}
|
||||
|
||||
def _error_response(self, status_code: int, error_type: str, message: str) -> JSONResponse:
|
||||
"""生成 Gemini 格式的错误响应"""
|
||||
# Gemini 错误响应格式
|
||||
return JSONResponse(
|
||||
status_code=status_code,
|
||||
content={
|
||||
"error": {
|
||||
"code": status_code,
|
||||
"message": message,
|
||||
"status": error_type.upper(),
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(
|
||||
cls,
|
||||
base_url: str,
|
||||
request_data: dict[str, Any] | None = None,
|
||||
model_name: str | None = None,
|
||||
*,
|
||||
provider_type: str | None = None,
|
||||
) -> str:
|
||||
"""构建Gemini API端点URL"""
|
||||
base_url = base_url.rstrip("/")
|
||||
if base_url.endswith("/v1beta"):
|
||||
return base_url # 子类需要处理model参数
|
||||
else:
|
||||
return f"{base_url}/v1beta"
|
||||
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI
|
||||
|
||||
@classmethod
|
||||
async def check_endpoint(
|
||||
cls,
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 端点规则参数
|
||||
body_rules: list[dict[str, Any]] | None = None,
|
||||
header_rules: list[dict[str, Any]] | None = None,
|
||||
# 用量计算参数
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
provider_name: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
# Provider 上下文
|
||||
auth_type: str | None = None,
|
||||
provider_type: str | None = None,
|
||||
decrypted_auth_config: dict[str, Any] | None = None,
|
||||
provider_endpoint: Any | None = None,
|
||||
provider_api_key: Any | None = None,
|
||||
# 代理配置
|
||||
proxy_config: dict[str, Any] | None = None,
|
||||
timeout_seconds: float | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""测试 Gemini API 模型连接性(非流式)"""
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
from src.api.handlers.base.request_builder import (
|
||||
apply_body_rules,
|
||||
evaluate_condition,
|
||||
)
|
||||
from src.core.api_format.headers import HeaderBuilder
|
||||
from src.services.provider.adapters.vertex_ai.transport import is_vertex_ai_context
|
||||
|
||||
# Gemini需要从request_data或model_name参数获取model名称
|
||||
effective_model_name = model_name or request_data.get("model", "")
|
||||
if not effective_model_name:
|
||||
return {
|
||||
"error": "Model name is required for Gemini API",
|
||||
"status_code": 400,
|
||||
}
|
||||
|
||||
is_antigravity = provider_type and provider_type.lower() == "antigravity"
|
||||
is_gemini_cli = provider_type and provider_type.lower() == "gemini_cli"
|
||||
is_vertex = is_vertex_ai_context(
|
||||
base_url=base_url,
|
||||
provider_type=provider_type,
|
||||
endpoint=provider_endpoint,
|
||||
key=provider_api_key,
|
||||
)
|
||||
is_oauth = auth_type == "oauth"
|
||||
vertex_auth_info: Any | None = None
|
||||
|
||||
# Antigravity provider 使用 v1internal 路径,而非标准 Gemini API 路径
|
||||
if is_antigravity:
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
V1INTERNAL_PATH_TEMPLATE,
|
||||
get_v1internal_extra_headers,
|
||||
)
|
||||
from src.services.provider.adapters.antigravity.url_availability import url_availability
|
||||
|
||||
ordered_urls = url_availability.get_ordered_urls(prefer_daily=True)
|
||||
ag_base = ordered_urls[0] if ordered_urls else base_url
|
||||
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
|
||||
url = f"{str(ag_base).rstrip('/')}{path}"
|
||||
elif is_gemini_cli:
|
||||
from src.services.provider.adapters.gemini_cli.constants import V1INTERNAL_PATH_TEMPLATE
|
||||
|
||||
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
|
||||
url = f"{str(base_url).rstrip('/')}{path}"
|
||||
elif is_vertex and provider_endpoint is not None and provider_api_key is not None:
|
||||
# Vertex AI: test-model 必须走统一 provider transport/auth,
|
||||
# 否则会错误命中普通 Gemini URL(导致 404)。
|
||||
from src.services.provider.auth import get_provider_auth
|
||||
from src.services.provider.transport import build_provider_url
|
||||
|
||||
vertex_auth_info = await get_provider_auth(provider_endpoint, provider_api_key)
|
||||
effective_auth_config = (
|
||||
vertex_auth_info.decrypted_auth_config
|
||||
if vertex_auth_info
|
||||
else decrypted_auth_config
|
||||
)
|
||||
if effective_auth_config:
|
||||
decrypted_auth_config = effective_auth_config
|
||||
url = build_provider_url(
|
||||
provider_endpoint,
|
||||
path_params={"model": effective_model_name},
|
||||
is_stream=bool(request_data.get("stream", False)),
|
||||
key=provider_api_key,
|
||||
decrypted_auth_config=effective_auth_config,
|
||||
)
|
||||
else:
|
||||
# 使用基类配置方法,但重写URL构建逻辑
|
||||
base_url_resolved = cls.build_endpoint_url(base_url)
|
||||
url = f"{base_url_resolved}/models/{effective_model_name}:generateContent"
|
||||
|
||||
# 构建请求组件
|
||||
# Antigravity 需要特定的 User-Agent
|
||||
merged_extra = dict(extra_headers) if extra_headers else {}
|
||||
if is_antigravity:
|
||||
merged_extra.update(get_v1internal_extra_headers())
|
||||
elif is_gemini_cli:
|
||||
from src.services.provider.adapters.gemini_cli.constants import (
|
||||
get_v1internal_extra_headers,
|
||||
)
|
||||
|
||||
merged_extra.update(get_v1internal_extra_headers())
|
||||
if is_vertex and provider_endpoint is not None and provider_api_key is not None:
|
||||
headers = dict(merged_extra)
|
||||
if (
|
||||
vertex_auth_info
|
||||
and getattr(vertex_auth_info, "auth_header", None)
|
||||
and getattr(vertex_auth_info, "auth_value", None)
|
||||
):
|
||||
headers[str(vertex_auth_info.auth_header)] = str(vertex_auth_info.auth_value)
|
||||
else:
|
||||
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
|
||||
|
||||
# OAuth 统一处理:替换端点默认认证头(x-goog-api-key)为 Authorization: Bearer
|
||||
if is_oauth:
|
||||
from src.core.api_format import get_auth_config_for_endpoint
|
||||
|
||||
default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
|
||||
if default_auth_header.lower() != "authorization":
|
||||
headers.pop(default_auth_header, None)
|
||||
auth_header_name = resolve_header_name_case(extra_headers, "Authorization")
|
||||
headers[auth_header_name] = f"Bearer {api_key}"
|
||||
|
||||
body = cls.build_request_body(request_data)
|
||||
|
||||
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
|
||||
if body_rules:
|
||||
body = apply_body_rules(
|
||||
body,
|
||||
body_rules,
|
||||
original_body=body,
|
||||
)
|
||||
|
||||
# Antigravity 需要将请求体包装为 v1internal 信封格式
|
||||
if is_antigravity:
|
||||
from src.services.provider.adapters.antigravity.envelope import wrap_v1internal_request
|
||||
|
||||
project_id = (decrypted_auth_config or {}).get("project_id", "")
|
||||
body = wrap_v1internal_request(
|
||||
body,
|
||||
project_id=project_id,
|
||||
model=effective_model_name,
|
||||
request_type="endpoint_test",
|
||||
)
|
||||
elif is_gemini_cli:
|
||||
from src.services.provider.adapters.gemini_cli.envelope import wrap_v1internal_request
|
||||
|
||||
project_id = (decrypted_auth_config or {}).get("project_id", "")
|
||||
body = wrap_v1internal_request(
|
||||
body,
|
||||
project_id=project_id,
|
||||
model=effective_model_name,
|
||||
)
|
||||
|
||||
# 应用请求头规则(在请求头构建后应用)
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config_for_endpoint
|
||||
|
||||
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
if vertex_auth_info and getattr(vertex_auth_info, "auth_header", None):
|
||||
protected_keys.add(str(vertex_auth_info.auth_header).lower())
|
||||
|
||||
header_builder = HeaderBuilder()
|
||||
header_builder.add_many(headers)
|
||||
header_builder.apply_rules(
|
||||
header_rules,
|
||||
protected_keys,
|
||||
body=body,
|
||||
original_body=body,
|
||||
condition_evaluator=evaluate_condition,
|
||||
)
|
||||
headers = header_builder.build()
|
||||
|
||||
return await run_endpoint_check(
|
||||
client=client,
|
||||
url=url,
|
||||
headers=headers,
|
||||
json_body=body,
|
||||
api_format=cls.FORMAT_ID,
|
||||
is_stream=bool(request_data.get("stream", False)),
|
||||
# 用量计算参数(现在强制记录)
|
||||
db=db,
|
||||
user=user,
|
||||
provider_name=provider_name,
|
||||
provider_id=provider_id,
|
||||
api_key_id=api_key_id,
|
||||
model_name=effective_model_name,
|
||||
proxy_config=proxy_config,
|
||||
timeout=timeout_seconds,
|
||||
)
|
||||
|
||||
|
||||
def build_gemini_adapter(x_app_header: str = "") -> GeminiChatAdapter: # noqa: ARG001
|
||||
"""
|
||||
根据请求头构建适当的 Gemini 适配器
|
||||
|
||||
Args:
|
||||
x_app_header: X-App 请求头值
|
||||
|
||||
Returns:
|
||||
GeminiChatAdapter 实例
|
||||
"""
|
||||
# 目前只有一种 Gemini 适配器
|
||||
# 未来可以根据 x_app_header 返回不同的适配器(如 CLI 模式)
|
||||
return GeminiChatAdapter()
|
||||
|
||||
|
||||
__all__ = ["GeminiChatAdapter", "build_gemini_adapter"]
|
||||
198
_deprecated_py_src/api/handlers/gemini/handler.py
Normal file
198
_deprecated_py_src/api/handlers/gemini/handler.py
Normal file
@@ -0,0 +1,198 @@
|
||||
"""
|
||||
Gemini Chat Handler
|
||||
|
||||
处理 Gemini API 格式的请求
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from starlette.requests import Request
|
||||
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
from src.core.api_format import ApiFamily, EndpointKind
|
||||
|
||||
|
||||
class GeminiChatHandler(ChatHandlerBase):
|
||||
"""
|
||||
Gemini Chat Handler - 处理 Google Gemini API 格式的请求
|
||||
|
||||
格式特点:
|
||||
- 使用 promptTokenCount / candidatesTokenCount
|
||||
- 支持 cachedContentTokenCount
|
||||
- 请求格式: GeminiRequest
|
||||
- 响应格式: JSON 数组流(非 SSE)
|
||||
"""
|
||||
|
||||
FORMAT_ID = "gemini:chat"
|
||||
API_FAMILY = ApiFamily.GEMINI
|
||||
ENDPOINT_KIND = EndpointKind.CHAT
|
||||
|
||||
async def _resolve_preferred_key_ids(
|
||||
self,
|
||||
model_name: str, # noqa: ARG002 - 仅做文件绑定
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> list[str] | None:
|
||||
"""
|
||||
从 files/xxx 绑定关系中解析优先 Key ID 列表。
|
||||
|
||||
Gemini 文件与上传它的 API Key 绑定,必须使用同一 Key 访问。
|
||||
此方法从缓存中查找文件→Key 映射,优先使用正确的 Key。
|
||||
|
||||
当同一源文件被上传到多个 Key 时,会返回所有可用的 Key ID,
|
||||
让系统能够选择任意可用的 Key。
|
||||
|
||||
注意事项:
|
||||
- 如果映射缺失(缓存过期/重启),会记录警告,请求可能失败
|
||||
- 优先返回所有支持该文件的 Key,让调度器选择可用的
|
||||
"""
|
||||
from src.core.logger import logger
|
||||
from src.services.gemini_files_mapping import (
|
||||
extract_file_names_from_request,
|
||||
get_all_key_ids_for_file,
|
||||
)
|
||||
|
||||
file_names = extract_file_names_from_request(request_body or {})
|
||||
if not file_names:
|
||||
return None
|
||||
|
||||
all_key_ids: set[str] = set()
|
||||
unmapped_files: list[str] = []
|
||||
|
||||
for file_name in file_names:
|
||||
# 获取所有支持该文件的 Key(包括通过 source_hash 关联的)
|
||||
key_ids = await get_all_key_ids_for_file(file_name)
|
||||
if key_ids:
|
||||
all_key_ids.update(key_ids)
|
||||
else:
|
||||
unmapped_files.append(file_name)
|
||||
|
||||
# 警告:映射缺失
|
||||
if unmapped_files:
|
||||
logger.warning(
|
||||
f"[{self.request_id}] Gemini 文件→Key 映射缺失: {unmapped_files},"
|
||||
"请求可能失败(文件属于其他 Key 或映射已过期)"
|
||||
)
|
||||
|
||||
if all_key_ids:
|
||||
logger.debug(f"[{self.request_id}] 文件引用可用的 Key: {list(all_key_ids)}")
|
||||
|
||||
return list(all_key_ids) if all_key_ids else None
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Gemini Chat 格式实现
|
||||
|
||||
Gemini Chat 模式下,model 在请求体中(经过转换后的 GeminiRequest)。
|
||||
与 Gemini CLI 不同,CLI 模式的 model 在 URL 路径中。
|
||||
|
||||
Args:
|
||||
request_body: 请求体
|
||||
path_params: URL 路径参数(Chat 模式通常不使用)
|
||||
|
||||
Returns:
|
||||
模型名
|
||||
"""
|
||||
# 优先从请求体获取,其次从 path_params
|
||||
model = request_body.get("model")
|
||||
if model:
|
||||
return str(model)
|
||||
if path_params and "model" in path_params:
|
||||
return str(path_params["model"])
|
||||
return "unknown"
|
||||
|
||||
async def _convert_request(self, request: Request) -> None:
|
||||
"""
|
||||
将请求转换为 Gemini 格式的 Pydantic 对象
|
||||
|
||||
注意:此方法只做类型转换(dict → Pydantic),不做跨格式转换。
|
||||
跨格式转换由调度/执行层(TaskService + RequestDispatcher)在选中候选后、发送请求前执行,
|
||||
并受全局开关和端点配置控制。
|
||||
|
||||
Args:
|
||||
request: 原始请求对象(应已是 Gemini 格式)
|
||||
|
||||
Returns:
|
||||
GeminiRequest 对象
|
||||
"""
|
||||
from src.models.gemini import GeminiRequest
|
||||
|
||||
# 如果已经是 Gemini 格式 Pydantic 对象,直接返回
|
||||
if isinstance(request, GeminiRequest):
|
||||
return request
|
||||
|
||||
# 如果是字典,转换为 Pydantic 对象(假设已是 Gemini 格式)
|
||||
if isinstance(request, dict):
|
||||
return GeminiRequest(**request)
|
||||
|
||||
return request
|
||||
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从 Gemini 响应中提取 token 使用情况
|
||||
|
||||
调用 GeminiStreamParser.extract_usage 作为单一实现源
|
||||
"""
|
||||
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
|
||||
|
||||
usage = GeminiStreamParser().extract_usage(response)
|
||||
|
||||
if not usage:
|
||||
return {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
}
|
||||
|
||||
return {
|
||||
"input_tokens": usage.get("input_tokens", 0),
|
||||
"output_tokens": usage.get("output_tokens", 0),
|
||||
"cache_creation_input_tokens": 0, # Gemini 不区分缓存创建
|
||||
"cache_read_input_tokens": usage.get("cached_tokens", 0),
|
||||
}
|
||||
|
||||
def finalize_provider_request(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: str | None,
|
||||
provider_api_format: str | None, # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
from src.api.handlers.gemini.image_gen import (
|
||||
adapt_request_for_image_gen,
|
||||
is_image_gen_model,
|
||||
)
|
||||
|
||||
# Sanitize Gemini contents: strip parts without a valid data-oneof
|
||||
# field and merge consecutive same-role entries. This catches cases
|
||||
# missed by the normalizer (passthrough) or the antigravity envelope.
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
if not is_image_gen_model(mapped_model):
|
||||
return request_body
|
||||
return adapt_request_for_image_gen(request_body)
|
||||
|
||||
def _normalize_response(self, response: dict) -> dict:
|
||||
"""
|
||||
规范化 Gemini 响应
|
||||
|
||||
Args:
|
||||
response: 原始响应
|
||||
|
||||
Returns:
|
||||
规范化后的响应
|
||||
"""
|
||||
# 作为中转站,直接透传响应,不做标准化处理
|
||||
return response
|
||||
41
_deprecated_py_src/api/handlers/gemini/image_gen.py
Normal file
41
_deprecated_py_src/api/handlers/gemini/image_gen.py
Normal file
@@ -0,0 +1,41 @@
|
||||
"""
|
||||
Gemini 图像生成模型请求适配
|
||||
|
||||
- 图像生成模型不支持 tools / system_instruction,需要移除
|
||||
- responseModalities / responseMimeType 与 imageConfig 冲突,需要移除
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.core.video_utils import is_image_gen_model
|
||||
|
||||
__all__ = ["is_image_gen_model", "adapt_request_for_image_gen"]
|
||||
|
||||
|
||||
def adapt_request_for_image_gen(body: dict[str, Any]) -> dict[str, Any]:
|
||||
"""为图像生成模型清理不兼容字段"""
|
||||
# 移除图像生成不支持的顶层字段
|
||||
for key in ("tools", "tool_config", "toolConfig", "system_instruction", "systemInstruction"):
|
||||
if key in body:
|
||||
body.pop(key)
|
||||
|
||||
# 处理 generationConfig
|
||||
gc_key = "generationConfig" if "generationConfig" in body else "generation_config"
|
||||
gc = body.get(gc_key)
|
||||
if not isinstance(gc, dict):
|
||||
gc = {}
|
||||
body[gc_key] = gc
|
||||
|
||||
# 移除与图像生成冲突的字段
|
||||
for key in (
|
||||
"responseMimeType",
|
||||
"response_mime_type",
|
||||
"responseModalities",
|
||||
"response_modalities",
|
||||
):
|
||||
gc.pop(key, None)
|
||||
|
||||
# 设置输出模态
|
||||
gc["responseModalities"] = ["TEXT", "IMAGE"]
|
||||
|
||||
return body
|
||||
316
_deprecated_py_src/api/handlers/gemini/stream_parser.py
Normal file
316
_deprecated_py_src/api/handlers/gemini/stream_parser.py
Normal file
@@ -0,0 +1,316 @@
|
||||
"""
|
||||
Gemini 流解析器(SSE + JSON-array 兼容)
|
||||
|
||||
Gemini streamGenerateContent 常见两种返回(与上游/代理实现有关):
|
||||
1) `?alt=sse`:SSE(`data: {GenerateContentResponse}`)
|
||||
2) 默认:JSON-array / JSON-chunks(`[{...},{...},...]`,可能跨 chunk/跨行)
|
||||
|
||||
本解析器提供:
|
||||
- parse_line(): 适用于 SSE data 行或逐行 JSON 对象
|
||||
- parse_chunk(): 适用于 JSON-array/chunks(可跨 chunk 拼接)
|
||||
|
||||
参考:
|
||||
- https://ai.google.dev/gemini-api/docs/text-generation?lang=python#generate-a-text-stream
|
||||
- https://generativelanguage.googleapis.com/$discovery/rest?version=v1beta
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
|
||||
class GeminiStreamParser:
|
||||
"""
|
||||
Gemini 流解析器
|
||||
|
||||
解析 Gemini streamGenerateContent API 的响应流。
|
||||
|
||||
Gemini 流式响应特点:
|
||||
- 每个事件块本质上都是一个 GenerateContentResponse JSON 对象(包含 candidates、usageMetadata 等)
|
||||
- 结束判定以 `candidates[].finishReason` 为准(存在且不为 FINISH_REASON_UNSPECIFIED)
|
||||
"""
|
||||
|
||||
# finishReason(官方枚举值很多,见 discovery;这里仅保留一个明确的“未结束”哨兵)
|
||||
FINISH_REASON_UNSPECIFIED = "FINISH_REASON_UNSPECIFIED"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._buffer = ""
|
||||
self._in_array = False
|
||||
self._brace_depth = 0
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置解析器状态"""
|
||||
self._buffer = ""
|
||||
self._in_array = False
|
||||
self._brace_depth = 0
|
||||
|
||||
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
解析流式数据块
|
||||
|
||||
Args:
|
||||
chunk: 原始数据(bytes 或 str)
|
||||
|
||||
Returns:
|
||||
解析后的事件列表
|
||||
"""
|
||||
if isinstance(chunk, bytes):
|
||||
text = chunk.decode("utf-8")
|
||||
else:
|
||||
text = chunk
|
||||
|
||||
events: list[dict[str, Any]] = []
|
||||
|
||||
for char in text:
|
||||
if char == "[" and not self._in_array:
|
||||
self._in_array = True
|
||||
continue
|
||||
|
||||
if char == "]" and self._in_array and self._brace_depth == 0:
|
||||
# 数组结束
|
||||
self._in_array = False
|
||||
if self._buffer.strip():
|
||||
try:
|
||||
obj = json.loads(self._buffer.strip().rstrip(","))
|
||||
events.append(obj)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
self._buffer = ""
|
||||
continue
|
||||
|
||||
if self._in_array:
|
||||
if char == "{":
|
||||
self._brace_depth += 1
|
||||
elif char == "}":
|
||||
self._brace_depth -= 1
|
||||
|
||||
self._buffer += char
|
||||
|
||||
# 当 brace_depth 回到 0 时,说明一个完整的 JSON 对象结束
|
||||
if self._brace_depth == 0 and self._buffer.strip():
|
||||
try:
|
||||
obj = json.loads(self._buffer.strip().rstrip(","))
|
||||
events.append(obj)
|
||||
self._buffer = ""
|
||||
except json.JSONDecodeError:
|
||||
# 可能还不完整,继续累积
|
||||
pass
|
||||
|
||||
return events
|
||||
|
||||
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||
"""
|
||||
解析单行 JSON 数据
|
||||
|
||||
Args:
|
||||
line: JSON 数据行
|
||||
|
||||
Returns:
|
||||
解析后的事件字典,如果无法解析返回 None
|
||||
"""
|
||||
if not line or line.strip() in ["[", "]", ","]:
|
||||
return None
|
||||
|
||||
try:
|
||||
result = json.loads(line.strip().rstrip(","))
|
||||
if isinstance(result, dict):
|
||||
return result
|
||||
return None
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
def is_done_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为结束事件
|
||||
|
||||
Args:
|
||||
event: 事件字典
|
||||
|
||||
Returns:
|
||||
True 如果是结束事件
|
||||
"""
|
||||
candidates = event.get("candidates", [])
|
||||
if not candidates:
|
||||
return False
|
||||
|
||||
for candidate in candidates:
|
||||
finish_reason = candidate.get("finishReason")
|
||||
if not finish_reason:
|
||||
continue
|
||||
# 只要出现非 UNSPECIFIED 的 finishReason,通常表示该 candidate 已结束。
|
||||
# 例如:STOP/MAX_TOKENS/SAFETY/RECITATION/MALFORMED_FUNCTION_CALL/...(枚举持续演进)
|
||||
if str(finish_reason) != self.FINISH_REASON_UNSPECIFIED:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def is_error_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为错误事件
|
||||
|
||||
检测多种 Gemini 错误格式:
|
||||
1. 顶层 error: {"error": {...}}
|
||||
2. chunks 内嵌套 error: {"chunks": [{"error": {...}}]}
|
||||
3. candidates 内的错误状态
|
||||
|
||||
Args:
|
||||
event: 事件字典
|
||||
|
||||
Returns:
|
||||
True 如果是错误事件
|
||||
"""
|
||||
# 顶层 error
|
||||
if "error" in event:
|
||||
return True
|
||||
|
||||
# chunks 内嵌套 error (某些 Gemini 响应格式)
|
||||
chunks = event.get("chunks", [])
|
||||
if chunks and isinstance(chunks, list):
|
||||
for chunk in chunks:
|
||||
if isinstance(chunk, dict) and "error" in chunk:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def extract_error_info(self, event: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""
|
||||
从事件中提取错误信息
|
||||
|
||||
Args:
|
||||
event: 事件字典
|
||||
|
||||
Returns:
|
||||
错误信息字典 {"code": int, "message": str, "status": str},无错误返回 None
|
||||
"""
|
||||
# 顶层 error
|
||||
if "error" in event:
|
||||
error = event["error"]
|
||||
if isinstance(error, dict):
|
||||
return {
|
||||
"code": error.get("code"),
|
||||
"message": error.get("message", str(error)),
|
||||
"status": error.get("status"),
|
||||
}
|
||||
return {"code": None, "message": str(error), "status": None}
|
||||
|
||||
# chunks 内嵌套 error
|
||||
chunks = event.get("chunks", [])
|
||||
if chunks and isinstance(chunks, list):
|
||||
for chunk in chunks:
|
||||
if isinstance(chunk, dict) and "error" in chunk:
|
||||
error = chunk["error"]
|
||||
if isinstance(error, dict):
|
||||
return {
|
||||
"code": error.get("code"),
|
||||
"message": error.get("message", str(error)),
|
||||
"status": error.get("status"),
|
||||
}
|
||||
return {"code": None, "message": str(error), "status": None}
|
||||
|
||||
return None
|
||||
|
||||
def get_finish_reason(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
获取结束原因
|
||||
|
||||
Args:
|
||||
event: 事件字典
|
||||
|
||||
Returns:
|
||||
结束原因字符串
|
||||
"""
|
||||
candidates = event.get("candidates", [])
|
||||
if candidates:
|
||||
reason = candidates[0].get("finishReason")
|
||||
return str(reason) if reason is not None else None
|
||||
return None
|
||||
|
||||
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从响应中提取文本内容
|
||||
|
||||
Args:
|
||||
event: 事件字典
|
||||
|
||||
Returns:
|
||||
文本内容,如果没有文本返回 None
|
||||
"""
|
||||
candidates = event.get("candidates", [])
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
content = candidates[0].get("content", {})
|
||||
parts = content.get("parts", [])
|
||||
|
||||
text_parts = []
|
||||
for part in parts:
|
||||
if "text" in part:
|
||||
text_parts.append(part["text"])
|
||||
|
||||
return "".join(text_parts) if text_parts else None
|
||||
|
||||
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
|
||||
"""
|
||||
从事件中提取 token 使用量
|
||||
|
||||
这是 Gemini token 提取的单一实现源,其他地方都应该调用此方法。
|
||||
|
||||
Args:
|
||||
event: 事件字典(包含 usageMetadata)
|
||||
|
||||
Returns:
|
||||
使用量字典,如果没有完整的使用量信息返回 None
|
||||
|
||||
注意:
|
||||
- 只有当 totalTokenCount 存在时才提取(确保是完整的 usage 数据)
|
||||
- 输出 token = thoughtsTokenCount + candidatesTokenCount
|
||||
"""
|
||||
usage_metadata = event.get("usageMetadata", {})
|
||||
if not usage_metadata or "totalTokenCount" not in usage_metadata:
|
||||
return None
|
||||
|
||||
# 输出 token = thoughtsTokenCount + candidatesTokenCount
|
||||
thoughts_tokens = usage_metadata.get("thoughtsTokenCount", 0)
|
||||
candidates_tokens = usage_metadata.get("candidatesTokenCount", 0)
|
||||
output_tokens = thoughts_tokens + candidates_tokens
|
||||
|
||||
return {
|
||||
"input_tokens": usage_metadata.get("promptTokenCount", 0),
|
||||
"output_tokens": output_tokens,
|
||||
"total_tokens": usage_metadata.get("totalTokenCount", 0),
|
||||
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0),
|
||||
}
|
||||
|
||||
def extract_model_version(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从响应中提取模型版本
|
||||
|
||||
Args:
|
||||
event: 事件字典
|
||||
|
||||
Returns:
|
||||
模型版本,如果没有返回 None
|
||||
"""
|
||||
version = event.get("modelVersion")
|
||||
return str(version) if version is not None else None
|
||||
|
||||
def extract_safety_ratings(self, event: dict[str, Any]) -> list[dict[str, Any]] | None:
|
||||
"""
|
||||
从响应中提取安全评级
|
||||
|
||||
Args:
|
||||
event: 事件字典
|
||||
|
||||
Returns:
|
||||
安全评级列表,如果没有返回 None
|
||||
"""
|
||||
candidates = event.get("candidates", [])
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
ratings = candidates[0].get("safetyRatings")
|
||||
if isinstance(ratings, list):
|
||||
return ratings
|
||||
return None
|
||||
|
||||
|
||||
__all__ = ["GeminiStreamParser"]
|
||||
24
_deprecated_py_src/api/handlers/gemini/video_adapter.py
Normal file
24
_deprecated_py_src/api/handlers/gemini/video_adapter.py
Normal file
@@ -0,0 +1,24 @@
|
||||
"""
|
||||
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
|
||||
from src.core.api_format import ApiFamily
|
||||
|
||||
|
||||
class GeminiVeoAdapter(VideoAdapterBase):
|
||||
FORMAT_ID = "gemini:video"
|
||||
API_FAMILY = ApiFamily.GEMINI
|
||||
name = "gemini.video"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> type[VideoHandlerBase]:
|
||||
from src.api.handlers.gemini.video_handler import GeminiVeoHandler
|
||||
|
||||
return GeminiVeoHandler
|
||||
|
||||
|
||||
__all__ = ["GeminiVeoAdapter"]
|
||||
870
_deprecated_py_src/api/handlers/gemini/video_handler.py
Normal file
870
_deprecated_py_src/api/handlers/gemini/video_handler.py
Normal file
@@ -0,0 +1,870 @@
|
||||
"""
|
||||
Gemini Video Handler - Veo 视频生成实现
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
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 (
|
||||
apply_body_rules,
|
||||
evaluate_condition,
|
||||
get_provider_auth,
|
||||
)
|
||||
from src.api.handlers.base.video_handler_base import (
|
||||
VideoHandlerBase,
|
||||
normalize_gemini_operation_id,
|
||||
sanitize_error_message,
|
||||
)
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import (
|
||||
ApiFamily,
|
||||
EndpointKind,
|
||||
build_upstream_headers_for_endpoint,
|
||||
get_extra_headers_from_endpoint,
|
||||
make_signature_key,
|
||||
)
|
||||
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.conversion.registry import format_conversion_registry
|
||||
from src.core.api_format.headers import HOP_BY_HOP_HEADERS
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
from src.core.logger import logger
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
|
||||
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
|
||||
from src.services.provider.provider_context import resolve_provider_proxy
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
class GeminiVeoHandler(VideoHandlerBase):
|
||||
FORMAT_ID = "gemini:video"
|
||||
API_FAMILY = ApiFamily.GEMINI
|
||||
ENDPOINT_KIND = EndpointKind.VIDEO
|
||||
|
||||
DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com"
|
||||
|
||||
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()
|
||||
|
||||
@staticmethod
|
||||
def _get_request_base_url(http_request: Request) -> str:
|
||||
"""从 HTTP 请求中获取基础 URL(协议 + 主机)"""
|
||||
# 优先使用 X-Forwarded-Proto 和 X-Forwarded-Host(代理场景)
|
||||
proto = http_request.headers.get("x-forwarded-proto") or http_request.url.scheme
|
||||
host = http_request.headers.get("x-forwarded-host") or http_request.headers.get("host")
|
||||
if host:
|
||||
return f"{proto}://{host}"
|
||||
# 回退到 request.url
|
||||
return f"{http_request.url.scheme}://{http_request.url.netloc}"
|
||||
|
||||
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))
|
||||
|
||||
# 异步任务:提前创建 pending usage,便于前端看到“处理中”
|
||||
try:
|
||||
UsageService.create_pending_usage(
|
||||
db=self.db,
|
||||
request_id=self.request_id,
|
||||
user=self.user,
|
||||
api_key=self.api_key,
|
||||
model=internal_request.model,
|
||||
is_stream=False,
|
||||
request_type="video",
|
||||
api_format=self.FORMAT_ID,
|
||||
request_headers=original_headers,
|
||||
request_body=original_request_body,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to create pending usage for video request_id={}: {}",
|
||||
self.request_id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
# 用于跟踪是否发生了格式转换
|
||||
format_conversion_info: dict[str, Any] = {
|
||||
"converted": False,
|
||||
"provider_format": None,
|
||||
}
|
||||
|
||||
async def _submit(candidate: ProviderCandidate) -> Any:
|
||||
upstream_key, endpoint, _key, auth_info = await self._resolve_upstream_key(candidate)
|
||||
|
||||
# 检测目标格式
|
||||
provider_format = make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
needs_conversion = provider_format.upper() != self.FORMAT_ID.upper()
|
||||
format_conversion_info["provider_format"] = provider_format
|
||||
format_conversion_info["converted"] = needs_conversion
|
||||
|
||||
# 应用端点的请求体规则
|
||||
endpoint_body_rules = getattr(endpoint, "body_rules", None)
|
||||
|
||||
if needs_conversion and provider_format.upper().startswith("OPENAI:"):
|
||||
# Gemini -> OpenAI 格式转换
|
||||
converted_body = format_conversion_registry.convert_video_request(
|
||||
original_request_body,
|
||||
self.FORMAT_ID,
|
||||
provider_format,
|
||||
)
|
||||
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string)
|
||||
if "seconds" in converted_body and converted_body["seconds"] is not None:
|
||||
converted_body["seconds"] = str(converted_body["seconds"])
|
||||
|
||||
if endpoint_body_rules:
|
||||
converted_body = apply_body_rules(
|
||||
converted_body,
|
||||
endpoint_body_rules,
|
||||
original_body=original_request_body,
|
||||
)
|
||||
|
||||
# 构建 OpenAI 风格的 URL
|
||||
upstream_url = self._build_openai_upstream_url(endpoint.base_url)
|
||||
|
||||
# 构建 OpenAI 风格的请求头
|
||||
headers = self._build_openai_upstream_headers(
|
||||
original_headers,
|
||||
upstream_key,
|
||||
endpoint,
|
||||
body=converted_body,
|
||||
original_body=original_request_body,
|
||||
)
|
||||
|
||||
return await self._try_rust_sync_http_response(
|
||||
method="POST",
|
||||
url=upstream_url,
|
||||
headers=headers,
|
||||
body=converted_body,
|
||||
provider_name=str(candidate.provider.name),
|
||||
provider_id=str(candidate.provider.id),
|
||||
endpoint_id=str(endpoint.id),
|
||||
key_id=str(_key.id),
|
||||
provider_api_format=provider_format,
|
||||
client_api_format=self.FORMAT_ID,
|
||||
model_name=internal_request.model,
|
||||
content_type=str(headers.get("content-type") or "").strip()
|
||||
or "application/json",
|
||||
log_label="GeminiVideoCreate",
|
||||
)
|
||||
else:
|
||||
# 原始 Gemini 格式
|
||||
request_body = (
|
||||
original_request_body.copy() if endpoint_body_rules else original_request_body
|
||||
)
|
||||
if endpoint_body_rules:
|
||||
request_body = apply_body_rules(
|
||||
request_body,
|
||||
endpoint_body_rules,
|
||||
original_body=original_request_body,
|
||||
)
|
||||
|
||||
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
|
||||
headers = self._build_upstream_headers(
|
||||
original_headers,
|
||||
upstream_key,
|
||||
endpoint,
|
||||
auth_info,
|
||||
body=request_body,
|
||||
original_body=original_request_body,
|
||||
)
|
||||
return await self._try_rust_sync_http_response(
|
||||
method="POST",
|
||||
url=upstream_url,
|
||||
headers=headers,
|
||||
body=request_body,
|
||||
provider_name=str(candidate.provider.name),
|
||||
provider_id=str(candidate.provider.id),
|
||||
endpoint_id=str(endpoint.id),
|
||||
key_id=str(_key.id),
|
||||
provider_api_format=provider_format,
|
||||
client_api_format=self.FORMAT_ID,
|
||||
model_name=internal_request.model,
|
||||
content_type=str(headers.get("content-type") or "").strip()
|
||||
or "application/json",
|
||||
log_label="GeminiVideoCreate",
|
||||
)
|
||||
|
||||
def _extract_task_id(payload: dict[str, Any]) -> str | None:
|
||||
# 根据响应格式提取 task ID
|
||||
# Gemini: {"name": "operations/..."}
|
||||
# OpenAI: {"id": "..."}
|
||||
if "name" in payload:
|
||||
value = payload.get("name")
|
||||
logger.debug(
|
||||
"[GeminiVeoHandler] Upstream response name={}, keys={}",
|
||||
value,
|
||||
list(payload.keys()) if isinstance(payload, dict) else type(payload),
|
||||
)
|
||||
if not value:
|
||||
return None
|
||||
return normalize_gemini_operation_id(str(value))
|
||||
if "id" in payload:
|
||||
# OpenAI 格式
|
||||
return str(payload["id"])
|
||||
return None
|
||||
|
||||
outcome_or_response = await self._submit_with_failover(
|
||||
api_format=self.FORMAT_ID,
|
||||
model_name=internal_request.model,
|
||||
task_type="video",
|
||||
submit_func=_submit,
|
||||
extract_external_task_id=_extract_task_id,
|
||||
supported_auth_types={"api_key", "service_account", "vertex_ai"},
|
||||
allow_format_conversion=True,
|
||||
max_candidates=10,
|
||||
)
|
||||
if isinstance(outcome_or_response, JSONResponse):
|
||||
return outcome_or_response
|
||||
outcome = outcome_or_response
|
||||
|
||||
# 冻结 billing_rule 配置(用于异步任务的成本一致性)
|
||||
# 复用 _select_candidate 中已查询的结果;billing_require_rule=false 时需补查
|
||||
rule_lookup = outcome.rule_lookup
|
||||
if rule_lookup is None:
|
||||
rule_lookup = BillingRuleService.find_rule(
|
||||
self.db,
|
||||
provider_id=outcome.candidate.provider.id,
|
||||
model_name=internal_request.model,
|
||||
task_type="video",
|
||||
)
|
||||
billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup)
|
||||
|
||||
external_task_id = outcome.external_task_id
|
||||
|
||||
# 如果发生了格式转换,记录转换后的请求体
|
||||
converted_request_body = original_request_body
|
||||
if format_conversion_info["converted"]:
|
||||
try:
|
||||
converted_request_body = format_conversion_registry.convert_video_request(
|
||||
original_request_body,
|
||||
self.FORMAT_ID,
|
||||
format_conversion_info["provider_format"],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"[GeminiVeoHandler] Failed to record converted request: {}",
|
||||
sanitize_error_message(str(e)),
|
||||
)
|
||||
|
||||
task = self._create_task_record(
|
||||
external_task_id=external_task_id,
|
||||
candidate=outcome.candidate,
|
||||
original_request_body=original_request_body,
|
||||
converted_request_body=converted_request_body,
|
||||
internal_request=internal_request,
|
||||
candidate_keys=outcome.candidate_keys,
|
||||
original_headers=original_headers,
|
||||
billing_rule_snapshot=billing_rule_snapshot,
|
||||
format_converted=format_conversion_info["converted"],
|
||||
)
|
||||
try:
|
||||
self.db.add(task)
|
||||
self.db.flush() # 先 flush 检测冲突
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
logger.debug(
|
||||
"[GeminiVeoHandler] Task created: id={}, external_task_id={}",
|
||||
task.id,
|
||||
task.external_task_id,
|
||||
)
|
||||
except IntegrityError:
|
||||
self.db.rollback()
|
||||
raise HTTPException(status_code=409, detail="Task already exists")
|
||||
|
||||
# 先构建返回给客户端的响应(使用短 ID 对外暴露)
|
||||
internal_task = InternalVideoTask(
|
||||
id=task.short_id,
|
||||
external_id=external_task_id,
|
||||
status=VideoStatus.SUBMITTED,
|
||||
created_at=task.created_at,
|
||||
original_request=internal_request,
|
||||
)
|
||||
base_url = self._get_request_base_url(http_request)
|
||||
response_body = self._normalizer.video_task_from_internal(internal_task, base_url=base_url)
|
||||
|
||||
# 提交成功后补齐 Usage 的 provider 上下文,真正结算留到轮询完成时
|
||||
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||
try:
|
||||
# 构建发送给上游的请求头(脱敏)
|
||||
upstream_request_headers = self._build_upstream_headers(
|
||||
original_headers,
|
||||
"", # key 不重要,只是用于记录
|
||||
outcome.candidate.endpoint,
|
||||
None, # auth_info
|
||||
body=converted_request_body,
|
||||
original_body=original_request_body,
|
||||
)
|
||||
|
||||
UsageService.finalize_submitted(
|
||||
self.db,
|
||||
request_id=self.request_id,
|
||||
provider_name=outcome.candidate.provider.name,
|
||||
provider_id=outcome.candidate.provider.id,
|
||||
provider_endpoint_id=outcome.candidate.endpoint.id,
|
||||
provider_api_key_id=outcome.candidate.key.id,
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=outcome.upstream_status_code or 200,
|
||||
endpoint_api_format=make_signature_key(
|
||||
str(getattr(outcome.candidate.endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(outcome.candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
),
|
||||
provider_request_headers=upstream_request_headers,
|
||||
response_headers=outcome.upstream_headers,
|
||||
response_body=response_body, # 使用我们转换后的响应(包含我们的 ID)
|
||||
)
|
||||
self.db.commit()
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to finalize submitted usage for video request_id={}: {}",
|
||||
self.request_id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
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)
|
||||
base_url = self._get_request_base_url(http_request)
|
||||
response_body = self._normalizer.video_task_from_internal(internal_task, base_url=base_url)
|
||||
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()
|
||||
)
|
||||
base_url = self._get_request_base_url(http_request)
|
||||
items = [
|
||||
self._normalizer.video_task_from_internal(self._task_to_internal(t), base_url=base_url)
|
||||
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:
|
||||
from src.services.task.service import TaskService
|
||||
|
||||
_ = (http_request, query_params, path_params) # reserved for future extensions
|
||||
err_resp = await TaskService(self.db).cancel(
|
||||
task_id,
|
||||
user_id=str(self.user.id),
|
||||
original_headers=original_headers,
|
||||
)
|
||||
if err_resp is not None:
|
||||
return self._build_error_response(err_resp)
|
||||
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")
|
||||
|
||||
# 获取 provider 的认证信息(Gemini 下载视频需要带 API Key)
|
||||
endpoint, key = self._get_endpoint_and_key(task)
|
||||
download_headers: dict[str, str] = {}
|
||||
if key.api_key:
|
||||
try:
|
||||
upstream_key = crypto_service.decrypt(key.api_key)
|
||||
# Gemini API 使用 x-goog-api-key 头进行认证
|
||||
download_headers["x-goog-api-key"] = upstream_key
|
||||
|
||||
# 如果是 Vertex AI,需要使用 OAuth Bearer token
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
if auth_info:
|
||||
download_headers.pop("x-goog-api-key", None)
|
||||
download_headers[auth_info.auth_header] = auth_info.auth_value
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[VideoDownload] Failed to get auth for download task={}: {}",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
# 继续尝试无认证下载(某些 URL 可能是预签名的)
|
||||
|
||||
# 代理下载而非直接重定向,避免暴露上游存储 URL
|
||||
# 使用 httpx 支持重定向(Gemini 视频 URL 会重定向到实际存储位置)
|
||||
return await self._try_rust_download_stream(
|
||||
url=task.video_url,
|
||||
headers=download_headers,
|
||||
task_id=str(task.id),
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
model_name=str(getattr(task, "model", "") or "") or None,
|
||||
)
|
||||
|
||||
async def _try_rust_download_stream(
|
||||
self,
|
||||
*,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
task_id: str,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
model_name: str | None = None,
|
||||
) -> Response | StreamingResponse:
|
||||
import httpx
|
||||
|
||||
from src.services.proxy_node.resolver import (
|
||||
build_proxy_url_async,
|
||||
get_system_proxy_config_async,
|
||||
resolve_delegate_config_async,
|
||||
resolve_effective_proxy,
|
||||
resolve_proxy_info_async,
|
||||
)
|
||||
from src.services.request.execution_runtime_plan import (
|
||||
ExecutionPlan,
|
||||
ExecutionPlanBody,
|
||||
ExecutionPlanTimeouts,
|
||||
ExecutionProxySnapshot,
|
||||
)
|
||||
from src.services.request.execution_runtime_client import (
|
||||
ExecutionRuntimeClient,
|
||||
ExecutionRuntimeClientError,
|
||||
)
|
||||
|
||||
if config.execution_runtime_backend != "rust":
|
||||
raise ProviderNotAvailableException(
|
||||
"Video 下载仅支持 Rust executor",
|
||||
provider_name="gemini",
|
||||
upstream_response=f"executor_backend={config.execution_runtime_backend}",
|
||||
)
|
||||
|
||||
try:
|
||||
effective_proxy = resolve_effective_proxy(
|
||||
resolve_provider_proxy(endpoint=endpoint, key=key),
|
||||
getattr(key, "proxy", None),
|
||||
)
|
||||
if not effective_proxy or not effective_proxy.get("enabled", True):
|
||||
effective_proxy = await get_system_proxy_config_async()
|
||||
|
||||
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
|
||||
proxy_url: str | None = None
|
||||
if effective_proxy and not (delegate_cfg and delegate_cfg.get("tunnel")):
|
||||
proxy_url = await build_proxy_url_async(effective_proxy)
|
||||
|
||||
proxy_info = await resolve_proxy_info_async(effective_proxy)
|
||||
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
|
||||
proxy_info,
|
||||
proxy_url=proxy_url,
|
||||
mode_override="tunnel" if delegate_cfg and delegate_cfg.get("tunnel") else None,
|
||||
node_id_override=(
|
||||
str(delegate_cfg.get("node_id") or "").strip() or None
|
||||
if delegate_cfg and delegate_cfg.get("tunnel")
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
plan = ExecutionPlan(
|
||||
request_id=str(self.request_id or ""),
|
||||
candidate_id=None,
|
||||
provider_name="gemini",
|
||||
provider_id=str(getattr(endpoint, "provider_id", "") or ""),
|
||||
endpoint_id=str(getattr(endpoint, "id", "") or ""),
|
||||
key_id=str(getattr(key, "id", "") or ""),
|
||||
method="GET",
|
||||
url=url,
|
||||
headers=dict(headers),
|
||||
body=ExecutionPlanBody(),
|
||||
stream=True,
|
||||
provider_api_format=self.FORMAT_ID,
|
||||
client_api_format=self.FORMAT_ID,
|
||||
model_name=str(model_name or ""),
|
||||
proxy=proxy_snapshot,
|
||||
timeouts=ExecutionPlanTimeouts(
|
||||
connect_ms=30_000,
|
||||
read_ms=300_000,
|
||||
write_ms=300_000,
|
||||
pool_ms=30_000,
|
||||
total_ms=None,
|
||||
),
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[VideoDownload] Rust plan build failed task={} url={}: {}",
|
||||
task_id,
|
||||
url,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
raise ProviderNotAvailableException(
|
||||
"Rust executor 请求计划构建失败",
|
||||
provider_name="gemini",
|
||||
upstream_response=sanitize_error_message(str(exc)),
|
||||
) from exc
|
||||
|
||||
try:
|
||||
rust_stream = await ExecutionRuntimeClient().execute_stream(plan)
|
||||
except (ExecutionRuntimeClientError, httpx.HTTPError, ValueError) as exc:
|
||||
logger.warning(
|
||||
"[VideoDownload] Rust executor unavailable task={} url={}: {}",
|
||||
task_id,
|
||||
url,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
raise ProviderNotAvailableException(
|
||||
"执行器暂时不可用,请稍后重试",
|
||||
provider_name="gemini",
|
||||
upstream_response=sanitize_error_message(str(exc)),
|
||||
) from exc
|
||||
|
||||
safe_headers = {
|
||||
k: v for k, v in rust_stream.headers.items() if k.lower() not in HOP_BY_HOP_HEADERS
|
||||
}
|
||||
|
||||
if rust_stream.status_code >= 400:
|
||||
try:
|
||||
async for _ in rust_stream.byte_iterator:
|
||||
pass
|
||||
finally:
|
||||
await rust_stream.response_ctx.__aexit__(None, None, None)
|
||||
raise HTTPException(status_code=rust_stream.status_code, detail="Upstream error")
|
||||
|
||||
async def _iter_bytes() -> AsyncIterator[bytes]:
|
||||
try:
|
||||
async for chunk in rust_stream.byte_iterator:
|
||||
yield chunk
|
||||
finally:
|
||||
await rust_stream.response_ctx.__aexit__(None, None, None)
|
||||
|
||||
return StreamingResponse(
|
||||
_iter_bytes(),
|
||||
status_code=rust_stream.status_code,
|
||||
headers=safe_headers,
|
||||
media_type=safe_headers.get("content-type", "video/mp4"),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
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={}: {}",
|
||||
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,
|
||||
*,
|
||||
body: dict[str, Any] | None = None,
|
||||
original_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, str]:
|
||||
extra_headers = get_extra_headers_from_endpoint(endpoint)
|
||||
endpoint_sig = make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
headers = build_upstream_headers_for_endpoint(
|
||||
original_headers,
|
||||
endpoint_sig,
|
||||
upstream_key,
|
||||
endpoint_headers=extra_headers,
|
||||
header_rules=getattr(endpoint, "header_rules", None),
|
||||
body=body,
|
||||
original_body=original_body,
|
||||
condition_evaluator=evaluate_condition,
|
||||
)
|
||||
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"),
|
||||
}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# OpenAI format conversion helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _build_openai_upstream_url(self, base_url: str | None) -> str:
|
||||
"""构建 OpenAI Sora API 的上游 URL"""
|
||||
base = (base_url or "https://api.openai.com").rstrip("/")
|
||||
if base.endswith("/v1"):
|
||||
return f"{base}/videos"
|
||||
return f"{base}/v1/videos"
|
||||
|
||||
def _build_openai_upstream_headers(
|
||||
self,
|
||||
original_headers: dict[str, str],
|
||||
upstream_key: str,
|
||||
endpoint: ProviderEndpoint,
|
||||
*,
|
||||
body: dict[str, Any] | None = None,
|
||||
original_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""构建 OpenAI 格式的请求头"""
|
||||
extra_headers = get_extra_headers_from_endpoint(endpoint)
|
||||
endpoint_sig = make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
return build_upstream_headers_for_endpoint(
|
||||
original_headers,
|
||||
endpoint_sig,
|
||||
upstream_key,
|
||||
endpoint_headers=extra_headers,
|
||||
header_rules=getattr(endpoint, "header_rules", None),
|
||||
body=body,
|
||||
original_body=original_body,
|
||||
condition_evaluator=evaluate_condition,
|
||||
)
|
||||
|
||||
def _create_task_record(
|
||||
self,
|
||||
*,
|
||||
external_task_id: str,
|
||||
candidate: ProviderCandidate,
|
||||
original_request_body: dict[str, Any],
|
||||
internal_request: Any,
|
||||
candidate_keys: list[dict[str, Any]] | None = None,
|
||||
original_headers: dict[str, str] | None = None,
|
||||
billing_rule_snapshot: dict[str, Any] | None = None,
|
||||
converted_request_body: dict[str, Any] | None = None,
|
||||
format_converted: bool = False,
|
||||
) -> VideoTask:
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# 构建请求元数据(使用追踪信息)
|
||||
request_metadata = {
|
||||
"candidate_keys": candidate_keys or [],
|
||||
"selected_key_id": candidate.key.id,
|
||||
"selected_endpoint_id": candidate.endpoint.id,
|
||||
"client_ip": self.client_ip,
|
||||
"user_agent": self.user_agent,
|
||||
"request_id": self.request_id,
|
||||
"billing_rule_snapshot": billing_rule_snapshot,
|
||||
}
|
||||
# 记录请求头(脱敏处理)
|
||||
if original_headers:
|
||||
safe_headers = {
|
||||
k: v
|
||||
for k, v in original_headers.items()
|
||||
if k.lower() not in {"authorization", "x-api-key", "x-goog-api-key", "cookie"}
|
||||
}
|
||||
request_metadata["request_headers"] = safe_headers
|
||||
|
||||
provider_api_format = make_signature_key(
|
||||
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
|
||||
return VideoTask(
|
||||
id=str(uuid4()),
|
||||
request_id=self.request_id,
|
||||
external_task_id=external_task_id,
|
||||
user_id=self.user.id,
|
||||
api_key_id=self.api_key.id,
|
||||
username=self.user.username,
|
||||
api_key_name=self.api_key.name,
|
||||
provider_id=candidate.provider.id,
|
||||
endpoint_id=candidate.endpoint.id,
|
||||
key_id=candidate.key.id,
|
||||
client_api_format=self.FORMAT_ID,
|
||||
provider_api_format=provider_api_format,
|
||||
format_converted=format_converted,
|
||||
model=internal_request.model,
|
||||
prompt=internal_request.prompt,
|
||||
original_request_body=original_request_body,
|
||||
converted_request_body=converted_request_body or 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=config.video_poll_interval_seconds,
|
||||
next_poll_at=now + timedelta(seconds=config.video_poll_interval_seconds),
|
||||
poll_count=0,
|
||||
max_poll_count=config.video_max_poll_count,
|
||||
submitted_at=now,
|
||||
request_metadata=request_metadata,
|
||||
)
|
||||
|
||||
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:
|
||||
"""覆盖父类方法,Gemini 使用 short_id 作为对外暴露的 ID"""
|
||||
try:
|
||||
status = VideoStatus(task.status)
|
||||
except ValueError:
|
||||
status = VideoStatus.PENDING
|
||||
return InternalVideoTask(
|
||||
id=task.short_id, # Gemini 使用短 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},
|
||||
)
|
||||
|
||||
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
|
||||
"""按 short_id 查找任务(我们对外暴露的 operation 格式是 models/{model}/operations/{short_id})"""
|
||||
from src.api.handlers.base.video_handler_base import extract_short_id_from_operation
|
||||
|
||||
short_id = extract_short_id_from_operation(external_id)
|
||||
|
||||
# 通过 short_id 查找任务
|
||||
task = (
|
||||
self.db.query(VideoTask)
|
||||
.filter(
|
||||
VideoTask.short_id == short_id,
|
||||
VideoTask.user_id == self.user.id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if not task:
|
||||
logger.debug("[GeminiVeoHandler] Task not found: short_id={}", short_id)
|
||||
raise HTTPException(status_code=404, detail="Video task not found")
|
||||
return task
|
||||
|
||||
|
||||
__all__ = ["GeminiVeoHandler"]
|
||||
Reference in New Issue
Block a user