Files
xiaoxia-saas/tests/unit/test_viral_video.py
T
CI Bot eda3a3a540
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 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m54s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 19s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 10s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m33s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker 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 / Retag skipped Staging API 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
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m4s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 35s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Successful in 2m38s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 4m15s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 4m49s
AI Code Review / AI Code Review (pull_request) Successful in 6m48s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 7m46s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 11m53s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m43s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 19m21s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
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
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 35s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 4m43s
style: auto-format with black + isort + ruff + prettier [skip ci-format-check]
2026-10-01 06:02:55 +00:00

571 lines
22 KiB
Python
Executable File

"""爆款视频模块单元测试。
覆盖范围:
- 领域实体状态机转换
- 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 == 15
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",
"script_generation",
"review",
"tts",
"rendering",
"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 == 15
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=15,
user_copy_text="这款产品超好用",
fusion_level="ai_polish",
video_ratio="9:16",
)
@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_script_generation_returns_copy_result(self, mock_llm, mock_job):
"""v1.6: _step_script_generation 返回 dict 形式的 CopyResult,含 voiceover_script + shots。"""
from apps.worker.worker_app.tasks.viral_video import _step_script_generation
mock_llm.return_value = {
"overview": {"theme": "口红推荐", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台,柔和自然光",
"shots": [
{
"time_range": "0-5秒",
"shot_type_angle_movement": "近景平视,缓慢推镜",
"scene_and_dialogue": "女主微笑展示口红:大家好,今天分享一款口红",
"action_details": "手持口红特写",
"audio_bgm": "轻快流行BGM",
"transition": "硬切",
"reference_image_index": 0,
},
{
"time_range": "5-15秒",
"shot_type_angle_movement": "特写,固定镜头",
"scene_and_dialogue": "涂抹口红:颜色特别好看很显白",
"action_details": "嘴唇涂抹特写",
"audio_bgm": "轻快BGM继续",
"transition": "结束",
"reference_image_index": 1,
},
],
"hard_constraints": ["无字幕无水印"],
"negative_prompts": ["字幕", "水印"],
"voiceover_script": "大家好,今天分享一款口红,颜色特别好看很显白。",
}
result = _step_script_generation(
mock_job, {"intent": "推广口红", "key_messages": [], "tone": "亲切"}, {"products": []}
)
assert isinstance(result, dict)
assert "voiceover_script" in result
assert "shots" in result
assert isinstance(result["shots"], list)
assert len(result["shots"]) == 2
assert result["overview"]["total_duration"] == 15
# final_copy 必须 = voiceover_script(向后兼容)
assert result.get("final_copy") == result["voiceover_script"]
@patch("packages.shared.ai_service.call_llm")
def test_script_generation_fallback(self, mock_llm, mock_job):
"""LLM 返回异常时使用兜底脚本(不会抛错)。"""
from apps.worker.worker_app.tasks.viral_video import _fallback_script
result = _fallback_script(mock_job)
assert isinstance(result, dict)
assert result["voiceover_script"]
assert len(result["shots"]) >= 1
@patch("packages.shared.ai_service.call_llm")
def test_review_pass_v16(self, mock_llm, mock_job):
"""v1.6 _step_review 接收 copy_result dict。"""
from apps.worker.worker_app.tasks.viral_video import _step_review
mock_llm.return_value = {"passed": True, "score": 90, "details": {}}
cr = {"voiceover_script": "大家好", "shots": []}
result = _step_review(mock_job, cr)
assert result["passed"] is True
def test_assemble_seedance_prompt(self, mock_job):
"""编导脚本必须能拼出完整的 Seedance prompt,含总览/场景/逐镜头/约束。"""
from apps.worker.worker_app.tasks.viral_video import _assemble_seedance_prompt
cr = {
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台",
"shots": [
{
"time_range": "0-15秒",
"shot_type_angle_movement": "中景平视",
"scene_and_dialogue": "你好分享",
"action_details": "展示",
"audio_bgm": "BGM",
"transition": "结束",
"reference_image_index": 0,
}
],
"hard_constraints": ["无字幕"],
"negative_prompts": ["水印"],
}
prompt = _assemble_seedance_prompt(cr, mock_job)
assert "【视频总览】" in prompt
assert "【逐镜头时间轴】" in prompt
assert "【硬性约束】" in prompt
assert "【负面提示词】" in prompt
assert "0-15秒" in prompt
# ── 端到端流水线集成测试 ────────────────────────────────────────────────
class TestPipelineIntegration:
"""v1.6 流水线端到端集成测试(mock 外部依赖):TTS+单次 Seedance+上传。"""
@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._upload_tts_to_oss")
@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_script_generation")
@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_script,
mock_review,
mock_tts,
mock_tts_upload,
mock_render,
mock_upload,
):
"""v1.6: TTS整段合成 → 上传TTS到OSS → 单次 Seedance → 上传成片。"""
from apps.worker.worker_app.tasks.viral_video import (
resume_viral_video_pipeline,
)
job = ViralVideoJob(
user_id="user-001",
images=["https://img.com/1.jpg"],
industry="美妆",
status=ViralVideoStatus.RUNNING,
intent_result={"intent": "推广"},
duration=15,
video_ratio="9:16",
)
mock_repo = MagicMock()
mock_session = MagicMock()
mock_get_repo.return_value = (mock_session, mock_repo, job)
# v1.6: 如果没有 copy_result 会现场补生成
mock_intent.return_value = {"intent": "推广", "key_messages": [], "tone": "亲切"}
mock_script.return_value = {
"overview": {"theme": "口红", "total_duration": 15, "aspect_ratio": "9:16"},
"scene_and_lighting": "明亮化妆台",
"shots": [],
"hard_constraints": [],
"negative_prompts": [],
"voiceover_script": "大家好,分享一款口红。",
"final_copy": "大家好,分享一款口红。",
}
mock_review.return_value = {"passed": True, "score": 90}
mock_tts.return_value = None # TTS 失败也能走下去(Seedance generate_audio=True 会自己合成音效)
mock_tts_upload.return_value = 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