Files
xiaoxia-saas/tests/unit/test_viral_video.py
T
xiaoxia 22e04d65a7
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 51s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m0s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 4m7s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 4m45s
AI Code Review / AI Code Review (pull_request) Successful in 7m14s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m39s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 20m32s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 20m41s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 24m10s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 26m22s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 31s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 5m18s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 1h9m19s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 2s
refactor(viral-video): #2106 删除 _step_musetalk 步骤
爆款视频由 Seedance 2.5 直接生成人物口型,不需要 MuseTalk 事后对口型。
MuseTalk 是 AI 数字人路线(上传人物视频+配音→对嘴型)用的,跟爆款视频是两条不同路线。

- 删除 _step_musetalk 函数
- resume_pipeline 直接把 render 输出传给 upload
- 更新流水线 docstring 为 9 步
- 删除/更新对应单测
- ViralVideoStage.MUSETALK 枚举值保留以避免前端 breaking change
2026-09-30 20:42:12 +08:00

518 lines
19 KiB
Python
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""爆款视频模块单元测试。
覆盖范围:
- 领域实体状态机转换
- Repository CRUD
- API 端点(6 个)
- Celery 编排器流水线
- Schema 校验
"""
from __future__ import annotations
from datetime import datetime, timezone
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from packages.domain.viral_video import (
CREDITS_VIRAL_VIDEO_COST,
STAGE_LABELS,
FusionLevel,
StyleStrength,
ViralVideoJob,
ViralVideoStage,
ViralVideoStatus,
)
# ── 领域模型测试 ─────────────────────────────────────────────────────────
class TestViralVideoStatus:
"""状态枚举测试。"""
def test_status_values(self):
assert ViralVideoStatus.PENDING == "pending"
assert ViralVideoStatus.RUNNING == "running"
assert ViralVideoStatus.WAIT_USER_CONFIRM == "wait_user_confirm"
assert ViralVideoStatus.COMPLETED == "completed"
assert ViralVideoStatus.FAILED == "failed"
assert ViralVideoStatus.CANCELLED == "cancelled"
def test_terminal_statuses(self):
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED).is_terminal
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.FAILED).is_terminal
assert ViralVideoJob(user_id="u1", status=ViralVideoStatus.CANCELLED).is_terminal
assert not ViralVideoJob(user_id="u1", status=ViralVideoStatus.PENDING).is_terminal
assert not ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING).is_terminal
class TestViralVideoJobStateTransitions:
"""状态机转换测试。"""
def test_mark_running_from_pending(self):
job = ViralVideoJob(user_id="u1")
job.mark_running()
assert job.status == ViralVideoStatus.RUNNING
assert job.started_at is not None
def test_mark_running_from_running(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
job.mark_running()
assert job.status == ViralVideoStatus.RUNNING
def test_mark_running_from_completed_raises(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED)
with pytest.raises(ValueError, match="Cannot transition"):
job.mark_running()
def test_mark_wait_user_confirm(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
intent = {"intent": "推广", "key_messages": ["卖点1"]}
job.mark_wait_user_confirm(intent)
assert job.status == ViralVideoStatus.WAIT_USER_CONFIRM
assert job.intent_result == intent
def test_mark_wait_user_confirm_from_non_running_raises(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.PENDING)
with pytest.raises(ValueError, match="Cannot transition"):
job.mark_wait_user_confirm({})
def test_resume_from_confirm(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
job.resume_from_confirm()
assert job.status == ViralVideoStatus.RUNNING
def test_resume_from_non_confirm_raises(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
with pytest.raises(ValueError, match="Cannot resume"):
job.resume_from_confirm()
def test_mark_completed(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
job.mark_completed("https://oss.example.com/video.mp4")
assert job.status == ViralVideoStatus.COMPLETED
assert job.result_video_url == "https://oss.example.com/video.mp4"
assert job.completed_at is not None
def test_mark_failed(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
job.mark_failed("渲染超时")
assert job.status == ViralVideoStatus.FAILED
assert job.error_msg == "渲染超时"
def test_mark_cancelled(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.RUNNING)
job.mark_cancelled()
assert job.status == ViralVideoStatus.CANCELLED
def test_mark_cancelled_from_terminal_raises(self):
job = ViralVideoJob(user_id="u1", status=ViralVideoStatus.COMPLETED)
with pytest.raises(ValueError, match="Cannot cancel"):
job.mark_cancelled()
class TestViralVideoJobDefaults:
"""默认值测试。"""
def test_default_values(self):
job = ViralVideoJob(user_id="u1")
assert job.images == []
assert job.industry == ""
assert job.duration == 30
assert job.fusion_level == FusionLevel.AI_POLISH
assert job.style_strength == StyleStrength.MEDIUM
assert job.status == ViralVideoStatus.PENDING
assert job.credits_cost == 0
assert job.retry_count == 0
assert job.result_video_url == ""
assert job.error_msg == ""
def test_credits_cost_constant(self):
assert CREDITS_VIRAL_VIDEO_COST == 50
class TestViralVideoStage:
"""阶段枚举测试。"""
def test_all_stages_have_labels(self):
for stage in ViralVideoStage:
assert stage in STAGE_LABELS, f"Stage {stage} missing label"
def test_stage_order(self):
expected_order = [
"image_analysis",
"video_analysis",
"intent_parsing",
"copy_fusion",
"storyboard",
"review",
"tts",
"bgm_select",
"rendering",
"musetalk",
"uploading",
]
actual_order = [s.value for s in ViralVideoStage]
assert actual_order == expected_order
# ── Schema 校验测试 ──────────────────────────────────────────────────────
class TestViralVideoSchemas:
"""Pydantic Schema 校验测试。"""
def test_create_request_valid(self):
from app.schemas.viral_video import CreateViralVideoRequest
req = CreateViralVideoRequest(images=["https://example.com/img.jpg"])
assert req.images == ["https://example.com/img.jpg"]
assert req.fusion_level == "ai_polish"
assert req.style_strength == "medium"
assert req.duration == 30
def test_create_request_empty_images_raises(self):
from app.schemas.viral_video import CreateViralVideoRequest
with pytest.raises(ValidationError):
CreateViralVideoRequest(images=[])
def test_create_request_invalid_fusion_level(self):
from app.schemas.viral_video import CreateViralVideoRequest
with pytest.raises(ValidationError):
CreateViralVideoRequest(
images=["https://example.com/img.jpg"],
fusion_level="invalid_level",
)
def test_create_request_invalid_style_strength(self):
from app.schemas.viral_video import CreateViralVideoRequest
with pytest.raises(ValidationError):
CreateViralVideoRequest(
images=["https://example.com/img.jpg"],
style_strength="ultra",
)
def test_confirm_intent_request_defaults(self):
from app.schemas.viral_video import ConfirmIntentRequest
req = ConfirmIntentRequest()
assert req.confirmed_copy == ""
assert req.adjustments == ""
def test_analyze_style_request(self):
from app.schemas.viral_video import AnalyzeStyleRequest
req = AnalyzeStyleRequest(reference_video_url="https://example.com/video.mp4")
assert req.reference_video_url == "https://example.com/video.mp4"
def test_ws_progress_event(self):
from app.schemas.viral_video import WSProgressEvent
event = WSProgressEvent(
job_id="abc123",
stage="image_analysis",
progress=10.0,
message="正在分析图片",
)
assert event.type == "viral_video:progress"
assert event.job_id == "abc123"
assert event.progress == 10.0
# ── Repository 测试 ─────────────────────────────────────────────────────
class TestViralVideoRepository:
"""SQLAlchemy Repository CRUD 测试(使用内存数据库)。"""
@pytest.fixture
def db_session(self):
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.models import Base
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(engine)
SessionLocal = sessionmaker(bind=engine)
session = SessionLocal()
yield session
session.close()
def test_save_and_get(self, db_session):
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
SQLAlchemyViralVideoJobRepository,
)
repo = SQLAlchemyViralVideoJobRepository(db_session)
job = ViralVideoJob(
user_id="user-001",
images=["https://img.com/1.jpg"],
industry="美妆",
duration=60,
)
repo.save(job)
fetched = repo.get(job.id)
assert fetched is not None
assert fetched.id == job.id
assert fetched.user_id == "user-001"
assert fetched.images == ["https://img.com/1.jpg"]
assert fetched.industry == "美妆"
assert fetched.duration == 60
def test_get_nonexistent(self, db_session):
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
SQLAlchemyViralVideoJobRepository,
)
repo = SQLAlchemyViralVideoJobRepository(db_session)
assert repo.get("nonexistent-id") is None
def test_list_by_user(self, db_session):
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
SQLAlchemyViralVideoJobRepository,
)
repo = SQLAlchemyViralVideoJobRepository(db_session)
for i in range(3):
job = ViralVideoJob(user_id="user-001", industry=f"行业{i}")
repo.save(job)
# 另一个用户的任务
other_job = ViralVideoJob(user_id="user-002", industry="其他")
repo.save(other_job)
jobs = repo.list_by_user("user-001")
assert len(jobs) == 3
assert all(j.user_id == "user-001" for j in jobs)
def test_update_status(self, db_session):
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
SQLAlchemyViralVideoJobRepository,
)
repo = SQLAlchemyViralVideoJobRepository(db_session)
job = ViralVideoJob(user_id="user-001")
repo.save(job)
job.mark_running()
repo.update(job)
fetched = repo.get(job.id)
assert fetched.status == ViralVideoStatus.RUNNING
assert fetched.started_at is not None
def test_count_pending_by_user(self, db_session):
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
SQLAlchemyViralVideoJobRepository,
)
repo = SQLAlchemyViralVideoJobRepository(db_session)
# 2 个 pending
for _ in range(2):
repo.save(ViralVideoJob(user_id="user-001"))
# 1 个 completed
completed = ViralVideoJob(user_id="user-001", status=ViralVideoStatus.COMPLETED)
repo.save(completed)
assert repo.count_pending_by_user("user-001") == 2
def test_style_template_repo(self, db_session):
from packages.adapters.sqlalchemy_impl.models import ViralVideoStyleTemplateModel
from packages.adapters.sqlalchemy_impl.viral_video_repository import (
SQLAlchemyViralVideoStyleTemplateRepository,
)
# 插入模板
tpl = ViralVideoStyleTemplateModel(
id="tpl-001",
name="快节奏",
description="适合快消品",
style_config={"cut_speed": "fast"},
is_system=True,
sort_order=1,
)
db_session.add(tpl)
db_session.commit()
repo = SQLAlchemyViralVideoStyleTemplateRepository(db_session)
templates = repo.list_all()
assert len(templates) == 1
assert templates[0]["name"] == "快节奏"
fetched = repo.get("tpl-001")
assert fetched is not None
assert fetched["style_config"] == {"cut_speed": "fast"}
# ── Celery 编排器测试 ───────────────────────────────────────────────────
class TestViralVideoPipeline:
"""编排器流水线测试。"""
@pytest.fixture
def mock_job(self):
return ViralVideoJob(
user_id="user-001",
images=["https://img.com/1.jpg", "https://img.com/2.jpg"],
industry="美妆",
target_customer="年轻女性",
marketing_purpose="品牌推广",
duration=30,
user_copy_text="这款产品超好用",
fusion_level="ai_polish",
)
@patch("packages.shared.ai_service.call_vision")
def test_image_analysis_step(self, mock_vision, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
mock_vision.return_value = {"name": "口红", "features": ["持久", "滋润"]}
result = _step_image_analysis(mock_job)
assert "products" in result
assert len(result["products"]) == 2 # 两张图片
@patch("packages.shared.ai_service.call_vision")
def test_image_analysis_fallback(self, mock_vision, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_image_analysis
# 模拟 call_vision 不存在
mock_vision.side_effect = ImportError("no module")
result = _step_image_analysis(mock_job)
assert "products" in result
def test_video_analysis_no_reference(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_video_analysis
# 没有参考视频
mock_job.reference_video_url = ""
result = _step_video_analysis(mock_job)
assert result is None
@patch("packages.shared.ai_service.call_llm")
def test_intent_parsing(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_intent_parsing
mock_llm.return_value = {"intent": "推广口红", "tone": "活泼"}
result = _step_intent_parsing(mock_job, {"products": []})
assert "intent" in result
@patch("packages.shared.ai_service.call_llm")
def test_copy_fusion_ai_polish(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_copy_fusion
mock_llm.return_value = "融合后的文案内容"
result = _step_copy_fusion(mock_job, {"intent": "推广"}, {"products": []})
assert isinstance(result, str)
assert len(result) > 0
@patch("packages.shared.ai_service.call_llm")
def test_storyboard_generation(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_storyboard
mock_llm.return_value = [
{"order": 0, "type": "product_shot", "duration": 10},
{"order": 1, "type": "closing", "duration": 5},
]
result = _step_storyboard(mock_job, "测试文案", {})
assert isinstance(result, list)
assert len(result) == 2
@patch("packages.shared.ai_service.call_llm")
def test_review_pass(self, mock_llm, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_review
mock_llm.return_value = {"passed": True, "score": 90, "details": {}}
result = _step_review(mock_job, "测试文案", [])
assert result["passed"] is True
def test_bgm_select(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
# P1: BGM 素材未就绪前 _step_bgm_select 统一返回 None(跳过 BGM 混音)
mock_job.bgm_preference = "upbeat"
bgm = _step_bgm_select(mock_job)
assert bgm is None
def test_bgm_select_default(self, mock_job):
from apps.worker.worker_app.tasks.viral_video import _step_bgm_select
mock_job.bgm_preference = ""
bgm = _step_bgm_select(mock_job)
assert bgm is None
# ── 端到端流水线集成测试 ────────────────────────────────────────────────
class TestPipelineIntegration:
"""流水线端到端集成测试(mock 外部依赖)。"""
@patch("apps.worker.worker_app.tasks.viral_video._step_upload")
@patch("apps.worker.worker_app.tasks.viral_video._step_render")
@patch("apps.worker.worker_app.tasks.viral_video._step_bgm_select")
@patch("apps.worker.worker_app.tasks.viral_video._step_tts")
@patch("apps.worker.worker_app.tasks.viral_video._step_review")
@patch("apps.worker.worker_app.tasks.viral_video._step_storyboard")
@patch("apps.worker.worker_app.tasks.viral_video._step_copy_fusion")
@patch("apps.worker.worker_app.tasks.viral_video._step_intent_parsing")
@patch("apps.worker.worker_app.tasks.viral_video._step_video_analysis")
@patch("apps.worker.worker_app.tasks.viral_video._step_image_analysis")
@patch("apps.worker.worker_app.tasks.viral_video._get_repo_and_job")
@patch("apps.worker.worker_app.tasks.viral_video._emit_progress")
def test_resume_pipeline_completes(
self,
mock_emit,
mock_get_repo,
mock_img_analysis,
mock_video_analysis,
mock_intent,
mock_copy_fusion,
mock_storyboard,
mock_review,
mock_tts,
mock_bgm,
mock_render,
mock_upload,
):
"""测试 resume 流水线能从确认状态走到完成。"""
from apps.worker.worker_app.tasks.viral_video import (
resume_viral_video_pipeline,
)
# 构造 mock job
job = ViralVideoJob(
user_id="user-001",
images=["https://img.com/1.jpg"],
industry="美妆",
status=ViralVideoStatus.RUNNING,
intent_result={"intent": "推广"},
)
mock_repo = MagicMock()
mock_session = MagicMock()
mock_get_repo.return_value = (mock_session, mock_repo, job)
# 设置各步骤返回值
mock_copy_fusion.return_value = "融合文案"
mock_storyboard.return_value = [{"order": 0, "duration": 10}]
mock_review.return_value = {"passed": True, "score": 90}
mock_tts.return_value = None # P1: TTS 返回 Path|None,mock 用 None 跳过混音
mock_bgm.return_value = None # P1: BGM 未就绪前返回 None
mock_render.return_value = "/tmp/video.mp4"
mock_upload.return_value = "https://oss.example.com/final.mp4"
result = resume_viral_video_pipeline.run("job-001")
assert result["ok"] is True
assert result["video_url"] == "https://oss.example.com/final.mp4"
assert job.status == ViralVideoStatus.COMPLETED
assert job.credits_cost == CREDITS_VIRAL_VIDEO_COST