mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: 添加视频生成功能增强和多维度计费系统适配
- 视频生成: 增强 video_handler,重构 task_poller,新增 telemetry - 计费系统: 适配新的 signature 格式,支持 video 任务类型回退 - 数据库迁移: 添加 api_family/endpoint_kind 字段和 video_formats
This commit is contained in:
@@ -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 keys(rate_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")
|
||||||
@@ -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(
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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 --------------------
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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 = (
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
45
tests/services/test_video_task_poller.py
Normal file
45
tests/services/test_video_task_poller.py
Normal 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
|
||||||
172
tests/services/test_video_telemetry.py
Normal file
172
tests/services/test_video_telemetry.py
Normal 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"
|
||||||
Reference in New Issue
Block a user