diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index fa8fee9c5..96fced8fc 100644 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -1058,6 +1058,7 @@ def _download_all_assets( task_asset_ids: list[str], voice_library_id: str, task_id: str, + voice_ids: list[str] | None = None, ) -> tuple[list[Path], str | None]: """下载视频素材和配音素材。 @@ -1066,6 +1067,9 @@ def _download_all_assets( Note: gen_task 不传入下载函数(session 已关闭), 主函数在下载前后已有汇总日志。 + + 配音下载逻辑:优先使用 voice_library_id(配音素材库资产); + 若为空则 fallback 到 voice_ids[0](前端选择的音频 asset_id)。 """ logger.info("[task_id=%s] [下载素材] 开始下载视频素材", task_id) download_start = time.monotonic() @@ -1085,11 +1089,24 @@ def _download_all_assets( ) audio_path: str | None = None - if voice_library_id: + # 配音下载:优先 voice_library_id,fallback 到 voice_ids[0] + effective_voice_id = voice_library_id + if not effective_voice_id and voice_ids: + effective_voice_id = voice_ids[0] + logger.info( + "[task_id=%s] [下载配音] voice_library_id 为空,fallback 到 voice_ids[0]=%s", + task_id, + effective_voice_id, + ) + if effective_voice_id: local_audio = temp_path / "voice.mp3" - if _download_voice_asset(voice_library_id, local_audio): + if _download_voice_asset(effective_voice_id, local_audio): audio_path = str(local_audio) - logger.info("[task_id=%s] [下载配音] 配音下载成功", task_id) + logger.info( + "[task_id=%s] [下载配音] 配音下载成功 (source=%s)", + task_id, + "voice_library_id" if voice_library_id else "voice_ids", + ) return downloaded_videos, audio_path @@ -1422,6 +1439,7 @@ def generate_video(self, task_id: str) -> dict: task_asset_ids=task_asset_ids, voice_library_id=voice_library_id, task_id=task_id, + voice_ids=task_info.get("voice_ids", []), ) if gen_task: diff --git a/tests/unit/test_preview_voice_ids_fallback.py b/tests/unit/test_preview_voice_ids_fallback.py new file mode 100644 index 000000000..acac09d1a --- /dev/null +++ b/tests/unit/test_preview_voice_ids_fallback.py @@ -0,0 +1,151 @@ +"""Tests for voice_ids fallback in _download_all_assets. + +When voice_library_id is empty but voice_ids is non-empty, the Worker +should fallback to voice_ids[0] as the audio asset_id. +""" + +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + + +class TestDownloadAllAssetsVoiceIdsFallback: + """_download_all_assets 配音下载 fallback 逻辑测试。""" + + @patch("worker_app.tasks.generation._download_voice_asset") + @patch("worker_app.tasks.generation._download_library_assets") + def test_voice_library_id_takes_priority(self, mock_download_videos, mock_download_voice, tmp_path): + """voice_library_id 存在时优先使用,不 fallback 到 voice_ids。""" + from worker_app.tasks.generation import _download_all_assets + + mock_download_videos.return_value = [tmp_path / "v1.mp4"] + mock_download_voice.return_value = True + + videos, audio = _download_all_assets( + temp_path=tmp_path, + asset_library_id="lib-1", + project_id="proj-1", + task_asset_ids=["a1"], + voice_library_id="voice-lib-123", + task_id="task-1", + voice_ids=["voice-ids-456"], + ) + + assert audio is not None + mock_download_voice.assert_called_once() + call_args = mock_download_voice.call_args + assert call_args[0][0] == "voice-lib-123" # first positional arg + + @patch("worker_app.tasks.generation._download_voice_asset") + @patch("worker_app.tasks.generation._download_library_assets") + def test_fallback_to_voice_ids_when_voice_library_id_empty( + self, mock_download_videos, mock_download_voice, tmp_path + ): + """voice_library_id 为空时 fallback 到 voice_ids[0]。""" + from worker_app.tasks.generation import _download_all_assets + + mock_download_videos.return_value = [tmp_path / "v1.mp4"] + mock_download_voice.return_value = True + + videos, audio = _download_all_assets( + temp_path=tmp_path, + asset_library_id="lib-1", + project_id="proj-1", + task_asset_ids=["a1"], + voice_library_id="", # 空字符串 + task_id="task-2", + voice_ids=["voice-asset-789"], + ) + + assert audio is not None + mock_download_voice.assert_called_once() + call_args = mock_download_voice.call_args + assert call_args[0][0] == "voice-asset-789" + + @patch("worker_app.tasks.generation._download_voice_asset") + @patch("worker_app.tasks.generation._download_library_assets") + def test_no_audio_when_both_empty(self, mock_download_videos, mock_download_voice, tmp_path): + """voice_library_id 和 voice_ids 都为空时,不下载音频。""" + from worker_app.tasks.generation import _download_all_assets + + mock_download_videos.return_value = [tmp_path / "v1.mp4"] + + videos, audio = _download_all_assets( + temp_path=tmp_path, + asset_library_id="lib-1", + project_id="proj-1", + task_asset_ids=["a1"], + voice_library_id="", + task_id="task-3", + voice_ids=[], + ) + + assert audio is None + mock_download_voice.assert_not_called() + + @patch("worker_app.tasks.generation._download_voice_asset") + @patch("worker_app.tasks.generation._download_library_assets") + def test_no_audio_when_voice_ids_none(self, mock_download_videos, mock_download_voice, tmp_path): + """voice_ids 为 None 时,不触发 fallback。""" + from worker_app.tasks.generation import _download_all_assets + + mock_download_videos.return_value = [tmp_path / "v1.mp4"] + + videos, audio = _download_all_assets( + temp_path=tmp_path, + asset_library_id="lib-1", + project_id="proj-1", + task_asset_ids=["a1"], + voice_library_id="", + task_id="task-4", + voice_ids=None, + ) + + assert audio is None + mock_download_voice.assert_not_called() + + @patch("worker_app.tasks.generation._download_voice_asset") + @patch("worker_app.tasks.generation._download_library_assets") + def test_voice_library_id_empty_string_fallback(self, mock_download_videos, mock_download_voice, tmp_path): + """voice_library_id 为空字符串且 voice_ids 有多个元素时,取第一个。""" + from worker_app.tasks.generation import _download_all_assets + + mock_download_videos.return_value = [tmp_path / "v1.mp4"] + mock_download_voice.return_value = True + + videos, audio = _download_all_assets( + temp_path=tmp_path, + asset_library_id="lib-1", + project_id="proj-1", + task_asset_ids=["a1"], + voice_library_id="", + task_id="task-5", + voice_ids=["first-id", "second-id", "third-id"], + ) + + assert audio is not None + call_args = mock_download_voice.call_args + assert call_args[0][0] == "first-id" + + @patch("worker_app.tasks.generation._download_voice_asset") + @patch("worker_app.tasks.generation._download_library_assets") + def test_backward_compat_no_voice_ids_param(self, mock_download_videos, mock_download_voice, tmp_path): + """不传 voice_ids 参数时,行为与之前一致(向后兼容)。""" + from worker_app.tasks.generation import _download_all_assets + + mock_download_videos.return_value = [tmp_path / "v1.mp4"] + mock_download_voice.return_value = True + + # 不传 voice_ids + videos, audio = _download_all_assets( + temp_path=tmp_path, + asset_library_id="lib-1", + project_id="proj-1", + task_asset_ids=["a1"], + voice_library_id="voice-lib-999", + task_id="task-6", + ) + + assert audio is not None + mock_download_voice.assert_called_once_with("voice-lib-999", tmp_path / "voice.mp3")