"""爆款视频模块单元测试。 覆盖范围: - 领域实体状态机转换 - 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 ( 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 == "" 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", {"completion_tokens": 1000000}) 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 isinstance(job.credits_cost, float)