"""create 端点兜底复用预览产物 + confirm 标题同步 — 单元测试.""" 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 # ── Stubs ──────────────────────────────────────────────────────────────────── class StubGenerationTaskRepository: 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: self._store[task.id] = task return task def list_by_user(self, user_id: str) -> list[Any]: return [t for t in self._store.values() if t.created_by_user_id == user_id] def count_pending_by_user(self, user_id: str) -> int: return len( [ t for t in self._store.values() if t.created_by_user_id == user_id and t.status == GenerationTaskStatus.PENDING ] ) def count_pending_total(self) -> int: return len([t for t in self._store.values() if t.status == GenerationTaskStatus.PENDING]) def count_by_user(self, user_id: str) -> int: return len([t for t in self._store.values() if t.created_by_user_id == user_id]) @dataclass class FakeProject: id: str = "project-001" owner_user_id: str = "user-001" shared_users: list[str] = field(default_factory=list) name: str = "Test Project" def can_access(self, user_id: str) -> bool: return user_id == self.owner_user_id or user_id in self.shared_users class StubProjectRepository: def __init__(self) -> None: self._projects: dict[str, FakeProject] = {} def add(self, project: FakeProject) -> None: self._projects[project.id] = project def find_by_id(self, project_id: str) -> Optional[FakeProject]: return self._projects.get(project_id) @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 def _make_preview_task(**kwargs: Any) -> GenerationTask: defaults = dict( id="preview-task-001", project_id="project-001", asset_library_id="library-001", strategy_id="one_take", voice_library_id="", template_id="tmpl-001", asset_ids=["asset-1"], title_ids=[], voice_ids=[], status=GenerationTaskStatus.COMPLETED, progress=100.0, result_count=1, error_message="", created_by_user_id="user-001", source_edit_plan_id="plan-001", asset_select_mode="all", is_preview=True, source_task_id="", output_width=1080, output_height=1920, cover_url="", video_title="", resolution="", bgm_config={}, title_config={"text": "预览标题"}, ) defaults.update(kwargs) return GenerationTask(**defaults) def _make_db_with_preview(preview: GenerationTask): """Create a mock DB that returns a model-like object for the preview.""" db = MagicMock() mock_model = MagicMock() mock_model.id = preview.id mock_model.output_width = preview.output_width mock_model.output_height = preview.output_height # Set up chain: db.query(...).filter(...).order_by(...).first() chain = db.query.return_value chain.filter.return_value = chain chain.order_by.return_value = chain chain.first.return_value = mock_model return db, mock_model def _make_db_empty(): """Create a mock DB that returns None (no preview found).""" db = MagicMock() chain = db.query.return_value chain.filter.return_value = chain chain.order_by.return_value = chain chain.first.return_value = None return db # ── Fixtures ───────────────────────────────────────────────────────────────── @pytest.fixture def gen_task_repo() -> StubGenerationTaskRepository: return StubGenerationTaskRepository() @pytest.fixture def project_repo() -> StubProjectRepository: repo = StubProjectRepository() repo.add(FakeProject()) return repo @pytest.fixture def app( gen_task_repo: StubGenerationTaskRepository, project_repo: StubProjectRepository, ) -> FastAPI: from app.api.routes.generation_tasks import router from app.auth import get_current_user from app.dependencies import ( get_asset_library_repository, get_asset_repository, get_db_session, get_generated_video_repository, get_generation_task_repository, get_project_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_project_repository] = lambda: project_repo test_app.dependency_overrides[get_asset_library_repository] = lambda: MagicMock() test_app.dependency_overrides[get_asset_repository] = lambda: MagicMock() test_app.dependency_overrides[get_generated_video_repository] = lambda: MagicMock() # db_session will be overridden per-test yield test_app test_app.dependency_overrides.clear() def _make_client(app: FastAPI, db: MagicMock) -> TestClient: from app.dependencies import get_db_session app.dependency_overrides[get_db_session] = lambda: db return TestClient(app) # ── Domain: mark_confirmed with title_config ───────────────────────────────── class TestMarkConfirmedTitleConfig: def test_mark_confirmed_sets_title_config(self): task = _make_preview_task() task.mark_confirmed(title_config={"text": "新标题"}) assert task.is_preview is False assert task.title_config["text"] == "新标题" def test_mark_confirmed_without_title_config_preserves_existing(self): task = _make_preview_task(title_config={"text": "原标题"}) task.mark_confirmed() assert task.title_config["text"] == "原标题" # ── Create endpoint fallback ───────────────────────────────────────────────── class TestCreateEndpointFallback: """create 端点兜底复用预览产物。""" def test_reuse_completed_preview( self, app: FastAPI, gen_task_repo: StubGenerationTaskRepository, ): """带 source_edit_plan_id + is_preview=False → 复用已完成预览""" preview = _make_preview_task() gen_task_repo.create(preview) db, _ = _make_db_with_preview(preview) with ( patch( "packages.adapters.sqlalchemy_impl.generation_task_repository._to_domain", return_value=preview, ), patch( "app.api.routes.generation_tasks._writeback_edit_plan_config", ), ): client = _make_client(app, db) resp = client.post( "/api/v1/generation/tasks", json={ "project_id": "project-001", "template_id": "tmpl-001", "asset_ids": ["asset-1"], "source_edit_plan_id": "plan-001", "is_preview": False, "output_width": 1080, "output_height": 1920, }, ) assert resp.status_code == 200 data = resp.json() assert data["total"] == 1 assert data["items"][0]["id"] == preview.id assert data["items"][0]["is_preview"] is False def test_no_preview_found_creates_new_task( self, app: FastAPI, gen_task_repo: StubGenerationTaskRepository, ): """没有已完成预览 → 正常创建新任务""" db = _make_db_empty() with patch("app.api.routes.generation_tasks.safe_enqueue_generation_task", return_value=True): client = _make_client(app, db) resp = client.post( "/api/v1/generation/tasks", json={ "project_id": "project-001", "template_id": "tmpl-001", "asset_ids": ["asset-1"], "source_edit_plan_id": "plan-002", "is_preview": False, }, ) assert resp.status_code == 200 data = resp.json() assert data["total"] == 1 assert data["items"][0]["id"] != "preview-task-001" assert data["items"][0]["is_preview"] is False def test_preview_request_does_not_use_fallback( self, app: FastAPI, ): """is_preview=True → 不走兜底,正常创建预览任务""" db = _make_db_empty() with patch("app.api.routes.generation_tasks.safe_enqueue_generation_task", return_value=True): client = _make_client(app, db) resp = client.post( "/api/v1/generation/tasks", json={ "project_id": "project-001", "template_id": "tmpl-001", "asset_ids": ["asset-1"], "source_edit_plan_id": "plan-001", "is_preview": True, }, ) assert resp.status_code == 200 assert resp.json()["items"][0]["is_preview"] is True # 兜底查询 GenerationTaskModel 不应被调用(只可能查 EditPlanModel 做自动关联) from packages.adapters.sqlalchemy_impl.models import GenerationTaskModel for call_args in db.query.call_args_list: assert ( call_args[0][0] is not GenerationTaskModel ), "fallback should not query GenerationTaskModel for preview requests" def test_fallback_resolution_mismatch_creates_new( self, app: FastAPI, gen_task_repo: StubGenerationTaskRepository, ): """兜底找到预览但分辨率不一致 → 跳过复用,创建新任务""" preview = _make_preview_task(output_width=1080, output_height=1920) gen_task_repo.create(preview) db, _ = _make_db_with_preview(preview) with ( patch( "packages.adapters.sqlalchemy_impl.generation_task_repository._to_domain", return_value=preview, ), patch("app.api.routes.generation_tasks.safe_enqueue_generation_task", return_value=True), ): client = _make_client(app, db) resp = client.post( "/api/v1/generation/tasks", json={ "project_id": "project-001", "template_id": "tmpl-001", "asset_ids": ["asset-1"], "source_edit_plan_id": "plan-001", "is_preview": False, "output_width": 1920, "output_height": 1080, }, ) assert resp.status_code == 200 assert resp.json()["items"][0]["id"] != preview.id # ── Confirm endpoint title sync ────────────────────────────────────────────── class TestConfirmTitleSync: """confirm 端点 custom_title 同步到 title_config。""" def test_confirm_with_custom_title_updates_title_config( self, app: FastAPI, gen_task_repo: StubGenerationTaskRepository, ): """传了 custom_title → title_config.text 被更新""" preview = _make_preview_task( title_config={"text": "旧标题", "font_size": 32}, source_edit_plan_id="plan-001", ) gen_task_repo.create(preview) db = _make_db_empty() with patch("app.api.routes.generation_tasks._writeback_edit_plan_config"): client = _make_client(app, db) resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={ "output_width": 1080, "output_height": 1920, "custom_title": "新标题", }, ) assert resp.status_code == 200 item = resp.json()["items"][0] assert item["title_config"]["text"] == "新标题" assert item["title_config"]["font_size"] == 32 def test_confirm_without_custom_title_preserves_title( self, app: FastAPI, gen_task_repo: StubGenerationTaskRepository, ): """没传 custom_title → title_config 不变""" preview = _make_preview_task(title_config={"text": "原标题"}) gen_task_repo.create(preview) db = _make_db_empty() client = _make_client(app, db) resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={"output_width": 1080, "output_height": 1920}, ) assert resp.status_code == 200 assert resp.json()["items"][0]["title_config"]["text"] == "原标题"