diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 04ca03599..d690da309 100644 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -1040,6 +1040,7 @@ def _load_task_info(task_id: str) -> dict | None: "resolution": getattr(gen_task, "resolution", "") or "", "bgm_config": dict(getattr(gen_task, "bgm_config", {}) or {}), "is_preview": bool(getattr(gen_task, "is_preview", False)), + "voice_ids": list(getattr(gen_task, "voice_ids", []) or []), } finally: session.close() @@ -1101,6 +1102,7 @@ def _render_video( resolution: str = "", bgm_config: dict | None = None, is_preview: bool = False, + voice_ids: list[str] | None = None, ) -> tuple[Path, float]: """渲染视频(含配音混音)。 @@ -1175,6 +1177,20 @@ def _render_video( plan_cfg["export"] = export_cfg virtual_plan.config = plan_cfg + # 注入用户选择的配音 voice_id(ASR 字幕对齐模式) + if voice_ids: + plan_cfg = dict(virtual_plan.config or {}) + plan_cfg["voice_id"] = voice_ids[0] + subtitle_cfg = plan_cfg.get("subtitle", {}) or {} + subtitle_cfg["auto_generated"] = True + plan_cfg["subtitle"] = subtitle_cfg + virtual_plan.config = plan_cfg + logger.info( + "[task_id=%s] [渲染] 预览配音已注入: voice_id=%s", + task_id, + voice_ids[0], + ) + total_duration = sum(c.duration for c in virtual_clips) logger.info( "[task_id=%s] [剪辑计划] 片段数=%d, 总时长=%.1fs", @@ -1428,6 +1444,7 @@ def generate_video(self, task_id: str) -> dict: resolution=task_info.get("resolution", ""), bgm_config=task_info.get("bgm_config", {}), is_preview=task_info.get("is_preview", False), + voice_ids=task_info.get("voice_ids", []), ) if gen_task: diff --git a/tests/unit/test_1294_preview_voice_injection.py b/tests/unit/test_1294_preview_voice_injection.py new file mode 100644 index 000000000..3b27160e7 --- /dev/null +++ b/tests/unit/test_1294_preview_voice_injection.py @@ -0,0 +1,259 @@ +"""测试 #1294 修复:预览视频配音注入。 + +验证: +1. _load_task_info 正确加载 voice_ids +2. _render_video 接受 voice_ids 参数 +3. voice_ids 正确注入到 plan config 中(实际执行代码路径,diff-cover 可达) +""" + +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + + +class TestLoadTaskInfoVoiceIds: + """验证 _load_task_info 包含 voice_ids""" + + def test_voice_ids_loaded_from_task(self): + """voice_ids 从 gen_task 正确加载""" + mock_task = MagicMock() + mock_task.project_id = "proj_1" + mock_task.asset_library_id = "lib_1" + mock_task.voice_library_id = "voice_lib_1" + mock_task.template_id = "tmpl_1" + mock_task.strategy_id = "one_take" + mock_task.asset_ids = ["a1", "a2"] + mock_task.batch_id = "batch_1" + mock_task.created_by_user_id = "user_1" + mock_task.video_title = "test" + mock_task.resolution = "854x480" + mock_task.bgm_config = {} + mock_task.is_preview = True + mock_task.voice_ids = ["voice_1", "voice_2"] + + with patch( + "packages.adapters.sqlalchemy_impl.generation_task_repository.SQLAlchemyGenerationTaskRepository" + ) as MockRepo: + mock_repo = MagicMock() + mock_repo.get.return_value = mock_task + MockRepo.return_value = mock_repo + + from worker_app.tasks.generation import _load_task_info + + result = _load_task_info("test_task_id") + + assert result is not None + assert result["voice_ids"] == ["voice_1", "voice_2"] + + def test_voice_ids_empty_when_none(self): + """voice_ids 为 None 时返回空列表""" + mock_task = MagicMock() + mock_task.project_id = "proj_1" + mock_task.asset_library_id = "lib_1" + mock_task.voice_library_id = "" + mock_task.template_id = "tmpl_1" + mock_task.strategy_id = "one_take" + mock_task.asset_ids = ["a1"] + mock_task.batch_id = "" + mock_task.created_by_user_id = "user_1" + mock_task.video_title = "" + mock_task.resolution = "" + mock_task.bgm_config = {} + mock_task.is_preview = False + mock_task.voice_ids = None + + with patch( + "packages.adapters.sqlalchemy_impl.generation_task_repository.SQLAlchemyGenerationTaskRepository" + ) as MockRepo: + mock_repo = MagicMock() + mock_repo.get.return_value = mock_task + MockRepo.return_value = mock_repo + + from worker_app.tasks.generation import _load_task_info + + result = _load_task_info("test_task_id") + assert result["voice_ids"] == [] + + +class TestRenderVideoVoiceInjection: + """验证 _render_video 正确注入 voice_id 到 plan config(实际执行代码路径)""" + + def test_render_video_accepts_voice_ids(self): + """_render_video 签名包含 voice_ids 参数""" + import inspect + + from worker_app.tasks.generation import _render_video + + sig = inspect.signature(_render_video) + assert "voice_ids" in sig.parameters + + def test_voice_ids_default_none(self): + """voice_ids 参数默认为 None""" + import inspect + + from worker_app.tasks.generation import _render_video + + sig = inspect.signature(_render_video) + param = sig.parameters["voice_ids"] + assert param.default is None + + def test_voice_ids_injected_into_plan_config(self): + """voice_ids 非空时,voice_id 和 subtitle.auto_generated 被注入到 plan config。 + + 此测试实际执行 _render_video 的配音注入代码路径,确保 diff-cover 覆盖新增行。 + """ + from dataclasses import dataclass, field + + @dataclass + class MockClip: + """模拟 VirtualClip,至少需要 duration 属性。""" + + id: str = "clip_1" + duration: float = 5.0 + config: dict = field(default_factory=dict) + + @dataclass + class MockPlan: + """模拟 VirtualPlan,至少需要 config 属性。""" + + id: str = "test_plan" + name: str = "test" + config: dict = field(default_factory=dict) + + mock_plan = MockPlan(config={"some_key": "some_value"}) + mock_clips = [MockClip(duration=5.0), MockClip(duration=3.0)] + mock_asset_path_map = {"asset_1": Path("/tmp/video1.mp4")} + + # Mock RenderAdapter 和 render 结果 + mock_render_result = MagicMock() + mock_render_result.success = True + mock_render_result.output_path = Path("/tmp/output.mp4") + mock_render_result.duration = 8.0 + + mock_adapter_cls = MagicMock(return_value=MagicMock()) + mock_adapter_cls.return_value.render_from_memory.return_value = mock_render_result + + mock_db = MagicMock() + + with ( + patch( + "worker_app.tasks.generation._build_plan_and_clips_from_task", + return_value=(mock_plan, mock_clips, mock_asset_path_map), + ), + patch( + "worker_app.tasks.generation._load_template_plan_config", + return_value=None, + ), + patch( + "video_processing.render_adapter.RenderAdapter", + mock_adapter_cls, + ), + patch( + "worker_app.tasks.generation.SessionLocal", + return_value=mock_db, + ), + ): + from worker_app.tasks.generation import _render_video + + from packages.domain import EditingMode + + output_path, render_duration = _render_video( + task_id="test_task_123", + downloaded_videos=[Path("/tmp/video1.mp4")], + voice_path=None, + editing_mode=EditingMode.ONE_TAKE, + project_id="proj_1", + template_id="tmpl_1", + user_id="user_1", + temp_path=Path("/tmp"), + output_name="test_output.mp4", + resolution="854x480", + is_preview=True, + voice_ids=["voice_abc"], + ) + + # 验证 voice_id 被注入到 plan config(覆盖新增代码行) + assert mock_plan.config.get("voice_id") == "voice_abc" + # 验证 subtitle.auto_generated 被设置为 True + assert mock_plan.config.get("subtitle", {}).get("auto_generated") is True + # 验证 RenderAdapter 被调用 + mock_adapter_cls.return_value.render_from_memory.assert_called_once() + # 验证返回值 + assert output_path == Path("/tmp/output.mp4") + assert render_duration == 8.0 + + def test_voice_ids_empty_skips_injection(self): + """voice_ids 为空时,不注入 voice_id 到 plan config""" + from dataclasses import dataclass, field + + @dataclass + class MockClip: + id: str = "clip_1" + duration: float = 5.0 + + @dataclass + class MockPlan: + id: str = "test_plan" + name: str = "test" + config: dict = field(default_factory=dict) + + mock_plan = MockPlan(config={"export": {"resolution": "854x480"}}) + mock_clips = [MockClip(duration=5.0)] + + mock_render_result = MagicMock() + mock_render_result.success = True + mock_render_result.output_path = Path("/tmp/output.mp4") + mock_render_result.duration = 5.0 + + mock_adapter_cls = MagicMock(return_value=MagicMock()) + mock_adapter_cls.return_value.render_from_memory.return_value = mock_render_result + + with ( + patch( + "worker_app.tasks.generation._build_plan_and_clips_from_task", + return_value=(mock_plan, mock_clips, {}), + ), + patch( + "worker_app.tasks.generation._load_template_plan_config", + return_value=None, + ), + patch( + "video_processing.render_adapter.RenderAdapter", + mock_adapter_cls, + ), + patch( + "worker_app.tasks.generation.SessionLocal", + return_value=MagicMock(), + ), + ): + from worker_app.tasks.generation import _render_video + + from packages.domain import EditingMode + + _render_video( + task_id="test_task_456", + downloaded_videos=[Path("/tmp/video1.mp4")], + voice_path=None, + editing_mode=EditingMode.ONE_TAKE, + project_id="proj_1", + template_id="", + user_id="user_1", + temp_path=Path("/tmp"), + output_name="test_output.mp4", + voice_ids=[], + ) + + # 验证 voice_id 没有被注入 + assert "voice_id" not in mock_plan.config + + +class TestGenerateVideoPassesVoiceIds: + """验证 generate_video 调用 _render_video 时传递 voice_ids""" + + def test_generate_video_passes_voice_ids(self): + """generate_video 中 _render_video 调用包含 voice_ids 参数""" + with open("apps/worker/worker_app/tasks/generation.py", "r") as f: + content = f.read() + + assert 'voice_ids=task_info.get("voice_ids", [])' in content