"""查重 API 集成测试。 覆盖端点: - GET /records — 列表查询(含分页) - GET /records/{record_id} — 详情查询 - DELETE /records/{record_id} — 删除记录 - POST /records/{record_id}/retry — 重试查重 使用 FastAPI TestClient + dependency_overrides 模式, 不依赖真实数据库。 """ from __future__ import annotations import os import sys import types from dataclasses import dataclass, field from datetime import datetime, timezone from typing import Any from unittest.mock import MagicMock from uuid import uuid4 import pytest from fastapi import FastAPI from fastapi.testclient import TestClient # --------------------------------------------------------------------------- # 1. 安装 mock 模块(复用 test_duplication_upload_error_handling 的模式) # --------------------------------------------------------------------------- # 保存被覆盖的原始模块,以便测试结束后恢复 _SAVED_MODULES: dict[str, Any] = {} def _install_mocks(): """安装所有必需的 mock 模块,使路由模块可导入。""" # 记录所有将被覆盖的模块 key,用于后续恢复 _keys_to_save = [ "packages.domain.entities", "packages.domain.duplication", "packages.ports.user_repository", "packages.ports.duplication_repository", "packages.adapters.sqlalchemy_impl.user_repository", "packages.adapters.sqlalchemy_impl.duplication_repository", "packages.adapters.sqlalchemy_impl.session", "packages.adapters.redis", "packages.adapters.smtp", "packages.application", "app.config", "app.auth", "app.dependencies", "app.core.storage", "app.schemas.duplication", ] for _k in _keys_to_save: if _k in sys.modules: _SAVED_MODULES[_k] = sys.modules[_k] # packages.domain.entities @dataclass(slots=True) class User: id: str = "user-test-001" email: str = "test@example.com" display_name: str = "Test User" username: str = "testuser" password_hash: str = "" email_verified: bool = False email_verification_token: str | None = None password_reset_token: str | None = None password_reset_expires_at: datetime | None = None last_login_at: datetime | None = None last_login_ip: str | None = None subscription_plan: str = "free" subscription_status: str = "active" subscription_expires_at: datetime | None = None max_projects: int = 3 max_storage_gb: int = 10 used_storage_gb: float = 0.0 created_at: datetime = field(default_factory=lambda: datetime(2026, 1, 1, tzinfo=timezone.utc)) entities_mod = types.ModuleType("packages.domain.entities") entities_mod.User = User sys.modules["packages.domain.entities"] = entities_mod # packages.domain.duplication — 使用真实域模型 @dataclass(slots=True) class DuplicateSegment: id: str source_start: float source_end: float matched_video_id: str matched_video_name: str matched_start: float matched_end: float similarity: float @dataclass(slots=True) class DuplicationRecord: id: str user_id: str filename: str file_size: int storage_key: str duration_seconds: float = 0.0 status: str = "pending" duplicate_rate: float | None = None duplicate_count: int = 0 video_fingerprint: dict | None = None error_message: str = "" segments: list = field(default_factory=list) created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) @classmethod def create(cls, user_id, filename, file_size, storage_key, **kwargs): return cls( id=uuid4().hex, user_id=user_id, filename=filename, file_size=file_size, storage_key=storage_key, **kwargs, ) def mark_processing(self): self.status = "processing" self.updated_at = datetime.now(timezone.utc) def mark_completed(self, duplicate_rate, duplicate_count, segments): if not 0 <= duplicate_rate <= 100: raise ValueError("duplicate_rate must be between 0 and 100") self.status = "completed" self.duplicate_rate = duplicate_rate self.duplicate_count = duplicate_count self.segments = segments self.updated_at = datetime.now(timezone.utc) def mark_failed(self, error_message): self.status = "failed" self.error_message = error_message self.updated_at = datetime.now(timezone.utc) def can_retry(self): return self.status == "failed" def reset_for_retry(self): self.status = "pending" self.error_message = "" self.duplicate_rate = None self.duplicate_count = 0 self.segments = [] self.video_fingerprint = None duplication_mod = types.ModuleType("packages.domain.duplication") duplication_mod.DuplicateSegment = DuplicateSegment duplication_mod.DuplicationRecord = DuplicationRecord sys.modules["packages.domain.duplication"] = duplication_mod # packages.ports for name in ["user_repository", "duplication_repository"]: mod = types.ModuleType(f"packages.ports.{name}") sys.modules[f"packages.ports.{name}"] = mod sys.modules["packages.ports.user_repository"].UserRepository = MagicMock sys.modules["packages.ports.duplication_repository"].DuplicationRecordRepository = MagicMock # packages namespace modules for name in [ "packages", "packages.domain", "packages.ports", "packages.adapters", "packages.adapters.sqlalchemy_impl", "packages.adapters.sqlalchemy_impl.user_repository", "packages.adapters.sqlalchemy_impl.duplication_repository", "packages.adapters.sqlalchemy_impl.session", "packages.adapters.redis", "packages.adapters.smtp", ]: if name not in sys.modules: sys.modules[name] = types.ModuleType(name) sys.modules["packages.adapters.sqlalchemy_impl.user_repository"].SQLAlchemyUserRepository = MagicMock sys.modules["packages.adapters.sqlalchemy_impl.duplication_repository"].SQLAlchemyDuplicationRecordRepository = ( MagicMock ) sys.modules["packages.adapters.sqlalchemy_impl.session"].build_session_factory = MagicMock( return_value=(MagicMock(), MagicMock()) ) sys.modules["packages.adapters.redis"].NoopSessionStore = MagicMock sys.modules["packages.adapters.redis"].SessionStore = MagicMock sys.modules["packages.adapters.smtp"].EmailConfig = MagicMock sys.modules["packages.adapters.smtp"].NoopEmailService = MagicMock sys.modules["packages.adapters.smtp"].get_email_service = MagicMock() # packages.application (UseCases) — 使用真实逻辑 app_mod = types.ModuleType("packages.application") @dataclass class UploadForDuplicationCommand: user_id: str filename: str file_size: int storage_key: str duration_seconds: float = 0.0 class UploadForDuplicationUseCase: def __init__(self, repo): self.repo = repo def execute(self, cmd): record = DuplicationRecord.create( user_id=cmd.user_id, filename=cmd.filename, file_size=cmd.file_size, storage_key=cmd.storage_key, ) return self.repo.create(record) class ListDuplicationRecordsUseCase: def __init__(self, repo): self.repo = repo def execute(self, user_id, *, offset=0, limit=50): if not user_id.strip(): raise ValueError("user_id 不能为空") return self.repo.list_by_user(user_id.strip(), offset=offset, limit=limit) class GetDuplicationDetailUseCase: def __init__(self, repo): self.repo = repo def execute(self, record_id): return self.repo.get(record_id) class DeleteDuplicationRecordUseCase: def __init__(self, repo): self.repo = repo def execute(self, record_id): return self.repo.delete(record_id) class RetryDuplicationUseCase: def __init__(self, repo): self.repo = repo def execute(self, record_id): record = self.repo.get(record_id) if record is None: return None record.status = "pending" record.error_message = "" record.duplicate_rate = None record.duplicate_count = 0 record.segments = [] return self.repo.update(record) app_mod.UploadForDuplicationCommand = UploadForDuplicationCommand app_mod.UploadForDuplicationUseCase = UploadForDuplicationUseCase app_mod.ListDuplicationRecordsUseCase = ListDuplicationRecordsUseCase app_mod.GetDuplicationDetailUseCase = GetDuplicationDetailUseCase app_mod.DeleteDuplicationRecordUseCase = DeleteDuplicationRecordUseCase app_mod.RetryDuplicationUseCase = RetryDuplicationUseCase sys.modules["packages.application"] = app_mod # app.config config_mod = types.ModuleType("app.config") class _Settings: JWT_SECRET_KEY = "test-secret-key-for-dup-api-tests" DATABASE_URL = "sqlite:///test.db" REDIS_URL = "redis://localhost:6379/0" ENABLE_REDIS_SESSIONS = False SMTP_HOST = "" SMTP_PORT = 587 SMTP_USER = "" SMTP_PASSWORD = "" SMTP_FROM_EMAIL = "" SMTP_FROM_NAME = "" SMTP_USE_TLS = False ENABLE_EMAIL_DELIVERY = False OSS_DIRECT_UPLOAD_MAX_MB = 100 OSS_BUCKET_NAME = "test-bucket" OSS_ENDPOINT = "oss-cn-hangzhou.aliyuncs.com" OSS_ACCESS_KEY_ID = "test-key" OSS_ACCESS_KEY_SECRET = "test-secret" config_mod.settings = _Settings() config_mod.get_settings = lambda: _Settings() sys.modules["app.config"] = config_mod # app.auth @dataclass(frozen=True, slots=True) class AuthenticatedUser: user: User session_id: str | None = None token_type: str | None = None async def _mock_get_current_user(): return AuthenticatedUser(user=User()) auth_mod = types.ModuleType("app.auth") auth_mod.AuthenticatedUser = AuthenticatedUser auth_mod.get_current_user = _mock_get_current_user sys.modules["app.auth"] = auth_mod # app.dependencies deps_mod = types.ModuleType("app.dependencies") deps_mod.get_db_session = MagicMock() deps_mod.get_duplication_repository = MagicMock() sys.modules["app.dependencies"] = deps_mod # app.core.storage storage_mod = types.ModuleType("app.core.storage") class OSSStorageService: def upload_file(self, content, key, content_type=None): pass def get_storage_service(): return OSSStorageService() storage_mod.OSSStorageService = OSSStorageService storage_mod.get_storage_service = get_storage_service sys.modules["app.core.storage"] = storage_mod for ns in ["app.core"]: if ns not in sys.modules: sys.modules[ns] = types.ModuleType(ns) sys.modules["app.core"].storage = storage_mod # app.schemas.duplication from pydantic import BaseModel, Field class DuplicateSegmentResponse(BaseModel): id: str source_start: float source_end: float matched_video_id: str matched_video_name: str matched_start: float matched_end: float similarity: float class DuplicationRecordResponse(BaseModel): id: str filename: str file_size: int duration_seconds: float = 0.0 status: str = "pending" duplicate_rate: float | None = None duplicate_count: int = 0 created_at: str updated_at: str class DuplicationDetailResponse(DuplicationRecordResponse): segments: list[DuplicateSegmentResponse] = Field(default_factory=list) class DuplicationUploadResponse(BaseModel): id: str status: str message: str dup_schemas_mod = types.ModuleType("app.schemas.duplication") dup_schemas_mod.DuplicateSegmentResponse = DuplicateSegmentResponse dup_schemas_mod.DuplicationRecordResponse = DuplicationRecordResponse dup_schemas_mod.DuplicationDetailResponse = DuplicationDetailResponse dup_schemas_mod.DuplicationUploadResponse = DuplicationUploadResponse sys.modules["app.schemas.duplication"] = dup_schemas_mod sys.modules.setdefault("app.schemas", types.ModuleType("app.schemas")) sys.modules["app.schemas"].duplication = dup_schemas_mod return User, AuthenticatedUser, DuplicationRecord, DuplicateSegment User, AuthenticatedUser, DuplicationRecord, DuplicateSegment = _install_mocks() # ---------- 导入被测路由模块 ---------- for ns in ["app", "app.api", "app.api.routes"]: if ns not in sys.modules: sys.modules[ns] = types.ModuleType(ns) import importlib.util _route_path = os.path.join( os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "apps", "api", "app", "api", "routes", "duplication.py", ) _spec = importlib.util.spec_from_file_location( "app.api.routes.duplication", _route_path, ) duplication = importlib.util.module_from_spec(_spec) sys.modules["app.api.routes.duplication"] = duplication _spec.loader.exec_module(duplication) # 路由模块已导入,立即恢复原始模块,避免污染后续测试文件的 collection for _k, _v in _SAVED_MODULES.items(): sys.modules[_k] = _v # 删除本文件新增的、原始不存在的 mock 模块 for _k in [ "packages.domain.entities", "packages.domain.duplication", "packages.ports.user_repository", "packages.ports.duplication_repository", "packages.adapters.sqlalchemy_impl.user_repository", "packages.adapters.sqlalchemy_impl.duplication_repository", "packages.adapters.sqlalchemy_impl.session", "packages.adapters.redis", "packages.adapters.smtp", "packages.application", "app.config", "app.auth", "app.dependencies", "app.core.storage", "app.schemas.duplication", "app.api.routes.duplication", ]: if _k not in _SAVED_MODULES and _k in sys.modules: del sys.modules[_k] # --------------------------------------------------------------------------- # 2. 内存 Repository + Fixtures # --------------------------------------------------------------------------- class InMemoryDuplicationRepo: """内存中的查重记录 Repository,模拟持久化行为。""" def __init__(self): self.records: dict[str, DuplicationRecord] = {} def create(self, record): self.records[record.id] = record return record def get(self, record_id): return self.records.get(record_id) def list_by_user(self, user_id, *, offset=0, limit=50): all_records = [r for r in self.records.values() if r.user_id == user_id] return all_records[offset : offset + limit] def update(self, record): self.records[record.id] = record return record def delete(self, record_id): if record_id in self.records: del self.records[record_id] return True return False 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_record(user_id="user-test-001", status="pending", filename="test.mp4", **kw): """创建测试用 DuplicationRecord 并设置状态。""" record = DuplicationRecord( id=uuid4().hex, user_id=user_id, filename=filename, file_size=kw.get("file_size", 1024), storage_key=kw.get("storage_key", "oss/key"), duration_seconds=kw.get("duration", 30.0), ) if status == "processing": record.mark_processing() elif status == "completed": record.mark_processing() record.mark_completed(duplicate_rate=15.0, duplicate_count=1, segments=[]) elif status == "failed": record.mark_processing() record.mark_failed("处理失败") return record @pytest.fixture(autouse=True, scope="module") def _restore_modules_after_tests(): """测试结束后恢复被 mock 覆盖的原始模块,避免污染其他测试文件。""" yield # 恢复原始模块 for _k, _v in _SAVED_MODULES.items(): sys.modules[_k] = _v # 删除本文件新增的 mock 模块(不在原始 sys.modules 中的) _mock_keys = [ "packages.domain.entities", "packages.domain.duplication", "packages.ports.user_repository", "packages.ports.duplication_repository", "packages.adapters.sqlalchemy_impl.user_repository", "packages.adapters.sqlalchemy_impl.duplication_repository", "packages.adapters.sqlalchemy_impl.session", "packages.adapters.redis", "packages.adapters.smtp", "packages.application", "app.config", "app.auth", "app.dependencies", "app.core.storage", "app.schemas.duplication", "app.api.routes.duplication", ] for _k in _mock_keys: if _k not in _SAVED_MODULES and _k in sys.modules: del sys.modules[_k] @pytest.fixture def repo(): return InMemoryDuplicationRepo() @pytest.fixture def client(repo): """创建带有依赖覆盖的 TestClient。""" app = FastAPI() app.include_router(duplication.router) def _override_current_user(): return AuthenticatedUser(user=_make_user()) def _override_dup_repo(): return repo def _override_storage(): from app.core.storage import OSSStorageService return OSSStorageService() app.dependency_overrides[duplication.get_current_user] = _override_current_user app.dependency_overrides[duplication.get_duplication_repository] = _override_dup_repo app.dependency_overrides[duplication.get_storage_service] = _override_storage return TestClient(app) # --------------------------------------------------------------------------- # 3. GET /records — 列表查询 # --------------------------------------------------------------------------- class TestListDuplicationRecords: """列表查询端点测试。""" def test_empty_list(self, client): """无记录时返回空列表。""" resp = client.get("/records") assert resp.status_code == 200 assert resp.json() == [] def test_returns_records(self, client, repo): """有记录时返回列表。""" r1 = _make_record(filename="a.mp4") r2 = _make_record(filename="b.mp4") repo.create(r1) repo.create(r2) resp = client.get("/records") assert resp.status_code == 200 data = resp.json() assert len(data) == 2 filenames = {item["filename"] for item in data} assert filenames == {"a.mp4", "b.mp4"} def test_record_response_fields(self, client, repo): """返回的字段应包含所有必需字段。""" record = _make_record() repo.create(record) resp = client.get("/records") assert resp.status_code == 200 data = resp.json() assert len(data) == 1 item = data[0] assert "id" in item assert "filename" in item assert "file_size" in item assert "status" in item assert "created_at" in item assert "updated_at" in item def test_only_returns_current_user_records(self, client, repo): """只返回当前用户的记录。""" # 当前用户 user-test-001 r1 = _make_record(user_id="user-test-001", filename="mine.mp4") # 其他用户 r2 = _make_record(user_id="other-user", filename="other.mp4") repo.create(r1) repo.create(r2) resp = client.get("/records") assert resp.status_code == 200 data = resp.json() assert len(data) == 1 assert data[0]["filename"] == "mine.mp4" # --------------------------------------------------------------------------- # 4. GET /records/{record_id} — 详情查询 # --------------------------------------------------------------------------- class TestGetDuplicationDetail: """详情查询端点测试。""" def test_returns_detail_with_segments(self, client, repo): """返回记录详情含片段列表。""" seg = DuplicateSegment( id=uuid4().hex, source_start=0.0, source_end=5.0, matched_video_id="vid-1", matched_video_name="existing.mp4", matched_start=0.0, matched_end=5.0, similarity=92.5, ) record = _make_record(status="completed") record.segments = [seg] repo.create(record) resp = client.get(f"/records/{record.id}") assert resp.status_code == 200 data = resp.json() assert data["id"] == record.id assert data["status"] == "completed" assert len(data["segments"]) == 1 assert data["segments"][0]["similarity"] == 92.5 def test_returns_404_for_nonexistent(self, client): """不存在的记录返回 404。""" resp = client.get("/records/nonexistent-id") assert resp.status_code == 404 def test_returns_404_for_other_user_record(self, client, repo): """其他用户的记录返回 404(安全隔离)。""" record = _make_record(user_id="other-user") repo.create(record) resp = client.get(f"/records/{record.id}") assert resp.status_code == 404 def test_detail_includes_all_segment_fields(self, client, repo): """片段响应包含所有必需字段。""" seg = DuplicateSegment( id="seg-1", source_start=1.0, source_end=10.0, matched_video_id="vid-1", matched_video_name="ref.mp4", matched_start=2.0, matched_end=11.0, similarity=85.0, ) record = _make_record(status="completed") record.segments = [seg] repo.create(record) resp = client.get(f"/records/{record.id}") assert resp.status_code == 200 seg_data = resp.json()["segments"][0] assert seg_data["id"] == "seg-1" assert seg_data["source_start"] == 1.0 assert seg_data["source_end"] == 10.0 assert seg_data["matched_video_id"] == "vid-1" assert seg_data["matched_video_name"] == "ref.mp4" assert seg_data["matched_start"] == 2.0 assert seg_data["matched_end"] == 11.0 assert seg_data["similarity"] == 85.0 # --------------------------------------------------------------------------- # 5. DELETE /records/{record_id} — 删除记录 # --------------------------------------------------------------------------- class TestDeleteDuplicationRecord: """删除端点测试。""" def test_delete_existing_record(self, client, repo): """删除存在的记录返回 204。""" record = _make_record() repo.create(record) resp = client.delete(f"/records/{record.id}") assert resp.status_code == 204 assert repo.get(record.id) is None def test_delete_nonexistent_returns_404(self, client): """删除不存在的记录返回 404。""" resp = client.delete("/records/nonexistent-id") assert resp.status_code == 404 def test_delete_other_user_record_returns_404(self, client, repo): """删除其他用户的记录返回 404(安全隔离)。""" record = _make_record(user_id="other-user") repo.create(record) resp = client.delete(f"/records/{record.id}") assert resp.status_code == 404 # 记录应仍然存在 assert repo.get(record.id) is not None def test_delete_idempotent(self, client, repo): """删除后再次删除返回 404。""" record = _make_record() repo.create(record) resp1 = client.delete(f"/records/{record.id}") assert resp1.status_code == 204 resp2 = client.delete(f"/records/{record.id}") assert resp2.status_code == 404 # --------------------------------------------------------------------------- # 6. POST /records/{record_id}/retry — 重试查重 # --------------------------------------------------------------------------- class TestRetryDuplication: """重试端点测试。""" def test_retry_failed_record(self, client, repo): """重试失败记录应重置状态为 pending。""" record = _make_record(status="failed") repo.create(record) resp = client.post(f"/records/{record.id}/retry") assert resp.status_code == 200 data = resp.json() assert data["id"] == record.id assert data["status"] == "pending" assert "重新提交" in data["message"] # 验证 repo 中的记录也被更新 updated = repo.get(record.id) assert updated.status == "pending" assert updated.error_message == "" def test_retry_nonexistent_returns_404(self, client): """重试不存在的记录返回 404。""" resp = client.post("/records/nonexistent-id/retry") assert resp.status_code == 404 def test_retry_other_user_record_returns_404(self, client, repo): """重试其他用户的记录返回 404。""" record = _make_record(user_id="other-user", status="failed") repo.create(record) resp = client.post(f"/records/{record.id}/retry") assert resp.status_code == 404 def test_retry_completed_record_still_resets(self, client, repo): """重试已完成记录 — 路由层不校验状态,直接重置。""" record = _make_record(status="completed") repo.create(record) resp = client.post(f"/records/{record.id}/retry") # 路由层允许重试(状态校验在用例层) assert resp.status_code == 200 assert resp.json()["status"] == "pending" def test_retry_pending_record(self, client, repo): """重试 pending 状态的记录。""" record = _make_record(status="pending") repo.create(record) resp = client.post(f"/records/{record.id}/retry") assert resp.status_code == 200 assert resp.json()["status"] == "pending" # --------------------------------------------------------------------------- # 7. 跨端点场景 # --------------------------------------------------------------------------- class TestCrossEndpointScenarios: """跨端点集成场景。""" def test_create_then_list_then_detail(self, client, repo): """创建 → 列表 → 详情 完整流程。""" record = _make_record(filename="flow.mp4") repo.create(record) # 列表 list_resp = client.get("/records") assert list_resp.status_code == 200 assert len(list_resp.json()) == 1 # 详情 detail_resp = client.get(f"/records/{record.id}") assert detail_resp.status_code == 200 assert detail_resp.json()["filename"] == "flow.mp4" def test_create_then_delete_then_404(self, client, repo): """创建 → 删除 → 详情 404 流程。""" record = _make_record() repo.create(record) # 删除 del_resp = client.delete(f"/records/{record.id}") assert del_resp.status_code == 204 # 详情应 404 detail_resp = client.get(f"/records/{record.id}") assert detail_resp.status_code == 404 def test_failed_record_retry_then_detail(self, client, repo): """失败记录 → 重试 → 查看详情状态已重置。""" record = _make_record(status="failed") repo.create(record) # 重试 retry_resp = client.post(f"/records/{record.id}/retry") assert retry_resp.status_code == 200 # 详情确认状态 detail_resp = client.get(f"/records/{record.id}") assert detail_resp.status_code == 200 assert detail_resp.json()["status"] == "pending"