"""确认生成 API 单元测试. 覆盖 POST /generation/tasks/{task_id}/confirm 端点: - 预览任务已完成 → 直接复用(mark_confirmed),秒出 - 预览任务未完成 → 创建新任务走渲染流程 - 预览任务不存在 → 404 - 权限不足 → 403 - cover_url 正确传递 """ from __future__ import annotations import os import sys from dataclasses import dataclass, field from datetime import datetime, timezone 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 list_by_project(self, project_id: str) -> list[Any]: return [t for t in self._store.values() if t.project_id == project_id] 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_by_user(self, user_id: str) -> int: return len([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 list_recent_by_user(self, user_id: str, limit: int = 5) -> list[Any]: items = [t for t in self._store.values() if t.created_by_user_id == user_id] items.sort(key=lambda t: t.created_at, reverse=True) return items[:limit] def list_by_source_edit_plan(self, plan_id: str) -> list[Any]: return [t for t in self._store.values() if (t.source_edit_plan_id or "") == plan_id] # ── Stub Project Repository ────────────────────────────────────────────────── @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) # ── Fixtures ───────────────────────────────────────────────────────────────── @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 @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: """构建测试 FastAPI 应用,注入 Stub Repository""" 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_generated_video_repository, get_generation_task_repository, get_project_repository, ) test_app = FastAPI() test_app.include_router(router, prefix="/api/v1/generation") def override_get_current_user(): return FakeAuthenticatedUser() def override_get_generation_task_repository(): return gen_task_repo def override_get_project_repository(): return project_repo test_app.dependency_overrides[get_current_user] = override_get_current_user test_app.dependency_overrides[get_generation_task_repository] = override_get_generation_task_repository test_app.dependency_overrides[get_project_repository] = override_get_project_repository 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() yield test_app test_app.dependency_overrides.clear() @pytest.fixture def client(app: FastAPI) -> TestClient: return TestClient(app) 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="", 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="", asset_select_mode="all", is_preview=True, source_task_id="", output_width=1080, output_height=1920, cover_url="", video_title="", resolution="", bgm_config={}, ) defaults.update(kwargs) return GenerationTask(**defaults) # ── Tests ──────────────────────────────────────────────────────────────────── class TestConfirmGenerationReuse: """确认生成复用预览产物。""" def test_confirm_reuses_completed_preview( self, client: TestClient, gen_task_repo: StubGenerationTaskRepository, ) -> None: """预览已完成 → 直接复用,返回同一个任务,不创建新任务""" preview = _make_preview_task() gen_task_repo.create(preview) initial_count = len(gen_task_repo._store) resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={ "output_width": 1080, "output_height": 1920, "cover_url": "https://example.com/cover.jpg", }, ) assert resp.status_code == 200 data = resp.json() assert data["total"] == 1 item = data["items"][0] # 返回的是同一个任务(复用) assert item["id"] == preview.id # is_preview 变为 False assert item["is_preview"] is False # 分辨率更新 assert item["output_width"] == 1080 assert item["output_height"] == 1920 # 封面和标题更新 assert item["cover_url"] == "https://example.com/cover.jpg" # 没有创建新任务 assert len(gen_task_repo._store) == initial_count def test_confirm_updates_task_in_repo( self, client: TestClient, gen_task_repo: StubGenerationTaskRepository, ) -> None: """确认后的任务在 repo 中被更新""" preview = _make_preview_task() gen_task_repo.create(preview) resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={"cover_url": "https://cdn.example.com/cover.png", "custom_title": "测试标题"}, ) assert resp.status_code == 200 # 验证 repo 中的任务已被更新 updated = gen_task_repo.get(preview.id) assert updated is not None assert updated.is_preview is False assert updated.cover_url == "https://cdn.example.com/cover.png" def test_confirm_creates_new_task_when_preview_not_completed( self, client: TestClient, gen_task_repo: StubGenerationTaskRepository, ) -> None: """预览任务未完成 → 创建新任务走渲染流程""" preview = _make_preview_task(status=GenerationTaskStatus.RUNNING, progress=50.0) gen_task_repo.create(preview) with patch("app.api.routes.generation_tasks.safe_enqueue_generation_task", return_value=True): resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={"output_width": 1080, "output_height": 1920}, ) assert resp.status_code == 200 item = resp.json()["items"][0] # 创建了新任务 assert item["id"] != preview.id assert item["is_preview"] is False assert item["source_task_id"] == preview.id class TestConfirmGenerationErrors: """确认生成的错误处理。""" def test_confirm_not_found(self, client: TestClient) -> None: """预览任务不存在 → 404""" resp = client.post( "/api/v1/generation/tasks/nonexistent-task/confirm", json={"output_width": 1080, "output_height": 1920}, ) assert resp.status_code == 404 assert "not found" in resp.json()["detail"] def test_confirm_access_denied( self, client: TestClient, gen_task_repo: StubGenerationTaskRepository, ) -> None: """权限不足 → 403""" preview = _make_preview_task(created_by_user_id="other-user-999") gen_task_repo.create(preview) resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={"output_width": 1080, "output_height": 1920}, ) assert resp.status_code == 403 assert "denied" in resp.json()["detail"].lower() or "Access" in resp.json()["detail"] def test_confirm_preserves_config( self, client: TestClient, gen_task_repo: StubGenerationTaskRepository, ) -> None: """确认后任务保留预览任务的全部配置""" preview = _make_preview_task( voice_library_id="voice-001", template_id="tmpl-001", title_ids=["title-1", "title-2"], voice_ids=["voice-a"], ) gen_task_repo.create(preview) resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={"output_width": 1920, "output_height": 1080}, ) assert resp.status_code == 200 item = resp.json()["items"][0] assert item["is_preview"] is False assert item["output_width"] == 1920 assert item["output_height"] == 1080 # 默认封面和标题为空 assert item["cover_url"] == "" # 配置保留 assert item["voice_library_id"] == "voice-001" assert item["template_id"] == "tmpl-001" assert item["title_ids"] == ["title-1", "title-2"] assert item["voice_ids"] == ["voice-a"] def test_confirm_cover_and_title( self, client: TestClient, gen_task_repo: StubGenerationTaskRepository, ) -> None: """cover_url 正确传递""" preview = _make_preview_task() gen_task_repo.create(preview) resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={ "output_width": 1080, "output_height": 1920, "cover_url": "https://cdn.example.com/my-cover.png", }, ) assert resp.status_code == 200 item = resp.json()["items"][0] assert item["cover_url"] == "https://cdn.example.com/my-cover.png" def test_confirm_default_resolution( self, client: TestClient, gen_task_repo: StubGenerationTaskRepository, ) -> None: """不传分辨率时使用 ConfirmGenerationRequest 默认值 1080x1920""" preview = _make_preview_task() gen_task_repo.create(preview) resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={}, ) assert resp.status_code == 200 item = resp.json()["items"][0] assert item["output_width"] == 1080 assert item["output_height"] == 1920 def test_confirm_skips_reuse_when_resolution_mismatch( self, client: TestClient, gen_task_repo: StubGenerationTaskRepository, ) -> None: """请求的分辨率与预览渲染的分辨率不一致时,跳过复用,创建新任务""" preview = _make_preview_task(output_width=1080, output_height=1920) gen_task_repo.create(preview) initial_count = len(gen_task_repo._store) with patch("app.api.routes.generation_tasks.safe_enqueue_generation_task", return_value=True): resp = client.post( f"/api/v1/generation/tasks/{preview.id}/confirm", json={"output_width": 720, "output_height": 1280}, ) assert resp.status_code == 200 item = resp.json()["items"][0] # 创建了新任务(而非复用) assert item["id"] != preview.id assert item["is_preview"] is False assert item["source_task_id"] == preview.id assert len(gen_task_repo._store) == initial_count + 1