""" 生成视频管理 API 集成测试。 覆盖端点: - GET /generated-videos — 列出生成视频 - GET /generated-videos/{video_id} — 获取生成视频详情 - PATCH /generated-videos/{video_id}/review — 更新审核状态 - GET /generated-videos/{video_id}/download-url — 获取下载地址 使用 FastAPI TestClient + dependency_overrides 模式, 导入真实路由模块,mock 所有外部依赖。 """ from __future__ import annotations import os import sys from dataclasses import replace from datetime import datetime, timezone from unittest.mock import MagicMock # ── 环境变量 & 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.generated_videos import router from app.auth import AuthenticatedUser, get_current_user from app.core.storage import get_storage_service from app.dependencies import get_generated_video_repository, get_project_repository from packages.domain.entities import Project, User from packages.domain.generated_video import GeneratedVideo # --------------------------------------------------------------------------- # 1. 内存 Repository + 辅助函数 # --------------------------------------------------------------------------- class InMemoryGeneratedVideoRepository: """内存中的生成视频 Repository。""" def __init__(self): self._items: dict[str, GeneratedVideo] = {} def create(self, video: GeneratedVideo) -> GeneratedVideo: self._items[video.id] = video return video def get(self, video_id: str) -> GeneratedVideo | None: return self._items.get(video_id) def update(self, video: GeneratedVideo) -> GeneratedVideo: self._items[video.id] = video return video def list_by_project(self, project_id: str) -> list[GeneratedVideo]: return [v for v in self._items.values() if v.project_id == project_id] def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]: return [v for v in self._items.values() if v.generation_task_id == generation_task_id] def list_by_batch(self, batch_id: str) -> list[GeneratedVideo]: return [] class InMemoryProjectRepository: """内存中的项目 Repository。""" 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) -> Project | None: return self._projects.get(project_id) def find_by_owner_user_id(self, owner_user_id: str) -> list[Project]: return [p for p in self._projects.values() if p.owner_user_id == owner_user_id] def find_accessible_projects(self, user_id: str) -> list[Project]: 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 MockStorageService: """Mock OSS 存储服务。""" def get_download_url(self, file_url: str) -> str: return f"https://cdn.example.com/download/{file_url}?token=abc123" 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 = "proj-1", 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_video( project_id: str = "proj-1", name: str = "output.mp4", status: str = "completed", review_status: str = "pending_review", **kwargs, ) -> GeneratedVideo: return GeneratedVideo.create( project_id=project_id, generation_task_id=kwargs.pop("generation_task_id", "task-1"), name=name, file_url=kwargs.pop("file_url", f"generated/{name}"), file_size=kwargs.pop("file_size", 1024000), duration=kwargs.pop("duration", 30.5), width=kwargs.pop("width", 1920), height=kwargs.pop("height", 1080), fps=kwargs.pop("fps", 30.0), thumbnail_url=kwargs.pop("thumbnail_url", None), generation_params=kwargs.pop("generation_params", {"resolution": "1080p"}), ) # --------------------------------------------------------------------------- # 2. Fixtures # --------------------------------------------------------------------------- @pytest.fixture def video_repo(): return InMemoryGeneratedVideoRepository() @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 storage_service(): return MockStorageService() @pytest.fixture def client(video_repo, project_repo, storage_service): """创建带有依赖覆盖的 TestClient。""" test_app = FastAPI() test_app.include_router(router, prefix="/generated-videos") def _override_current_user(): return AuthenticatedUser(user=_make_user()) def _override_video_repo(): return video_repo def _override_project_repo(): return project_repo def _override_storage(): return storage_service test_app.dependency_overrides[get_current_user] = _override_current_user test_app.dependency_overrides[get_generated_video_repository] = _override_video_repo test_app.dependency_overrides[get_project_repository] = _override_project_repo test_app.dependency_overrides[get_storage_service] = _override_storage yield TestClient(test_app) test_app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 3. GET / — 列出生成视频 # --------------------------------------------------------------------------- class TestListGeneratedVideos: """列出生成视频端点测试。""" def test_empty_list(self, client): """无视频时返回空列表。""" resp = client.get("/generated-videos") assert resp.status_code == 200 data = resp.json() assert data["items"] == [] def test_list_all_user_videos(self, client, video_repo, project_repo): """列出当前用户所有项目的视频。""" v1 = _make_video(project_id="proj-1", name="video1.mp4") v2 = _make_video(project_id="proj-2", name="video2.mp4") v3 = _make_video(project_id="proj-other", name="other.mp4") # 其他用户 video_repo.create(v1) video_repo.create(v2) video_repo.create(v3) resp = client.get("/generated-videos") assert resp.status_code == 200 data = resp.json() assert len(data["items"]) == 2 names = {item["name"] for item in data["items"]} assert names == {"video1.mp4", "video2.mp4"} def test_filter_by_project_id(self, client, video_repo): """按 project_id 筛选视频。""" v1 = _make_video(project_id="proj-1", name="a.mp4") v2 = _make_video(project_id="proj-2", name="b.mp4") video_repo.create(v1) video_repo.create(v2) resp = client.get("/generated-videos?project_id=proj-1") assert resp.status_code == 200 data = resp.json() assert len(data["items"]) == 1 assert data["items"][0]["name"] == "a.mp4" def test_filter_by_nonexistent_project_returns_404(self, client): """筛选不存在的项目返回 404。""" resp = client.get("/generated-videos?project_id=nonexistent") assert resp.status_code == 404 def test_list_includes_download_url(self, client, video_repo): """列表响应应包含下载地址。""" v = _make_video(file_url="generated/test.mp4") video_repo.create(v) resp = client.get("/generated-videos") assert resp.status_code == 200 item = resp.json()["items"][0] assert "download_url" in item assert item["download_url"] is not None assert "cdn.example.com" in item["download_url"] def test_list_response_fields(self, client, video_repo): """列表响应包含所有必需字段。""" v = _make_video() video_repo.create(v) resp = client.get("/generated-videos") item = resp.json()["items"][0] for field in [ "id", "project_id", "generation_task_id", "name", "file_url", "file_size", "duration", "width", "height", "fps", "status", "review_status", "generation_params", "download_url", ]: assert field in item, f"缺少字段: {field}" def test_unauthorized_returns_401(self, video_repo, project_repo, storage_service): """未授权访问返回 401/403。""" test_app = FastAPI() test_app.include_router(router, prefix="/generated-videos") # 不覆盖 get_current_user,使用默认(会拒绝无 token 请求) test_app.dependency_overrides[get_generated_video_repository] = lambda: video_repo test_app.dependency_overrides[get_project_repository] = lambda: project_repo test_app.dependency_overrides[get_storage_service] = lambda: storage_service c = TestClient(test_app) resp = c.get("/generated-videos") # 无 token 时 fastapi HTTPBearer auto_error=False 会返回 None, # get_current_user 会抛 401 assert resp.status_code in (401, 403) test_app.dependency_overrides.clear() # --------------------------------------------------------------------------- # 4. GET /{video_id} — 获取生成视频详情 # --------------------------------------------------------------------------- class TestGetGeneratedVideo: """获取生成视频详情端点测试。""" def test_get_existing_video(self, client, video_repo): """获取存在的视频返回详情。""" v = _make_video(name="detail.mp4", duration=45.0) video_repo.create(v) resp = client.get(f"/generated-videos/{v.id}") assert resp.status_code == 200 data = resp.json() assert data["id"] == v.id assert data["name"] == "detail.mp4" assert data["duration"] == 45.0 assert data["status"] == "completed" def test_get_includes_download_url(self, client, video_repo): """详情响应包含下载地址。""" v = _make_video(file_url="generated/detail.mp4") video_repo.create(v) resp = client.get(f"/generated-videos/{v.id}") data = resp.json() assert "download_url" in data assert "cdn.example.com" in data["download_url"] def test_get_nonexistent_returns_404(self, client): """获取不存在的视频返回 404。""" resp = client.get("/generated-videos/nonexistent-video-id") assert resp.status_code == 404 assert "not found" in resp.json()["detail"].lower() def test_get_thumbnail_url(self, client, video_repo): """有缩略图时返回缩略图 URL。""" v = _make_video(thumbnail_url="thumbs/test.jpg") video_repo.create(v) resp = client.get(f"/generated-videos/{v.id}") data = resp.json() assert data["thumbnail_url"] == "thumbs/test.jpg" def test_get_generation_params(self, client, video_repo): """返回生成参数。""" params = {"resolution": "4k", "style": "cinematic"} v = _make_video(generation_params=params) video_repo.create(v) resp = client.get(f"/generated-videos/{v.id}") data = resp.json() assert data["generation_params"]["resolution"] == "4k" assert data["generation_params"]["style"] == "cinematic" # --------------------------------------------------------------------------- # 5. PATCH /{video_id}/review — 更新审核状态 # --------------------------------------------------------------------------- class TestUpdateReviewStatus: """更新审核状态端点测试。""" def test_approve_video(self, client, video_repo): """审核通过。""" v = _make_video(review_status="pending_review") video_repo.create(v) resp = client.patch( f"/generated-videos/{v.id}/review", json={"review_status": "approved"}, ) assert resp.status_code == 200 data = resp.json() assert data["review_status"] == "approved" # 验证 repository 已更新 updated = video_repo.get(v.id) assert updated.review_status == "approved" def test_reject_video(self, client, video_repo): """审核拒绝。""" v = _make_video(review_status="pending_review") video_repo.create(v) resp = client.patch( f"/generated-videos/{v.id}/review", json={"review_status": "rejected"}, ) assert resp.status_code == 200 assert resp.json()["review_status"] == "rejected" def test_set_pending_review(self, client, video_repo): """设置为待审核。""" v = _make_video(review_status="approved") video_repo.create(v) resp = client.patch( f"/generated-videos/{v.id}/review", json={"review_status": "pending_review"}, ) assert resp.status_code == 200 assert resp.json()["review_status"] == "pending_review" def test_nonexistent_video_returns_404(self, client): """更新不存在的视频返回 404。""" resp = client.patch( "/nonexistent-id/review", json={"review_status": "approved"}, ) assert resp.status_code == 404 def test_invalid_status_returns_422(self, client, video_repo): """无效审核状态返回 422。""" v = _make_video() video_repo.create(v) resp = client.patch( f"/generated-videos/{v.id}/review", json={"review_status": "invalid_status"}, ) assert resp.status_code == 422 def test_missing_status_returns_422(self, client, video_repo): """缺少 review_status 字段返回 422。""" v = _make_video() video_repo.create(v) resp = client.patch(f"/generated-videos/{v.id}/review", json={}) assert resp.status_code == 422 def test_update_returns_updated_fields(self, client, video_repo): """更新后返回完整的视频信息。""" v = _make_video(name="review_test.mp4") video_repo.create(v) resp = client.patch( f"/generated-videos/{v.id}/review", json={"review_status": "approved"}, ) data = resp.json() assert data["name"] == "review_test.mp4" assert "id" in data assert "download_url" in data # --------------------------------------------------------------------------- # 6. GET /{video_id}/download-url — 获取下载地址 # --------------------------------------------------------------------------- class TestGetDownloadUrl: """获取下载地址端点测试。""" def test_get_download_url_success(self, client, video_repo): """获取下载地址成功。""" v = _make_video(file_url="generated/video.mp4") video_repo.create(v) resp = client.get(f"/generated-videos/{v.id}/download-url") assert resp.status_code == 200 data = resp.json() assert data["video_id"] == v.id assert "download_url" in data assert "cdn.example.com" in data["download_url"] def test_nonexistent_video_returns_404(self, client): """获取不存在视频的下载地址返回 404。""" resp = client.get("/generated-videos/nonexistent-id/download-url") assert resp.status_code == 404 def test_download_url_format(self, client, video_repo): """下载地址格式正确。""" v = _make_video(file_url="my-video.mp4") video_repo.create(v) resp = client.get(f"/generated-videos/{v.id}/download-url") url = resp.json()["download_url"] assert url.startswith("https://") assert "token=" in url # --------------------------------------------------------------------------- # 7. 跨端点场景 # --------------------------------------------------------------------------- class TestCrossEndpointScenarios: """跨端点集成场景。""" def test_create_list_detail_review_flow(self, client, video_repo): """列表 → 详情 → 审核 完整流程。""" # 准备数据 v = _make_video(name="flow.mp4", review_status="pending_review") video_repo.create(v) # 1. 列表 list_resp = client.get("/generated-videos") assert list_resp.status_code == 200 assert len(list_resp.json()["items"]) == 1 # 2. 详情 detail_resp = client.get(f"/generated-videos/{v.id}") assert detail_resp.status_code == 200 assert detail_resp.json()["name"] == "flow.mp4" assert detail_resp.json()["review_status"] == "pending_review" # 3. 审核通过 review_resp = client.patch( f"/generated-videos/{v.id}/review", json={"review_status": "approved"}, ) assert review_resp.status_code == 200 assert review_resp.json()["review_status"] == "approved" # 4. 再次查看详情确认 detail_resp2 = client.get(f"/generated-videos/{v.id}") assert detail_resp2.json()["review_status"] == "approved" # 5. 获取下载地址 dl_resp = client.get(f"/generated-videos/{v.id}/download-url") assert dl_resp.status_code == 200 assert dl_resp.json()["video_id"] == v.id def test_multiple_videos_pagination_simulation(self, client, video_repo): """多个视频时列表正确返回所有视频。""" for i in range(5): v = _make_video(project_id="proj-1", name=f"video_{i}.mp4") video_repo.create(v) resp = client.get("/generated-videos") assert resp.status_code == 200 items = resp.json()["items"] assert len(items) == 5 names = {item["name"] for item in items} assert len(names) == 5 # 全部不同 if __name__ == "__main__": pytest.main([__file__, "-v"])