diff --git a/tests/unit/test_preview_edit_plan_association.py b/tests/unit/test_preview_edit_plan_association.py new file mode 100644 index 000000000..20351f84d --- /dev/null +++ b/tests/unit/test_preview_edit_plan_association.py @@ -0,0 +1,292 @@ +# -*- coding: utf-8 -*- +"""测试预览任务自动关联 edit_plan(generation_preview.py 增量覆盖率补充)。 + +覆盖 generation_preview.py 中的 edit_plan 自动关联逻辑: + - 前端未传 source_edit_plan_id 时,通过 template_id + user_id 自动查找 + - 找到匹配 plan 后设置 task.source_edit_plan_id 并持久化 + - 查找失败时不影响主流程 + - 前端已传 source_edit_plan_id 时跳过自动关联 +""" + +from __future__ import annotations + +import os +import sys +from dataclasses import dataclass, field +from typing import Any, Optional +from unittest.mock import MagicMock, patch + +os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing") +os.environ.setdefault("DATABASE_URL", "sqlite:///test.db") + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api")) + +from packages.domain.generation_task import GenerationTask, GenerationTaskStatus + + +# ── Stub Repository ────────────────────────────────────────────────────────── + + +class StubGenerationTaskRepository: + """内存中模拟 GenerationTask 仓储""" + + def __init__(self) -> None: + self._store: dict[str, Any] = {} + + def create(self, task: Any) -> Any: + self._store[task.id] = task + return task + + def get(self, task_id: str) -> Optional[Any]: + return self._store.get(task_id) + + def update(self, task: Any) -> Any: + if task.id not in self._store: + raise ValueError(f"GenerationTask {task.id} not found") + self._store[task.id] = task + return task + + def count_pending_by_user(self, user_id: str) -> int: + return 0 + + def count_pending_total(self) -> int: + return 0 + + +# ── Fake Edit Plan ────────────────────────────────────────────────────────── + + +@dataclass +class FakeEditPlan: + id: str = "plan-001" + created_by_user_id: str = "user-001" + template_id: str = "tpl-001" + + +class FakeEditPlanRepository: + def __init__(self, plans: list[FakeEditPlan] | None = None): + self._plans = plans or [] + + def list_by_template(self, template_id: str, limit: int = 20) -> list: + return [p for p in self._plans if p.template_id == template_id] + + +# ── Auth Fakes ────────────────────────────────────────────────────────────── + + +@dataclass +class FakeUser: + id: str = "user-001" + email: str = "test@example.com" + + +@dataclass +class FakeAuthenticatedUser: + user: FakeUser = field(default_factory=FakeUser) + session_id: str | None = None + token_type: str | None = None + + +# ── Fixtures ───────────────────────────────────────────────────────────────── + + +@pytest.fixture +def gen_task_repo() -> StubGenerationTaskRepository: + return StubGenerationTaskRepository() + + +@pytest.fixture +def mock_db() -> MagicMock: + return MagicMock() + + +@pytest.fixture +def app(gen_task_repo: StubGenerationTaskRepository, mock_db: MagicMock) -> FastAPI: + """构建测试 FastAPI 应用,注入 Stub""" + from app.api.routes.generation_preview import router + from app.auth import get_current_user + from app.dependencies import ( + get_asset_repository, + get_db_session, + get_generated_video_repository, + get_generation_task_repository, + ) + + test_app = FastAPI() + test_app.include_router(router, prefix="/api/v1/generation") + + test_app.dependency_overrides[get_current_user] = lambda: FakeAuthenticatedUser() + test_app.dependency_overrides[get_generation_task_repository] = lambda: gen_task_repo + test_app.dependency_overrides[get_db_session] = lambda: mock_db + test_app.dependency_overrides[get_asset_repository] = lambda: MagicMock() + test_app.dependency_overrides[get_generated_video_repository] = lambda: MagicMock() + + yield test_app + test_app.dependency_overrides.clear() + + +@pytest.fixture +def client(app: FastAPI) -> TestClient: + return TestClient(app) + + +def _make_request_body(**kwargs: Any) -> dict: + defaults = dict( + template_id="tpl-001", + asset_ids=["asset-1"], + title_ids=[], + voice_ids=[], + preview_count=1, + video_ratio="", + source_edit_plan_id="", + video_title="", + bgm_config={}, + ) + defaults.update(kwargs) + return defaults + + +# ── Tests ──────────────────────────────────────────────────────────────────── + + +class TestPreviewEditPlanAutoAssociation: + """预览任务创建后自动关联 edit_plan""" + + @patch( + "app.api.routes.generation_preview._resolve_strategy_id_from_template", + return_value="one_take", + ) + @patch( + "app.api.routes.generation_preview._infer_video_ratio_from_template", + return_value="9:16", + ) + @patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True) + def test_auto_associate_when_source_plan_empty( + self, + mock_enqueue, + mock_ratio, + mock_strategy, + client: TestClient, + gen_task_repo: StubGenerationTaskRepository, + ): + """前端未传 source_edit_plan_id 时,应通过 template_id+user_id 自动查找并关联""" + fake_plan = FakeEditPlan(id="plan-auto-001", created_by_user_id="user-001", template_id="tpl-001") + fake_plan_repo = FakeEditPlanRepository(plans=[fake_plan]) + + with patch( + "packages.adapters.sqlalchemy_impl.edit_plan_repository.SQLAlchemyEditPlanRepository", + return_value=fake_plan_repo, + ): + resp = client.post( + "/api/v1/generation/preview", + json=_make_request_body(source_edit_plan_id=""), + ) + + assert resp.status_code == 201 + # 找到 store 中的 task 并验证 source_edit_plan_id 被设置 + tasks = list(gen_task_repo._store.values()) + assert len(tasks) == 1 + task = tasks[0] + assert task.source_edit_plan_id == "plan-auto-001" + + @patch( + "app.api.routes.generation_preview._resolve_strategy_id_from_template", + return_value="one_take", + ) + @patch( + "app.api.routes.generation_preview._infer_video_ratio_from_template", + return_value="9:16", + ) + @patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True) + def test_skip_associate_when_source_plan_provided( + self, + mock_enqueue, + mock_ratio, + mock_strategy, + client: TestClient, + gen_task_repo: StubGenerationTaskRepository, + ): + """前端已传 source_edit_plan_id 时,不应触发自动关联""" + resp = client.post( + "/api/v1/generation/preview", + json=_make_request_body(source_edit_plan_id="plan-explicit-001"), + ) + + assert resp.status_code == 201 + tasks = list(gen_task_repo._store.values()) + assert len(tasks) == 1 + assert tasks[0].source_edit_plan_id == "plan-explicit-001" + + @patch( + "app.api.routes.generation_preview._resolve_strategy_id_from_template", + return_value="one_take", + ) + @patch( + "app.api.routes.generation_preview._infer_video_ratio_from_template", + return_value="9:16", + ) + @patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True) + def test_association_failure_does_not_break_main_flow( + self, + mock_enqueue, + mock_ratio, + mock_strategy, + client: TestClient, + gen_task_repo: StubGenerationTaskRepository, + ): + """edit_plan 查找异常时不影响任务创建和入队""" + with patch( + "packages.adapters.sqlalchemy_impl.edit_plan_repository.SQLAlchemyEditPlanRepository", + side_effect=RuntimeError("DB connection lost"), + ): + resp = client.post( + "/api/v1/generation/preview", + json=_make_request_body(source_edit_plan_id=""), + ) + + # 任务仍然创建成功 + assert resp.status_code == 201 + tasks = list(gen_task_repo._store.values()) + assert len(tasks) == 1 + # source_edit_plan_id 保持为空(关联失败) + assert tasks[0].source_edit_plan_id == "" + + @patch( + "app.api.routes.generation_preview._resolve_strategy_id_from_template", + return_value="one_take", + ) + @patch( + "app.api.routes.generation_preview._infer_video_ratio_from_template", + return_value="9:16", + ) + @patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True) + def test_auto_associate_skips_when_no_matching_user( + self, + mock_enqueue, + mock_ratio, + mock_strategy, + client: TestClient, + gen_task_repo: StubGenerationTaskRepository, + ): + """模板下有 plan 但 created_by_user_id 不匹配时,不关联""" + fake_plan = FakeEditPlan(id="plan-other-user", created_by_user_id="user-999", template_id="tpl-001") + fake_plan_repo = FakeEditPlanRepository(plans=[fake_plan]) + + with patch( + "packages.adapters.sqlalchemy_impl.edit_plan_repository.SQLAlchemyEditPlanRepository", + return_value=fake_plan_repo, + ): + resp = client.post( + "/api/v1/generation/preview", + json=_make_request_body(source_edit_plan_id=""), + ) + + assert resp.status_code == 201 + tasks = list(gen_task_repo._store.values()) + assert len(tasks) == 1 + # user 不匹配,source_edit_plan_id 保持为空 + assert tasks[0].source_edit_plan_id == ""