""" 仪表盘 API 集成测试。 覆盖端点: - GET /dashboard/overview — 仪表盘概览 验证返回数据结构、空数据场景、数据汇总正确性。 """ from __future__ import annotations import os import sys from datetime import datetime, timezone # ── 环境变量 & 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, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api")) from app.api.routes.dashboard import router from app.auth import AuthenticatedUser, get_current_user from app.dependencies import ( get_asset_repository, get_generation_task_repository, get_project_repository, get_title_library_repository, get_voice_library_repository, ) from packages.domain.entities import Project, User from packages.domain.generation_task import GenerationTask, GenerationTaskStatus # --------------------------------------------------------------------------- # 1. 内存 Repository # --------------------------------------------------------------------------- class InMemoryProjectRepository: def __init__(self): self._projects: dict[str, Project] = {} def save(self, project: Project) -> None: self._projects[project.id] = project def find_by_id(self, project_id: str): return self._projects.get(project_id) def find_by_owner_user_id(self, owner_user_id: str): return [p for p in self._projects.values() if p.owner_user_id == owner_user_id] def find_accessible_projects(self, user_id: str): return [p for p in self._projects.values() if p.owner_user_id == user_id] def count_by_owner(self, owner_user_id: str) -> int: return len(self.find_by_owner_user_id(owner_user_id)) def delete(self, project_id: str) -> bool: if project_id in self._projects: del self._projects[project_id] return True return False class InMemoryAssetRepository: def __init__(self): self._assets = [] def add_asset(self, project_id: str, storage_size: int = 0): self._assets.append({"project_id": project_id, "storage_size": storage_size}) def count_by_project_ids(self, project_ids: list[str]) -> int: return sum(1 for a in self._assets if a["project_id"] in project_ids) def sum_storage_by_project_ids(self, project_ids: list[str]) -> int: return sum(a["storage_size"] for a in self._assets if a["project_id"] in project_ids) # 其他方法占位 def create(self, asset): return asset def find_by_id(self, asset_id): return None def find_by_project(self, project_id, **kwargs): return [] def find_by_library(self, library_id, **kwargs): return [] def update(self, asset): return asset def delete(self, asset_id): return False def batch_delete(self, asset_ids): return 0 def search_candidates(self, **kwargs): return [] def find_by_tag_ids(self, tag_ids): return [] def count_by_project(self, project_id): return 0 def find_by_library_and_file_type(self, library_id, file_type): return [] def find_by_library_and_file_hash(self, library_id, file_hash): return None class InMemoryGenerationTaskRepository: def __init__(self): self._tasks = {} def add_task(self, task: GenerationTask): self._tasks[task.id] = 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 list_recent_by_user(self, user_id: str, limit: int = 5) -> list: user_tasks = [t for t in self._tasks.values() if t.created_by_user_id == user_id] # 按 created_at 倒序 user_tasks.sort(key=lambda t: t.created_at, reverse=True) return user_tasks[:limit] # 其他方法占位 def create(self, task): return task def get(self, task_id): return None def list_by_project(self, project_id): return [] def list_by_user(self, user_id): return [] def list_by_source_edit_plan(self, plan_id): return [] def update(self, task): return task class InMemoryTitleLibraryRepository: def __init__(self): self._items = {} def add_item(self, user_id: str): from uuid import uuid4 item_id = uuid4().hex self._items[item_id] = {"id": item_id, "user_id": user_id} return item_id def count_by_user(self, user_id: str, is_active: bool = True) -> int: return len([i for i in self._items.values() if i["user_id"] == user_id]) # 其他方法占位 def list_by_user(self, user_id, **kwargs): return [] def get(self, title_id, user_id): return None def create(self, item): return item def update(self, item): return item def delete(self, title_id, user_id): return False class InMemoryVoiceLibraryRepository: def __init__(self): self._items = {} def add_item(self, user_id: str): from uuid import uuid4 item_id = uuid4().hex self._items[item_id] = {"id": item_id, "user_id": user_id} return item_id def count_by_user(self, user_id: str) -> int: return len([i for i in self._items.values() if i["user_id"] == user_id]) # 其他方法占位 def list_by_user(self, user_id, **kwargs): return [] def get(self, voice_id, user_id): return None def create(self, item): return item def update(self, item): return item def delete(self, voice_id, user_id): return False # --------------------------------------------------------------------------- # 2. 辅助函数 # --------------------------------------------------------------------------- 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(project_id: str, owner_user_id: str = "user-test-001") -> Project: return Project( id=project_id, name=f"Project {project_id}", owner_user_id=owner_user_id, description="", created_at=datetime(2026, 1, 1, tzinfo=timezone.utc), ) def _make_generation_task( task_id: str, user_id: str = "user-test-001", status: GenerationTaskStatus = GenerationTaskStatus.COMPLETED, created_at: datetime | None = None, ) -> GenerationTask: return GenerationTask( id=task_id, project_id="proj-1", asset_library_id="lib-1", created_by_user_id=user_id, status=status, error_message="", created_at=created_at or datetime.now(timezone.utc), started_at=datetime.now(timezone.utc) if status != GenerationTaskStatus.PENDING else None, completed_at=datetime.now(timezone.utc) if status == GenerationTaskStatus.COMPLETED else None, ) # --------------------------------------------------------------------------- # 3. Fixtures # --------------------------------------------------------------------------- @pytest.fixture def project_repo(): repo = InMemoryProjectRepository() repo.save(_make_project("proj-1", "user-test-001")) repo.save(_make_project("proj-2", "user-test-001")) repo.save(_make_project("proj-other", "other-user")) return repo @pytest.fixture def asset_repo(): return InMemoryAssetRepository() @pytest.fixture def generation_task_repo(): return InMemoryGenerationTaskRepository() @pytest.fixture def title_library_repo(): return InMemoryTitleLibraryRepository() @pytest.fixture def voice_library_repo(): return InMemoryVoiceLibraryRepository() @pytest.fixture def client(project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo): """创建带有依赖覆盖的 TestClient。""" test_app = FastAPI() test_app.include_router(router, prefix="/dashboard") def _override_current_user(): return AuthenticatedUser(user=_make_user()) 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_asset_repository] = lambda: asset_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo yield TestClient(test_app) test_app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 4. GET /overview — 仪表盘概览 # --------------------------------------------------------------------------- class TestDashboardOverview: """仪表盘概览端点测试。""" def test_empty_data_returns_zeros(self, client): """空数据时所有计数为 0。""" resp = client.get("/dashboard/overview") assert resp.status_code == 200 data = resp.json() assert data["total_assets"] == 0 assert data["used_storage_bytes"] == 0 assert data["total_titles"] == 0 assert data["total_voices"] == 0 assert data["total_tasks"] == 0 assert data["total_products"] == 2 # fixture 中有 2 个项目 assert data["recent_tasks"] == [] def test_assets_count_and_storage(self, client, asset_repo): """素材统计正确。""" asset_repo.add_asset("proj-1", 1024) asset_repo.add_asset("proj-1", 2048) asset_repo.add_asset("proj-2", 4096) # 其他用户的不计入 asset_repo.add_asset("proj-other", 9999) resp = client.get("/dashboard/overview") data = resp.json() assert data["total_assets"] == 3 assert data["used_storage_bytes"] == 1024 + 2048 + 4096 def test_title_library_count(self, client, title_library_repo): """标题库统计正确。""" title_library_repo.add_item("user-test-001") title_library_repo.add_item("user-test-001") title_library_repo.add_item("user-test-001") title_library_repo.add_item("other-user") resp = client.get("/dashboard/overview") data = resp.json() assert data["total_titles"] == 3 def test_voice_library_count(self, client, voice_library_repo): """配音库统计正确。""" voice_library_repo.add_item("user-test-001") voice_library_repo.add_item("other-user") resp = client.get("/dashboard/overview") data = resp.json() assert data["total_voices"] == 1 def test_generation_tasks_count(self, client, generation_task_repo): """生成任务统计正确。""" generation_task_repo.add_task(_make_generation_task("task-1")) generation_task_repo.add_task(_make_generation_task("task-2")) generation_task_repo.add_task(_make_generation_task("task-other", user_id="other-user")) resp = client.get("/dashboard/overview") data = resp.json() assert data["total_tasks"] == 2 def test_recent_tasks_limited_to_5(self, client, generation_task_repo): """最近任务最多返回 5 个。""" for i in range(10): task = _make_generation_task(f"task-{i}") generation_task_repo.add_task(task) resp = client.get("/dashboard/overview") data = resp.json() assert len(data["recent_tasks"]) <= 5 def test_recent_tasks_have_correct_fields(self, client, generation_task_repo): """最近任务包含正确字段。""" task = _make_generation_task("task-1", status=GenerationTaskStatus.COMPLETED) generation_task_repo.add_task(task) resp = client.get("/dashboard/overview") data = resp.json() assert len(data["recent_tasks"]) == 1 item = data["recent_tasks"][0] for field in ["id", "task_type", "status", "current_step", "error_message", "updated_at"]: assert field in item, f"缺少字段: {field}" assert item["task_type"] == "generation" def test_subscription_info(self, client): """订阅信息正确。""" resp = client.get("/dashboard/overview") data = resp.json() assert "subscription" in data sub = data["subscription"] assert "plan" in sub assert "is_active" in sub assert sub["plan"] == "free" assert sub["is_active"] is True def test_pro_user_subscription( self, project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo ): """Pro 用户订阅信息正确。""" test_app = FastAPI() test_app.include_router(router, prefix="/dashboard") test_app.dependency_overrides[get_current_user] = lambda: AuthenticatedUser( user=_make_user(subscription_plan="pro", subscription_status="active") ) test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_asset_repository] = lambda: asset_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo c = TestClient(test_app) resp = c.get("/dashboard/overview") assert resp.status_code == 200 assert resp.json()["subscription"]["plan"] == "pro" assert resp.json()["subscription"]["is_active"] is True test_app.dependency_overrides.clear() def test_total_products_count(self, client, project_repo): """项目(产品)数量正确。""" resp = client.get("/dashboard/overview") data = resp.json() assert data["total_products"] == 2 # 新增一个项目后 project_repo.save(_make_project("proj-3", "user-test-001")) resp2 = client.get("/dashboard/overview") assert resp2.json()["total_products"] == 3 def test_unauthorized_returns_401( self, project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo ): """未授权访问返回 401/403。""" test_app = FastAPI() test_app.include_router(router, prefix="/dashboard") test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_asset_repository] = lambda: asset_repo test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo c = TestClient(test_app) resp = c.get("/dashboard/overview") assert resp.status_code in (401, 403) test_app.dependency_overrides.clear() def test_recent_tasks_status_mapping(self, client, generation_task_repo): """不同状态的任务显示正确的当前步骤。""" # 已完成任务 completed_task = _make_generation_task("task-completed", status=GenerationTaskStatus.COMPLETED) generation_task_repo.add_task(completed_task) resp = client.get("/dashboard/overview") tasks = resp.json()["recent_tasks"] completed = [t for t in tasks if t["id"] == "task-completed"][0] assert completed["status"] == "completed" assert "完成" in completed["current_step"] or "completed" in completed["current_step"].lower() if __name__ == "__main__": pytest.main([__file__, "-v"])