"""Issue #1677 多视频批量生成 — 变体独立配置与批量预览/批量生成测试。 覆盖: - 批量预览:preview_count=N 一次创建 N 个独立任务,返回变体数组 - 变体克隆链路:N 个预览/正式任务各自关联独立克隆 plan - 变体独立配置:titles[]/voice_library_ids[]/cover_urls[] 按变体注入 - 长度校验:数组长度必须为 1 或 N(共用或独立),非法长度报错 - N=1 向后兼容:旧字段单值行为不变 """ from datetime import datetime, timezone from unittest.mock import MagicMock, patch import pytest from app.core.task_enqueue import GlobalQueueFull, UserPendingLimitExceeded from app.schemas.generation_task import ( BatchPreviewGenerationTaskResponse, CreateGenerationTaskRequest, CreatePreviewGenerationTaskRequest, ) from packages.domain import GenerationTask from packages.domain.generation_task import GenerationTaskStatus # ════════════════════════════════════════════════════════════════════════════ # 辅助构造 # ════════════════════════════════════════════════════════════════════════════ def _make_user(user_id="test_user_001"): mock_user = MagicMock() mock_user.id = user_id auth = MagicMock() auth.user = mock_user return auth def _make_task(task_id=None, status=GenerationTaskStatus.PENDING, source_plan_id=None): task = GenerationTask.create( project_id="", asset_library_id="", template_id="tpl_001", asset_ids=["asset_1"], ) if task_id: task.id = task_id task.status = status task.is_preview = True task.source_edit_plan_id = source_plan_id or "" task.voice_library_id = "" task.title_config = {} task.cover_url = "" return task def _make_preview_request(**kwargs): defaults = { "template_id": "tpl_001", "asset_ids": ["asset_1", "asset_2"], } defaults.update(kwargs) return CreatePreviewGenerationTaskRequest(**defaults) def _repo_mock(): repo = MagicMock() repo.count_pending_by_user.return_value = 0 repo.count_pending_total.return_value = 0 repo.get.side_effect = lambda tid: None return repo # ════════════════════════════════════════════════════════════════════════════ # Schema 校验:变体数组长度 # ════════════════════════════════════════════════════════════════════════════ class TestVariantArrayValidation: """变体数组字段长度校验。""" def test_preview_titles_length_matches_count(self): """titles 长度 = preview_count 合法""" req = _make_preview_request(preview_count=3, titles=["标题A", "标题B", "标题C"]) assert len(req.titles) == 3 def test_preview_titles_single_shared(self): """titles 长度 1 = 所有变体共用,合法""" req = _make_preview_request(preview_count=3, titles=["共用标题"]) assert req.titles == ["共用标题"] def test_preview_titles_wrong_length_raises(self): """titles 长度 2 与 preview_count=3 不匹配 → 报错""" with pytest.raises(ValueError, match="titles"): _make_preview_request(preview_count=3, titles=["A", "B"]) def test_preview_voice_ids_wrong_length_raises(self): """voice_library_ids 长度非法 → 报错""" from pydantic import ValidationError with pytest.raises(ValidationError, match="voice_library_ids"): _make_preview_request(preview_count=4, voice_library_ids=["v1", "v2"]) def test_preview_empty_arrays_ok(self): """空数组(回退单值字段)合法""" req = _make_preview_request(preview_count=3) assert req.titles == [] assert req.voice_library_ids == [] assert req.cover_urls == [] def test_generation_titles_length_matches_count(self): """正式生成 titles 长度 = count 合法""" req = CreateGenerationTaskRequest( template_id="tpl_1", asset_ids=["a1"], count=3, titles=["A", "B", "C"], ) assert len(req.titles) == 3 def test_generation_arrays_wrong_length_raises(self): """正式生成 cover_urls 长度与 count 不匹配 → 报错""" from pydantic import ValidationError with pytest.raises(ValidationError, match="cover_urls"): CreateGenerationTaskRequest( template_id="tpl_1", asset_ids=["a1"], count=3, cover_urls=["c1", "c2"], ) def test_generation_single_count_no_arrays(self): """N=1 且不传数组:完全旧行为""" req = CreateGenerationTaskRequest(template_id="tpl_1", asset_ids=["a1"]) assert req.count == 1 assert req.titles == [] assert req.voice_library_ids == [] assert req.cover_urls == [] # ════════════════════════════════════════════════════════════════════════════ # 批量预览路由 # ════════════════════════════════════════════════════════════════════════════ class TestBatchPreviewRoute: """POST /preview 批量变体。""" def test_preview_count_1_returns_single_item_array(self): """N=1 返回 items 长度 1 的批量响应(结构统一)""" from app.api.routes.generation_preview import create_preview_generation_task task = _make_task(task_id="task_1") repo = _repo_mock() with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.return_value = task with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): resp = create_preview_generation_task( _make_preview_request(preview_count=1), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) assert isinstance(resp, BatchPreviewGenerationTaskResponse) assert resp.total == 1 assert len(resp.items) == 1 assert resp.items[0].task_id == "task_1" assert resp.items[0].variant_index == 0 def test_preview_count_3_creates_three_independent_tasks(self): """N=3 创建 3 个独立任务,返回 3 个变体,task_id 各不相同""" from app.api.routes.generation_preview import create_preview_generation_task tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] repo = _repo_mock() with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.side_effect = tasks with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): resp = create_preview_generation_task( _make_preview_request(preview_count=3), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) assert resp.total == 3 task_ids = [item.task_id for item in resp.items] assert task_ids == ["task_0", "task_1", "task_2"] assert len(set(task_ids)) == 3 for i, item in enumerate(resp.items): assert item.variant_index == i def test_preview_count_3_clones_three_variant_plans(self): """有源 plan 时,N=3 克隆 3 个独立变体 plan(预览全部克隆,不用源 plan)""" from app.api.routes.generation_preview import create_preview_generation_task tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)] repo = _repo_mock() cloned_plan_ids = ["clone_1", "clone_2", "clone_3"] with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.side_effect = tasks with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc: clone_results = [MagicMock(id=pid) for pid in cloned_plan_ids] MockPlanSvc.return_value.clone_plan_for_variant.side_effect = clone_results create_preview_generation_task( _make_preview_request(preview_count=3), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) # 克隆被调用 3 次 assert MockPlanSvc.return_value.clone_plan_for_variant.call_count == 3 # 每个任务关联到不同的克隆 plan for i, task in enumerate(tasks): assert task.source_edit_plan_id == cloned_plan_ids[i] def test_preview_variant_titles_injected_per_variant(self): """titles[] 按变体注入 title_config.text""" from app.api.routes.generation_preview import create_preview_generation_task tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] repo = _repo_mock() captured_commands = [] with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: def _execute(cmd): captured_commands.append(cmd) return tasks[len(captured_commands) - 1] MockUC.return_value.execute.side_effect = _execute with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): create_preview_generation_task( _make_preview_request( preview_count=3, title_config={"font": "黑体", "position": "bottom"}, titles=["标题A", "标题B", "标题C"], ), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) assert len(captured_commands) == 3 assert captured_commands[0].title_config["text"] == "标题A" assert captured_commands[1].title_config["text"] == "标题B" assert captured_commands[2].title_config["text"] == "标题C" # 样式全局共用 assert all(c.title_config["font"] == "黑体" for c in captured_commands) def test_preview_shared_title_when_single_length(self): """titles 长度 1 = 所有变体共用同一标题""" from app.api.routes.generation_preview import create_preview_generation_task tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] repo = _repo_mock() captured = [] with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: def _execute(cmd): captured.append(cmd) return tasks[len(captured) - 1] MockUC.return_value.execute.side_effect = _execute with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): create_preview_generation_task( _make_preview_request(preview_count=3, titles=["共用标题"]), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) assert all(c.title_config["text"] == "共用标题" for c in captured) def test_preview_independent_voice_per_variant(self): """voice_library_ids[] 按变体注入独立配音""" from app.api.routes.generation_preview import create_preview_generation_task tasks = [_make_task(task_id=f"task_{i}") for i in range(3)] repo = _repo_mock() captured = [] with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: def _execute(cmd): captured.append(cmd) return tasks[len(captured) - 1] MockUC.return_value.execute.side_effect = _execute with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): create_preview_generation_task( _make_preview_request( preview_count=3, voice_library_ids=["voice_a", "voice_b", "voice_c"], ), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) assert [c.voice_library_id for c in captured] == ["voice_a", "voice_b", "voice_c"] def test_preview_voice_fallback_to_single_field(self): """voice_library_ids 为空时回退 voice_library_id 单值字段(向后兼容)""" from app.api.routes.generation_preview import create_preview_generation_task task = _make_task(task_id="task_1") repo = _repo_mock() captured = [] with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: def _execute(cmd): captured.append(cmd) return task MockUC.return_value.execute.side_effect = _execute with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True): create_preview_generation_task( _make_preview_request(voice_library_id="legacy_voice"), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) assert captured[0].voice_library_id == "legacy_voice" def test_preview_queue_limit_checks_total_count(self): """限流预检查按变体总数计:用户 pending + N 超限 → 429""" from app.api.routes.generation_preview import create_preview_generation_task from fastapi import HTTPException repo = MagicMock() repo.count_pending_by_user.return_value = 3 repo.count_pending_total.return_value = 0 with pytest.raises(HTTPException) as exc: create_preview_generation_task( _make_preview_request(preview_count=5), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) assert exc.value.status_code == 429 def test_preview_clone_failure_marks_all_failed(self): """克隆变体 plan 失败 → 已创建任务全部标记 failed 并 500""" from app.api.routes.generation_preview import create_preview_generation_task from fastapi import HTTPException tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)] repo = _repo_mock() with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC: MockUC.return_value.execute.side_effect = tasks with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc: MockPlanSvc.return_value.clone_plan_for_variant.side_effect = RuntimeError("db down") with pytest.raises(HTTPException) as exc: create_preview_generation_task( _make_preview_request(preview_count=3), authenticated_user=_make_user(), generation_task_repository=repo, db=MagicMock(), ) assert exc.value.status_code == 500 # 所有已创建任务都被标记 failed assert all(t.status == GenerationTaskStatus.FAILED for t in tasks) # ════════════════════════════════════════════════════════════════════════════ # 批量正式生成:变体配置注入 # ════════════════════════════════════════════════════════════════════════════ class TestBatchGenerationVariantConfig: """POST /tasks count=N 时变体独立配置。""" def _call_create_tasks(self, request, repo=None): from app.api.routes.generation_tasks import create_generation_task repo = repo or MagicMock() repo.count_pending_by_user.return_value = 0 repo.count_pending_total.return_value = 0 repo.update.return_value = None # 模板模式:asset_repository.find_by_id 返回 None(无 project 关联, # 纯模板模式 project_id/library_id 都为空),避免 MagicMock 属性污染 asset_repo = MagicMock() asset_repo.find_by_id.return_value = None # db.query().filter()...first() 返回 None:不走兜底关联编辑计划 db = MagicMock() db.query.return_value.filter.return_value.order_by.return_value.first.return_value = None return create_generation_task( request, authenticated_user=_make_user(), generation_task_repository=repo, project_repository=MagicMock(), asset_library_repository=MagicMock(), asset_repository=asset_repo, db=db, ) def test_count_3_variant_titles_voices_covers_injected(self): """count=3:titles/voice_library_ids/cover_urls 按变体注入""" from app.api.routes import generation_tasks as routes tasks = [_make_task(task_id=f"gen_{i}") for i in range(3)] captured = [] with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: def _execute(cmd): captured.append(cmd) t = tasks[len(captured) - 1] t.title_config = cmd.title_config t.voice_library_id = cmd.voice_library_id t.cover_url = cmd.cover_url return t MockUC.return_value.execute.side_effect = _execute with patch.object(routes, "safe_enqueue_generation_task", return_value=True): req = CreateGenerationTaskRequest( template_id="tpl_1", asset_ids=["a1"], count=3, title_config={"font": "宋体"}, titles=["成片标题1", "成片标题2", "成片标题3"], voice_library_ids=["v1", "v2", "v3"], cover_urls=["http://c1", "http://c2", "http://c3"], ) resp = self._call_create_tasks(req) assert resp.total == 3 assert [c.title_config["text"] for c in captured] == ["成片标题1", "成片标题2", "成片标题3"] assert [c.voice_library_id for c in captured] == ["v1", "v2", "v3"] assert [c.cover_url for c in captured] == ["http://c1", "http://c2", "http://c3"] # 样式共用 assert all(c.title_config["font"] == "宋体" for c in captured) def test_count_1_legacy_fields_unchanged(self): """N=1 不传数组:旧字段 voice_library_id/cover_url/title_config 行为不变""" from app.api.routes import generation_tasks as routes task = _make_task(task_id="gen_1") task.is_preview = False captured = [] with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: def _execute(cmd): captured.append(cmd) return task MockUC.return_value.execute.side_effect = _execute with patch.object(routes, "safe_enqueue_generation_task", return_value=True): req = CreateGenerationTaskRequest( template_id="tpl_1", asset_ids=["a1"], count=1, voice_library_id="legacy_voice", cover_url="http://legacy-cover", title_config={"text": "旧标题", "font": "黑体"}, ) resp = self._call_create_tasks(req) assert resp.total == 1 assert captured[0].voice_library_id == "legacy_voice" assert captured[0].cover_url == "http://legacy-cover" assert captured[0].title_config["text"] == "旧标题" def test_count_3_shared_single_value_arrays(self): """数组长度 1:3 个变体共用同一配音/封面""" from app.api.routes import generation_tasks as routes tasks = [_make_task(task_id=f"gen_{i}") for i in range(3)] captured = [] with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC: def _execute(cmd): captured.append(cmd) return tasks[len(captured) - 1] MockUC.return_value.execute.side_effect = _execute with patch.object(routes, "safe_enqueue_generation_task", return_value=True): req = CreateGenerationTaskRequest( template_id="tpl_1", asset_ids=["a1"], count=3, voice_library_ids=["shared_voice"], cover_urls=["http://shared"], ) self._call_create_tasks(req) assert all(c.voice_library_id == "shared_voice" for c in captured) assert all(c.cover_url == "http://shared" for c in captured) class TestVariantValueHelper: """_variant_value 取值逻辑。""" def test_empty_returns_fallback(self): from app.api.routes.generation_preview import _variant_value assert _variant_value([], 0, fallback="fb") == "fb" def test_single_length_shared(self): from app.api.routes.generation_preview import _variant_value assert _variant_value(["only"], 5) == "only" def test_indexed_access(self): from app.api.routes.generation_preview import _variant_value assert _variant_value(["a", "b", "c"], 1) == "b" def test_index_out_of_range_fallback(self): from app.api.routes.generation_preview import _variant_value assert _variant_value(["a", "b"], 9, fallback="x") == "x"