""" 任务中心 API 集成测试 覆盖端点: - GET /tasks — 列出用户任务 - POST /tasks/{task_id}/retry — 重试用户任务 - GET /projects/{project_id}/tasks — 列出项目任务 - POST /tasks/{task_type}/{source_id}/retry — 重试项目任务 使用 FastAPI TestClient + dependency_overrides 模式, 导入真实模块,mock 外部依赖(Celery任务)。 """ from __future__ import annotations import os import sys from datetime import datetime, timezone from pathlib import Path from unittest.mock import MagicMock, patch # ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ────────────────────────── 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, str(Path(__file__).resolve().parents[2] / "apps" / "api")) from app.api.routes.task_center import router from app.auth import AuthenticatedUser, get_current_user from app.dependencies import ( get_generation_task_repository, get_ingest_job_repository, get_project_repository, ) from packages.domain import ( GenerationTask, GenerationTaskStatus, IngestJob, IngestJobStatus, Project, User, ) # --------------------------------------------------------------------------- # 1. Stub Repository 实现 # --------------------------------------------------------------------------- class StubProjectRepository: def __init__(self, projects: dict[str, Project] | None = None): self._projects = projects or {} def find_by_id(self, project_id: str) -> Project | None: return self._projects.get(project_id) class StubGenerationTaskRepository: def __init__(self, tasks: dict[str, GenerationTask] | None = None): self._tasks = tasks or {} def create(self, task: GenerationTask) -> GenerationTask: self._tasks[task.id] = task return task def get(self, task_id: str) -> GenerationTask | None: return self._tasks.get(task_id) def list_by_project(self, project_id: str) -> list[GenerationTask]: return [t for t in self._tasks.values() if t.project_id == project_id] def list_by_user(self, user_id: str) -> list[GenerationTask]: return [t for t in self._tasks.values() if t.created_by_user_id == user_id] def update(self, task: GenerationTask) -> GenerationTask: self._tasks[task.id] = task return task def count_by_user(self, user_id: str) -> int: return len([t for t in self._tasks.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._tasks.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._tasks.values() if t.status == GenerationTaskStatus.PENDING]) def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: items = [t for t in self._tasks.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[GenerationTask]: return [t for t in self._tasks.values() if getattr(t, "source_edit_plan_id", "") == plan_id] def list_by_user_filtered( self, user_id: str, *, status: str | None = None, limit: int | None = None, offset: int = 0, ) -> list: """按用户+状态筛选任务列表(stub实现)。""" items = [t for t in self._tasks.values() if t.created_by_user_id == user_id] if status: items = [t for t in items if str(t.status) == status] # 按创建时间倒序 items.sort(key=lambda t: t.created_at or "", reverse=True) if offset: items = items[offset:] if limit is not None: items = items[:limit] return items def count_by_user_filtered( self, user_id: str, *, status: str | None = None, ) -> int: """按用户+状态筛选计数(stub实现)。""" items = [t for t in self._tasks.values() if t.created_by_user_id == user_id] if status: items = [t for t in items if str(t.status) == status] return len(items) def list_by_project_filtered( self, project_id: str, *, status: str | None = None, limit: int | None = None, offset: int = 0, ) -> list: """按项目+状态筛选任务列表(stub实现)。""" items = [t for t in self._tasks.values() if t.project_id == project_id] if status: items = [t for t in items if str(t.status) == status] # 按创建时间倒序 items.sort(key=lambda t: t.created_at or "", reverse=True) if offset: items = items[offset:] if limit is not None: items = items[:limit] return items def count_by_project_filtered( self, project_id: str, *, status: str | None = None, ) -> int: """按项目+状态筛选计数(stub实现)。""" items = [t for t in self._tasks.values() if t.project_id == project_id] if status: items = [t for t in items if str(t.status) == status] return len(items) class StubIngestJobRepository: def __init__(self, jobs: dict[str, IngestJob] | None = None): self._jobs = jobs or {} def create(self, job: IngestJob) -> IngestJob: self._jobs[job.id] = job return job def get(self, job_id: str) -> IngestJob | None: return self._jobs.get(job_id) def update(self, job: IngestJob) -> IngestJob: self._jobs[job.id] = job return job def update_status(self, job_id: str, status, **kwargs): job = self._jobs.get(job_id) if job: job.status = status def list_by_project(self, project_id: str, skip: int = 0, limit: int = 50) -> list[IngestJob]: return [j for j in self._jobs.values() if j.project_id == project_id][skip : skip + limit] def list_by_library(self, library_id: str, skip: int = 0, limit: int = 50) -> list[IngestJob]: return [j for j in self._jobs.values() if j.library_id == library_id][skip : skip + limit] # --------------------------------------------------------------------------- # 2. Helpers & Fixtures # --------------------------------------------------------------------------- def _make_user(**overrides) -> User: defaults = dict( id="user-test-001", email="test@example.com", display_name="Test User", username="testuser", subscription_plan="free", subscription_status="active", max_projects=3, max_storage_gb=10, created_at=datetime(2026, 1, 1, tzinfo=timezone.utc), ) defaults.update(overrides) return User(**defaults) def _make_project(id: str = "proj-1", owner_user_id: str = "user-test-001") -> Project: return Project(id=id, name="Test Project", owner_user_id=owner_user_id) def _make_generation_task( task_id: str = "gen-task-1", project_id: str = "proj-1", user_id: str = "user-test-001", status: GenerationTaskStatus = GenerationTaskStatus.PENDING, ) -> GenerationTask: task = GenerationTask( id=task_id, project_id=project_id, asset_library_id="lib-1", strategy_id="s1", voice_library_id="v1", created_by_user_id=user_id, ) task.status = status return task def _make_ingest_job( job_id: str = "ingest-job-1", project_id: str = "proj-1", library_id: str = "lib-1", status: IngestJobStatus = IngestJobStatus.PENDING, ) -> IngestJob: job = IngestJob( id=job_id, project_id=project_id, library_id=library_id, storage_key="uploads/test.mp4", ) job.status = status return job @pytest.fixture def client(): """创建带有依赖覆盖的 TestClient。""" test_app = FastAPI() test_app.include_router(router) project = _make_project() project_repo = StubProjectRepository({project.id: project}) task_repo = StubGenerationTaskRepository() ingest_repo = StubIngestJobRepository() def _override_current_user(): mock_auth = MagicMock(spec=AuthenticatedUser) mock_auth.user = _make_user() return mock_auth test_app.dependency_overrides[get_current_user] = _override_current_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo yield TestClient(test_app) test_app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 3. GET /tasks — 列出用户任务 # --------------------------------------------------------------------------- class TestListUserTasks: """列出用户任务端点测试。""" def test_empty_list(self, client): """无任务时返回空列表。""" resp = client.get("/tasks") assert resp.status_code == 200 data = resp.json() assert "items" in data assert data["items"] == [] def test_list_returns_generation_tasks(self, client): """返回当前用户的 generation 任务。""" # 直接在 repository 中注入任务 from app.dependencies import get_generation_task_repository task_repo = StubGenerationTaskRepository() task1 = _make_generation_task("gen-1", status=GenerationTaskStatus.PENDING) task2 = _make_generation_task("gen-2", status=GenerationTaskStatus.COMPLETED) task_repo.create(task1) task_repo.create(task2) test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _override_current_user(): mock_auth = MagicMock(spec=AuthenticatedUser) mock_auth.user = _make_user() return mock_auth test_app.dependency_overrides[get_current_user] = _override_current_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: StubIngestJobRepository() tc = TestClient(test_app) resp = tc.get("/tasks") assert resp.status_code == 200 data = resp.json() assert len(data["items"]) == 2 # 验证响应字段 for item in data["items"]: assert "id" in item assert "task_type" in item assert item["task_type"] == "generation" assert "status" in item assert "current_step" in item assert "retryable" in item test_app.dependency_overrides.clear() def test_tasks_sorted_by_updated_time(self, client): """任务按更新时间倒序排列。""" # 由于两个任务同时创建,验证它们都出现在列表中 resp = client.get("/tasks") assert resp.status_code == 200 data = resp.json() assert isinstance(data["items"], list) def test_task_response_fields(self, client): """任务响应包含所有必需字段。""" resp = client.get("/tasks") assert resp.status_code == 200 # 空列表也应该返回正确的结构 assert resp.json()["items"] == [] # --------------------------------------------------------------------------- # 4. POST /tasks/{task_id}/retry — 重试用户任务 # --------------------------------------------------------------------------- class TestRetryUserTask: """重试用户任务端点测试。""" def test_retry_nonexistent_task_returns_404(self, client): """重试不存在的任务返回 404。""" resp = client.post("/tasks/nonexistent-task-id/retry") assert resp.status_code == 404 assert "not found" in resp.json()["detail"].lower() @patch("app.api.routes.task_center.celery_app") def test_retry_pending_task_returns_409(self, mock_celery, client): """重试 pending 状态的任务返回 409(只有 failed 任务才能重试)。""" mock_celery.send_task = MagicMock() # 在 repository 中创建一个 pending 任务 task_repo = StubGenerationTaskRepository() task = _make_generation_task("gen-pending", status=GenerationTaskStatus.PENDING) task_repo.create(task) test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: StubIngestJobRepository() tc = TestClient(test_app) resp = tc.post("/tasks/gen-pending/retry") assert resp.status_code == 409 assert "Only failed" in resp.json()["detail"] test_app.dependency_overrides.clear() @patch("app.api.routes.task_center.celery_app") def test_retry_completed_task_returns_409(self, mock_celery, client): """重试 completed 状态的任务返回 409。""" mock_celery.send_task = MagicMock() task_repo = StubGenerationTaskRepository() task = _make_generation_task("gen-completed", status=GenerationTaskStatus.COMPLETED) task_repo.create(task) test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: StubIngestJobRepository() tc = TestClient(test_app) resp = tc.post("/tasks/gen-completed/retry") assert resp.status_code == 409 test_app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 5. GET /projects/{project_id}/tasks — 列出项目任务 # --------------------------------------------------------------------------- class TestListProjectTasks: """列出项目任务端点测试。""" def test_empty_project_tasks(self, client): """项目无任务时返回空列表。""" resp = client.get("/projects/proj-1/tasks") assert resp.status_code == 200 data = resp.json() assert "items" in data assert data["items"] == [] def test_project_not_found(self, client): """项目不存在返回 404。""" resp = client.get("/projects/nonexistent-project/tasks") assert resp.status_code == 404 assert "Project not found" in resp.json()["detail"] def test_returns_ingest_and_generation_tasks(self, client): """返回项目中 ingest 和 generation 两种任务。""" # 在 repository 中注入任务 task_repo = StubGenerationTaskRepository() gen_task = _make_generation_task("gen-proj-1", status=GenerationTaskStatus.PENDING) task_repo.create(gen_task) ingest_repo = StubIngestJobRepository() ingest_job = _make_ingest_job("ingest-proj-1", status=IngestJobStatus.PENDING) ingest_repo.create(ingest_job) test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo tc = TestClient(test_app) resp = tc.get("/projects/proj-1/tasks") assert resp.status_code == 200 data = resp.json() assert len(data["items"]) == 2 task_types = {item["task_type"] for item in data["items"]} assert "generation" in task_types assert "ingest" in task_types test_app.dependency_overrides.clear() def test_project_task_response_fields(self, client): """项目任务响应包含所有必需字段。""" resp = client.get("/projects/proj-1/tasks") assert resp.status_code == 200 data = resp.json() assert isinstance(data["items"], list) # --------------------------------------------------------------------------- # 6. POST /tasks/{task_type}/{source_id}/retry — 重试项目任务 # --------------------------------------------------------------------------- class TestRetryProjectTask: """重试项目任务端点测试。""" def test_retry_unsupported_task_type_returns_400(self, client): """不支持的任务类型返回 400。""" resp = client.post("/tasks/unknown/some-source-id/retry") assert resp.status_code == 400 assert "Unsupported" in resp.json()["detail"] @patch("app.core.task_enqueue.celery_app") def test_retry_failed_generation_task(self, mock_celery, client): """重试失败的 generation 任务成功。""" mock_celery.send_task = MagicMock() task_repo = StubGenerationTaskRepository() task = _make_generation_task("gen-failed-1", status=GenerationTaskStatus.FAILED) task_repo.create(task) ingest_repo = StubIngestJobRepository() test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo tc = TestClient(test_app) resp = tc.post("/tasks/generation/gen-failed-1/retry") assert resp.status_code == 200 data = resp.json() assert data["task_type"] == "generation" assert data["status"] == "pending" assert "current_step" in data # 原地重试:source_id 保持不变(复用同一个任务) assert data["source_id"] == "gen-failed-1" # 验证 Celery 任务被发送 assert mock_celery.send_task.called test_app.dependency_overrides.clear() @patch("app.api.routes.task_center.celery_app") def test_retry_failed_ingest_task(self, mock_celery, client): """重试失败的 ingest 任务成功。""" mock_celery.send_task = MagicMock() ingest_repo = StubIngestJobRepository() job = _make_ingest_job("ingest-failed-1", status=IngestJobStatus.FAILED) ingest_repo.create(job) task_repo = StubGenerationTaskRepository() test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo tc = TestClient(test_app) resp = tc.post("/tasks/ingest/ingest-failed-1/retry") assert resp.status_code == 200 data = resp.json() assert data["task_type"] == "ingest" assert data["status"] == "pending" assert mock_celery.send_task.called assert mock_celery.send_task.call_args[0][0] == "worker.ingest_asset" test_app.dependency_overrides.clear() def test_retry_pending_generation_task_returns_409(self, client): """重试 pending 状态的 generation 任务返回 409。""" task_repo = StubGenerationTaskRepository() task = _make_generation_task("gen-pending-proj", status=GenerationTaskStatus.PENDING) task_repo.create(task) ingest_repo = StubIngestJobRepository() test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo tc = TestClient(test_app) resp = tc.post("/tasks/generation/gen-pending-proj/retry") assert resp.status_code == 409 assert "Only failed" in resp.json()["detail"] test_app.dependency_overrides.clear() def test_retry_nonexistent_generation_task_returns_404(self, client): """重试不存在的 generation 任务返回 404。""" task_repo = StubGenerationTaskRepository() ingest_repo = StubIngestJobRepository() test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo tc = TestClient(test_app) resp = tc.post("/tasks/generation/nonexistent-id/retry") assert resp.status_code == 404 test_app.dependency_overrides.clear() def test_retry_nonexistent_ingest_task_returns_404(self, client): """重试不存在的 ingest 任务返回 404。""" task_repo = StubGenerationTaskRepository() ingest_repo = StubIngestJobRepository() test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo tc = TestClient(test_app) resp = tc.post("/tasks/ingest/nonexistent-id/retry") assert resp.status_code == 404 test_app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 7. 跨端点集成场景 # --------------------------------------------------------------------------- class TestTaskCenterCrossEndpoint: """任务中心跨端点集成测试。""" @patch("app.core.task_enqueue.celery_app") def test_list_then_retry_then_list(self, mock_celery, client): """列出任务 → 重试失败任务 → 再列出验证新任务。""" mock_celery.send_task = MagicMock() task_repo = StubGenerationTaskRepository() failed_task = _make_generation_task("gen-fail-cross", status=GenerationTaskStatus.FAILED) failed_task.error_message = "ffmpeg error" task_repo.create(failed_task) ingest_repo = StubIngestJobRepository() test_app = FastAPI() test_app.include_router(router) project_repo = StubProjectRepository({_make_project().id: _make_project()}) def _get_user(): return MagicMock(spec=AuthenticatedUser, user=_make_user()) test_app.dependency_overrides[get_current_user] = _get_user test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: task_repo test_app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo tc = TestClient(test_app) # 1. 列出任务 list_resp = tc.get("/tasks") assert list_resp.status_code == 200 items = list_resp.json()["items"] assert len(items) == 1 assert items[0]["retryable"] is True # failed 任务应可重试 # 2. 重试失败任务 retry_resp = tc.post("/tasks/gen-fail-cross/retry") assert retry_resp.status_code == 200 # 3. 再次列出:原地重试,任务数不变(仍是1个),但状态变为 pending list_resp2 = tc.get("/tasks") assert list_resp2.status_code == 200 items2 = list_resp2.json()["items"] assert len(items2) == 1 assert items2[0]["status"] == "pending" assert items2[0]["source_id"] == "gen-fail-cross" test_app.dependency_overrides.clear() if __name__ == "__main__": pytest.main([__file__, "-v"])