"""生成任务 UseCase 单元测试.""" from __future__ import annotations from unittest.mock import MagicMock import pytest from packages.application.generation_tasks import ( CreateGenerationTaskCommand, CreateGenerationTaskUseCase, GetGenerationTaskUseCase, ListGenerationTasksResult, ListTasksFilter, ListUserTasksFilteredUseCase, RetryGenerationTaskUseCase, ) from packages.domain import GenerationTask @pytest.fixture def mock_repo(): return MagicMock() @pytest.fixture def sample_task(): task = MagicMock(spec=GenerationTask) task.id = "task_001" task.project_id = "proj_001" task.status = "pending" return task class TestCreateGenerationTaskUseCase: """CreateGenerationTaskUseCase 测试""" def test_create_task_success(self, mock_repo): """正常创建生成任务""" mock_repo.create.side_effect = lambda t: t use_case = CreateGenerationTaskUseCase(mock_repo) command = CreateGenerationTaskCommand( project_id="proj_001", template_id="tpl_001", asset_library_id="lib_001", voice_library_id="voice_lib_001", created_by_user_id="user_001", ) result = use_case.execute(command) assert isinstance(result, GenerationTask) assert result.project_id == "proj_001" assert result.template_id == "tpl_001" assert result.status == "pending" assert result.progress == 0.0 assert result.result_count == 0 mock_repo.create.assert_called_once() def test_create_task_generates_id(self, mock_repo): """创建任务时生成 id""" mock_repo.create.side_effect = lambda t: t use_case = CreateGenerationTaskUseCase(mock_repo) command = CreateGenerationTaskCommand(project_id="proj_001") result = use_case.execute(command) assert result.id is not None assert len(result.id) > 0 def test_create_task_with_asset_ids(self, mock_repo): """创建带 asset_ids 的任务""" mock_repo.create.side_effect = lambda t: t use_case = CreateGenerationTaskUseCase(mock_repo) command = CreateGenerationTaskCommand( project_id="proj_001", asset_ids=["asset_1", "asset_2", "asset_3"], title_ids=["title_1", "title_2"], voice_ids=["voice_1"], ) result = use_case.execute(command) assert len(result.asset_ids) == 3 assert len(result.title_ids) == 2 assert len(result.voice_ids) == 1 def test_create_task_with_auto_retry(self, mock_repo): """创建带自动重试配置的任务""" mock_repo.create.side_effect = lambda t: t use_case = CreateGenerationTaskUseCase(mock_repo) command = CreateGenerationTaskCommand( project_id="proj_001", auto_retry_enabled=True, auto_retry_max=3, ) result = use_case.execute(command) assert result.auto_retry_enabled is True assert result.auto_retry_max == 3 def test_create_task_with_bgm_config(self, mock_repo): """创建带 BGM 配置的任务""" mock_repo.create.side_effect = lambda t: t use_case = CreateGenerationTaskUseCase(mock_repo) bgm = {"enabled": True, "volume": 0.5, "library_id": "bgm_lib"} command = CreateGenerationTaskCommand( project_id="proj_001", bgm_config=bgm, resolution="1080p", video_title="测试视频", ) result = use_case.execute(command) assert result.bgm_config == bgm assert result.resolution == "1080p" assert result.video_title == "测试视频" def test_create_task_defaults(self, mock_repo): """默认参数的任务""" mock_repo.create.side_effect = lambda t: t use_case = CreateGenerationTaskUseCase(mock_repo) command = CreateGenerationTaskCommand() result = use_case.execute(command) assert result.project_id == "" assert result.asset_ids == [] assert result.auto_retry_enabled is False assert result.auto_retry_max == 0 class TestGetGenerationTaskUseCase: """GetGenerationTaskUseCase 测试""" def test_get_task_success(self, mock_repo, sample_task): """获取任务成功""" mock_repo.get.return_value = sample_task use_case = GetGenerationTaskUseCase(mock_repo) result = use_case.execute("task_001") assert result is sample_task mock_repo.get.assert_called_once_with("task_001") def test_get_task_not_found(self, mock_repo): """任务不存在返回 None""" mock_repo.get.return_value = None use_case = GetGenerationTaskUseCase(mock_repo) result = use_case.execute("nonexistent") assert result is None class TestListUserTasksFilteredUseCase: """ListUserTasksFilteredUseCase 测试""" def test_list_without_filter(self, mock_repo, sample_task): """不带筛选条件查询""" mock_repo.list_by_user_filtered.return_value = [sample_task] mock_repo.count_by_user_filtered.return_value = 1 use_case = ListUserTasksFilteredUseCase(mock_repo) result = use_case.execute("user_001") assert isinstance(result, ListGenerationTasksResult) assert len(result.items) == 1 assert result.total == 1 mock_repo.list_by_user_filtered.assert_called_once_with("user_001", status=None, limit=None, offset=0) def test_list_with_status_filter(self, mock_repo): """按状态筛选""" mock_repo.list_by_user_filtered.return_value = [] mock_repo.count_by_user_filtered.return_value = 0 use_case = ListUserTasksFilteredUseCase(mock_repo) result = use_case.execute("user_001", status="completed") assert result.total == 0 mock_repo.list_by_user_filtered.assert_called_once_with("user_001", status="completed", limit=None, offset=0) def test_list_with_pagination(self, mock_repo): """带分页参数查询""" mock_repo.list_by_user_filtered.return_value = [] mock_repo.count_by_user_filtered.return_value = 50 use_case = ListUserTasksFilteredUseCase(mock_repo) result = use_case.execute("user_001", limit=10, offset=20) assert result.total == 50 mock_repo.list_by_user_filtered.assert_called_once_with("user_001", status=None, limit=10, offset=20) def test_list_with_all_params(self, mock_repo): """带所有筛选和分页参数""" mock_repo.list_by_user_filtered.return_value = [] mock_repo.count_by_user_filtered.return_value = 5 use_case = ListUserTasksFilteredUseCase(mock_repo) use_case.execute("user_001", status="failed", limit=20, offset=0) mock_repo.list_by_user_filtered.assert_called_once_with("user_001", status="failed", limit=20, offset=0) mock_repo.count_by_user_filtered.assert_called_once_with("user_001", status="failed") class TestRetryGenerationTaskUseCase: """RetryGenerationTaskUseCase 测试""" def test_retry_failed_task(self, mock_repo): """重试失败的任务""" task = MagicMock(spec=GenerationTask) task.is_failed = True mock_repo.get.return_value = task mock_repo.update.return_value = task use_case = RetryGenerationTaskUseCase(mock_repo) result = use_case.execute("task_001") task.mark_pending_from_failed.assert_called_once() mock_repo.update.assert_called_once_with(task) assert result is task def test_retry_not_found(self, mock_repo): """任务不存在抛出 ValueError""" mock_repo.get.return_value = None use_case = RetryGenerationTaskUseCase(mock_repo) with pytest.raises(ValueError, match="任务不存在"): use_case.execute("nonexistent") mock_repo.update.assert_not_called() def test_retry_non_failed_task(self, mock_repo): """非失败状态的任务不能重试""" task = MagicMock(spec=GenerationTask) task.is_failed = False task.status = MagicMock() task.status.value = "running" mock_repo.get.return_value = task use_case = RetryGenerationTaskUseCase(mock_repo) with pytest.raises(ValueError, match="只有失败状态的任务才能重试"): use_case.execute("task_001") mock_repo.update.assert_not_called() task.mark_pending_from_failed.assert_not_called()