feat(viral-video): 多模型支持(4 Seedance+Wan 3.0)+DashScope接入+GET /models端点 (#2159) (#2161)
CI/CD Pipeline / Check if frontend-only change (push) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (push) Successful in 2s
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 3s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 5s
CI/CD Pipeline / Frontend Lint (push) Has been skipped
CI/CD Pipeline / PR Build API Image (push) Has been skipped
CI/CD Pipeline / PR Build Web Image (push) Has been skipped
CI/CD Pipeline / PR Build Worker Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been cancelled
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been cancelled
CI/CD Pipeline / PR Build API Image (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Lint (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Web Image (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
AI Code Review / AI Code Review (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
PR Automation / Auto Approve on CI Green (pull_request) Has been cancelled
Preview Deploy / Deploy Preview Environment (pull_request) Has been cancelled
CI/CD Pipeline / Check push changed paths (push) Successful in 28s
CI/CD Pipeline / Frontend Unit Tests (push) Successful in 3m25s
CI/CD Pipeline / Integration Tests (push) Successful in 5m29s
CI/CD Pipeline / Validate - Python (mypy + alembic) (push) Successful in 5m46s
CI/CD Pipeline / Build Staging API Image (push) Successful in 1m21s
CI/CD Pipeline / Validate - Style (push) Successful in 6m22s
CI/CD Pipeline / Build Staging Web Image (push) Successful in 1m34s
CI/CD Pipeline / Build Staging Worker Image (push) Successful in 1m16s
CI/CD Pipeline / Retag skipped Staging API Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (push) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Successful in 1m19s
CI/CD Pipeline / ACR Image Cleanup (push) Successful in 1m35s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 1m55s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m14s
CI/CD Pipeline / Validate - Security (push) Successful in 12m29s
CI/CD Pipeline / Unit Tests (push) Successful in 13m34s
CI/CD Pipeline / Build Production API Image (push) Has been skipped
CI/CD Pipeline / Build Production Web Image (push) Has been skipped
CI/CD Pipeline / CI Gate (push) Has been skipped
CI/CD Pipeline / Build Production Worker Image (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Canary Release to Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped

Co-authored-by: xiaoxia <dev@xiaoxiajianji.com>
Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
This commit was merged in pull request #2161.
This commit is contained in:
2026-10-03 17:42:38 +08:00
committed by auto-approve-bot
parent 76e11cb2f9
commit 1b02df4d4a
10 changed files with 779 additions and 159 deletions
+27 -4
View File
@@ -49,7 +49,9 @@ from packages.adapters.sqlalchemy_impl.viral_video_repository import (
SQLAlchemyViralVideoJobRepository,
SQLAlchemyViralVideoStyleTemplateRepository,
)
from packages.domain.points_rules import list_viral_video_models
from packages.domain.viral_video import ViralVideoStatus
from packages.shared.dashscope_client import get_dashscope_client
logger = logging.getLogger(__name__)
@@ -411,7 +413,10 @@ def estimate_credits(
w, h = resolve_video_dimensions(resolution, ratio)
credits, bd = calculate_viral_video_credits_with_breakdown(
duration, w, h, model,
duration,
w,
h,
model,
)
breakdown = CreditsFormulaBreakdown(**bd)
return EstimateCreditsResponse(estimated_credits=credits, formula_breakdown=breakdown)
@@ -451,6 +456,17 @@ def list_style_templates(
return StyleTemplateListResponse(items=items)
@router.get("/models")
def list_available_models() -> dict:
"""返回爆款视频可用模型列表(供前端模型选择器使用)。"""
dashscope_available = get_dashscope_client() is not None
models = list_viral_video_models(
include_placeholder=False,
dashscope_available=dashscope_available,
)
return {"models": models}
@router.get("/{job_id}", response_model=ViralVideoJobResponse)
def get_viral_video_job(
job_id: str,
@@ -574,7 +590,9 @@ def retry_viral_video_job(
job.credits_prepaid = round(old_prepaid + diff, 2)
logger.info(
"[爆款视频][retry] 补扣差额 job_id=%s diff=%.2f new_prepaid=%.2f",
job.id, diff, job.credits_prepaid,
job.id,
diff,
job.credits_prepaid,
)
else:
# 新预扣更少:退还差额
@@ -591,7 +609,9 @@ def retry_viral_video_job(
job.credits_prepaid = round(old_prepaid - refund, 2)
logger.info(
"[爆款视频][retry] 退还差额 job_id=%s refund=%.2f new_prepaid=%.2f",
job.id, refund, job.credits_prepaid,
job.id,
refund,
job.credits_prepaid,
)
# 差额为 0 则不调整
@@ -611,7 +631,10 @@ def retry_viral_video_job(
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
logger.info(
"[爆款视频] 重试入队: job_id=%s retry_count=%d stale=%s params_changed=%s",
job.id, job.retry_count, is_stale_running, param_changed,
job.id,
job.retry_count,
is_stale_running,
param_changed,
)
except Exception as e:
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
+8 -2
View File
@@ -1126,6 +1126,7 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
返回 (本地视频路径, usage dict|None)。失败抛异常。
"""
from packages.domain.points_rules import get_viral_video_model_config
from packages.shared.ai_service import call_video_generation
prompt = _assemble_seedance_prompt(copy_result, job)
@@ -1133,6 +1134,9 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
ratio = getattr(job, "video_ratio", None) or "9:16"
model = getattr(job, "video_model", "") or None
resolution = getattr(job, "video_resolution", "720p") or "720p"
# 按模型配置决定是否开启音频生成(#2159 多模型支持)
_mcfg = get_viral_video_model_config(model)
gen_audio = bool(_mcfg.get("supports_audio", True))
# reference_audios: TTS 音频驱动口型
ref_audios = [tts_audio_url] if tts_audio_url else []
@@ -1145,10 +1149,12 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
tmpdir = Path(tempfile.mkdtemp(prefix=f"viral_{job.id}_"))
logger.info(
"[爆款视频] 开始单次 Seedance 生成 dur=%ds ratio=%s model=%s ref_imgs=%d ref_audios=%d ref_videos=%d tmpdir=%s",
"[爆款视频] 开始单次视频生成 dur=%ds ratio=%s model=%s provider=%s gen_audio=%s ref_imgs=%d ref_audios=%d ref_videos=%d tmpdir=%s",
dur,
ratio if not first_image else "(follow-image)",
model or "default",
_mcfg.get("provider", "doubao"),
gen_audio,
len(rest_images) + (1 if first_image else 0),
len(ref_audios),
len(ref_videos),
@@ -1164,7 +1170,7 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
resolution=resolution,
output_dir=str(tmpdir),
model=model,
generate_audio=True, # Seedance 原生生成环境音效/BGM;口型由 reference_audios 的 TTS 驱动
generate_audio=gen_audio, # 按模型能力:有声模型走原生音画同生;Wan 等需后配 TTS
reference_images=rest_images,
reference_audios=ref_audios,
reference_videos=ref_videos,
+6
View File
@@ -103,6 +103,12 @@ class SharedSettings(BaseSettings):
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
dashscope_api_key: str = ""
dashscope_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
dashscope_video_timeout: int = 900 # Wan 视频任务轮询总超时(秒)
dashscope_video_poll_interval: int = 10
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
mediakit_api_key: str = ""
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
+136 -13
View File
@@ -10,7 +10,9 @@ from __future__ import annotations
import math
# ============ 爆款视频动态定价 (#2151) ============
# key = (model_id, resolution, has_video_input),单位:元/百万token
# key = (model_id, resolution, has_video_input),单位:
# - billing_mode=token: 元/百万tokens(输出)
# - billing_mode=per_second: 元/秒(视频时长)
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("seedance-2.5", "480p", False): 70.0,
("seedance-2.5", "720p", False): 70.0,
@@ -21,6 +23,14 @@ VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("seedance-2.0", "480p", False): 46.0,
("seedance-2.0", "720p", False): 46.0,
("seedance-2.0", "1080p", False): 51.0,
("seedance-2.0", "4k", False): 80.0,
("seedance-2.0-fast", "480p", False): 28.0,
("seedance-2.0-fast", "720p", False): 28.0,
("seedance-2.0-mini", "480p", False): 9.2,
("seedance-2.0-mini", "720p", False): 9.2,
("wan-3.0", "480p", False): 0.3,
("wan-3.0", "720p", False): 0.6,
("wan-3.0", "1080p", False): 1.2,
}
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器
@@ -49,7 +59,10 @@ _RESOLUTION_ALIASES: dict[str, str] = {
"fhd": "1080p",
}
# 分辨率 -> 短边像素数(p 值代表短边,不是 height)
_RESOLUTION_SHORT_SIDE: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080}
_RESOLUTION_SHORT_SIDE: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080, "4k": 2160}
_RESOLUTION_ALIASES["4k"] = "4k"
_RESOLUTION_ALIASES["2160p"] = "4k"
_RESOLUTION_ALIASES["uhd"] = "4k"
def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
@@ -81,18 +94,116 @@ def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
return int(w), int(h)
def _match_model_prefix(model: str) -> str:
"""匹配 model 前缀。"""
m = (model or "").strip().lower()
for prefix in ("seedance-2.5", "seedance-2.0"):
if m.startswith(prefix):
return prefix
# ── 爆款视频多模型元数据 (#2159) ──────────────────────────────────────
VIRAL_VIDEO_MODEL_CONFIG: dict[str, dict] = {
"seedance-2.5": {
"key": "seedance-2.5",
"display_name": "Seedance 2.5(最新模型)",
"model_id": "doubao-seedance-2-5-260628",
"provider": "doubao",
"supports_audio": True,
"supported_resolutions": ["480p", "720p", "1080p"],
"max_duration": 15,
"billing_mode": "token",
"is_default": True,
},
"seedance-2.0": {
"key": "seedance-2.0",
"display_name": "Seedance 2.0 完整版(正式投放首选)",
"model_id": "doubao-seedance-2-0-260128",
"provider": "doubao",
"supports_audio": True,
"supported_resolutions": ["480p", "720p", "1080p", "4k"],
"max_duration": 15,
"billing_mode": "token",
"is_default": False,
},
"seedance-2.0-fast": {
"key": "seedance-2.0-fast",
"display_name": "Seedance 2.0 Fast — 快速低成本",
"model_id": "doubao-seedance-2-0-fast-260128",
"provider": "doubao",
"supports_audio": True,
"supported_resolutions": ["480p", "720p"],
"max_duration": 10,
"billing_mode": "token",
"is_default": False,
},
"seedance-2.0-mini": {
"key": "seedance-2.0-mini",
"display_name": "Seedance 2.0 Mini — 低成本",
"model_id": "doubao-seedance-2-0-mini-260615",
"provider": "doubao",
"supports_audio": True,
"supported_resolutions": ["480p", "720p"],
"max_duration": 10,
"billing_mode": "token",
"is_default": False,
},
"wan-3.0": {
"key": "wan-3.0",
"display_name": "Wan 3.0 — 低成本(DashScope)",
"model_id": "wan3.0-video",
"provider": "dashscope",
"supports_audio": True,
"supported_resolutions": ["480p", "720p", "1080p"],
"max_duration": 10,
"billing_mode": "per_second",
"is_default": False,
},
}
def get_viral_video_model_config(model_key: str | None) -> dict:
"""获取模型配置,未知 key 回落到默认 seedance-2.5。"""
key = (model_key or "").strip().lower()
if key and key in VIRAL_VIDEO_MODEL_CONFIG:
return VIRAL_VIDEO_MODEL_CONFIG[key]
return VIRAL_VIDEO_MODEL_CONFIG["seedance-2.5"]
def list_viral_video_models(
include_placeholder: bool = False,
dashscope_available: bool = False,
) -> list[dict]:
"""返回前端可用的模型列表(供 GET /api/v1/viral-video/models 端点用)。"""
out: list[dict] = []
for _k, cfg in VIRAL_VIDEO_MODEL_CONFIG.items():
if cfg.get("_placeholder") and not include_placeholder:
continue
if cfg.get("provider") == "dashscope" and not dashscope_available:
continue
out.append(
{
"key": cfg["key"],
"display_name": cfg["display_name"],
"supports_audio": bool(cfg.get("supports_audio", True)),
"supported_resolutions": list(cfg.get("supported_resolutions", ["720p"])),
"max_duration": int(cfg.get("max_duration", 15)),
"billing_mode": cfg.get("billing_mode", "token"),
"is_default": bool(cfg.get("is_default", False)),
}
)
return out
def _match_model_prefix(model: str | None) -> str:
"""匹配 model key(支持全部内部别名,未知回落到 seedance-2.5)。
按 key 长度从长到短匹配,避免 "seedance-2.0-fast" 被 "seedance-2.0" 前缀命中。
"""
mm = (model or "").strip().lower()
for k in sorted(VIRAL_VIDEO_MODEL_CONFIG.keys(), key=len, reverse=True):
if mm == k or mm.startswith(k):
return k
return "seedance-2.5"
def _infer_resolution_key(width: int, height: int) -> str:
"""从实际 (width, height) 用短边推断 resolution key。"""
short = min(int(width or 720), int(height or 720))
if short >= 1900:
return "4k"
if short >= 1000:
return "1080p"
if short >= 650:
@@ -128,18 +239,26 @@ def calculate_viral_video_credits_with_breakdown(
effective_fps = int(fps or VIRAL_VIDEO_FPS)
prefix = _match_model_prefix(model)
cfg = get_viral_video_model_config(prefix)
res_key = _infer_resolution_key(w, h)
billing = cfg.get("billing_mode", "token")
key = (prefix, res_key, bool(has_video_input))
price = VIRAL_VIDEO_MODEL_PRICES.get(key)
if price is None:
price = VIRAL_VIDEO_MODEL_PRICES.get(("seedance-2.5", res_key, False), 70.0)
if actual_tokens is not None and actual_tokens > 0:
tokens = float(actual_tokens)
dur = max(1, int(duration_seconds or 15))
if billing == "per_second":
tokens = 0.0
video_cost = dur * float(price)
billing_unit = "second"
else:
dur = max(1, int(duration_seconds or 15))
tokens = dur * w * h * effective_fps / 1024.0
if actual_tokens is not None and actual_tokens > 0:
tokens = float(actual_tokens)
else:
tokens = dur * w * h * effective_fps / 1024.0
video_cost = tokens / 1_000_000.0 * float(price)
billing_unit = "token"
video_cost = tokens / 1_000_000.0 * float(price)
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
credits = round(float(total), 2)
breakdown = {
@@ -148,9 +267,13 @@ def calculate_viral_video_credits_with_breakdown(
"fixed_cost": float(VIRAL_VIDEO_FIXED_COST),
"profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER),
"model_price": float(price),
"model_key": prefix,
"billing_mode": billing,
"billing_unit": billing_unit,
"width": int(w),
"height": int(h),
"fps": int(effective_fps),
"duration": dur,
}
return credits, breakdown
+59 -32
View File
@@ -25,39 +25,46 @@ from packages.shared.config import get_shared_settings
logger = logging.getLogger(__name__)
# 内部 pricing key → 火山方舟实际模型 ID
# 前端/后端内部用简短别名 seedance-2.5/seedance-2.0 做 VIRAL_VIDEO_MODEL_PRICES 配置 key,
# 但传给方舟 API 时必须用真实模型 ID(如 doubao-seedance-2-5-260628),否则会 404 InvalidEndpointOrModel.NotFound。
VIRAL_VIDEO_MODEL_ID_MAP: dict[str, str] = {
"seedance-2.5": "doubao-seedance-2-5-260628",
"seedance-2.0": "doubao-seedance-2-0-250628",
}
# 视频模型 ID 解析逻辑(#2159 多模型支持)。
# 内部使用简短别名(seedance-2.5 / seedance-2.0 / wan-3.0 等)做 PRICING key 和前端选择值;
# 实际调用时按 VIRAL_VIDEO_MODEL_CONFIG(domain/points_rules.py)的 provider/model_id 分发。
# - provider=doubao → 火山方舟
# - provider=dashscope → 阿里云 DashScope(Wan 系列)
def _resolve_video_provider_and_id(model: str | None) -> tuple[str, str, dict]:
"""把内部 model key 解析成 (provider, model_id, cfg)。
- provider: "doubao" | "dashscope"
- model_id: 对应 API 的真实模型 ID
- cfg: VIRAL_VIDEO_MODEL_CONFIG 条目
未识别或空值回落到默认 seedance-2.5。已经是 doubao-/ep- 开头的完整 ID 视为 doubao provider。
"""
from packages.domain.points_rules import get_viral_video_model_config
settings = get_shared_settings()
default_id = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
m = (model or "").strip()
if not m:
cfg = get_viral_video_model_config("seedance-2.5")
return "doubao", default_id, cfg
# 已经是 doubao-/ep- 开头:直接透传,默认视为 doubao provider
if m.startswith("doubao-") or m.startswith("ep-"):
return "doubao", m, {"provider": "doubao", "model_id": m, "supports_audio": True}
# 别名 → 从 domain config 查
cfg = get_viral_video_model_config(m)
provider = cfg.get("provider", "doubao")
resolved_id = cfg.get("model_id", "")
if not resolved_id:
logger.warning("[ai_client] model %r 无 model_id,回落到默认 %s", m, default_id)
return "doubao", default_id, cfg
return provider, resolved_id, cfg
def _resolve_video_model_id(model: str | None) -> str:
"""把内部 pricing key(seedance-2.5/seedance-2.0)解析成方舟实际模型 ID。
未识别的别名或空值回落到 settings.doubao_video_model 默认值。
"""
settings = get_shared_settings()
default_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
m = (model or "").strip()
if not m:
return default_model
# 已经是 doubao- 前缀的完整 ID,直接用
if m.startswith("doubao-") or m.startswith("ep-"):
return m
# 别名 → 完整 ID
resolved = VIRAL_VIDEO_MODEL_ID_MAP.get(m.lower())
if resolved:
return resolved
# 未识别:带 seedance 字样且不是 doubao-/ep- 开头的,拼前缀兜底
if "seedance" in m.lower():
# seedance-2.5 → doubao-seedance-2-5-260628 等兜底逻辑
normalized = m.lower().replace(".", "-")
if normalized in VIRAL_VIDEO_MODEL_ID_MAP:
return VIRAL_VIDEO_MODEL_ID_MAP[normalized]
logger.warning("[ai_client] 未识别的 video_model=%r,使用默认模型 %s", m, default_model)
return default_model
"""兼容旧调用:只返回 doubao model_id。wan/dashscope 调用方应直接用 _resolve_video_provider_and_id。"""
_provider, mid, _cfg = _resolve_video_provider_and_id(model)
return mid
class DoubaoClient:
@@ -317,8 +324,28 @@ class DoubaoClient:
poll_interval = getattr(settings, "doubao_video_poll_interval", 10) or 10
# 收紧总超时:轮询 8min + 下载 2min = 最长 ~10min,防止出现 20min 卡死
total_timeout = getattr(settings, "doubao_video_timeout", 480) or 480
# 内部 pricing key (seedance-2.5/seedance-2.0) → 方舟实际模型 ID
video_model = _resolve_video_model_id(model)
# 内部 key → (provider, 实际模型 ID, cfg),按 provider 分发
provider, video_model, model_cfg = _resolve_video_provider_and_id(model)
if provider == "dashscope":
from packages.shared.dashscope_client import get_dashscope_client
ds = get_dashscope_client()
if ds is None:
logger.error("DashScope client 不可用(未配置 DASHSCOPE_API_KEY),video_model=%s", model)
return None
try:
return ds.video_generation(
prompt=prompt,
image_url=image_url,
duration=duration,
ratio=ratio,
resolution=resolution,
output_dir=output_dir,
model=video_model,
)
except Exception as de:
logger.error("DashScope video_generation 异常: %s", de, exc_info=True)
return None
ref_audios = [u for u in (reference_audios or [])[:10] if u and isinstance(u, str)]
ref_videos = [u for u in (reference_videos or [])[:3] if u and isinstance(u, str)]
+221
View File
@@ -0,0 +1,221 @@
"""DashScope 客户端(阿里云百炼 Wan 3.0 等非方舟模型)。
#2159: 新增 Wan 3.0 视频生成支持。DashScope 异步协议:
- POST {base_url}/services/aigc/video-generation/video-synthesis (X-DashScope-Async: enable)
→ 返回 output.task_id
- GET {base_url}/tasks/{task_id} 轮询状态
→ SUCCEEDED 时 output.video_url 可下载
认证:Authorization: Bearer {DASHSCOPE_API_KEY}
"""
from __future__ import annotations
import logging
import os
import time
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
import httpx
from packages.shared.config import get_shared_settings
logger = logging.getLogger(__name__)
_DASHSCOPE_CLIENT_SINGLETON: "DashScopeClient | None" = None
class DashScopeClient:
"""阿里云 DashScope 异步 API 客户端(Wan 3.0 等视频生成)。"""
def __init__(self) -> None:
settings = get_shared_settings()
self.api_key: str = getattr(settings, "dashscope_api_key", "") or os.getenv("DASHSCOPE_API_KEY", "")
self.base_url: str = (
getattr(settings, "dashscope_base_url", "") or "https://dashscope.aliyuncs.com/api/v1"
).rstrip("/")
self.poll_interval: int = int(getattr(settings, "dashscope_video_poll_interval", 10) or 10)
self.total_timeout: int = int(getattr(settings, "dashscope_video_timeout", 900) or 900)
self.max_retries: int = 2
@property
def is_available(self) -> bool:
return bool(self.api_key)
def video_generation(
self,
prompt: str,
*,
image_url: str | None = None,
duration: int = 5,
ratio: str | None = "9:16",
resolution: str = "720p",
watermark: bool = False,
output_dir: str | None = None,
model: str = "wan3.0-video",
) -> dict | None:
"""调用 DashScope 异步视频合成接口,轮询完成后下载到本地。
返回 {"video_path": str, "usage": dict | None};失败返回 None。
"""
if not self.is_available:
logger.error("[dashscope] API key 未配置,无法调用视频生成")
return None
if not prompt or not prompt.strip():
return None
# DashScope 分辨率参数:720P / 1080P / 480P(大写 P)
res_upper = (resolution or "720p").upper().replace("P", "P")
if res_upper == "480P":
ds_res = "480P"
elif res_upper == "1080P":
ds_res = "1080P"
else:
ds_res = "720P"
# 构造 input+parameters
input_obj: dict[str, Any] = {"prompt": prompt.strip()}
if image_url:
input_obj["img_url"] = image_url
params: dict[str, Any] = {
"resolution": ds_res,
"duration": str(float(duration)),
"watermark": bool(watermark),
}
# 比例透传:Wan 支持 "9:16" / "16:9" / "1:1" 等
if ratio and ratio != "adaptive":
params["aspect_ratio"] = ratio
payload: dict[str, Any] = {
"model": model,
"input": input_obj,
"parameters": params,
}
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
"X-DashScope-Async": "enable",
}
create_url = f"{self.base_url}/services/aigc/video-generation/video-synthesis"
logger.info(
"[dashscope] 创建任务: model=%s dur=%ds ratio=%s res=%s img=%s",
model,
duration,
ratio,
ds_res,
bool(image_url),
)
# 创建任务
task_id: str | None = None
last_err: Exception | None = None
for attempt in range(self.max_retries + 1):
try:
resp = httpx.post(create_url, headers=headers, json=payload, timeout=60)
sc = int(getattr(resp, "status_code", 0) or 0)
body_text = (getattr(resp, "text", "") or "")[:1500]
if sc >= 400:
logger.error("[dashscope] 创建任务 HTTP %d: %s", sc, body_text)
resp.raise_for_status()
data = resp.json()
tid = (data.get("output") or {}).get("task_id")
if tid:
task_id = tid
break
# 部分情况下 code != 错误
code = data.get("code")
if code and code != "":
last_err = RuntimeError(f"dashscope create failed: {body_text[:300]}")
else:
last_err = RuntimeError(f"create ok but no task_id: {str(data)[:300]}")
except Exception as e:
last_err = e
if attempt < self.max_retries:
time.sleep(0.5 * (2**attempt))
continue
logger.error("[dashscope] 创建任务最终失败: %s", last_err)
return None
if not task_id:
return None
# 轮询任务
poll_url = f"{self.base_url}/tasks/{task_id}"
deadline = time.time() + self.total_timeout
video_url: str | None = None
usage: dict | None = None
while time.time() < deadline:
try:
r = httpx.get(poll_url, headers=headers, timeout=30)
if int(getattr(r, "status_code", 0) or 0) >= 400:
logger.warning("[dashscope] 轮询 HTTP %d", r.status_code)
time.sleep(self.poll_interval)
continue
d = r.json()
out = d.get("output") or {}
task_status = out.get("task_status") or d.get("task_status") or ""
if task_status == "SUCCEEDED":
video_url = out.get("video_url") or ""
usage = d.get("usage")
if not video_url:
# 结果在 results 数组
results = out.get("results") or []
if results and isinstance(results, list):
video_url = results[0].get("url") or results[0].get("video_url")
if video_url:
logger.info("[dashscope] 任务 %s 完成: %s", task_id, video_url[:120])
break
logger.error("[dashscope] 任务 %s SUCCEEDED 但无 video_url", task_id)
return None
if task_status in ("FAILED", "FAILED_WITH_ERROR", "ERROR"):
msg = out.get("message") or d.get("message") or "unknown error"
logger.error("[dashscope] 任务 %s 失败: %s", task_id, msg)
return None
if task_status in ("CANCELED", "CANCELLED"):
logger.warning("[dashscope] 任务 %s 被取消", task_id)
return None
# PENDING / RUNNING / SUSPENDED → 继续轮询
logger.debug("[dashscope] 任务 %s 状态 %s,继续轮询", task_id, task_status)
except Exception as e:
logger.warning("[dashscope] 轮询异常: %s", e)
time.sleep(self.poll_interval)
if not video_url:
logger.error("[dashscope] 任务 %s 轮询超时(%ds)", task_id, self.total_timeout)
return None
# 下载视频
out_dir = output_dir or os.path.join(os.getcwd(), "seedance_outputs")
os.makedirs(out_dir, exist_ok=True)
suffix = Path(urlparse(video_url).path).suffix or ".mp4"
if suffix.lower() not in (".mp4", ".mov", ".webm"):
suffix = ".mp4"
safe_tid = "".join(c if c.isalnum() or c in "-_" else "_" for c in task_id)[:40]
out_path = os.path.join(out_dir, f"wan_{safe_tid}{suffix}")
try:
with httpx.stream("GET", video_url, timeout=300, follow_redirects=True) as resp:
if int(getattr(resp, "status_code", 0) or 0) >= 400:
logger.error("[dashscope] 下载 HTTP %d", resp.status_code)
return None
with open(out_path, "wb") as f:
for chunk in resp.iter_bytes(chunk_size=1024 * 256):
if chunk:
f.write(chunk)
except Exception as e:
logger.error("[dashscope] 下载视频失败: %s", e)
return None
size = os.path.getsize(out_path) if os.path.exists(out_path) else 0
if size < 1024:
logger.error("[dashscope] 下载文件过小: %d bytes", size)
return None
logger.info("[dashscope] 视频已下载: %s (%d bytes)", out_path, size)
return {"video_path": out_path, "usage": usage}
def get_dashscope_client() -> DashScopeClient | None:
"""返回 DashScope 客户端单例;未配置 API key 时返回 None。"""
global _DASHSCOPE_CLIENT_SINGLETON
if _DASHSCOPE_CLIENT_SINGLETON is None:
_DASHSCOPE_CLIENT_SINGLETON = DashScopeClient()
if not _DASHSCOPE_CLIENT_SINGLETON.is_available:
return None
return _DASHSCOPE_CLIENT_SINGLETON
+23 -17
View File
@@ -522,7 +522,23 @@ class TestResolveVideoModelId:
def test_seedance_2_0_alias(self):
fn = self._import_target()
assert fn("seedance-2.0") == "doubao-seedance-2-0-250628"
assert fn("seedance-2.0") == "doubao-seedance-2-0-260128"
def test_seedance_2_0_fast_alias(self):
fn = self._import_target()
assert fn("seedance-2.0-fast") == "doubao-seedance-2-0-fast-260128"
def test_seedance_2_0_mini_alias(self):
fn = self._import_target()
assert fn("seedance-2.0-mini") == "doubao-seedance-2-0-mini-260615"
def test_wan_3_0_returns_dashscope_provider(self):
from packages.shared.ai_client import _resolve_video_provider_and_id
prov, mid, cfg = _resolve_video_provider_and_id("wan-3.0")
assert prov == "dashscope"
assert mid == "wan3.0-video"
assert cfg.get("billing_mode") == "per_second"
def test_seedance_2_5_uppercase(self):
fn = self._import_target()
@@ -530,23 +546,13 @@ class TestResolveVideoModelId:
def test_seedance_dot_normalize(self):
fn = self._import_target()
# 传 "seedance-2.5" 带 dot 走 MAP.get 已命中;
# 构造带 seedance 但 key 变体的兜底场景
with patch(
"packages.shared.ai_client.VIRAL_VIDEO_MODEL_ID_MAP",
{
"seedance-2-5": "doubao-seedance-2-5-260628",
},
):
assert fn("seedance-2.5") == "doubao-seedance-2-5-260628"
# dot 形式 "seedance-2.5" 直接命中 domain config 的 key(与 2-5 同等)
assert fn("seedance-2.5") == "doubao-seedance-2-5-260628"
def test_unknown_seedance_falls_back_to_default_and_warns(self, caplog):
def test_unknown_model_falls_back_to_default_seedance_2_5(self, caplog):
fn = self._import_target()
import logging
with patch("packages.shared.ai_client.get_shared_settings") as ms:
ms.return_value = MagicMock(doubao_video_model="doubao-seedance-2-5-260628")
with caplog.at_level(logging.WARNING, logger="shared.ai_client"):
# 任何不认识的别名
assert fn("seedance-9.9") == "doubao-seedance-2-5-260628"
assert any("未识别" in r.message for r in caplog.records if "ai_client" in r.name)
# 未知 model key 会通过 get_viral_video_model_config 回落到 seedance-2.5
with caplog.at_level(logging.WARNING, logger="shared.ai_client"):
assert fn("some-random-model") == "doubao-seedance-2-5-260628"
+159
View File
@@ -0,0 +1,159 @@
"""tests for packages/shared/dashscope_client.py (#2159 Wan 3.0 DashScope client)."""
from __future__ import annotations
from unittest.mock import MagicMock, mock_open, patch
import pytest
_SINGLETON = "_DASHSCOPE_CLIENT_SINGLETON"
@pytest.fixture(autouse=True)
def reset_singleton():
import packages.shared.dashscope_client as d
# 兼容实际 singleton 名
for name in ("_DASHSCOPE_CLIENT_SINGLETON", "_dashscope_client"):
if hasattr(d, name):
setattr(d, name, None)
yield
for name in ("_DASHSCOPE_CLIENT_SINGLETON", "_dashscope_client"):
if hasattr(d, name):
setattr(d, name, None)
def _make_settings(api_key="test-key"):
return MagicMock(
dashscope_api_key=api_key,
dashscope_base_url="https://dashscope.aliyuncs.com/api/v1",
dashscope_video_timeout=10,
dashscope_video_poll_interval=0,
video_dir="/tmp/videos",
)
class TestDashScopeAvailability:
def test_unavailable_without_key(self):
from packages.shared.dashscope_client import get_dashscope_client
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
ms.return_value = _make_settings(api_key="")
assert get_dashscope_client() is None
def test_available_with_key(self):
from packages.shared.dashscope_client import get_dashscope_client
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = get_dashscope_client()
assert c is not None
assert c.is_available is True
def _mock_stream_response(min_size=2048):
"""构造 httpx.stream 上下文返回值,模拟返回若干字节的 mp4 内容。"""
m = MagicMock()
m.status_code = 200
chunk = b"x" * min_size
m.iter_bytes.return_value = [chunk]
ctx = MagicMock()
ctx.__enter__.return_value = m
return ctx
class TestDashScopeVideoGeneration:
def test_happy_path_returns_video_path(self):
"""POST create → GET poll (SUCCEEDED) → download → returns path + correct payload."""
import packages.shared.dashscope_client as d
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = d.DashScopeClient()
create_resp = MagicMock(status_code=200)
create_resp.json.return_value = {"output": {"task_id": "task-abc"}}
poll_resp = MagicMock(status_code=200)
poll_resp.json.return_value = {
"output": {"task_status": "SUCCEEDED", "video_url": "http://x/y.mp4"},
"usage": {"billed_duration": 10},
}
# fake file: write enough bytes to pass the size>=1024 check
m_open = mock_open()
m_open.return_value.write.return_value = None
fake_size = {"/tmp/videos/wan_task-abc.mp4": 4096}
def fake_getsize(p):
return fake_size.get(p, 0)
def fake_exists(p):
return p in fake_size
with (
patch.object(d.httpx, "post", return_value=create_resp) as mock_post,
patch.object(d.httpx, "get", return_value=poll_resp),
patch.object(d.httpx, "stream", return_value=_mock_stream_response()),
patch("packages.shared.dashscope_client.time.sleep"),
patch("packages.shared.dashscope_client.os.makedirs"),
patch("builtins.open", m_open),
patch("packages.shared.dashscope_client.os.path.getsize", side_effect=fake_getsize),
patch("packages.shared.dashscope_client.os.path.exists", side_effect=fake_exists),
):
res = c.video_generation(
prompt="test",
duration=5,
ratio="9:16",
resolution="720p",
output_dir="/tmp/videos",
)
assert res is not None, "expected success"
assert res["video_path"] == "/tmp/videos/wan_task-abc.mp4"
_, kwargs = mock_post.call_args
body = kwargs["json"]
assert body["parameters"]["resolution"] == "720P"
assert body["model"] == "wan3.0-video"
def test_create_http_error_returns_none(self):
import packages.shared.dashscope_client as d
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = d.DashScopeClient()
err_resp = MagicMock(status_code=400, text="bad")
err_resp.raise_for_status.side_effect = RuntimeError("bad")
with patch.object(d.httpx, "post", return_value=err_resp):
res = c.video_generation(
prompt="test", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos"
)
assert res is None
def test_poll_failed_returns_none(self):
import packages.shared.dashscope_client as d
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = d.DashScopeClient()
create_resp = MagicMock(status_code=200)
create_resp.json.return_value = {"output": {"task_id": "task-abc"}}
poll_resp = MagicMock(status_code=200)
poll_resp.json.return_value = {"output": {"task_status": "FAILED", "message": "nope"}}
with (
patch.object(d.httpx, "post", return_value=create_resp),
patch.object(d.httpx, "get", return_value=poll_resp),
patch("packages.shared.dashscope_client.time.sleep"),
):
res = c.video_generation(
prompt="test", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos"
)
assert res is None
def test_empty_prompt_returns_none(self):
import packages.shared.dashscope_client as d
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = d.DashScopeClient()
assert (
c.video_generation(prompt=" ", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos")
is None
)
+111 -1
View File
@@ -196,10 +196,24 @@ class TestResolveVideoDimensions:
"""未知分辨率字符串兜底到 720p。"""
from packages.domain.points_rules import resolve_video_dimensions
w, h = resolve_video_dimensions("2160p", "1:1")
w, h = resolve_video_dimensions("garbage-xxx", "1:1")
assert h == 720
assert w == 720
def test_4k_16_9(self):
"""#2159 4k 横屏:短边=height=2160,width=3840。"""
from packages.domain.points_rules import resolve_video_dimensions
w, h = resolve_video_dimensions("4k", "16:9")
assert (w, h) == (3840, 2160)
def test_2160p_alias(self):
"""2160p 别名→4k。"""
from packages.domain.points_rules import resolve_video_dimensions
w, h = resolve_video_dimensions("2160p", "9:16")
assert (w, h) == (2160, 3840)
def test_empty_resolution_defaults_to_720p_9_16(self):
"""空 resolution + 空 ratio → 默认 720p + 9:16 竖屏 (720×1280)。"""
from packages.domain.points_rules import resolve_video_dimensions
@@ -531,3 +545,99 @@ class TestViralVideoCreditsWithBreakdown:
assert bd["tokens"] == 1_000_000.0
# video_cost = 1M/1M * 70 = 70; total = (70+0.15)*1.3 = 91.195 → 91.20
assert c == 91.20
# ============ #2159 多模型定价单测 ============
class TestMultiModelCredits:
"""#2159 多模型积分估算正确性(含 token/second 两种计费模式)。"""
def test_seedance_2_5_15s_720p_9x16(self):
from packages.domain.points_rules import calculate_viral_video_credits
# 15s/720p/9:16 → 720×1280
# tokens = 15*720*1280*24/1024 = 324000
# video_cost = 324000/1M*70 = 22.68
# total = (22.68+0.15)*1.3 = 29.679 ≈ 29.68
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.5")
assert c == 29.68, f"got {c}"
def test_seedance_2_0_30s_1080p_9x16(self):
# 30s/1080p/9:16 → 1080×1920
# tokens = 30*1080*1920*24/1024 = 1,458,000
# video_cost = 1.458M/1M*51 = 74.358
# total = (74.358+0.15)*1.3 = 96.86
from packages.domain.points_rules import calculate_viral_video_credits
c = calculate_viral_video_credits(30, 1080, 1920, model="seedance-2.0")
assert c == 96.86, f"got {c}"
def test_seedance_2_0_fast_15s_720p_9x16(self):
# 15s/720p/9:16 tokens=324000, price=28
# video_cost = 0.324*28 = 9.072
# total = (9.072+0.15)*1.3 = 11.99
from packages.domain.points_rules import calculate_viral_video_credits
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.0-fast")
assert c == 11.99, f"got {c}"
def test_seedance_2_0_mini_15s_720p_9x16(self):
# price=9.2, tokens=324000
# video_cost = 0.324*9.2 = 2.9808
# total = (2.9808+0.15)*1.3 = 4.07
from packages.domain.points_rules import calculate_viral_video_credits
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.0-mini")
assert c == 4.07, f"got {c}"
def test_wan_3_0_per_second_billing(self):
# per_second: 10s/720p price=0.6元/秒
# video_cost = 10*0.6 = 6.0
# total = (6.0+0.15)*1.3 = 7.995 ≈ 8.00
from packages.domain.points_rules import calculate_viral_video_credits
c = calculate_viral_video_credits(10, 720, 1280, model="wan-3.0")
assert c == 8.0, f"got {c}"
def test_seedance_2_0_4k_16x9(self):
# 5s/4k/16:9 → 3840×2160, price=80
# tokens = 5*3840*2160*24/1024 = 972000
# video_cost = 0.972*80 = 77.76
# total = (77.76+0.15)*1.3 = 101.28
from packages.domain.points_rules import calculate_viral_video_credits
c = calculate_viral_video_credits(5, 3840, 2160, model="seedance-2.0")
assert c == 101.28, f"got {c}"
def test_model_config_has_all_6_models(self):
from packages.domain.points_rules import VIRAL_VIDEO_MODEL_CONFIG
expected = {"seedance-2.5", "seedance-2.0", "seedance-2.0-fast", "seedance-2.0-mini", "wan-3.0"}
assert expected.issubset(set(VIRAL_VIDEO_MODEL_CONFIG.keys()))
def test_list_models_hides_wan_when_dashscope_unavailable(self):
from packages.domain.points_rules import list_viral_video_models
all_models = list_viral_video_models(include_placeholder=False, dashscope_available=False)
keys = {m["key"] for m in all_models}
assert "wan-3.0" not in keys
assert "seedance-2.5" in keys
# is_default
defaults = [m for m in all_models if m["is_default"]]
assert len(defaults) == 1
assert defaults[0]["key"] == "seedance-2.5"
def test_list_models_includes_wan_when_dashscope_available(self):
from packages.domain.points_rules import list_viral_video_models
models = list_viral_video_models(include_placeholder=False, dashscope_available=True)
keys = {m["key"] for m in models}
assert "wan-3.0" in keys
def test_infer_4k(self):
from packages.domain.points_rules import _infer_resolution_key
assert _infer_resolution_key(3840, 2160) == "4k"
assert _infer_resolution_key(2160, 3840) == "4k"
assert _infer_resolution_key(1920, 1080) == "1080p"
+29 -90
View File
@@ -1,99 +1,38 @@
"""爆款视频 DB 模型单元测试(#2039 PR1:DB + migration)。
验证:
- 3 张新表可在内存 SQLite 上创建
- 默认值与基本 CRUD 正常
"""
"""tests for GET /api/v1/viral-video/models route function (#2159)."""
from __future__ import annotations
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.models import (
Base,
ViralVideoJobModel,
ViralVideoPromptTemplateModel,
ViralVideoStyleTemplateModel,
)
from unittest.mock import MagicMock, patch
def _make_session():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(engine)
return sessionmaker(bind=engine)()
class TestModelsRoute:
def _call(self):
from apps.api.app.api.routes import viral_video as routes
return routes.list_available_models()
class TestViralVideoJobModel:
def test_create_and_get(self):
session = _make_session()
job = ViralVideoJobModel(
id="job-001",
user_id="user-001",
images=["https://img.com/1.jpg"],
industry="美妆",
duration=60,
)
session.add(job)
session.commit()
def test_without_dashscope_hides_wan(self):
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=None):
result = self._call()
assert "models" in result
keys = {m["key"] for m in result["models"]}
assert "seedance-2.5" in keys
assert "wan-3.0" not in keys
for m in result["models"]:
if m["key"].startswith("seedance"):
assert m["supports_audio"] is True
fetched = session.query(ViralVideoJobModel).filter_by(id="job-001").one()
assert fetched.user_id == "user-001"
assert fetched.images == ["https://img.com/1.jpg"]
assert fetched.industry == "美妆"
assert fetched.duration == 60
def test_with_dashscope_includes_wan(self):
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=MagicMock()):
result = self._call()
keys = {m["key"] for m in result["models"]}
assert "wan-3.0" in keys
wan = next(m for m in result["models"] if m["key"] == "wan-3.0")
assert wan["billing_mode"] == "per_second"
def test_default_values(self):
session = _make_session()
job = ViralVideoJobModel(id="job-002", user_id="user-002")
session.add(job)
session.commit()
fetched = session.get(ViralVideoJobModel, "job-002")
assert fetched.images == []
assert fetched.fusion_level == "ai_polish"
assert fetched.style_strength == "medium"
assert fetched.status == "pending"
assert fetched.credits_cost == 0
assert fetched.retry_count == 0
assert fetched.style_guide is None
assert fetched.intent_result is None
class TestViralVideoStyleTemplateModel:
def test_create_and_get(self):
session = _make_session()
tpl = ViralVideoStyleTemplateModel(
id="tpl-001",
name="快节奏",
style_config={"cut_speed": "fast"},
sort_order=1,
)
session.add(tpl)
session.commit()
fetched = session.get(ViralVideoStyleTemplateModel, "tpl-001")
assert fetched.name == "快节奏"
assert fetched.style_config == {"cut_speed": "fast"}
assert fetched.sort_order == 1
class TestViralVideoPromptTemplateModel:
def test_create_and_get(self):
session = _make_session()
tpl = ViralVideoPromptTemplateModel(
id="pt-001",
prompt_type="image_analysis",
name="图片分析模板",
content="请分析图片:{image_url}",
variables=["image_url"],
)
session.add(tpl)
session.commit()
fetched = session.get(ViralVideoPromptTemplateModel, "pt-001")
assert fetched.prompt_type == "image_analysis"
assert fetched.content == "请分析图片:{image_url}"
assert fetched.variables == ["image_url"]
assert fetched.version == 1
assert fetched.is_active is True
def test_exactly_one_default(self):
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=None):
result = self._call()
defaults = [m for m in result["models"] if m["is_default"]]
assert len(defaults) == 1
assert defaults[0]["key"] == "seedance-2.5"