feat: 添加视频生成功能增强和多维度计费系统适配

- 视频生成: 增强 video_handler,重构 task_poller,新增 telemetry
- 计费系统: 适配新的 signature 格式,支持 video 任务类型回退
- 数据库迁移: 添加 api_family/endpoint_kind 字段和 video_formats
This commit is contained in:
fawney19
2026-02-01 17:28:27 +08:00
parent 7b66505634
commit 4ac8e63c94
15 changed files with 1514 additions and 527 deletions

View File

@@ -0,0 +1,570 @@
"""Add api_family/endpoint_kind and migrate api_format to endpoint signature keys
Revision ID: cf40e6a5c5b1
Revises: c8d2e4f6a1b3
Create Date: 2026-01-31 15:30:00.000000
"""
from __future__ import annotations
import json
from typing import Sequence, Union
from uuid import uuid4
import sqlalchemy as sa
from sqlalchemy import inspect, text
from alembic import op
# revision identifiers, used by Alembic.
revision: str = "cf40e6a5c5b1"
down_revision: Union[str, None] = "c8d2e4f6a1b3"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def _json_loads(val):
if val is None:
return None
if isinstance(val, (dict, list)):
return val
if isinstance(val, str):
try:
return json.loads(val)
except Exception:
return None
return None
def _json_dumps(val):
"""将 dict/list 转为 JSON 字符串None 保持 None"""
if val is None:
return None
if isinstance(val, str):
return val
return json.dumps(val)
def _normalize_signature(value: str | None) -> str | None:
"""
Normalize legacy api_format / signature-ish strings to canonical signature key.
- canonical: `<family>:<kind>` (lowercase)
- legacy examples: "OPENAI", "OPENAI_CLI", "GEMINI_VIDEO"
"""
if value is None:
return None
raw = str(value).strip()
if not raw:
return None
if ":" in raw:
fam, kind = raw.split(":", 1)
fam = fam.strip().lower()
kind = kind.strip().lower()
if not fam or not kind:
return None
return f"{fam}:{kind}"
upper = raw.upper()
if upper.startswith("CLAUDE"):
fam = "claude"
elif upper.startswith("OPENAI"):
fam = "openai"
elif upper.startswith("GEMINI"):
fam = "gemini"
else:
return None
kind = "chat"
if upper.endswith("_CLI"):
kind = "cli"
elif upper.endswith("_VIDEO"):
kind = "video"
return f"{fam}:{kind}"
def _normalize_signature_list(values) -> list[str] | None:
if values is None:
return None
if isinstance(values, str):
values = _json_loads(values)
if not isinstance(values, list):
return None
out: list[str] = []
seen: set[str] = set()
for v in values:
sig = _normalize_signature(str(v) if v is not None else None)
if not sig:
continue
if sig in seen:
continue
seen.add(sig)
out.append(sig)
return out
def _normalize_signature_dict(values) -> dict | None:
if values is None:
return None
if isinstance(values, str):
values = _json_loads(values)
if not isinstance(values, dict):
return None
out: dict = {}
for k, v in values.items():
sig = _normalize_signature(str(k) if k is not None else None)
if not sig:
continue
out[sig] = v
return out
def _add_video_variants(formats: list[str]) -> list[str]:
"""
迁移策略:如果 key/限制里包含 openai/gemini 的 chat/cli则自动补齐 video 变体。
这是为了兼容旧数据:历史上 video 复用了 chat 的 api_format。
"""
if any(f.startswith("openai:") for f in formats) and "openai:video" not in formats:
formats.append("openai:video")
if any(f.startswith("gemini:") for f in formats) and "gemini:video" not in formats:
formats.append("gemini:video")
return formats
def table_exists(table_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
return table_name in inspector.get_table_names()
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [col["name"] for col in inspector.get_columns(table_name)]
return column_name in columns
def index_exists(table_name: str, index_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
try:
indexes = inspector.get_indexes(table_name)
except Exception:
return False
return any(idx.get("name") == index_name for idx in indexes)
def _migrate_format_acceptance_config(cfg) -> dict | None:
cfg_obj = _json_loads(cfg)
if not isinstance(cfg_obj, dict):
return cfg_obj if cfg_obj is None else None
for key in ("accept_formats", "reject_formats"):
raw = cfg_obj.get(key)
if not isinstance(raw, list):
continue
normalized = _normalize_signature_list(raw) or []
cfg_obj[key] = normalized
return cfg_obj
def migrate_provider_endpoints(connection) -> None:
"""
- 将 provider_endpoints.api_format 统一迁移为 signature key小写
- 填充/校准 api_family / endpoint_kind
- 迁移 format_acceptance_config 中的 accept/reject formats
"""
rows = connection.execute(text("""
SELECT
id,
api_format,
api_family,
endpoint_kind,
format_acceptance_config
FROM provider_endpoints
""")).fetchall()
for row in rows:
sig = _normalize_signature(row.api_format)
if not sig:
continue
fam, kind = sig.split(":", 1)
cfg = _migrate_format_acceptance_config(row.format_acceptance_config)
connection.execute(
text("""
UPDATE provider_endpoints
SET
api_format = :api_format,
api_family = :api_family,
endpoint_kind = :endpoint_kind,
format_acceptance_config = CAST(:format_acceptance_config AS json)
WHERE id = :id
"""),
{
"id": row.id,
"api_format": sig,
"api_family": fam,
"endpoint_kind": kind,
"format_acceptance_config": _json_dumps(cfg),
},
)
def create_video_endpoints(connection) -> None:
"""
为已有 openai:chat / gemini:chat endpoint 的 provider 创建对应的 *:video endpoint。
重要custom_path 必须置空。源 endpoint 的 custom_path 大概率是 Chat 路径,
复制过去会导致 Video 请求发到错误路径;置空后走 *:video 的默认路径。
"""
result = connection.execute(text("""
SELECT
e1.provider_id,
e1.api_family,
e1.base_url,
e1.is_active,
e1.header_rules,
e1.max_retries,
e1.config,
e1.format_acceptance_config,
e1.proxy,
e1.api_format
FROM provider_endpoints e1
WHERE e1.api_format IN ('openai:chat', 'gemini:chat')
AND NOT EXISTS (
SELECT 1 FROM provider_endpoints e2
WHERE e2.provider_id = e1.provider_id
AND e2.api_format = CASE
WHEN e1.api_format = 'openai:chat' THEN 'openai:video'
ELSE 'gemini:video'
END
)
"""))
for row in result:
base_format = str(row.api_format or "").strip().lower()
if base_format == "openai:chat":
new_format = "openai:video"
new_family = "openai"
elif base_format == "gemini:chat":
new_format = "gemini:video"
new_family = "gemini"
else:
continue
connection.execute(
text("""
INSERT INTO provider_endpoints
(
id,
provider_id,
api_format,
api_family,
endpoint_kind,
base_url,
is_active,
header_rules,
max_retries,
custom_path,
config,
format_acceptance_config,
proxy
)
VALUES
(
:id,
:provider_id,
:api_format,
:api_family,
'video',
:base_url,
:is_active,
:header_rules,
:max_retries,
NULL,
:config,
:format_acceptance_config,
:proxy
)
"""),
{
"id": str(uuid4()),
"provider_id": row.provider_id,
"api_format": new_format,
"api_family": new_family,
"base_url": row.base_url,
"is_active": row.is_active,
"header_rules": _json_dumps(_json_loads(row.header_rules)),
"max_retries": row.max_retries,
"config": _json_dumps(_json_loads(row.config)),
"format_acceptance_config": _json_dumps(_json_loads(row.format_acceptance_config)),
"proxy": _json_dumps(_json_loads(row.proxy)),
},
)
def migrate_provider_api_keys(connection) -> None:
"""
迁移 provider_api_keys:
- api_formats -> signature keys并补齐 video 变体)
- dict 字段 key -> signature keysrate_multipliers/global_priority/health/circuit_breaker
- rate_multipliers/global_priority_by_format 复制 chat -> video如 openai:chat -> openai:video
"""
rows = connection.execute(text("""
SELECT
id,
api_formats,
rate_multipliers,
global_priority_by_format,
health_by_format,
circuit_breaker_by_format
FROM provider_api_keys
""")).fetchall()
for row in rows:
api_formats = _normalize_signature_list(row.api_formats)
if api_formats is not None:
api_formats = _add_video_variants(api_formats)
rate_multipliers = _normalize_signature_dict(row.rate_multipliers)
if isinstance(rate_multipliers, dict):
if "openai:chat" in rate_multipliers and "openai:video" not in rate_multipliers:
rate_multipliers["openai:video"] = rate_multipliers["openai:chat"]
if "gemini:chat" in rate_multipliers and "gemini:video" not in rate_multipliers:
rate_multipliers["gemini:video"] = rate_multipliers["gemini:chat"]
global_priority_by_format = _normalize_signature_dict(row.global_priority_by_format)
if isinstance(global_priority_by_format, dict):
if (
"openai:chat" in global_priority_by_format
and "openai:video" not in global_priority_by_format
):
global_priority_by_format["openai:video"] = global_priority_by_format["openai:chat"]
if (
"gemini:chat" in global_priority_by_format
and "gemini:video" not in global_priority_by_format
):
global_priority_by_format["gemini:video"] = global_priority_by_format["gemini:chat"]
health_by_format = _normalize_signature_dict(row.health_by_format)
circuit_breaker_by_format = _normalize_signature_dict(row.circuit_breaker_by_format)
connection.execute(
text("""
UPDATE provider_api_keys
SET
api_formats = CAST(:api_formats AS json),
rate_multipliers = CAST(:rate_multipliers AS json),
global_priority_by_format = CAST(:global_priority_by_format AS json),
health_by_format = CAST(:health_by_format AS json),
circuit_breaker_by_format = CAST(:circuit_breaker_by_format AS json)
WHERE id = :id
"""),
{
"id": row.id,
"api_formats": _json_dumps(api_formats),
"rate_multipliers": _json_dumps(rate_multipliers),
"global_priority_by_format": _json_dumps(global_priority_by_format),
"health_by_format": _json_dumps(health_by_format),
"circuit_breaker_by_format": _json_dumps(circuit_breaker_by_format),
},
)
def migrate_allowed_api_formats(connection, *, table_name: str) -> None:
"""迁移 users/api_keys.allowed_api_formats 为 signature keys并补齐 video 变体)。"""
if not table_exists(table_name):
return
rows = connection.execute(text(f"""
SELECT id, allowed_api_formats
FROM {table_name}
""")).fetchall()
for row in rows:
allowed = _normalize_signature_list(row.allowed_api_formats)
if allowed is None:
continue
allowed = _add_video_variants(allowed)
connection.execute(
text(f"""
UPDATE {table_name}
SET allowed_api_formats = CAST(:allowed_api_formats AS json)
WHERE id = :id
"""),
{"id": row.id, "allowed_api_formats": _json_dumps(allowed)},
)
def migrate_video_tasks(connection) -> None:
"""
video_tasks.*_api_format 迁移为 signature keys。
注意video_tasks 表天然是 video 任务,因此将 openai/gemini 的 kind 强制归一为 video
以兼容历史上复用 chat 格式存储的旧记录。
"""
if not table_exists("video_tasks"):
return
rows = connection.execute(text("""
SELECT id, client_api_format, provider_api_format
FROM video_tasks
""")).fetchall()
for row in rows:
client_sig = _normalize_signature(row.client_api_format) or ""
provider_sig = _normalize_signature(row.provider_api_format) or ""
def _force_video(sig: str) -> str:
if not sig or ":" not in sig:
return sig
fam, _kind = sig.split(":", 1)
fam = fam.strip().lower()
if fam in ("openai", "gemini"):
return f"{fam}:video"
return sig
client_sig = _force_video(client_sig)
provider_sig = _force_video(provider_sig)
if not client_sig or not provider_sig:
continue
connection.execute(
text("""
UPDATE video_tasks
SET client_api_format = :client_api_format,
provider_api_format = :provider_api_format
WHERE id = :id
"""),
{
"id": row.id,
"client_api_format": client_sig,
"provider_api_format": provider_sig,
},
)
def migrate_model_provider_mappings(connection) -> None:
"""迁移 models.provider_model_mappings[*].api_formats 为 signature keys。"""
if not table_exists("models"):
return
rows = connection.execute(text("""
SELECT id, provider_model_mappings
FROM models
WHERE provider_model_mappings IS NOT NULL
""")).fetchall()
for row in rows:
mappings = _json_loads(row.provider_model_mappings)
if not isinstance(mappings, list):
continue
changed = False
new_mappings: list = []
for item in mappings:
if not isinstance(item, dict):
new_mappings.append(item)
continue
raw_formats = item.get("api_formats")
if isinstance(raw_formats, list):
normalized = _normalize_signature_list(raw_formats) or []
# 内容比较(而非引用比较),避免已迁移数据被无意义地重复 UPDATE
if set(normalized) != set(raw_formats):
changed = True
item = dict(item)
item["api_formats"] = normalized
new_mappings.append(item)
if not changed:
continue
connection.execute(
text("""
UPDATE models
SET provider_model_mappings = CAST(:provider_model_mappings AS json)
WHERE id = :id
"""),
{"id": row.id, "provider_model_mappings": _json_dumps(new_mappings)},
)
def migrate_dimension_collectors(connection) -> None:
"""迁移 dimension_collectors.api_format 为 signature keys如果存在历史数据"""
if not table_exists("dimension_collectors"):
return
rows = connection.execute(text("""
SELECT id, api_format
FROM dimension_collectors
WHERE api_format IS NOT NULL
""")).fetchall()
for row in rows:
sig = _normalize_signature(row.api_format)
if not sig:
continue
connection.execute(
text("""
UPDATE dimension_collectors
SET api_format = :api_format
WHERE id = :id
"""),
{"id": row.id, "api_format": sig},
)
def upgrade() -> None:
if not table_exists("provider_endpoints"):
return
# ==================== provider_endpoints.api_family / endpoint_kind ====================
if not column_exists("provider_endpoints", "api_family"):
op.add_column("provider_endpoints", sa.Column("api_family", sa.String(50), nullable=True))
if not column_exists("provider_endpoints", "endpoint_kind"):
op.add_column(
"provider_endpoints", sa.Column("endpoint_kind", sa.String(50), nullable=True)
)
# ==================== idx_provider_family_kind ====================
if not index_exists("provider_endpoints", "idx_provider_family_kind"):
op.create_index(
"idx_provider_family_kind",
"provider_endpoints",
["provider_id", "api_family", "endpoint_kind"],
)
# ==================== data migrations (idempotent) ====================
conn = op.get_bind()
migrate_provider_endpoints(conn)
create_video_endpoints(conn)
if table_exists("provider_api_keys"):
migrate_provider_api_keys(conn)
migrate_allowed_api_formats(conn, table_name="users")
migrate_allowed_api_formats(conn, table_name="api_keys")
migrate_video_tasks(conn)
migrate_model_provider_mappings(conn)
migrate_dimension_collectors(conn)
def downgrade() -> None:
# Drop index/columns only; data changes are intentionally kept (safe rollback strategy).
if table_exists("provider_endpoints"):
if index_exists("provider_endpoints", "idx_provider_family_kind"):
op.drop_index("idx_provider_family_kind", table_name="provider_endpoints")
if column_exists("provider_endpoints", "endpoint_kind"):
op.drop_column("provider_endpoints", "endpoint_kind")
if column_exists("provider_endpoints", "api_family"):
op.drop_column("provider_endpoints", "api_family")

View File

@@ -6,7 +6,7 @@ Video Adapter 通用基类
from __future__ import annotations from __future__ import annotations
from typing import Any from typing import Any, ClassVar
from fastapi import HTTPException, Request from fastapi import HTTPException, Request
from fastapi.responses import Response from fastapi.responses import Response
@@ -14,7 +14,12 @@ from fastapi.responses import Response
from src.api.base.adapter import ApiAdapter, ApiMode from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext from src.api.base.context import ApiRequestContext
from src.api.handlers.base.video_handler_base import VideoHandlerBase from src.api.handlers.base.video_handler_base import VideoHandlerBase
from src.core.api_format import APIFormat, get_auth_handler, get_default_auth_method from src.core.api_format import (
ApiFamily,
EndpointKind,
get_auth_handler,
get_default_auth_method_for_endpoint,
)
from src.core.logger import logger from src.core.logger import logger
@@ -24,21 +29,18 @@ class VideoAdapterBase(ApiAdapter):
FORMAT_ID: str = "UNKNOWN" FORMAT_ID: str = "UNKNOWN"
HANDLER_CLASS: type[VideoHandlerBase] HANDLER_CLASS: type[VideoHandlerBase]
# 新架构:结构化标识(逐步替代直接依赖 FORMAT_ID 的语义)
API_FAMILY: ClassVar[ApiFamily | None] = None
ENDPOINT_KIND: ClassVar[EndpointKind] = EndpointKind.VIDEO
name: str = "video.base" name: str = "video.base"
mode = ApiMode.STANDARD mode = ApiMode.STANDARD
def __init__(self, allowed_api_formats: list[str] | None = None): def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID] self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
@classmethod
def _get_api_format(cls) -> APIFormat:
try:
return APIFormat[cls.FORMAT_ID]
except KeyError:
return APIFormat.OPENAI
def extract_api_key(self, request: Request) -> str | None: def extract_api_key(self, request: Request) -> str | None:
auth_method = get_default_auth_method(self._get_api_format()) auth_method = get_default_auth_method_for_endpoint(self.FORMAT_ID)
handler = get_auth_handler(auth_method) handler = get_auth_handler(auth_method)
return handler.extract_credentials(request) return handler.extract_credentials(request)
@@ -93,6 +95,17 @@ class VideoAdapterBase(ApiAdapter):
path_params=path_params, path_params=path_params,
) )
# Remix task
if method == "POST" and path.endswith("/remix") and task_id:
return await handler.handle_remix_task(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
original_request_body=original_request_body,
query_params=context.query_params,
path_params=path_params,
)
# Get task # Get task
if method == "GET" and task_id: if method == "GET" and task_id:
return await handler.handle_get_task( return await handler.handle_get_task(

View File

@@ -8,19 +8,26 @@ from __future__ import annotations
import re import re
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from fastapi import HTTPException, Request from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, Response, StreamingResponse from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.api_format.conversion.internal_video import InternalVideoTask, VideoStatus from src.core.api_format.conversion.internal_video import InternalVideoTask, VideoStatus
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult from src.services.billing.rule_service import BillingRuleLookupResult
from src.services.cache.aware_scheduler import ProviderCandidate
if TYPE_CHECKING: if TYPE_CHECKING:
import httpx import httpx
from src.services.task.orchestrator import SubmitOutcome
# 敏感信息匹配正则(预编译提升性能) # 敏感信息匹配正则(预编译提升性能)
_SENSITIVE_PATTERN = re.compile( _SENSITIVE_PATTERN = re.compile(
r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+", r"(api[_-]?key|token|bearer|authorization)[=:\s]+\S+",
@@ -136,6 +143,19 @@ class VideoHandlerBase(ABC):
) -> JSONResponse: ) -> JSONResponse:
"""取消任务""" """取消任务"""
async def handle_remix_task(
self,
*,
task_id: str,
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:
"""Remix 任务(基于已完成视频创建新视频)- 可选实现"""
raise HTTPException(status_code=501, detail="Remix not supported for this provider")
@abstractmethod @abstractmethod
async def handle_download_content( async def handle_download_content(
self, self,
@@ -164,7 +184,7 @@ class VideoHandlerBase(ABC):
status_code=response.status_code, status_code=response.status_code,
content={"error": payload}, content={"error": payload},
) )
except ValueError, KeyError, TypeError: except (ValueError, KeyError, TypeError):
pass pass
message = sanitize_error_message(response.text or "Upstream error") message = sanitize_error_message(response.text or "Upstream error")
fallback_payload = self._format_error_payload({"message": message}, response.status_code) fallback_payload = self._format_error_payload({"message": message}, response.status_code)
@@ -246,5 +266,74 @@ class VideoHandlerBase(ABC):
"dimension_mappings": rule.dimension_mappings, "dimension_mappings": rule.dimension_mappings,
} }
async def _submit_with_failover(
self,
*,
api_format: str,
model_name: str,
task_type: str,
submit_func: Callable[[ProviderCandidate], Awaitable["httpx.Response"]],
extract_external_task_id: Callable[[dict[str, Any]], str | None],
supported_auth_types: set[str] | None,
allow_format_conversion: bool = False,
capability_requirements: dict[str, bool] | None = None,
max_candidates: int = 10,
) -> "SubmitOutcome | JSONResponse":
"""
提交阶段故障转移(只负责拿到 external_task_id
返回:
- 成功SubmitOutcome
- 上游客户端错误:直接返回脱敏后的 JSONResponse保留 API 格式差异)
失败时:
- 无可用候选 / 全部失败:抛 HTTPException(503)
"""
# 延迟导入,避免 handler 基类层引入过多依赖导致循环
from src.services.task.orchestrator import (
AllCandidatesFailedError,
AsyncTaskOrchestrator,
SubmitOutcome,
UpstreamClientRequestError,
)
orchestrator = AsyncTaskOrchestrator(self.db)
try:
return await orchestrator.submit_with_failover(
api_format=api_format,
model_name=model_name,
affinity_key=str(self.api_key.id),
user_api_key=self.api_key,
request_id=self.request_id,
task_type=task_type,
submit_func=submit_func,
extract_external_task_id=extract_external_task_id,
supported_auth_types=supported_auth_types,
allow_format_conversion=allow_format_conversion,
capability_requirements=capability_requirements,
max_candidates=max_candidates,
)
except UpstreamClientRequestError as exc:
return self._build_error_response(exc.response)
except AllCandidatesFailedError as exc:
detail = "No available provider for video generation"
if config.billing_require_rule:
detail = "No available provider with billing rule for video generation"
# 记录候选信息到日志
logger.warning(
"[VideoHandler] All candidates failed: reason=%s, candidate_keys=%s",
exc.reason,
exc.candidate_keys,
)
# 创建带有 candidate_keys 的 HTTPException
http_exc = HTTPException(status_code=503, detail=detail)
http_exc.candidate_keys = exc.candidate_keys # type: ignore[attr-defined]
raise http_exc
except ProviderNotAvailableException:
detail = "No available provider for video generation"
if config.billing_require_rule:
detail = "No available provider with billing rule for video generation"
raise HTTPException(status_code=503, detail=detail)
__all__ = ["VideoHandlerBase", "normalize_gemini_operation_id", "sanitize_error_message"] __all__ = ["VideoHandlerBase", "normalize_gemini_operation_id", "sanitize_error_message"]

View File

@@ -6,10 +6,12 @@ from __future__ import annotations
from src.api.handlers.base.video_adapter_base import VideoAdapterBase from src.api.handlers.base.video_adapter_base import VideoAdapterBase
from src.api.handlers.base.video_handler_base import VideoHandlerBase from src.api.handlers.base.video_handler_base import VideoHandlerBase
from src.core.api_format import ApiFamily
class GeminiVeoAdapter(VideoAdapterBase): class GeminiVeoAdapter(VideoAdapterBase):
FORMAT_ID = "GEMINI" FORMAT_ID = "gemini:video"
API_FAMILY = ApiFamily.GEMINI
name = "gemini.video" name = "gemini.video"
@property @property

View File

@@ -21,7 +21,13 @@ from src.api.handlers.base.video_handler_base import (
) )
from src.clients.http_client import HTTPClientPool from src.clients.http_client import HTTPClientPool
from src.config.settings import config from src.config.settings import config
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint 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 ( from src.core.api_format.conversion.internal_video import (
InternalVideoRequest, InternalVideoRequest,
InternalVideoTask, InternalVideoTask,
@@ -33,11 +39,13 @@ from src.core.crypto import crypto_service
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate from src.services.cache.aware_scheduler import ProviderCandidate
class GeminiVeoHandler(VideoHandlerBase): class GeminiVeoHandler(VideoHandlerBase):
FORMAT_ID = "GEMINI" FORMAT_ID = "gemini:video"
API_FAMILY = ApiFamily.GEMINI
ENDPOINT_KIND = EndpointKind.VIDEO
DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com" DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com"
@@ -84,56 +92,55 @@ class GeminiVeoHandler(VideoHandlerBase):
except ValueError as e: except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) raise HTTPException(status_code=400, detail=str(e))
candidate, candidate_keys, rule_lookup = await self._select_candidate( async def _submit(candidate: ProviderCandidate) -> Any:
internal_request.model upstream_key, endpoint, _key, auth_info = await self._resolve_upstream_key(candidate)
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
headers = self._build_upstream_headers(
original_headers, upstream_key, endpoint, auth_info
)
client = await HTTPClientPool.get_default_client_async()
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))
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", "vertex_ai"},
allow_format_conversion=False,
max_candidates=10,
) )
if not candidate: if isinstance(outcome_or_response, JSONResponse):
detail = "No available provider for video generation" return outcome_or_response
if config.billing_require_rule: outcome = outcome_or_response
detail = "No available provider with billing rule for video generation"
raise HTTPException(status_code=503, detail=detail)
# 冻结 billing_rule 配置(用于异步任务的成本一致性) # 冻结 billing_rule 配置(用于异步任务的成本一致性)
# 复用 _select_candidate 中已查询的结果billing_require_rule=false 时需补查 # 复用 _select_candidate 中已查询的结果billing_require_rule=false 时需补查
rule_lookup = outcome.rule_lookup
if rule_lookup is None: if rule_lookup is None:
rule_lookup = BillingRuleService.find_rule( rule_lookup = BillingRuleService.find_rule(
self.db, self.db,
provider_id=candidate.provider.id, provider_id=outcome.candidate.provider.id,
model_name=internal_request.model, model_name=internal_request.model,
task_type="video", task_type="video",
) )
billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup) billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup)
upstream_key, endpoint, key, auth_info = await self._resolve_upstream_key(candidate) external_task_id = outcome.external_task_id
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint, auth_info)
api_key_header = (
headers.get("x-goog-api-key", "")[:10] + "..."
if headers.get("x-goog-api-key")
else "MISSING"
)
logger.info(
f"[GeminiVeoHandler] Create task: endpoint_id={endpoint.id}, base_url={endpoint.base_url}, upstream_url={upstream_url}, api_key_prefix={api_key_header}"
)
client = await HTTPClientPool.get_default_client_async()
response = await client.post(upstream_url, headers=headers, json=original_request_body)
if response.status_code >= 400:
return self._build_error_response(response)
payload = response.json()
external_task_id = str(payload.get("name") or "")
if not external_task_id:
raise HTTPException(status_code=502, detail="Upstream returned empty task id")
external_task_id = normalize_gemini_operation_id(external_task_id)
task = self._create_task_record( task = self._create_task_record(
external_task_id=external_task_id, external_task_id=external_task_id,
candidate=candidate, candidate=outcome.candidate,
original_request_body=original_request_body, original_request_body=original_request_body,
internal_request=internal_request, internal_request=internal_request,
candidate_keys=candidate_keys, candidate_keys=outcome.candidate_keys,
original_headers=original_headers, original_headers=original_headers,
billing_rule_snapshot=billing_rule_snapshot, billing_rule_snapshot=billing_rule_snapshot,
) )
@@ -308,61 +315,6 @@ class GeminiVeoHandler(VideoHandlerBase):
# Helpers # Helpers
# ------------------------------------------------------------------ # ------------------------------------------------------------------
async def _select_candidate(
self, model_name: str
) -> tuple[ProviderCandidate | None, list[dict[str, Any]], BillingRuleLookupResult | None]:
"""选择候选 key返回 (选中的候选, 所有候选列表, 选中候选的 billing rule lookup)"""
scheduler = CacheAwareScheduler()
candidates, _ = await scheduler.list_all_candidates(
db=self.db,
api_format=APIFormat.GEMINI,
model_name=model_name,
affinity_key=str(self.api_key.id),
user_api_key=self.api_key,
max_candidates=10,
)
# 记录所有候选 key 信息
candidate_keys = []
selected_candidate = None
selected_index = -1
selected_rule_lookup: BillingRuleLookupResult | None = None
for idx, candidate in enumerate(candidates):
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
has_billing_rule = True
rule_lookup: BillingRuleLookupResult | None = None
if config.billing_require_rule:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=candidate.provider.id,
model_name=model_name,
task_type="video",
)
has_billing_rule = rule_lookup is not None
candidate_info = {
"index": idx,
"provider_id": candidate.provider.id,
"provider_name": candidate.provider.name,
"endpoint_id": candidate.endpoint.id,
"key_id": candidate.key.id,
"key_name": candidate.key.name,
"auth_type": auth_type,
"has_billing_rule": has_billing_rule,
"priority": getattr(candidate.key, "priority", 0) or 0,
}
candidate_keys.append(candidate_info)
if (
selected_candidate is None
and auth_type in {"api_key", "vertex_ai"}
and has_billing_rule
):
selected_candidate = candidate
selected_index = idx
selected_rule_lookup = rule_lookup
# 标记选中的候选
if selected_index >= 0:
candidate_keys[selected_index]["selected"] = True
return selected_candidate, candidate_keys, selected_rule_lookup
async def _resolve_upstream_key( async def _resolve_upstream_key(
self, candidate: ProviderCandidate self, candidate: ProviderCandidate
) -> tuple[str, ProviderEndpoint, ProviderAPIKey, Any | None]: ) -> tuple[str, ProviderEndpoint, ProviderAPIKey, Any | None]:
@@ -399,9 +351,13 @@ class GeminiVeoHandler(VideoHandlerBase):
auth_info: Any | None, auth_info: Any | None,
) -> dict[str, str]: ) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint) extra_headers = get_extra_headers_from_endpoint(endpoint)
headers = build_upstream_headers( 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, original_headers,
APIFormat.GEMINI, endpoint_sig,
upstream_key, upstream_key,
endpoint_headers=extra_headers, endpoint_headers=extra_headers,
) )
@@ -459,8 +415,11 @@ class GeminiVeoHandler(VideoHandlerBase):
provider_id=candidate.provider.id, provider_id=candidate.provider.id,
endpoint_id=candidate.endpoint.id, endpoint_id=candidate.endpoint.id,
key_id=candidate.key.id, key_id=candidate.key.id,
client_api_format="GEMINI", client_api_format=self.FORMAT_ID,
provider_api_format=str(candidate.endpoint.api_format), 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, format_converted=False,
model=internal_request.model, model=internal_request.model,
prompt=internal_request.prompt, prompt=internal_request.prompt,

View File

@@ -6,10 +6,12 @@ from __future__ import annotations
from src.api.handlers.base.video_adapter_base import VideoAdapterBase from src.api.handlers.base.video_adapter_base import VideoAdapterBase
from src.api.handlers.base.video_handler_base import VideoHandlerBase from src.api.handlers.base.video_handler_base import VideoHandlerBase
from src.core.api_format import ApiFamily
class OpenAIVideoAdapter(VideoAdapterBase): class OpenAIVideoAdapter(VideoAdapterBase):
FORMAT_ID = "OPENAI" FORMAT_ID = "openai:video"
API_FAMILY = ApiFamily.OPENAI
name = "openai.video" name = "openai.video"
@property @property

View File

@@ -17,7 +17,13 @@ from sqlalchemy.orm import Session
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
from src.clients.http_client import HTTPClientPool from src.clients.http_client import HTTPClientPool
from src.config.settings import config from src.config.settings import config
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint 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 ( from src.core.api_format.conversion.internal_video import (
InternalVideoRequest, InternalVideoRequest,
InternalVideoTask, InternalVideoTask,
@@ -29,11 +35,14 @@ from src.core.crypto import crypto_service
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.usage.service import UsageService
class OpenAIVideoHandler(VideoHandlerBase): class OpenAIVideoHandler(VideoHandlerBase):
FORMAT_ID = "OPENAI" FORMAT_ID = "openai:video"
API_FAMILY = ApiFamily.OPENAI
ENDPOINT_KIND = EndpointKind.VIDEO
DEFAULT_BASE_URL = "https://api.openai.com" DEFAULT_BASE_URL = "https://api.openai.com"
@@ -72,47 +81,87 @@ class OpenAIVideoHandler(VideoHandlerBase):
try: try:
internal_request = self._normalizer.video_request_to_internal(original_request_body) internal_request = self._normalizer.video_request_to_internal(original_request_body)
except ValueError as e: except ValueError as e:
# 请求解析失败,记录失败的使用记录
await self._record_failed_usage(
model="unknown",
error_message=str(e),
status_code=400,
original_request_body=original_request_body,
original_headers=original_headers,
)
raise HTTPException(status_code=400, detail=str(e)) raise HTTPException(status_code=400, detail=str(e))
candidate, candidate_keys, rule_lookup = await self._select_candidate(
internal_request.model async def _submit(candidate: ProviderCandidate) -> Any:
) upstream_key, endpoint, _provider_key = await self._resolve_upstream_key(candidate)
if not candidate: upstream_url = self._build_upstream_url(endpoint.base_url)
detail = "No available provider for video generation" headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
if config.billing_require_rule: client = await HTTPClientPool.get_default_client_async()
detail = "No available provider with billing rule for video generation" return await client.post(upstream_url, headers=headers, json=original_request_body)
raise HTTPException(status_code=503, detail=detail)
def _extract_task_id(payload: dict[str, Any]) -> str | None:
value = payload.get("id")
return str(value) if value else None
# 捕获提交阶段的所有错误,记录失败任务
try:
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"},
allow_format_conversion=False,
max_candidates=10,
)
except HTTPException as exc:
# 提交失败(如无可用 provider创建失败任务记录
# 尝试获取 candidate_keys如果是 AllCandidatesFailedError 转换来的)
candidate_keys = getattr(exc, "candidate_keys", None)
await self._create_failed_task_and_usage(
internal_request=internal_request,
original_request_body=original_request_body,
original_headers=original_headers,
error_code="provider_unavailable",
error_message=exc.detail if isinstance(exc.detail, str) else str(exc.detail),
status_code=exc.status_code,
candidate_keys=candidate_keys,
)
raise
if isinstance(outcome_or_response, JSONResponse):
# 上游返回客户端错误,也要记录
await self._create_failed_task_and_usage(
internal_request=internal_request,
original_request_body=original_request_body,
original_headers=original_headers,
error_code="upstream_client_error",
error_message="Upstream rejected the request",
status_code=outcome_or_response.status_code,
)
return outcome_or_response
outcome = outcome_or_response
# 冻结 billing_rule 配置(用于异步任务的成本一致性) # 冻结 billing_rule 配置(用于异步任务的成本一致性)
# 复用 _select_candidate 中已查询的结果billing_require_rule=false 时需补查 # 复用 _select_candidate 中已查询的结果billing_require_rule=false 时需补查
rule_lookup = outcome.rule_lookup
if rule_lookup is None: if rule_lookup is None:
rule_lookup = BillingRuleService.find_rule( rule_lookup = BillingRuleService.find_rule(
self.db, self.db,
provider_id=candidate.provider.id, provider_id=outcome.candidate.provider.id,
model_name=internal_request.model, model_name=internal_request.model,
task_type="video", task_type="video",
) )
billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup) billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup)
upstream_key, endpoint, provider_key = await self._resolve_upstream_key(candidate) external_task_id = outcome.external_task_id
upstream_url = self._build_upstream_url(endpoint.base_url)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
response = await client.post(upstream_url, headers=headers, json=original_request_body)
if response.status_code >= 400:
return self._build_error_response(response)
payload = response.json()
external_task_id = str(payload.get("id") or "")
if not external_task_id:
raise HTTPException(status_code=502, detail="Upstream returned empty task id")
task = self._create_task_record( task = self._create_task_record(
external_task_id=external_task_id, external_task_id=external_task_id,
candidate=candidate, candidate=outcome.candidate,
original_request_body=original_request_body, original_request_body=original_request_body,
internal_request=internal_request, internal_request=internal_request,
candidate_keys=candidate_keys, candidate_keys=outcome.candidate_keys,
original_headers=original_headers, original_headers=original_headers,
billing_rule_snapshot=billing_rule_snapshot, billing_rule_snapshot=billing_rule_snapshot,
) )
@@ -199,6 +248,117 @@ class OpenAIVideoHandler(VideoHandlerBase):
self.db.commit() self.db.commit()
return JSONResponse({}) return JSONResponse({})
async def handle_remix_task(
self,
*,
task_id: str,
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:
# 获取原始任务以验证所有权和状态
original_task = self._get_task(task_id)
if original_task.status != VideoStatus.COMPLETED.value:
raise HTTPException(
status_code=400,
detail=f"Can only remix completed videos (current status: {original_task.status})",
)
if not original_task.external_task_id:
raise HTTPException(status_code=500, detail="Original task missing external_task_id")
# 获取原始任务的 endpoint 和 key
endpoint, key = self._get_endpoint_and_key(original_task)
if not key.api_key:
raise HTTPException(status_code=500, detail="Provider key not configured")
upstream_key = crypto_service.decrypt(key.api_key)
# 构建 remix 请求的上游 URL
upstream_url = self._build_upstream_url(
endpoint.base_url, f"{original_task.external_task_id}/remix"
)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async()
response = await client.post(upstream_url, headers=headers, json=original_request_body)
if response.status_code >= 400:
return self._build_error_response(response)
# 解析上游响应
try:
response_data = response.json()
except (ValueError, TypeError):
raise HTTPException(status_code=502, detail="Invalid response from upstream")
external_task_id = response_data.get("id")
if not external_task_id:
raise HTTPException(status_code=502, detail="Upstream did not return task ID")
# 解析 remix 请求
try:
internal_request = self._normalizer.video_request_to_internal(
{
"prompt": original_request_body.get("prompt", ""),
"model": original_task.model,
"size": original_task.size,
"seconds": original_task.duration_seconds,
}
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
# 复用原始任务的 billing rule snapshot
billing_rule_snapshot = None
if original_task.request_metadata:
billing_rule_snapshot = original_task.request_metadata.get("billing_rule_snapshot")
# 构建 ProviderCandidate复用原始任务的 provider 配置)
from src.models.database import Provider
provider = self.db.query(Provider).filter(Provider.id == original_task.provider_id).first()
if not provider:
raise HTTPException(status_code=500, detail="Provider not found")
candidate = ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
)
# 创建新任务记录
task = self._create_task_record(
external_task_id=external_task_id,
candidate=candidate,
original_request_body={
**original_request_body,
"remix_video_id": task_id,
},
internal_request=internal_request,
original_headers=original_headers,
billing_rule_snapshot=billing_rule_snapshot,
)
try:
self.db.add(task)
self.db.flush()
self.db.commit()
self.db.refresh(task)
except IntegrityError:
self.db.rollback()
raise HTTPException(status_code=409, detail="Task already exists")
internal_task = InternalVideoTask(
id=task.id,
external_id=external_task_id,
status=VideoStatus.SUBMITTED,
created_at=task.created_at,
original_request=internal_request,
)
response_body = self._normalizer.video_task_from_internal(internal_task)
return JSONResponse(response_body)
async def handle_download_content( async def handle_download_content(
self, self,
*, *,
@@ -236,9 +396,19 @@ class OpenAIVideoHandler(VideoHandlerBase):
raise HTTPException(status_code=500, detail="Provider key not configured") raise HTTPException(status_code=500, detail="Provider key not configured")
upstream_key = crypto_service.decrypt(key.api_key) upstream_key = crypto_service.decrypt(key.api_key)
upstream_url = self._build_upstream_url( # 支持 variant 查询参数: video (默认), thumbnail, spritesheet
endpoint.base_url, f"{task.external_task_id}/content" variant = (query_params or {}).get("variant", "video")
) if variant not in {"video", "thumbnail", "spritesheet"}:
raise HTTPException(
status_code=400,
detail=f"Invalid variant '{variant}'. Must be one of: video, thumbnail, spritesheet",
)
# 构建上游 URL透传 variant 参数
content_path = f"{task.external_task_id}/content"
if variant != "video":
content_path = f"{content_path}?variant={variant}"
upstream_url = self._build_upstream_url(endpoint.base_url, content_path)
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint) headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async() client = await HTTPClientPool.get_default_client_async()
@@ -299,57 +469,6 @@ class OpenAIVideoHandler(VideoHandlerBase):
# Helpers # Helpers
# ------------------------------------------------------------------ # ------------------------------------------------------------------
async def _select_candidate(
self, model_name: str
) -> tuple[ProviderCandidate | None, list[dict[str, Any]], BillingRuleLookupResult | None]:
"""选择候选 key返回 (选中的候选, 所有候选列表, 选中候选的 billing rule lookup)"""
scheduler = CacheAwareScheduler()
candidates, _ = await scheduler.list_all_candidates(
db=self.db,
api_format=APIFormat.OPENAI,
model_name=model_name,
affinity_key=str(self.api_key.id),
user_api_key=self.api_key,
max_candidates=10,
)
# 记录所有候选 key 信息
candidate_keys = []
selected_candidate = None
selected_index = -1
selected_rule_lookup: BillingRuleLookupResult | None = None
for idx, candidate in enumerate(candidates):
auth_type = getattr(candidate.key, "auth_type", "api_key") or "api_key"
has_billing_rule = True
rule_lookup: BillingRuleLookupResult | None = None
if config.billing_require_rule:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=candidate.provider.id,
model_name=model_name,
task_type="video",
)
has_billing_rule = rule_lookup is not None
candidate_info = {
"index": idx,
"provider_id": candidate.provider.id,
"provider_name": candidate.provider.name,
"endpoint_id": candidate.endpoint.id,
"key_id": candidate.key.id,
"key_name": candidate.key.name,
"auth_type": auth_type,
"has_billing_rule": has_billing_rule,
"priority": getattr(candidate.key, "priority", 0) or 0,
}
candidate_keys.append(candidate_info)
if selected_candidate is None and auth_type == "api_key" and has_billing_rule:
selected_candidate = candidate
selected_index = idx
selected_rule_lookup = rule_lookup
# 标记选中的候选
if selected_index >= 0:
candidate_keys[selected_index]["selected"] = True
return selected_candidate, candidate_keys, selected_rule_lookup
async def _resolve_upstream_key( async def _resolve_upstream_key(
self, candidate: ProviderCandidate self, candidate: ProviderCandidate
) -> tuple[str, ProviderEndpoint, ProviderAPIKey]: ) -> tuple[str, ProviderEndpoint, ProviderAPIKey]:
@@ -378,9 +497,13 @@ class OpenAIVideoHandler(VideoHandlerBase):
self, original_headers: dict[str, str], upstream_key: str, endpoint: ProviderEndpoint self, original_headers: dict[str, str], upstream_key: str, endpoint: ProviderEndpoint
) -> dict[str, str]: ) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint) extra_headers = get_extra_headers_from_endpoint(endpoint)
return build_upstream_headers( 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, original_headers,
APIFormat.OPENAI, endpoint_sig,
upstream_key, upstream_key,
endpoint_headers=extra_headers, endpoint_headers=extra_headers,
) )
@@ -428,8 +551,11 @@ class OpenAIVideoHandler(VideoHandlerBase):
provider_id=candidate.provider.id, provider_id=candidate.provider.id,
endpoint_id=candidate.endpoint.id, endpoint_id=candidate.endpoint.id,
key_id=candidate.key.id, key_id=candidate.key.id,
client_api_format="OPENAI", client_api_format=self.FORMAT_ID,
provider_api_format=str(candidate.endpoint.api_format), 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, format_converted=False,
model=internal_request.model, model=internal_request.model,
prompt=internal_request.prompt, prompt=internal_request.prompt,
@@ -473,5 +599,184 @@ class OpenAIVideoHandler(VideoHandlerBase):
extra={"model": task.model, "size": task.size, "seconds": task.duration_seconds}, extra={"model": task.model, "size": task.size, "seconds": task.duration_seconds},
) )
async def _record_failed_usage(
self,
*,
model: str,
error_message: str,
status_code: int,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
) -> None:
"""记录失败请求的使用记录(无任务记录)"""
import time
response_time_ms = int((time.time() - self.start_time) * 1000)
safe_headers = {
k: v
for k, v in original_headers.items()
if k.lower() not in {"authorization", "x-api-key", "cookie"}
}
try:
await UsageService.record_usage_with_custom_cost(
db=self.db,
user=self.user,
api_key=self.api_key,
provider="unknown",
model=model,
request_type="video",
total_cost_usd=0.0,
request_cost_usd=0.0,
input_tokens=0,
output_tokens=0,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
api_format=self.FORMAT_ID,
endpoint_api_format=None,
has_format_conversion=False,
is_stream=False,
response_time_ms=response_time_ms,
first_byte_time_ms=None,
status_code=status_code,
error_message=error_message,
metadata={"failure_stage": "request_parsing"},
request_headers=safe_headers,
request_body=original_request_body,
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=self.request_id,
provider_id=None,
provider_endpoint_id=None,
provider_api_key_id=None,
status="failed",
target_model=None,
)
except Exception as exc:
logger.warning("Failed to record failed usage: %s", sanitize_error_message(str(exc)))
async def _create_failed_task_and_usage(
self,
*,
internal_request: InternalVideoRequest,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
error_code: str,
error_message: str,
status_code: int,
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)
# 构建请求元数据
safe_headers = {
k: v
for k, v in original_headers.items()
if k.lower() not in {"authorization", "x-api-key", "cookie"}
}
request_metadata: dict[str, Any] = {
"client_ip": self.client_ip,
"user_agent": self.user_agent,
"request_id": self.request_id,
"request_headers": safe_headers,
"failure_stage": "submit",
}
# 添加候选链路追踪信息
if candidate_keys:
request_metadata["candidate_keys"] = candidate_keys
size = internal_request.extra.get("original_size")
# 创建失败的任务记录
task = VideoTask(
id=str(uuid4()),
external_task_id=None,
user_id=self.user.id,
api_key_id=self.api_key.id,
provider_id=None,
endpoint_id=None,
key_id=None,
client_api_format=self.FORMAT_ID,
provider_api_format=self.FORMAT_ID,
format_converted=False,
model=internal_request.model,
prompt=internal_request.prompt,
original_request_body=original_request_body,
converted_request_body=original_request_body,
duration_seconds=internal_request.duration_seconds,
resolution=internal_request.resolution,
aspect_ratio=internal_request.aspect_ratio,
size=size,
status=VideoStatus.FAILED.value,
progress_percent=0,
error_code=error_code,
error_message=error_message,
submitted_at=now,
completed_at=now,
request_metadata=request_metadata,
)
try:
self.db.add(task)
self.db.commit()
self.db.refresh(task)
except Exception as exc:
self.db.rollback()
logger.warning(
"Failed to create failed task record: %s", sanitize_error_message(str(exc))
)
# 即使任务记录失败,仍然尝试记录使用记录
task = None
# 记录使用记录
try:
await UsageService.record_usage_with_custom_cost(
db=self.db,
user=self.user,
api_key=self.api_key,
provider="unknown",
model=internal_request.model,
request_type="video",
total_cost_usd=0.0,
request_cost_usd=0.0,
input_tokens=0,
output_tokens=0,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
api_format=self.FORMAT_ID,
endpoint_api_format=None,
has_format_conversion=False,
is_stream=False,
response_time_ms=response_time_ms,
first_byte_time_ms=None,
status_code=status_code,
error_message=error_message,
metadata={
"failure_stage": "submit",
"error_code": error_code,
"video_task_id": task.id if task else None,
},
request_headers=safe_headers,
request_body=original_request_body,
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=self.request_id,
provider_id=None,
provider_endpoint_id=None,
provider_api_key_id=None,
status="failed",
target_model=None,
)
except Exception as exc:
logger.warning("Failed to record failed usage: %s", sanitize_error_message(str(exc)))
__all__ = ["OpenAIVideoHandler"] __all__ = ["OpenAIVideoHandler"]

View File

@@ -90,6 +90,21 @@ async def download_video_content_sora(
) )
@router.post("/v1/videos/{task_id}/remix")
async def remix_video_sora(
task_id: str, http_request: Request, db: Session = Depends(get_db)
) -> Any:
adapter = OpenAIVideoAdapter()
return await pipeline.run(
adapter=adapter,
http_request=http_request,
db=db,
mode=adapter.mode,
api_format_hint=adapter.allowed_api_formats[0],
path_params={"task_id": task_id},
)
# -------------------- Gemini Veo compatible -------------------- # -------------------- Gemini Veo compatible --------------------

View File

@@ -11,7 +11,7 @@
from src.services.billing import BillingCalculator, UsageMapper, StandardizedUsage from src.services.billing import BillingCalculator, UsageMapper, StandardizedUsage
# 1. 将原始 usage 映射为标准格式 # 1. 将原始 usage 映射为标准格式
usage = UsageMapper.map(raw_usage, api_format="OPENAI") usage = UsageMapper.map(raw_usage, api_format="openai:chat")
# 2. 使用计费计算器计算费用 # 2. 使用计费计算器计算费用
calculator = BillingCalculator(template="openai") calculator = BillingCalculator(template="openai")

View File

@@ -29,7 +29,11 @@ ValueType = Literal["float", "int", "string"]
def _normalize_api_format(api_format: str | None) -> str: def _normalize_api_format(api_format: str | None) -> str:
return (api_format or "").upper() if not api_format:
return ""
from src.core.api_format.signature import normalize_signature_key
return normalize_signature_key(api_format)
def _normalize_task_type(task_type: str | None) -> str: def _normalize_task_type(task_type: str | None) -> str:
@@ -312,6 +316,45 @@ class DimensionCollectorService:
task = _normalize_task_type(task_type) task = _normalize_task_type(task_type)
api_variants = list({api, api.lower()}) api_variants = list({api, api.lower()})
if task == "video":
# VIDEO → base 回退:优先使用 family:video 专用 collector
# 缺失的维度再回退到 family:chat。
from src.core.api_format.signature import parse_signature_key
base_api = api
try:
sig = parse_signature_key(api)
if sig.endpoint_kind.value == "video":
base_api = f"{sig.api_family.value}:chat"
except Exception:
base_api = api
base_variants = list({base_api, base_api.lower()})
video_collectors = (
self.db.query(DimensionCollector)
.filter(
DimensionCollector.api_format.in_(api_variants),
DimensionCollector.task_type == "video",
DimensionCollector.is_enabled == True, # noqa: E712
)
.all()
)
base_collectors = (
self.db.query(DimensionCollector)
.filter(
DimensionCollector.api_format.in_(base_variants),
DimensionCollector.task_type == "video",
DimensionCollector.is_enabled == True, # noqa: E712
)
.all()
)
video_dims: set[str] = {c.dimension_name for c in video_collectors}
result: list[DimensionCollector] = list(video_collectors)
for c in base_collectors:
if c.dimension_name not in video_dims:
result.append(c)
return result
if task == "cli": if task == "cli":
# CLI → chat按维度回退维度存在 cli collector 则用 cli否则用 chat # CLI → chat按维度回退维度存在 cli collector 则用 cli否则用 chat
cli_collectors = ( cli_collectors = (

View File

@@ -4,9 +4,9 @@ Usage 字段映射器
将不同 API 格式的原始 usage 数据映射为标准化格式。 将不同 API 格式的原始 usage 数据映射为标准化格式。
支持的格式: 支持的格式:
- OPENAI / OPENAI_CLI: OpenAI Chat Completions API - openai:*: OpenAI compatible (Chat/CLI)
- CLAUDE / CLAUDE_CLI: Anthropic Messages API - claude:*: Anthropic Messages (Chat/CLI)
- GEMINI / GEMINI_CLI: Google Gemini API - gemini:*: Google Gemini (Chat/CLI)
""" """
from typing import Any from typing import Any
@@ -73,16 +73,6 @@ class UsageMapper:
"usageMetadata.cachedContentTokenCount": "cache_read_tokens", "usageMetadata.cachedContentTokenCount": "cache_read_tokens",
} }
# 格式名称到映射的对应关系
FORMAT_MAPPINGS: dict[str, dict[str, str]] = {
"OPENAI": OPENAI_MAPPING,
"OPENAI_CLI": OPENAI_MAPPING,
"CLAUDE": CLAUDE_MAPPING,
"CLAUDE_CLI": CLAUDE_MAPPING,
"GEMINI": GEMINI_MAPPING,
"GEMINI_CLI": GEMINI_MAPPING,
}
@classmethod @classmethod
def map( def map(
cls, cls,
@@ -142,12 +132,13 @@ class UsageMapper:
Returns: Returns:
标准化的 usage 对象 标准化的 usage 对象
""" """
format_upper = api_format.upper() if api_format else "" format_norm = (api_format or "").strip().lower()
api_family = format_norm.split(":", 1)[0] if ":" in format_norm else format_norm
# 提取 usage 部分 # 提取 usage 部分
usage_data: dict[str, Any] = {} usage_data: dict[str, Any] = {}
if format_upper.startswith("GEMINI"): if api_family == "gemini":
# Gemini: usageMetadata # Gemini: usageMetadata
usage_data = response.get("usageMetadata", {}) usage_data = response.get("usageMetadata", {})
if not usage_data: if not usage_data:
@@ -164,21 +155,14 @@ class UsageMapper:
@classmethod @classmethod
def _get_mapping(cls, api_format: str) -> dict[str, str]: def _get_mapping(cls, api_format: str) -> dict[str, str]:
"""获取对应格式的字段映射""" """获取对应格式的字段映射"""
if not api_format: format_norm = (api_format or "").strip().lower()
return cls.CLAUDE_MAPPING api_family = format_norm.split(":", 1)[0] if ":" in format_norm else format_norm
format_upper = api_format.upper() if api_family == "openai":
return cls.OPENAI_MAPPING
# 精确匹配 if api_family == "gemini":
if format_upper in cls.FORMAT_MAPPINGS: return cls.GEMINI_MAPPING
return cls.FORMAT_MAPPINGS[format_upper] # 默认 Claude也覆盖未知/空值)
# 前缀匹配
for key, mapping in cls.FORMAT_MAPPINGS.items():
if format_upper.startswith(key.split("_")[0]):
return mapping
# 默认使用 Claude 映射
return cls.CLAUDE_MAPPING return cls.CLAUDE_MAPPING
@classmethod @classmethod

View File

@@ -19,19 +19,20 @@ from src.api.handlers.base.video_handler_base import (
from src.clients.http_client import HTTPClientPool from src.clients.http_client import HTTPClientPool
from src.clients.redis_client import get_redis_client from src.clients.redis_client import get_redis_client
from src.config.settings import config from src.config.settings import config
from src.core.api_format import APIFormat, build_upstream_headers, get_extra_headers_from_endpoint from src.core.api_format import (
build_upstream_headers_for_endpoint,
get_extra_headers_from_endpoint,
make_signature_key,
)
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
from src.core.crypto import crypto_service from src.core.crypto import crypto_service
from src.core.logger import logger from src.core.logger import logger
from src.database import create_session from src.database import create_session
from src.models.database import ApiKey, Provider, ProviderAPIKey, ProviderEndpoint, User, VideoTask from src.models.database import ProviderAPIKey, ProviderEndpoint, VideoTask
from src.services.billing.dimension_collector_service import DimensionCollectorService
from src.services.billing.formula_engine import BillingIncompleteError, FormulaEngine
from src.services.billing.rule_service import BillingRuleService
from src.services.system.scheduler import get_scheduler from src.services.system.scheduler import get_scheduler
from src.services.usage.service import UsageService from src.services.task.impl.video_telemetry import VideoTelemetry
# 永久性错误指示词(用于降级判断,不应重试) # 永久性错误指示词(用于降级判断,不应重试)
_PERMANENT_ERROR_INDICATORS = frozenset( _PERMANENT_ERROR_INDICATORS = frozenset(
@@ -71,7 +72,6 @@ class VideoTaskPollerService:
self.redis = None self.redis = None
self._openai_normalizer = OpenAINormalizer() self._openai_normalizer = OpenAINormalizer()
self._gemini_normalizer = GeminiNormalizer() self._gemini_normalizer = GeminiNormalizer()
self._formula_engine = FormulaEngine()
# 追踪连续失败次数(用于告警) # 追踪连续失败次数(用于告警)
self._consecutive_failures = 0 self._consecutive_failures = 0
# 从配置读取参数 # 从配置读取参数
@@ -248,7 +248,7 @@ class VideoTaskPollerService:
# 终态写入 Usage复用外层 per-task session # 终态写入 Usage复用外层 per-task session
if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value): if task.status in (VideoStatus.COMPLETED.value, VideoStatus.FAILED.value):
try: try:
await self._record_terminal_usage(db, task) await VideoTelemetry(db, redis_client=self.redis).record_terminal_usage(task)
except Exception as exc: except Exception as exc:
logger.exception( logger.exception(
"Failed to record video usage for task=%s: %s", "Failed to record video usage for task=%s: %s",
@@ -264,289 +264,6 @@ class VideoTaskPollerService:
# 仅在终态写一次,避免污染 request_metadata # 仅在终态写一次,避免污染 request_metadata
task.request_metadata["poll_raw_response"] = result.raw_response task.request_metadata["poll_raw_response"] = result.raw_response
async def _record_terminal_usage(self, db: Session, task: VideoTask) -> None:
"""
为视频任务终态写入 Usage
- COMPLETED: 使用 FormulaEngine 计算 cost或 no_rule / incomplete -> cost=0
- FAILED: cost=0
"""
request_id = None
if isinstance(task.request_metadata, dict):
request_id = task.request_metadata.get("request_id")
request_id = request_id or task.id
# 计算异步任务总耗时ms
response_time_ms = None
if task.submitted_at and task.completed_at:
delta = task.completed_at - task.submitted_at
response_time_ms = int(delta.total_seconds() * 1000)
# 基础维度(无需 collectors 也可计费)
base_dimensions: dict[str, Any] = {
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size or "",
"retry_count": task.retry_count,
}
# collectors 可用的 metadata结构稳定便于配置 path
collector_metadata: dict[str, Any] = {
"task": {
"id": task.id,
"external_task_id": task.external_task_id,
"model": task.model,
"duration_seconds": task.duration_seconds,
"resolution": task.resolution,
"aspect_ratio": task.aspect_ratio,
"size": task.size,
"retry_count": task.retry_count,
"video_size_bytes": task.video_size_bytes,
},
"result": {
"video_url": task.video_url,
"video_urls": task.video_urls or [],
},
}
# 维度采集base + collectors 覆盖/补全
dims = DimensionCollectorService(db).collect_dimensions(
api_format=task.provider_api_format,
task_type="video",
request=task.original_request_body or {},
response=(
(task.request_metadata or {}).get("poll_raw_response")
if isinstance(task.request_metadata, dict)
else None
),
metadata=collector_metadata,
base_dimensions=base_dimensions,
)
# 取冻结的 rule_snapshot若缺失则回退 DB 查找(兼容旧任务)
rule_snapshot = None
if isinstance(task.request_metadata, dict):
rule_snapshot = task.request_metadata.get("billing_rule_snapshot")
# 构建 billing_snapshot写入 Usage.request_metadata
billing_snapshot: dict[str, Any] = {
"status": "complete",
"missing_required": [],
"strict_mode": config.billing_strict_mode,
}
cost = 0.0
if task.status == VideoStatus.FAILED.value:
billing_snapshot["billed_reason"] = "task_failed"
else:
# COMPLETED计算成本
expression = None
variables = None
dimension_mappings = None
rule_id = None
rule_name = None
rule_scope = None
if isinstance(rule_snapshot, dict) and rule_snapshot.get("status") == "ok":
rule_id = rule_snapshot.get("rule_id")
rule_name = rule_snapshot.get("rule_name")
rule_scope = rule_snapshot.get("scope")
expression = rule_snapshot.get("expression")
variables = rule_snapshot.get("variables")
dimension_mappings = rule_snapshot.get("dimension_mappings")
else:
lookup = BillingRuleService.find_rule(
db,
provider_id=task.provider_id,
model_name=task.model,
task_type="video",
)
if lookup:
rule = lookup.rule
rule_id = rule.id
rule_name = rule.name
rule_scope = lookup.scope
expression = rule.expression
variables = rule.variables
dimension_mappings = rule.dimension_mappings
if not expression:
billing_snapshot["status"] = "no_rule"
billing_snapshot["cost_breakdown"] = {"total": 0.0}
logger.warning(
"No billing rule for video task (request_id=%s, model=%s, provider_id=%s)",
request_id,
task.model,
task.provider_id,
)
else:
billing_snapshot.update(
{
"rule_id": rule_id,
"rule_name": rule_name,
"rule_scope": rule_scope,
"expression": expression,
"variables": variables or {},
}
)
try:
result = self._formula_engine.evaluate(
expression=expression,
variables=variables or {},
dimensions=dims,
dimension_mappings=dimension_mappings or {},
strict_mode=config.billing_strict_mode,
)
billing_snapshot["status"] = result.status
billing_snapshot["missing_required"] = result.missing_required
billing_snapshot["resolved_values"] = result.resolved_values
if result.status == "complete":
cost = result.cost
else:
logger.error(
"Billing incomplete due to missing required dimensions "
"(request_id=%s, model=%s, missing=%s)",
request_id,
task.model,
result.missing_required,
)
cost = 0.0
await self._maybe_alert_missing_required(
model=task.model,
missing_required=result.missing_required,
)
if result.error:
billing_snapshot["error"] = result.error
except BillingIncompleteError as exc:
logger.error(
"Billing strict mode triggered (request_id=%s, model=%s, missing=%s)",
request_id,
task.model,
exc.missing_required,
)
billing_snapshot["status"] = "incomplete"
billing_snapshot["missing_required"] = exc.missing_required
billing_snapshot["resolved_values"] = {}
billing_snapshot["error"] = "strict_mode_missing_required"
cost = 0.0
# strict_mode=true标记任务失败并隐藏产物避免"免费放行"
task.status = VideoStatus.FAILED.value
task.error_code = "billing_incomplete"
task.error_message = f"Missing required dimensions: {exc.missing_required}"
task.video_url = None
task.video_urls = None
await self._maybe_alert_missing_required(
model=task.model,
missing_required=exc.missing_required,
)
billing_snapshot["cost_breakdown"] = {"total": cost}
# Usage 元数据(包含 snapshot + dimensions + raw_response_ref
usage_metadata: dict[str, Any] = {
"billing_snapshot": billing_snapshot,
"dimensions": dims,
"raw_response_ref": {
"video_task_id": task.id,
"field": "video_tasks.request_metadata.poll_raw_response",
},
}
# 查询关联对象(用于写入 usage.user_id/api_key_id 等)
user_obj = db.query(User).filter(User.id == task.user_id).first()
api_key_obj = (
db.query(ApiKey).filter(ApiKey.id == task.api_key_id).first()
if task.api_key_id
else None
)
provider_obj = (
db.query(Provider).filter(Provider.id == task.provider_id).first()
if task.provider_id
else None
)
provider_name = provider_obj.name if provider_obj else "unknown"
await UsageService.record_usage_with_custom_cost(
db=db,
user=user_obj,
api_key=api_key_obj,
provider=provider_name,
model=task.model,
request_type="video",
total_cost_usd=cost,
request_cost_usd=cost,
input_tokens=0,
output_tokens=0,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
api_format=task.client_api_format,
endpoint_api_format=task.provider_api_format,
has_format_conversion=bool(task.format_converted),
is_stream=False,
response_time_ms=response_time_ms,
first_byte_time_ms=None,
status_code=200 if task.status == VideoStatus.COMPLETED.value else 500,
error_message=(
None
if task.status == VideoStatus.COMPLETED.value
else (task.error_message or task.error_code or "video_task_failed")
),
metadata=usage_metadata,
request_headers=(
(task.request_metadata or {}).get("request_headers")
if isinstance(task.request_metadata, dict)
else None
),
request_body=task.original_request_body,
provider_request_headers=None,
response_headers=None,
client_response_headers=None,
response_body=None,
request_id=request_id,
provider_id=task.provider_id,
provider_endpoint_id=task.endpoint_id,
provider_api_key_id=task.key_id,
status="completed" if task.status == VideoStatus.COMPLETED.value else "failed",
target_model=None,
)
async def _maybe_alert_missing_required(
self, *, model: str, missing_required: list[str]
) -> None:
"""required 维度缺失告警:同一 (model, dimension) 1 小时内 >= 10 次触发升级告警。"""
if not missing_required:
return
if not self.redis:
# Redis 不可用:降级为日志
logger.error(
"Missing required billing dimensions (model=%s): %s", model, missing_required
)
return
# 按小时 bucket 聚合
now = datetime.now(timezone.utc)
hour_bucket = now.strftime("%Y%m%d%H")
for dim in missing_required:
key = f"billing:missing_required:{model}:{dim}:{hour_bucket}"
try:
count = await self.redis.incr(key)
# TTL 略大于 1h避免边界抖动
if count == 1:
await self.redis.expire(key, 3700)
if count >= 10:
logger.warning(
"Billing required dimension missing frequently (model=%s, dim=%s, count=%s/hour)",
model,
dim,
count,
)
except Exception as exc:
logger.warning(
"Failed to record billing alert counter: %s", sanitize_error_message(str(exc))
)
def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool: def _is_permanent_error(self, exc: Exception, status_code: int | None = None) -> bool:
"""判断是否为永久性错误(不应重试)""" """判断是否为永久性错误(不应重试)"""
# 优先使用 HTTP 状态码判断 # 优先使用 HTTP 状态码判断
@@ -583,7 +300,14 @@ class VideoTaskPollerService:
error_message="Failed to decrypt provider key", error_message="Failed to decrypt provider key",
) )
if (task.provider_api_format or "").upper() == "GEMINI": provider_format = (task.provider_api_format or "").strip().lower()
if not provider_format:
provider_format = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
if provider_format.startswith("gemini:"):
auth_info = await get_provider_auth(endpoint, key) auth_info = await get_provider_auth(endpoint, key)
return await self._poll_gemini(task, endpoint, upstream_key, auth_info) return await self._poll_gemini(task, endpoint, upstream_key, auth_info)
return await self._poll_openai(task, endpoint, upstream_key) return await self._poll_openai(task, endpoint, upstream_key)
@@ -601,7 +325,11 @@ class VideoTaskPollerService:
error_message="Task missing external_task_id", error_message="Task missing external_task_id",
) )
url = self._build_openai_url(endpoint.base_url, task.external_task_id) url = self._build_openai_url(endpoint.base_url, task.external_task_id)
headers = self._build_headers(APIFormat.OPENAI, upstream_key, endpoint) endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
headers = self._build_headers(endpoint_sig, upstream_key, endpoint)
client = await HTTPClientPool.get_default_client_async() client = await HTTPClientPool.get_default_client_async()
response = await client.get(url, headers=headers) response = await client.get(url, headers=headers)
@@ -629,7 +357,11 @@ class VideoTaskPollerService:
) )
operation_name = normalize_gemini_operation_id(task.external_task_id) operation_name = normalize_gemini_operation_id(task.external_task_id)
url = self._build_gemini_url(endpoint.base_url, operation_name) url = self._build_gemini_url(endpoint.base_url, operation_name)
headers = self._build_headers(APIFormat.GEMINI, upstream_key, endpoint, auth_info) endpoint_sig = (task.provider_api_format or "").strip().lower() or make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
headers = self._build_headers(endpoint_sig, upstream_key, endpoint, auth_info)
client = await HTTPClientPool.get_default_client_async() client = await HTTPClientPool.get_default_client_async()
response = await client.get(url, headers=headers) response = await client.get(url, headers=headers)
@@ -656,15 +388,15 @@ class VideoTaskPollerService:
def _build_headers( def _build_headers(
self, self,
api_format: APIFormat, endpoint_sig: str,
upstream_key: str, upstream_key: str,
endpoint: ProviderEndpoint, endpoint: ProviderEndpoint,
auth_info: ProviderAuthInfo | None = None, auth_info: ProviderAuthInfo | None = None,
) -> dict[str, str]: ) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint) extra_headers = get_extra_headers_from_endpoint(endpoint)
headers = build_upstream_headers( headers = build_upstream_headers_for_endpoint(
{}, {},
api_format, endpoint_sig,
upstream_key, upstream_key,
endpoint_headers=extra_headers, endpoint_headers=extra_headers,
) )

View File

@@ -1,7 +1,10 @@
from unittest.mock import MagicMock
from src.models.database import DimensionCollector from src.models.database import DimensionCollector
from src.services.billing.dimension_collector_service import ( from src.services.billing.dimension_collector_service import (
DimensionCollectInput, DimensionCollectInput,
DimensionCollectorRuntime, DimensionCollectorRuntime,
DimensionCollectorService,
) )
@@ -10,7 +13,7 @@ class TestDimensionCollectorRuntime:
runtime = DimensionCollectorRuntime() runtime = DimensionCollectorRuntime()
collectors = [ collectors = [
DimensionCollector( DimensionCollector(
api_format="OPENAI", api_format="openai:chat",
task_type="chat", task_type="chat",
dimension_name="input_tokens", dimension_name="input_tokens",
source_type="response", source_type="response",
@@ -20,7 +23,7 @@ class TestDimensionCollectorRuntime:
is_enabled=True, is_enabled=True,
), ),
DimensionCollector( DimensionCollector(
api_format="OPENAI", api_format="openai:chat",
task_type="chat", task_type="chat",
dimension_name="input_tokens", dimension_name="input_tokens",
source_type="response", source_type="response",
@@ -42,7 +45,7 @@ class TestDimensionCollectorRuntime:
runtime = DimensionCollectorRuntime() runtime = DimensionCollectorRuntime()
collectors = [ collectors = [
DimensionCollector( DimensionCollector(
api_format="GEMINI", api_format="gemini:video",
task_type="video", task_type="video",
dimension_name="file_size_mb", dimension_name="file_size_mb",
source_type="metadata", source_type="metadata",
@@ -63,7 +66,7 @@ class TestDimensionCollectorRuntime:
runtime = DimensionCollectorRuntime() runtime = DimensionCollectorRuntime()
collectors = [ collectors = [
DimensionCollector( DimensionCollector(
api_format="CLAUDE", api_format="claude:chat",
task_type="chat", task_type="chat",
dimension_name="input_tokens", dimension_name="input_tokens",
source_type="request", source_type="request",
@@ -73,7 +76,7 @@ class TestDimensionCollectorRuntime:
is_enabled=True, is_enabled=True,
), ),
DimensionCollector( DimensionCollector(
api_format="CLAUDE", api_format="claude:chat",
task_type="chat", task_type="chat",
dimension_name="cache_read_tokens", dimension_name="cache_read_tokens",
source_type="request", source_type="request",
@@ -83,7 +86,7 @@ class TestDimensionCollectorRuntime:
is_enabled=True, is_enabled=True,
), ),
DimensionCollector( DimensionCollector(
api_format="CLAUDE", api_format="claude:chat",
task_type="chat", task_type="chat",
dimension_name="total_input_tokens", dimension_name="total_input_tokens",
source_type="computed", source_type="computed",
@@ -103,3 +106,56 @@ class TestDimensionCollectorRuntime:
assert dims["input_tokens"] == 100 assert dims["input_tokens"] == 100
assert dims["cache_read_tokens"] == 20 assert dims["cache_read_tokens"] == 20
assert dims["total_input_tokens"] == 120 assert dims["total_input_tokens"] == 120
class TestDimensionCollectorService:
def test_video_fallback_merges_base_collectors(self) -> None:
db = MagicMock()
video_collectors = [
DimensionCollector(
api_format="openai:video",
task_type="video",
dimension_name="duration_seconds",
source_type="metadata",
source_path="task.duration_seconds",
value_type="int",
priority=0,
is_enabled=True,
)
]
base_collectors = [
# Should be kept (dimension not present in video_collectors)
DimensionCollector(
api_format="openai:chat",
task_type="video",
dimension_name="resolution",
source_type="metadata",
source_path="task.resolution",
value_type="string",
priority=0,
is_enabled=True,
),
# Should be ignored (dimension already present in video_collectors)
DimensionCollector(
api_format="openai:chat",
task_type="video",
dimension_name="duration_seconds",
source_type="metadata",
source_path="task.duration_seconds",
value_type="int",
priority=0,
is_enabled=True,
),
]
q1 = MagicMock()
q1.filter.return_value.all.return_value = video_collectors
q2 = MagicMock()
q2.filter.return_value.all.return_value = base_collectors
db.query.side_effect = [q1, q2]
svc = DimensionCollectorService(db)
result = svc.list_enabled_collectors(api_format="openai:video", task_type="video")
assert [c.dimension_name for c in result] == ["duration_seconds", "resolution"]

View File

@@ -0,0 +1,45 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
from src.services.video.task_poller import VideoTaskPollerService
@pytest.mark.asyncio
async def test_poll_task_status_routes_gemini_video_to_gemini(
monkeypatch: pytest.MonkeyPatch,
) -> None:
poller = VideoTaskPollerService()
task = SimpleNamespace(
endpoint_id="e1",
key_id="k1",
provider_api_format="gemini:video",
external_task_id="operations/123",
)
endpoint = SimpleNamespace(id="e1", base_url="https://example.com", api_format="gemini:video")
key = SimpleNamespace(id="k1", api_key="enc")
monkeypatch.setattr(poller, "_get_endpoint", lambda _db, _id: endpoint)
monkeypatch.setattr(poller, "_get_key", lambda _db, _id: key)
monkeypatch.setattr(
"src.services.video.task_poller.crypto_service.decrypt", lambda _v: "decrypted"
)
auth_info = SimpleNamespace(auth_header="authorization", auth_value="Bearer x")
monkeypatch.setattr(
"src.services.video.task_poller.get_provider_auth",
AsyncMock(return_value=auth_info),
)
poll_gemini = AsyncMock(return_value=InternalVideoPollResult(status=VideoStatus.PROCESSING))
poll_openai = AsyncMock(return_value=InternalVideoPollResult(status=VideoStatus.PROCESSING))
monkeypatch.setattr(poller, "_poll_gemini", poll_gemini)
monkeypatch.setattr(poller, "_poll_openai", poll_openai)
result = await poller._poll_task_status(MagicMock(), task)
assert result.status == VideoStatus.PROCESSING
assert poll_gemini.await_count == 1
assert poll_openai.await_count == 0

View File

@@ -0,0 +1,172 @@
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import pytest
from src.config.settings import config
from src.core.api_format.conversion.internal_video import VideoStatus
from src.services.billing.formula_engine import BillingIncompleteError
from src.services.task.impl.video_telemetry import VideoTelemetry
def _make_task(**overrides: Any) -> SimpleNamespace:
task = SimpleNamespace(
id="t1",
user_id="u1",
api_key_id="ak1",
provider_id="p1",
endpoint_id="e1",
key_id="k1",
external_task_id="ext-1",
client_api_format="openai:video",
provider_api_format="openai:video",
format_converted=False,
model="sora",
original_request_body={"model": "sora"},
duration_seconds=4,
resolution="720p",
aspect_ratio="16:9",
size="1024x1024",
retry_count=0,
video_size_bytes=None,
video_url="https://example.com/v.mp4",
video_urls=["https://example.com/v.mp4"],
submitted_at=None,
completed_at=None,
error_code=None,
error_message=None,
status=VideoStatus.COMPLETED.value,
request_metadata={"request_id": "req-1", "poll_raw_response": {"foo": "bar"}},
)
for k, v in overrides.items():
setattr(task, k, v)
return task
def _make_db() -> MagicMock:
db = MagicMock()
user_obj = SimpleNamespace(id="u1")
api_key_obj = SimpleNamespace(id="ak1")
provider_obj = SimpleNamespace(id="p1", name="prov1")
q_user = MagicMock()
q_user.filter.return_value.first.return_value = user_obj
q_key = MagicMock()
q_key.filter.return_value.first.return_value = api_key_obj
q_provider = MagicMock()
q_provider.filter.return_value.first.return_value = provider_obj
db.query.side_effect = [q_user, q_key, q_provider]
return db
@pytest.mark.asyncio
async def test_video_telemetry_failed_records_cost_zero(monkeypatch: pytest.MonkeyPatch) -> None:
db = _make_db()
task = _make_task(status=VideoStatus.FAILED.value, error_message="boom")
monkeypatch.setattr(
"src.services.task.impl.video_telemetry.DimensionCollectorService.collect_dimensions",
lambda _self, **_kwargs: {"duration_seconds": 4},
)
record = AsyncMock()
monkeypatch.setattr(
"src.services.task.impl.video_telemetry.UsageService.record_usage_with_custom_cost",
record,
)
telemetry = VideoTelemetry(db)
await telemetry.record_terminal_usage(task)
# billing_snapshot should be written back to task.request_metadata
assert task.request_metadata["billing_snapshot"]["billed_reason"] == "task_failed"
assert record.await_count == 1
kwargs = record.await_args.kwargs
assert kwargs["total_cost_usd"] == 0.0
assert kwargs["status"] == "failed"
@pytest.mark.asyncio
async def test_video_telemetry_completed_no_rule(monkeypatch: pytest.MonkeyPatch) -> None:
db = _make_db()
task = _make_task(status=VideoStatus.COMPLETED.value)
monkeypatch.setattr(
"src.services.task.impl.video_telemetry.DimensionCollectorService.collect_dimensions",
lambda _self, **_kwargs: {"duration_seconds": 4},
)
monkeypatch.setattr(
"src.services.task.impl.video_telemetry.BillingRuleService.find_rule",
lambda *_args, **_kwargs: None,
)
record = AsyncMock()
monkeypatch.setattr(
"src.services.task.impl.video_telemetry.UsageService.record_usage_with_custom_cost",
record,
)
telemetry = VideoTelemetry(db)
await telemetry.record_terminal_usage(task)
assert task.request_metadata["billing_snapshot"]["status"] == "no_rule"
kwargs = record.await_args.kwargs
assert kwargs["total_cost_usd"] == 0.0
assert kwargs["status"] == "completed"
@pytest.mark.asyncio
async def test_video_telemetry_strict_mode_missing_required_marks_failed(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = _make_db()
task = _make_task(status=VideoStatus.COMPLETED.value)
# provide a billing rule so formula path is taken
rule = SimpleNamespace(
id="r1",
name="video",
expression="duration_seconds",
variables={},
dimension_mappings={},
)
lookup = SimpleNamespace(rule=rule, scope="model", effective_task_type="video")
monkeypatch.setattr(
"src.services.task.impl.video_telemetry.DimensionCollectorService.collect_dimensions",
lambda _self, **_kwargs: {"duration_seconds": None},
)
monkeypatch.setattr(
"src.services.task.impl.video_telemetry.BillingRuleService.find_rule",
lambda *_args, **_kwargs: lookup,
)
record = AsyncMock()
monkeypatch.setattr(
"src.services.task.impl.video_telemetry.UsageService.record_usage_with_custom_cost",
record,
)
old = config.billing_strict_mode
try:
config.billing_strict_mode = True
telemetry = VideoTelemetry(db)
telemetry._formula_engine.evaluate = MagicMock(
side_effect=BillingIncompleteError(
"Missing required dimensions", missing_required=["duration_seconds"]
)
)
await telemetry.record_terminal_usage(task)
finally:
config.billing_strict_mode = old
assert task.status == VideoStatus.FAILED.value
assert task.video_url is None
assert task.video_urls is None
assert "billing_incomplete" in (task.error_code or "")
kwargs = record.await_args.kwargs
assert kwargs["total_cost_usd"] == 0.0
assert kwargs["status"] == "failed"