Files
xiaoxia-saas/tests/integration/test_task_center_api.py
xiaoxia 368baf683b
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
fix: StubGenerationTaskRepository补list_by_user_filtered方法 (#306)
2026-07-14 11:02:57 +08:00

710 lines
26 KiB
Python
Executable File

"""
任务中心 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"])