mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
refactor: 重构异步任务系统和计费服务架构
- 重构任务系统:新增 lifecycle (TaskStatus/BillingStatus)、context、application 模块 - 将 video tasks 泛化为 async tasks,支持更通用的异步任务管理 - 新增 Gemini Files 管理模块和管理界面 - 重构 billing 服务:拆分 schema.py 和 service.py - 新增 candidate 服务模块用于请求候选管理 - 数据库迁移:添加 billing_status、request_id、gemini_file_mappings 表和索引 - 移除废弃的 video_telemetry、task orchestrator 等模块
This commit is contained in:
@@ -79,10 +79,8 @@ class VideoAdapterBase(ApiAdapter):
|
||||
path_params=path_params,
|
||||
)
|
||||
|
||||
# Cancel task
|
||||
if method in {"DELETE", "POST"} and (
|
||||
path.endswith("/cancel") or path_params.get("action") == "cancel"
|
||||
):
|
||||
# Cancel task (POST /videos/{id}/cancel or explicit action=cancel)
|
||||
if (method == "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"
|
||||
@@ -95,6 +93,16 @@ class VideoAdapterBase(ApiAdapter):
|
||||
path_params=path_params,
|
||||
)
|
||||
|
||||
# Delete task (DELETE /videos/{id})
|
||||
if method == "DELETE" and task_id:
|
||||
return await handler.handle_delete_task(
|
||||
task_id=task_id,
|
||||
http_request=http_request,
|
||||
original_headers=context.original_headers,
|
||||
query_params=context.query_params,
|
||||
path_params=path_params,
|
||||
)
|
||||
|
||||
# Remix task
|
||||
if method == "POST" and path.endswith("/remix") and task_id:
|
||||
return await handler.handle_remix_task(
|
||||
|
||||
@@ -26,7 +26,7 @@ from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
if TYPE_CHECKING:
|
||||
import httpx
|
||||
|
||||
from src.services.task.orchestrator import SubmitOutcome
|
||||
from src.services.candidate.submit import SubmitOutcome
|
||||
|
||||
# 敏感信息匹配正则(预编译提升性能)
|
||||
_SENSITIVE_PATTERN = re.compile(
|
||||
@@ -53,21 +53,39 @@ def sanitize_error_message(message: str, max_length: int = 200) -> str:
|
||||
return sanitized[:max_length]
|
||||
|
||||
|
||||
def extract_short_id_from_operation(operation_id: str) -> str:
|
||||
"""
|
||||
从 operation ID 中提取短 ID
|
||||
|
||||
我们对外暴露的 operation name 格式是:
|
||||
- models/{model}/operations/{short_id}
|
||||
|
||||
此函数提取最后一部分作为 short_id,用于在数据库中查找任务。
|
||||
|
||||
Args:
|
||||
operation_id: 原始 operation ID(如 "models/veo-3.1/operations/abc123")
|
||||
|
||||
Returns:
|
||||
short_id(如 "abc123")
|
||||
"""
|
||||
# 格式: models/{model}/operations/{short_id}
|
||||
# 或者直接是 short_id
|
||||
if "/" in operation_id:
|
||||
# 提取最后一部分
|
||||
return operation_id.rsplit("/", 1)[-1]
|
||||
return operation_id
|
||||
|
||||
|
||||
def normalize_gemini_operation_id(operation_id: str) -> str:
|
||||
"""
|
||||
规范化 Gemini operation ID,确保以 "operations/" 开头
|
||||
|
||||
Gemini API 返回的任务 ID 格式可能是 "operations/xxx" 或 "xxx",
|
||||
此函数统一规范化为 "operations/xxx" 格式。
|
||||
规范化 Gemini operation ID(保留用于向后兼容)
|
||||
|
||||
Args:
|
||||
operation_id: 原始 operation ID
|
||||
|
||||
Returns:
|
||||
规范化后的 operation ID
|
||||
规范化后的 operation ID(原样返回)
|
||||
"""
|
||||
if not operation_id.startswith("operations/"):
|
||||
return f"operations/{operation_id}"
|
||||
return operation_id
|
||||
|
||||
|
||||
@@ -143,6 +161,18 @@ class VideoHandlerBase(ABC):
|
||||
) -> JSONResponse:
|
||||
"""取消任务"""
|
||||
|
||||
async def handle_delete_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:
|
||||
"""删除已完成或失败的视频任务 - 可选实现"""
|
||||
raise HTTPException(status_code=501, detail="Delete not supported for this provider")
|
||||
|
||||
async def handle_remix_task(
|
||||
self,
|
||||
*,
|
||||
@@ -205,6 +235,7 @@ class VideoHandlerBase(ABC):
|
||||
}
|
||||
|
||||
def _get_task(self, task_id: str) -> VideoTask:
|
||||
"""通过 UUID 查找任务(OpenAI Sora 风格)"""
|
||||
task = (
|
||||
self.db.query(VideoTask)
|
||||
.filter(VideoTask.id == task_id, VideoTask.user_id == self.user.id)
|
||||
@@ -229,7 +260,7 @@ class VideoHandlerBase(ABC):
|
||||
except ValueError:
|
||||
status = VideoStatus.PENDING
|
||||
return InternalVideoTask(
|
||||
id=task.id,
|
||||
id=task.id, # OpenAI Sora 使用 UUID
|
||||
external_id=task.external_task_id,
|
||||
status=status,
|
||||
progress_percent=task.progress_percent or 0,
|
||||
@@ -243,6 +274,52 @@ class VideoHandlerBase(ABC):
|
||||
extra={"model": task.model},
|
||||
)
|
||||
|
||||
def _finalize_usage_on_submit_failure(
|
||||
self,
|
||||
candidate_keys: list[dict[str, Any]],
|
||||
status_code: int | None,
|
||||
) -> None:
|
||||
"""
|
||||
提交失败时结算 pending usage(避免遗留 pending 状态)。
|
||||
|
||||
从 candidate_keys 中提取最后尝试的 provider 信息,更新 Usage 记录。
|
||||
"""
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
# 提取 provider 信息:优先取最后一个有 attempt 的候选
|
||||
provider_name = "unknown"
|
||||
provider_id = None
|
||||
endpoint_id = None
|
||||
key_id = None
|
||||
|
||||
for ck in reversed(candidate_keys):
|
||||
if ck.get("attempt_status") or ck.get("selected"):
|
||||
provider_name = ck.get("provider_name") or "unknown"
|
||||
provider_id = ck.get("provider_id")
|
||||
endpoint_id = ck.get("endpoint_id")
|
||||
key_id = ck.get("key_id")
|
||||
break
|
||||
|
||||
try:
|
||||
# 更新 usage 状态并设置 provider 信息
|
||||
UsageService.update_usage_status(
|
||||
self.db,
|
||||
request_id=self.request_id,
|
||||
status="failed",
|
||||
error_message=f"submit_failed (status_code={status_code or 'unknown'})",
|
||||
provider=provider_name,
|
||||
provider_id=provider_id,
|
||||
provider_endpoint_id=endpoint_id,
|
||||
provider_api_key_id=key_id,
|
||||
status_code=status_code,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to finalize usage on submit failure: request_id=%s, error=%s",
|
||||
self.request_id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
def _build_billing_rule_snapshot(
|
||||
self, rule_lookup: BillingRuleLookupResult | None
|
||||
) -> dict[str, Any]:
|
||||
@@ -290,16 +367,16 @@ class VideoHandlerBase(ABC):
|
||||
- 无可用候选 / 全部失败:抛 HTTPException(503)
|
||||
"""
|
||||
# 延迟导入,避免 handler 基类层引入过多依赖导致循环
|
||||
from src.services.task.orchestrator import (
|
||||
from src.services.candidate.service import CandidateService
|
||||
from src.services.candidate.submit import (
|
||||
AllCandidatesFailedError,
|
||||
AsyncTaskOrchestrator,
|
||||
SubmitOutcome,
|
||||
UpstreamClientRequestError,
|
||||
)
|
||||
|
||||
orchestrator = AsyncTaskOrchestrator(self.db)
|
||||
candidate_service = CandidateService(self.db)
|
||||
try:
|
||||
return await orchestrator.submit_with_failover(
|
||||
return await candidate_service.submit_with_failover(
|
||||
api_format=api_format,
|
||||
model_name=model_name,
|
||||
affinity_key=str(self.api_key.id),
|
||||
@@ -314,8 +391,12 @@ class VideoHandlerBase(ABC):
|
||||
max_candidates=max_candidates,
|
||||
)
|
||||
except UpstreamClientRequestError as exc:
|
||||
# 将 pending usage 结算为 failed,并记录 provider 信息
|
||||
self._finalize_usage_on_submit_failure(exc.candidate_keys, exc.response.status_code)
|
||||
return self._build_error_response(exc.response)
|
||||
except AllCandidatesFailedError as exc:
|
||||
# 将 pending usage 结算为 failed
|
||||
self._finalize_usage_on_submit_failure(exc.candidate_keys, exc.last_status_code)
|
||||
detail = "No available provider for video generation"
|
||||
if config.billing_require_rule:
|
||||
detail = "No available provider with billing rule for video generation"
|
||||
|
||||
@@ -18,7 +18,6 @@ from src.core.api_format import ApiFamily, 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
|
||||
from src.services.provider.transport import redact_url_for_log
|
||||
|
||||
|
||||
@@ -63,11 +62,9 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
def detect_capability_requirements(
|
||||
self,
|
||||
headers: dict[str, str], # noqa: ARG002 - 预留
|
||||
request_body: dict[str, Any] | None = None,
|
||||
request_body: dict[str, Any] | None = None, # noqa: ARG002 - 预留
|
||||
) -> dict[str, bool]:
|
||||
"""检测是否需要 Gemini Files API 能力"""
|
||||
if request_body and extract_file_names_from_request(request_body):
|
||||
return {"gemini_files_api": True}
|
||||
"""Gemini API 无特殊能力要求"""
|
||||
return {}
|
||||
|
||||
def _merge_path_params(
|
||||
|
||||
@@ -40,28 +40,31 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
Gemini 文件与上传它的 API Key 绑定,必须使用同一 Key 访问。
|
||||
此方法从缓存中查找文件→Key 映射,优先使用正确的 Key。
|
||||
|
||||
当同一源文件被上传到多个 Key 时,会返回所有可用的 Key ID,
|
||||
让系统能够选择任意可用的 Key。
|
||||
|
||||
注意事项:
|
||||
- 如果映射缺失(缓存过期/重启),会记录警告,请求可能失败
|
||||
- 如果多个文件属于不同 Key,只能使用其中一个,其他文件可能无法访问
|
||||
- 优先返回所有支持该文件的 Key,让调度器选择可用的
|
||||
"""
|
||||
from src.core.logger import logger
|
||||
from src.services.gemini_files_mapping import (
|
||||
extract_file_names_from_request,
|
||||
get_file_key_mapping,
|
||||
get_all_key_ids_for_file,
|
||||
)
|
||||
|
||||
file_names = extract_file_names_from_request(request_body or {})
|
||||
if not file_names:
|
||||
return None
|
||||
|
||||
preferred_key_ids: list[str] = []
|
||||
unmapped_files: list[str] = [] # 记录找不到映射的文件
|
||||
all_key_ids: set[str] = set()
|
||||
unmapped_files: list[str] = []
|
||||
|
||||
for file_name in file_names:
|
||||
key_id = await get_file_key_mapping(file_name)
|
||||
if key_id:
|
||||
if key_id not in preferred_key_ids:
|
||||
preferred_key_ids.append(key_id)
|
||||
# 获取所有支持该文件的 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)
|
||||
|
||||
@@ -72,14 +75,10 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
"请求可能失败(文件属于其他 Key 或映射已过期)"
|
||||
)
|
||||
|
||||
# 警告:多个文件属于不同 Key
|
||||
if len(preferred_key_ids) > 1:
|
||||
logger.warning(
|
||||
f"[{self.request_id}] 请求使用了多个文件,但它们属于不同的 Key: "
|
||||
f"{preferred_key_ids},只能使用第一个 Key,其他文件可能无法访问"
|
||||
)
|
||||
if all_key_ids:
|
||||
logger.debug(f"[{self.request_id}] 文件引用可用的 Key: {list(all_key_ids)}")
|
||||
|
||||
return preferred_key_ids or None
|
||||
return list(all_key_ids) if all_key_ids else None
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
|
||||
@@ -4,6 +4,7 @@ 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
|
||||
@@ -34,12 +35,14 @@ from src.core.api_format.conversion.internal_video import (
|
||||
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.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.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
class GeminiVeoHandler(VideoHandlerBase):
|
||||
@@ -92,20 +95,93 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
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=%s: %s",
|
||||
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)
|
||||
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
|
||||
headers = self._build_upstream_headers(
|
||||
original_headers, upstream_key, endpoint, auth_info
|
||||
|
||||
# 检测目标格式
|
||||
provider_format = make_signature_key(
|
||||
str(getattr(endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
)
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
return await client.post(upstream_url, headers=headers, json=original_request_body)
|
||||
needs_conversion = provider_format.upper() != self.FORMAT_ID.upper()
|
||||
format_conversion_info["provider_format"] = provider_format
|
||||
format_conversion_info["converted"] = needs_conversion
|
||||
|
||||
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"])
|
||||
|
||||
# 构建 OpenAI 风格的 URL
|
||||
upstream_url = self._build_openai_upstream_url(endpoint.base_url)
|
||||
|
||||
# 构建 OpenAI 风格的请求头
|
||||
headers = self._build_openai_upstream_headers(
|
||||
original_headers, upstream_key, endpoint
|
||||
)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
return await client.post(upstream_url, headers=headers, json=converted_body)
|
||||
else:
|
||||
# 原始 Gemini 格式
|
||||
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()
|
||||
return await client.post(upstream_url, headers=headers, json=original_request_body)
|
||||
|
||||
def _extract_task_id(payload: dict[str, Any]) -> str | None:
|
||||
value = payload.get("name")
|
||||
if not value:
|
||||
return None
|
||||
return normalize_gemini_operation_id(str(value))
|
||||
# 根据响应格式提取 task ID
|
||||
# Gemini: {"name": "operations/..."}
|
||||
# OpenAI: {"id": "..."}
|
||||
if "name" in payload:
|
||||
value = payload.get("name")
|
||||
logger.debug(
|
||||
"[GeminiVeoHandler] Upstream response name=%s, keys=%s",
|
||||
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,
|
||||
@@ -114,7 +190,7 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
submit_func=_submit,
|
||||
extract_external_task_id=_extract_task_id,
|
||||
supported_auth_types={"api_key", "vertex_ai"},
|
||||
allow_format_conversion=False,
|
||||
allow_format_conversion=True,
|
||||
max_candidates=10,
|
||||
)
|
||||
if isinstance(outcome_or_response, JSONResponse):
|
||||
@@ -135,35 +211,92 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
|
||||
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: %s",
|
||||
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.info(
|
||||
f"[GeminiVeoHandler] Task created: id={task.id}, external_task_id={task.external_task_id}, user_id={task.user_id}"
|
||||
logger.debug(
|
||||
"[GeminiVeoHandler] Task created: id=%s, external_task_id=%s",
|
||||
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.id,
|
||||
id=task.short_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)
|
||||
|
||||
# 提交成功后立即结算 Usage(费用暂时为 0,轮询完成后更新)
|
||||
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
|
||||
)
|
||||
|
||||
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=%s: %s",
|
||||
self.request_id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
return JSONResponse(response_body)
|
||||
|
||||
async def handle_get_task(
|
||||
@@ -234,6 +367,28 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
|
||||
task.status = VideoStatus.CANCELLED.value
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
# 将 Usage 作废(不收费)
|
||||
# 尝试 finalize_void(处理 pending)和 void_settled(处理已 settled)
|
||||
try:
|
||||
voided = UsageService.finalize_void(
|
||||
self.db,
|
||||
request_id=task.request_id,
|
||||
reason="cancelled_by_user",
|
||||
)
|
||||
if not voided:
|
||||
# pending 状态未找到,尝试处理已 settled 的记录
|
||||
UsageService.void_settled(
|
||||
self.db,
|
||||
request_id=task.request_id,
|
||||
reason="cancelled_by_user",
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to void usage for cancelled task=%s: %s",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
self.db.commit()
|
||||
return JSONResponse({})
|
||||
|
||||
@@ -275,12 +430,38 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
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=%s: %s",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
# 继续尝试无认证下载(某些 URL 可能是预签名的)
|
||||
|
||||
# 代理下载而非直接重定向,避免暴露上游存储 URL
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
# 使用 httpx 支持重定向(Gemini 视频 URL 会重定向到实际存储位置)
|
||||
import httpx
|
||||
|
||||
try:
|
||||
request = client.build_request("GET", task.video_url)
|
||||
# 视频下载可能较大,设置 5 分钟超时
|
||||
response = await client.send(request, stream=True, timeout=300.0)
|
||||
# 使用 follow_redirects=True 跟随重定向
|
||||
async with httpx.AsyncClient(
|
||||
follow_redirects=True, timeout=httpx.Timeout(300.0)
|
||||
) as client:
|
||||
response = await client.get(task.video_url, headers=download_headers)
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
"[VideoDownload] Upstream fetch failed user=%s task=%s: %s",
|
||||
@@ -291,21 +472,14 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
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()
|
||||
|
||||
# 返回完整的视频内容(非 streaming,因为需要跟随重定向)
|
||||
safe_headers = {
|
||||
k: v for k, v in response.headers.items() if k.lower() not in HOP_BY_HOP_HEADERS
|
||||
}
|
||||
return StreamingResponse(
|
||||
_iter_bytes(),
|
||||
return Response(
|
||||
content=response.content,
|
||||
status_code=response.status_code,
|
||||
headers=safe_headers,
|
||||
media_type=response.headers.get("content-type", "video/mp4"),
|
||||
@@ -375,6 +549,36 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
"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,
|
||||
) -> 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,
|
||||
)
|
||||
|
||||
def _create_task_record(
|
||||
self,
|
||||
*,
|
||||
@@ -385,6 +589,8 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
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)
|
||||
|
||||
@@ -407,8 +613,14 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
}
|
||||
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,
|
||||
@@ -416,15 +628,12 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
endpoint_id=candidate.endpoint.id,
|
||||
key_id=candidate.key.id,
|
||||
client_api_format=self.FORMAT_ID,
|
||||
provider_api_format=make_signature_key(
|
||||
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
),
|
||||
format_converted=False,
|
||||
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=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,
|
||||
@@ -438,30 +647,45 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
request_metadata=request_metadata,
|
||||
)
|
||||
|
||||
def _get_task_by_external_id(self, external_id: str) -> VideoTask:
|
||||
"""按 external_task_id 查找任务(Gemini 使用 operations/{id} 格式)"""
|
||||
normalized_id = normalize_gemini_operation_id(external_id)
|
||||
|
||||
logger.info(
|
||||
f"[GeminiVeoHandler] Looking for task: normalized_id={normalized_id}, user_id={self.user.id}"
|
||||
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.external_task_id == normalized_id,
|
||||
VideoTask.short_id == short_id,
|
||||
VideoTask.user_id == self.user.id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if not task:
|
||||
logger.warning(
|
||||
f"[GeminiVeoHandler] Task not found: normalized_id={normalized_id}, user_id={self.user.id}"
|
||||
)
|
||||
logger.debug("[GeminiVeoHandler] Task not found: short_id=%s", short_id)
|
||||
raise HTTPException(status_code=404, detail="Video task not found")
|
||||
logger.info(
|
||||
f"[GeminiVeoHandler] Task found: id={task.id}, external_task_id={task.external_task_id}"
|
||||
)
|
||||
return task
|
||||
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ OpenAI Video Handler - Sora 视频生成实现
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, AsyncIterator
|
||||
from uuid import uuid4
|
||||
@@ -14,6 +15,7 @@ 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.config.settings import config
|
||||
@@ -30,6 +32,7 @@ from src.core.api_format.conversion.internal_video import (
|
||||
VideoStatus,
|
||||
)
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
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.logger import logger
|
||||
@@ -91,16 +94,91 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
)
|
||||
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=%s: %s",
|
||||
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, _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()
|
||||
return await client.post(upstream_url, headers=headers, json=original_request_body)
|
||||
|
||||
# 检测目标格式
|
||||
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
|
||||
|
||||
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string)
|
||||
request_body = original_request_body.copy()
|
||||
if "seconds" in request_body and request_body["seconds"] is not None:
|
||||
request_body["seconds"] = str(request_body["seconds"])
|
||||
|
||||
if needs_conversion and provider_format.upper().startswith("GEMINI:"):
|
||||
# OpenAI -> Gemini 格式转换
|
||||
converted_body = format_conversion_registry.convert_video_request(
|
||||
request_body,
|
||||
self.FORMAT_ID,
|
||||
provider_format,
|
||||
)
|
||||
# 如果 model 不在请求体中,从路径或内部请求中获取
|
||||
if "model" not in converted_body:
|
||||
converted_body["model"] = internal_request.model
|
||||
|
||||
# 构建 Gemini 风格的 URL
|
||||
upstream_url = self._build_gemini_upstream_url(
|
||||
endpoint.base_url, internal_request.model
|
||||
)
|
||||
|
||||
# 构建 Gemini 风格的请求头
|
||||
auth_info = await get_provider_auth(endpoint, _provider_key)
|
||||
headers = self._build_gemini_upstream_headers(
|
||||
original_headers, upstream_key, endpoint, auth_info
|
||||
)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
return await client.post(upstream_url, headers=headers, json=converted_body)
|
||||
else:
|
||||
# 原始 OpenAI 格式
|
||||
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()
|
||||
return await client.post(upstream_url, headers=headers, json=request_body)
|
||||
|
||||
def _extract_task_id(payload: dict[str, Any]) -> str | None:
|
||||
value = payload.get("id")
|
||||
return str(value) if value else None
|
||||
# 根据响应格式提取 task ID
|
||||
# OpenAI: {"id": "..."}
|
||||
# Gemini: {"name": "operations/..."}
|
||||
if "id" in payload:
|
||||
return str(payload["id"])
|
||||
if "name" in payload:
|
||||
# Gemini 格式
|
||||
return str(payload["name"])
|
||||
return None
|
||||
|
||||
# 捕获提交阶段的所有错误,记录失败任务
|
||||
try:
|
||||
@@ -110,8 +188,8 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
task_type="video",
|
||||
submit_func=_submit,
|
||||
extract_external_task_id=_extract_task_id,
|
||||
supported_auth_types={"api_key"},
|
||||
allow_format_conversion=False,
|
||||
supported_auth_types={"api_key", "vertex_ai"},
|
||||
allow_format_conversion=True,
|
||||
max_candidates=10,
|
||||
)
|
||||
except HTTPException as exc:
|
||||
@@ -156,14 +234,31 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
|
||||
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(
|
||||
"[OpenAIVideoHandler] Failed to record converted request: %s",
|
||||
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)
|
||||
@@ -174,6 +269,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
self.db.rollback()
|
||||
raise HTTPException(status_code=409, detail="Task already exists")
|
||||
|
||||
# 先构建返回给客户端的响应(OpenAI Sora 使用 UUID)
|
||||
internal_task = InternalVideoTask(
|
||||
id=task.id,
|
||||
external_id=external_task_id,
|
||||
@@ -182,6 +278,42 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
original_request=internal_request,
|
||||
)
|
||||
response_body = self._normalizer.video_task_from_internal(internal_task)
|
||||
|
||||
# 提交成功后立即结算 Usage(费用暂时为 0,轮询完成后更新)
|
||||
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||
try:
|
||||
# 构建发送给上游的请求头(脱敏)
|
||||
upstream_request_headers = self._build_upstream_headers(
|
||||
original_headers,
|
||||
"", # key 不重要,只是用于记录
|
||||
outcome.candidate.endpoint,
|
||||
)
|
||||
|
||||
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=%s: %s",
|
||||
self.request_id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
|
||||
return JSONResponse(response_body)
|
||||
|
||||
async def handle_get_task(
|
||||
@@ -206,17 +338,59 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
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()
|
||||
)
|
||||
params = query_params or {}
|
||||
|
||||
# 解析分页参数
|
||||
after = params.get("after")
|
||||
try:
|
||||
limit = min(int(params.get("limit") or 20), 100) # 默认 20,最大 100
|
||||
except (ValueError, TypeError):
|
||||
limit = 20
|
||||
order = params.get("order", "desc").lower()
|
||||
if order not in ("asc", "desc"):
|
||||
order = "desc"
|
||||
|
||||
# 构建查询
|
||||
query = self.db.query(VideoTask).filter(VideoTask.user_id == self.user.id)
|
||||
|
||||
# 处理游标分页(after 参数使用 UUID)
|
||||
if after:
|
||||
after_task = (
|
||||
self.db.query(VideoTask)
|
||||
.filter(VideoTask.id == after, VideoTask.user_id == self.user.id)
|
||||
.first()
|
||||
)
|
||||
if after_task and after_task.created_at:
|
||||
if order == "desc":
|
||||
query = query.filter(VideoTask.created_at < after_task.created_at)
|
||||
else:
|
||||
query = query.filter(VideoTask.created_at > after_task.created_at)
|
||||
|
||||
# 排序
|
||||
if order == "asc":
|
||||
query = query.order_by(VideoTask.created_at.asc())
|
||||
else:
|
||||
query = query.order_by(VideoTask.created_at.desc())
|
||||
|
||||
# 获取 limit + 1 条记录以判断是否有更多数据
|
||||
tasks = query.limit(limit + 1).all()
|
||||
has_more = len(tasks) > limit
|
||||
tasks = tasks[:limit]
|
||||
|
||||
items = [
|
||||
self._normalizer.video_task_from_internal(self._task_to_internal(t)) for t in tasks
|
||||
]
|
||||
return JSONResponse({"object": "list", "data": items})
|
||||
|
||||
response_data: dict[str, Any] = {
|
||||
"object": "list",
|
||||
"data": items,
|
||||
"has_more": has_more,
|
||||
}
|
||||
# 如果有更多数据,返回最后一条的 ID 作为下一页游标
|
||||
if has_more and tasks:
|
||||
response_data["last_id"] = tasks[-1].id
|
||||
|
||||
return JSONResponse(response_data)
|
||||
|
||||
async def handle_cancel_task(
|
||||
self,
|
||||
@@ -245,9 +419,80 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
|
||||
task.status = VideoStatus.CANCELLED.value
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
# 将 Usage 作废(不收费)
|
||||
# 尝试 finalize_void(处理 pending)和 void_settled(处理已 settled)
|
||||
try:
|
||||
voided = UsageService.finalize_void(
|
||||
self.db,
|
||||
request_id=task.request_id,
|
||||
reason="cancelled_by_user",
|
||||
)
|
||||
if not voided:
|
||||
# pending 状态未找到,尝试处理已 settled 的记录
|
||||
UsageService.void_settled(
|
||||
self.db,
|
||||
request_id=task.request_id,
|
||||
reason="cancelled_by_user",
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to void usage for cancelled task=%s: %s",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
self.db.commit()
|
||||
return JSONResponse({})
|
||||
|
||||
async def handle_delete_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 task.status not in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Can only delete completed or failed videos (current status: {task.status})",
|
||||
)
|
||||
|
||||
# 如果有 external_task_id,向上游发送删除请求
|
||||
if task.external_task_id:
|
||||
try:
|
||||
endpoint, key = self._get_endpoint_and_key(task)
|
||||
if key.api_key:
|
||||
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 and response.status_code != 404:
|
||||
# 404 表示上游已删除,不算错误
|
||||
return self._build_error_response(response)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Failed to delete video from upstream task=%s: %s",
|
||||
task.id,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
# 继续删除本地记录
|
||||
|
||||
# 删除本地任务记录
|
||||
self.db.delete(task)
|
||||
self.db.commit()
|
||||
|
||||
return JSONResponse({"id": task_id, "object": "video", "deleted": True})
|
||||
|
||||
async def handle_remix_task(
|
||||
self,
|
||||
*,
|
||||
@@ -280,8 +525,13 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
)
|
||||
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
||||
|
||||
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string)
|
||||
request_body = original_request_body.copy()
|
||||
if "seconds" in request_body and request_body["seconds"] is not None:
|
||||
request_body["seconds"] = str(request_body["seconds"])
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.post(upstream_url, headers=headers, json=original_request_body)
|
||||
response = await client.post(upstream_url, headers=headers, json=request_body)
|
||||
|
||||
if response.status_code >= 400:
|
||||
return self._build_error_response(response)
|
||||
@@ -350,7 +600,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
raise HTTPException(status_code=409, detail="Task already exists")
|
||||
|
||||
internal_task = InternalVideoTask(
|
||||
id=task.id,
|
||||
id=task.id, # OpenAI Sora 使用 UUID
|
||||
external_id=external_task_id,
|
||||
status=VideoStatus.SUBMITTED,
|
||||
created_at=task.created_at,
|
||||
@@ -389,15 +639,25 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
if task.status == VideoStatus.CANCELLED.value:
|
||||
raise HTTPException(status_code=404, detail="Video task was cancelled")
|
||||
|
||||
# 支持 variant 查询参数: video (默认), thumbnail, spritesheet
|
||||
variant = (query_params or {}).get("variant", "video")
|
||||
|
||||
# 如果 video_url 是完整的 HTTP URL,直接代理该 URL(适用于不支持 /content 端点的上游如 API易)
|
||||
# 保持流式代理而非重定向,确保客户端行为与官方 OpenAI 一致
|
||||
if variant == "video" and task.video_url and task.video_url.startswith("http"):
|
||||
logger.debug(
|
||||
"[VideoDownload] Proxying direct URL task=%s url=%s",
|
||||
task_id,
|
||||
task.video_url,
|
||||
)
|
||||
return await self._proxy_direct_url(task.video_url, 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)
|
||||
|
||||
# 支持 variant 查询参数: video (默认), thumbnail, spritesheet
|
||||
variant = (query_params or {}).get("variant", "video")
|
||||
if variant not in {"video", "thumbnail", "spritesheet"}:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
@@ -412,6 +672,12 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
logger.debug(
|
||||
"[VideoDownload] Requesting upstream url=%s task=%s external_task_id=%s",
|
||||
upstream_url,
|
||||
task_id,
|
||||
task.external_task_id,
|
||||
)
|
||||
try:
|
||||
# 使用 httpx 的 stream 方法并正确管理上下文
|
||||
# 视频下载可能较大,设置 5 分钟超时
|
||||
@@ -419,8 +685,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
response = await client.send(request, stream=True, timeout=300.0)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[VideoDownload] Upstream connection failed task=%s: %s",
|
||||
"[VideoDownload] Upstream connection failed task=%s url=%s: %s",
|
||||
task_id,
|
||||
upstream_url,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
raise HTTPException(status_code=502, detail="Upstream connection failed") from exc
|
||||
@@ -483,6 +750,46 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
|
||||
return upstream_key, candidate.endpoint, candidate.key
|
||||
|
||||
async def _proxy_direct_url(self, url: str, task_id: str) -> Response | StreamingResponse:
|
||||
"""代理直接的视频 URL(如 CDN URL),保持与官方 API 一致的流式返回行为"""
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
try:
|
||||
request = client.build_request("GET", url)
|
||||
response = await client.send(request, stream=True, timeout=300.0)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"[VideoDownload] Direct URL connection failed task=%s url=%s: %s",
|
||||
task_id,
|
||||
url,
|
||||
sanitize_error_message(str(exc)),
|
||||
)
|
||||
raise HTTPException(status_code=502, detail="Video download failed") from exc
|
||||
|
||||
if response.status_code >= 400:
|
||||
await response.aread() # consume body before closing
|
||||
await response.aclose()
|
||||
return JSONResponse(
|
||||
status_code=response.status_code,
|
||||
content={"error": {"type": "upstream_error", "message": "Video not available"}},
|
||||
)
|
||||
|
||||
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"),
|
||||
)
|
||||
|
||||
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"):
|
||||
@@ -508,6 +815,42 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
endpoint_headers=extra_headers,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Gemini format conversion helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _build_gemini_upstream_url(self, base_url: str | None, model: str) -> str:
|
||||
"""构建 Gemini Veo API 的上游 URL"""
|
||||
base = (base_url or "https://generativelanguage.googleapis.com").rstrip("/")
|
||||
if base.endswith("/v1beta"):
|
||||
base = base[: -len("/v1beta")]
|
||||
return f"{base}/v1beta/models/{model}:predictLongRunning"
|
||||
|
||||
def _build_gemini_upstream_headers(
|
||||
self,
|
||||
original_headers: dict[str, str],
|
||||
upstream_key: str,
|
||||
endpoint: ProviderEndpoint,
|
||||
auth_info: Any | None,
|
||||
) -> dict[str, str]:
|
||||
"""构建 Gemini 格式的请求头"""
|
||||
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,
|
||||
)
|
||||
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
|
||||
|
||||
# _build_error_response 继承自基类 VideoHandlerBase
|
||||
|
||||
def _create_task_record(
|
||||
@@ -520,6 +863,8 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
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)
|
||||
size = internal_request.extra.get("original_size")
|
||||
@@ -543,8 +888,14 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
}
|
||||
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,
|
||||
@@ -552,15 +903,12 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
endpoint_id=candidate.endpoint.id,
|
||||
key_id=candidate.key.id,
|
||||
client_api_format=self.FORMAT_ID,
|
||||
provider_api_format=make_signature_key(
|
||||
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
|
||||
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
|
||||
),
|
||||
format_converted=False,
|
||||
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=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,
|
||||
@@ -580,8 +928,23 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
status = VideoStatus(task.status)
|
||||
except ValueError:
|
||||
status = VideoStatus.PENDING
|
||||
|
||||
# 构建 extra 字段
|
||||
extra: dict[str, Any] = {
|
||||
"model": task.model,
|
||||
"size": task.size,
|
||||
"seconds": str(task.duration_seconds) if task.duration_seconds else None,
|
||||
"prompt": task.prompt,
|
||||
}
|
||||
|
||||
# 检查是否是 remix 视频
|
||||
if task.original_request_body and isinstance(task.original_request_body, dict):
|
||||
remixed_from = task.original_request_body.get("remix_video_id")
|
||||
if remixed_from:
|
||||
extra["remixed_from_video_id"] = remixed_from
|
||||
|
||||
return InternalVideoTask(
|
||||
id=task.id,
|
||||
id=task.id, # OpenAI Sora 使用 UUID
|
||||
external_id=task.external_task_id,
|
||||
status=status,
|
||||
progress_percent=task.progress_percent or 0,
|
||||
@@ -596,7 +959,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
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},
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
async def _record_failed_usage(
|
||||
@@ -609,8 +972,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
original_headers: dict[str, str],
|
||||
) -> None:
|
||||
"""记录失败请求的使用记录(无任务记录)"""
|
||||
import time
|
||||
|
||||
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||
safe_headers = {
|
||||
k: v
|
||||
@@ -669,8 +1030,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
candidate_keys: list[dict[str, Any]] | None = None,
|
||||
) -> None:
|
||||
"""创建失败的任务记录和使用记录"""
|
||||
import time
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
response_time_ms = int((time.time() - self.start_time) * 1000)
|
||||
|
||||
@@ -696,6 +1055,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
# 创建失败的任务记录
|
||||
task = VideoTask(
|
||||
id=str(uuid4()),
|
||||
request_id=self.request_id,
|
||||
external_task_id=None,
|
||||
user_id=self.user.id,
|
||||
api_key_id=self.api_key.id,
|
||||
|
||||
Reference in New Issue
Block a user