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
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:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user