refactor(dedup): 查重模块优化 — XOR 汉明距离 + 简化判定逻辑 #148

Merged
xiaoxia merged 5 commits from feature/duplication-optimization into develop 2026-07-01 16:18:05 +08:00
10 changed files with 1886 additions and 30 deletions
+6
View File
@@ -6,6 +6,7 @@ from app.api.routes.chunked_upload import router as chunked_upload_router
from app.api.routes.classification_jobs import router as classification_jobs_router
from app.api.routes.dashboard import router as dashboard_router
from app.api.routes.duplication import router as duplication_router
from app.api.routes.edit_plans import router as edit_plans_router
from app.api.routes.edit_templates import router as edit_templates_router
from app.api.routes.generated_videos import router as generated_videos_router
from app.api.routes.generation_tasks import router as generation_tasks_router
@@ -122,3 +123,8 @@ api_router.include_router(
prefix="/edit-templates",
tags=["EditTemplate"],
)
api_router.include_router(
edit_plans_router,
prefix="/edit-plans",
tags=["EditPlan"],
)
+25 -6
View File
@@ -15,7 +15,7 @@ from app.schemas.duplication import (
DuplicationRecordResponse,
DuplicationUploadResponse,
)
from fastapi import APIRouter, Depends, File, HTTPException, Response, UploadFile, status
from fastapi import APIRouter, Depends, File, HTTPException, Query, Response, UploadFile, status
from packages.application import (
DeleteDuplicationRecordUseCase,
@@ -29,7 +29,9 @@ from packages.domain.duplication import DuplicationRecord
logger = logging.getLogger(__name__)
router = APIRouter()
router = APIRouter(
tags=["查重"],
)
# 查重功能只接受视频文件
ALLOWED_VIDEO_MIME_TYPES = frozenset(
@@ -199,12 +201,19 @@ async def upload_for_duplication(
@router.get("/records", response_model=list[DuplicationRecordResponse])
def list_duplication_records(
offset: int = Query(0, ge=0, description="分页偏移量"),
limit: int = Query(50, ge=1, le=200, description="每页数量,最大 200"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
duplication_repository: Any = Depends(get_duplication_repository),
) -> list[DuplicationRecordResponse]:
"""获取当前用户的查重记录列表。"""
"""
获取当前用户的查重记录列表。
支持分页:通过 offset 和 limit 参数控制。
返回按创建时间倒序排列的记录。
"""
use_case = ListDuplicationRecordsUseCase(duplication_repository)
records = use_case.execute(authenticated_user.user.id)
records = use_case.execute(user_id=authenticated_user.user.id, offset=offset, limit=limit)
return [_to_record_response(r) for r in records]
@@ -257,7 +266,11 @@ def retry_duplication(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
duplication_repository: Any = Depends(get_duplication_repository),
) -> DuplicationUploadResponse:
"""重新提交查重。"""
"""
重新提交查重。
仅 failed 状态的记录允许重试,其他状态返回 400。
"""
# 检查记录存在且属于当前用户
detail_uc = GetDuplicationDetailUseCase(duplication_repository)
record = detail_uc.execute(record_id)
@@ -268,7 +281,13 @@ def retry_duplication(
)
use_case = RetryDuplicationUseCase(duplication_repository)
updated = use_case.execute(record_id)
try:
updated = use_case.execute(record_id)
except ValueError as e:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=str(e),
)
if updated is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
+104 -17
View File
@@ -23,7 +23,22 @@ logger = logging.getLogger(__name__)
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
"""Compute perceptual hash of an image using DCT."""
"""计算图像的感知哈希(pHash),基于 DCT(离散余弦变换)。
算法步骤:
1. 将图像缩放到 hash_size*4 × hash_size*4(默认 32×32
2. 转为灰度图,应用 2D DCT 提取频率分量
3. 取左上角 hash_size×hash_size 的低频分量(默认 8×8 = 64 bit)
4. 排除 DC 分量([0,0] 位置),计算中位数
5. 每个分量与中位数比较,生成二值 hash
Args:
image: BGR 格式的 numpy 图像数组
hash_size: 哈希边长,默认 8(生成 64-bit hash
Returns:
十六进制字符串表示的感知哈希
"""
# Resize to 32x32 for DCT
resized = cv2.resize(image, (hash_size * 4, hash_size * 4))
gray = cv2.cvtColor(resized, cv2.COLOR_BGR2GRAY).astype(np.float32)
@@ -41,7 +56,18 @@ def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
def hamming_distance(hash1: str, hash2: str) -> int:
"""Calculate Hamming distance between two hex hashes."""
"""计算两个十六进制哈希之间的汉明距离(不同 bit 位数)。
使用 XOR 异或 + bit 计数:bin(h1 ^ h2).count("1")。
例如:hamming_distance("00", "ff") = 88 个 bit 全不同)。
Args:
hash1: 十六进制字符串
hash2: 十六进制字符串
Returns:
不同 bit 的数量
"""
h1, h2 = int(hash1, 16), int(hash2, 16)
return bin(h1 ^ h2).count("1")
@@ -122,7 +148,25 @@ class VideoDeduplicator:
)
def check_duplicate(self, fingerprint: VideoFingerprint, project_id: str, session: Session) -> Optional[dict]:
"""Check if video is duplicate of existing one. Returns duplicate info if found."""
"""检查视频是否与项目中已有视频重复。
判定逻辑(按优先级):
1. MD5 精确匹配:完全一致则 similarity=1.0,立即返回
2. pHash 相似度:计算新视频每帧 phash 与已有视频每帧 phash 的最小汉明距离,
取所有帧的平均值 avg_distance。若 avg_distance < PHASH_THRESHOLD(10)
则判定为重复,similarity = 1.0 - (avg_distance / 64)
注意:返回第一个通过阈值的匹配(非最优匹配)。
Args:
fingerprint: 待检测视频的指纹
project_id: 项目 ID,仅在同一项目内搜索
session: 数据库会话
Returns:
重复信息字典(含 duplicate, duplicate_of, reason, similarity),
或 None 表示未找到重复。
"""
video_repo = SQLAlchemyGeneratedVideoRepository(session)
existing_videos = video_repo.list_by_project(project_id)
@@ -132,28 +176,71 @@ class VideoDeduplicator:
ef = existing.video_fingerprint
# 精确匹配:MD5 完全一致
if fingerprint.md5 == ef.get("md5"):
return {"duplicate": True, "duplicate_of": existing.id, "reason": "exact_md5_match", "similarity": 1.0}
# 感知哈希相似度
existing_phashes = ef.get("keyframe_phashes", [])
if existing_phashes:
total_distance = 0
min_distances = []
for phash in fingerprint.keyframe_phashes:
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
min_distances.append(min(distances))
avg_distance = sum(min_distances) / len(min_distances) if min_distances else 100
if not existing_phashes:
continue
if avg_distance < self.PHASH_THRESHOLD:
return {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "phash_similar",
"similarity": 1.0 - (avg_distance / 64),
}
# 计算每个新关键帧到已有关键帧的最小汉明距离,取平均
min_distances = []
for phash in fingerprint.keyframe_phashes:
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
min_distances.append(min(distances))
avg_distance = sum(min_distances) / len(min_distances) if min_distances else 100
if avg_distance >= self.PHASH_THRESHOLD:
continue
phash_similarity = 1.0 - (avg_distance / 64)
return {"duplicate": True, "duplicate_of": existing.id, "reason": "phash_similar", "similarity": phash_similarity}
return None
@staticmethod
def _average_histogram_similarity(
histograms_a: list[list[float]], histograms_b: list[list[float]]
) -> float:
"""
计算两组颜色直方图之间的平均余弦相似度。
对每组直方图对取最小长度对齐,计算余弦相似度后取平均。
Args:
histograms_a: 第一组直方图(每帧一个 list)
histograms_b: 第二组直方图
Returns:
平均余弦相似度,范围 [0, 1]
"""
if not histograms_a or not histograms_b:
return 0.0
similarities = []
for ha in histograms_a:
best = 0.0
vec_a = np.array(ha, dtype=np.float64)
norm_a = np.linalg.norm(vec_a)
if norm_a == 0:
continue
for hb in histograms_b:
vec_b = np.array(hb, dtype=np.float64)
# 对齐长度
min_len = min(len(vec_a), len(vec_b))
va, vb = vec_a[:min_len], vec_b[:min_len]
norm_b = np.linalg.norm(vb)
if norm_b == 0:
continue
sim = float(np.dot(va, vb) / (norm_a * norm_b))
best = max(best, sim)
similarities.append(best)
return sum(similarities) / len(similarities) if similarities else 0.0
@celery_app.task(bind=True, max_retries=3, name="worker.check_duplicate")
def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
@@ -82,8 +82,14 @@ class SQLAlchemyDuplicationRecordRepository:
return record
def delete(self, record_id: str) -> bool:
count = self.session.query(DuplicationRecordModel).filter(DuplicationRecordModel.id == record_id).delete()
"""
删除查重记录及其关联片段。
注意:必须先删片段再删记录,防止进程崩溃时产生孤儿片段数据。
"""
# 先删关联片段,再删主记录(安全顺序)
self.session.query(DuplicationSegmentModel).filter(DuplicationSegmentModel.record_id == record_id).delete()
count = self.session.query(DuplicationRecordModel).filter(DuplicationRecordModel.id == record_id).delete()
self.session.commit()
return count > 0
+20 -6
View File
@@ -72,20 +72,34 @@ class DeleteDuplicationRecordUseCase:
class RetryDuplicationUseCase:
"""重新提交查重 — 将记录状态重置为 pending。"""
"""
重新提交查重 — 将 failed 状态的记录重置为 pending。
仅 failed 状态允许重试,其他状态抛出 ValueError。
"""
def __init__(self, duplication_repository: DuplicationRecordRepository):
self.duplication_repository = duplication_repository
def execute(self, record_id: str) -> DuplicationRecord | None:
"""
重试查重。
Args:
record_id: 查重记录 ID
Returns:
重置后的记录;不存在返回 None
Raises:
ValueError: 记录状态不允许重试(非 failed)
"""
record = self.duplication_repository.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 = []
if not record.can_retry():
raise ValueError(f"只有 failed 状态的记录可以重试,当前状态: {record.status}")
record.reset_for_retry()
record = self.duplication_repository.update(record)
logger.info("Duplication record %s reset to pending for retry", record_id)
return record
+26
View File
@@ -111,3 +111,29 @@ class DuplicationRecord:
self.status = "failed"
self.error_message = error_message
self.updated_at = datetime.now(timezone.utc)
def can_retry(self) -> bool:
"""
判断是否可重试。
仅 failed 状态的记录允许重试,防止误操作已完成的记录。
Returns:
True 表示可以重试
"""
return self.status == "failed"
def reset_for_retry(self) -> None:
"""
重置记录以重新查重。
清空查重结果和错误信息,状态回到 pending。
调用前应先通过 can_retry() 确认状态合法。
"""
self.status = "pending"
self.duplicate_rate = None
self.duplicate_count = 0
self.error_message = ""
self.segments = []
self.video_fingerprint = None
self.updated_at = datetime.now(timezone.utc)
+751
View File
@@ -0,0 +1,751 @@
"""查重 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 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 的模式)
# ---------------------------------------------------------------------------
def _install_mocks():
"""安装所有必需的 mock 模块,使路由模块可导入。"""
# 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
_spec = importlib.util.spec_from_file_location(
"app.api.routes.duplication",
"/tmp/xiaoxia-saas/apps/api/app/api/routes/duplication.py",
)
duplication = importlib.util.module_from_spec(_spec)
sys.modules["app.api.routes.duplication"] = duplication
_spec.loader.exec_module(duplication)
# ---------------------------------------------------------------------------
# 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
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"
+440
View File
@@ -0,0 +1,440 @@
"""查重引擎单元测试。
覆盖:
- hamming_distance() 汉明距离计算(XOR bit 计数)
- compute_phash() 感知哈希算法(需真实 cv2,无则跳过)
- compute_color_histogram() 颜色直方图(需真实 cv2,无则跳过)
- VideoDeduplicator.check_duplicate() 相似度判定逻辑
"""
from __future__ import annotations
import sys
from unittest.mock import MagicMock
# ---------------------------------------------------------------------------
# 保存 sys.modules 原始状态,测试结束后恢复,避免污染其他测试文件
# ---------------------------------------------------------------------------
_ORIGINAL_MODULES = dict(sys.modules)
_MOCKED_MODULE_NAMES: list[str] = []
def _mock_if_absent(name: str, mock_obj=None):
"""仅在模块不在 sys.modules 中时注入 mock,并记录以便清理。"""
if name not in sys.modules:
sys.modules[name] = mock_obj if mock_obj is not None else MagicMock()
_MOCKED_MODULE_NAMES.append(name)
# Mock heavy deps before importing dedup module
_mock_if_absent("ffmpeg")
# Mock worker_app (celery) and its submodules
for mod_name in ["worker_app", "worker_app.celery_app", "worker_app.db"]:
_mock_if_absent(mod_name)
if "worker_app.celery_app" in sys.modules and isinstance(sys.modules["worker_app.celery_app"], MagicMock):
sys.modules["worker_app.celery_app"].celery_app = MagicMock()
if "worker_app.db" in sys.modules and isinstance(sys.modules["worker_app.db"], MagicMock):
sys.modules["worker_app.db"].SessionLocal = MagicMock()
# Mock celery.Task base class
_mock_if_absent("celery", MagicMock())
if "celery" in sys.modules and isinstance(sys.modules["celery"], MagicMock):
sys.modules["celery"].Task = object
# Mock packages.shared.storage
_mock_if_absent("packages.shared")
_mock_if_absent("packages.shared.storage")
# Mock packages.adapters.sqlalchemy_impl.generated_video_repository
_mock_if_absent("packages.adapters.sqlalchemy_impl.generated_video_repository")
# Check if cv2 is available as a real module (not mocked)
_HAS_CV2 = False
try:
import cv2 as _cv2
if not isinstance(_cv2, MagicMock):
_HAS_CV2 = True
except (ImportError, ModuleNotFoundError):
pass
import numpy as np # noqa: E402
import pytest # noqa: E402
# Mock cv2 if not available (so dedup module can import)
if not _HAS_CV2:
_mock_if_absent("cv2")
from apps.worker.video_processing.dedup import ( # noqa: E402
VideoDeduplicator,
VideoFingerprint,
compute_color_histogram,
compute_phash,
hamming_distance,
)
@pytest.fixture(autouse=True, scope="session")
def _cleanup_mocks():
"""测试结束后恢复 sys.modules,防止 mock 污染其他测试文件。"""
yield
# 移除本次新增的 mock 模块
for name in _MOCKED_MODULE_NAMES:
sys.modules.pop(name, None)
# 恢复被覆盖的模块
for name, mod in _ORIGINAL_MODULES.items():
if sys.modules.get(name) is not mod:
sys.modules[name] = mod
class TestHammingDistance:
"""hamming_distance() 测试。
实现使用 XOR + bit 计数:bin(h1 ^ h2).count("1")。
空字符串会触发 ValueErrorint("", 16) 失败),属于边界行为。
"""
def test_identical_hashes_zero_distance(self):
assert hamming_distance("abcdef01", "abcdef01") == 0
def test_completely_different_bytes(self):
# 0x00 XOR 0xFF = 0xFF → 8 bits
assert hamming_distance("00", "ff") == 8
def test_single_bit_difference(self):
# 0x0 XOR 0x1 = 0x1 → 1 bit
assert hamming_distance("0", "1") == 1
def test_unequal_length_leading_zeros(self):
# int("abc", 16) == int("0abc", 16) → XOR = 0 → 0 bits
dist = hamming_distance("abc", "0abc")
assert dist == 0
def test_unequal_length_with_leading_zeros_ff(self):
# int("ff", 16) == int("00ff", 16) → XOR = 0 → 0 bits
dist = hamming_distance("ff", "00ff")
assert dist == 0
def test_all_bits_different_64bit(self):
# 16 hex chars = 64 bits, all different → 64
dist = hamming_distance("0000000000000000", "ffffffffffffffff")
assert dist == 64
def test_partial_difference(self):
# 0x0F = 00001111, 0xF0 = 11110000 → XOR = 0xFF → 8 bits
assert hamming_distance("0f", "f0") == 8
def test_one_bit_in_second_byte(self):
# 0x0000 XOR 0x0001 = 0x0001 → 1 bit
assert hamming_distance("0000", "0001") == 1
@pytest.mark.skipif(not _HAS_CV2, reason="需要真实 cv2 模块")
class TestComputePhash:
"""compute_phash() 测试(需真实 cv2)。"""
def test_returns_hex_string(self):
image = np.random.randint(0, 256, (64, 64, 3), dtype=np.uint8)
result = compute_phash(image)
assert isinstance(result, str)
int(result, 16) # 不应抛出异常
def test_identical_images_same_hash(self):
image = np.full((64, 64, 3), 128, dtype=np.uint8)
hash1 = compute_phash(image)
hash2 = compute_phash(image)
assert hash1 == hash2
def test_different_images_different_hash(self):
img1 = np.zeros((64, 64, 3), dtype=np.uint8)
img2 = np.full((64, 64, 3), 255, dtype=np.uint8)
hash1 = compute_phash(img1)
hash2 = compute_phash(img2)
assert hash1 != hash2
def test_custom_hash_size(self):
image = np.random.randint(0, 256, (64, 64, 3), dtype=np.uint8)
result = compute_phash(image, hash_size=16)
assert isinstance(result, str)
int(result, 16)
@pytest.mark.skipif(not _HAS_CV2, reason="需要真实 cv2 模块")
class TestComputeColorHistogram:
"""compute_color_histogram() 测试(需真实 cv2)。"""
def test_returns_correct_length(self):
image = np.random.randint(0, 256, (64, 64, 3), dtype=np.uint8)
hist = compute_color_histogram(image, bins=32)
assert len(hist) == 96 # 3 channels × 32 bins
def test_custom_bins(self):
image = np.random.randint(0, 256, (64, 64, 3), dtype=np.uint8)
hist = compute_color_histogram(image, bins=16)
assert len(hist) == 48 # 3 channels × 16 bins
def test_normalized_values(self):
image = np.random.randint(0, 256, (64, 64, 3), dtype=np.uint8)
hist = compute_color_histogram(image)
for v in hist:
assert 0.0 <= v <= 1.0 + 1e-6
def test_identical_images_same_histogram(self):
image = np.full((64, 64, 3), 100, dtype=np.uint8)
hist1 = compute_color_histogram(image)
hist2 = compute_color_histogram(image)
assert hist1 == hist2
class TestVideoDeduplicatorCheckDuplicate:
"""VideoDeduplicator.check_duplicate() 测试。
当前实现仅使用 MD5 精确匹配和 pHash 距离判定,
不包含颜色直方图相似度计算。
"""
@pytest.fixture
def deduplicator(self):
return VideoDeduplicator()
@pytest.fixture
def mock_session(self):
return MagicMock()
def _make_existing_video(self, video_id, md5, phashes=None):
"""创建模拟已有视频的 mock 对象。"""
video = MagicMock()
video.id = video_id
video.video_fingerprint = {
"md5": md5,
"keyframe_phashes": phashes or [],
"color_histograms": [],
}
return video
def _patch_repo(self, mock_repo):
"""Patch SQLAlchemyGeneratedVideoRepository。"""
import apps.worker.video_processing.dedup as dedup_module
original = dedup_module.SQLAlchemyGeneratedVideoRepository
dedup_module.SQLAlchemyGeneratedVideoRepository = MagicMock(return_value=mock_repo)
return original, dedup_module
def _restore_repo(self, dedup_module, original):
dedup_module.SQLAlchemyGeneratedVideoRepository = original
def test_exact_md5_match(self, deduplicator, mock_session):
"""MD5 完全匹配应返回 similarity=1.0。"""
existing = self._make_existing_video("vid-1", "abc123")
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = [existing]
fingerprint = VideoFingerprint(
md5="abc123",
keyframe_phashes=["ff"],
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
assert result["similarity"] == 1.0
assert result["reason"] == "exact_md5_match"
finally:
self._restore_repo(mod, orig)
def test_phash_similar_match(self, deduplicator, mock_session):
"""pHash 距离 < 阈值时应判定为重复。"""
existing = self._make_existing_video("vid-1", "different_md5", phashes=["abcdef01"])
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = [existing]
fingerprint = VideoFingerprint(
md5="different_md5_new",
keyframe_phashes=["abcdef01"], # 完全相同,距离=0
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
assert result["similarity"] == 1.0 # distance=0 → 1.0
assert result["reason"] == "phash_similar"
finally:
self._restore_repo(mod, orig)
def test_no_match_returns_none(self, deduplicator, mock_session):
"""pHash 平均距离 >= PHASH_THRESHOLD(10) 时应返回 None。"""
# 使用 16 字符 phash64 bit),全部不同 → 距离=64 >= 10
existing = self._make_existing_video("vid-1", "md5_a", phashes=["0000000000000000"])
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = [existing]
fingerprint = VideoFingerprint(
md5="md5_b",
keyframe_phashes=["ffffffffffffffff"], # 64 bits 全不同 → 距离=64
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is None
finally:
self._restore_repo(mod, orig)
def test_empty_project_returns_none(self, deduplicator, mock_session):
"""项目中没有视频时应返回 None。"""
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = []
fingerprint = VideoFingerprint(
md5="abc",
keyframe_phashes=["ff"],
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is None
finally:
self._restore_repo(mod, orig)
def test_skip_videos_without_fingerprint(self, deduplicator, mock_session):
"""没有指纹的视频应被跳过。"""
existing = MagicMock()
existing.id = "vid-1"
existing.video_fingerprint = None
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = [existing]
fingerprint = VideoFingerprint(
md5="abc",
keyframe_phashes=["ff"],
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is None
finally:
self._restore_repo(mod, orig)
def test_first_match_returned(self, deduplicator, mock_session):
"""返回第一个通过阈值的匹配(非最优匹配)。"""
# vid-1: 距离=2 bits0x03 XOR 0x01 = 0x02 → 1 bit),通过阈值
vid1 = self._make_existing_video("vid-1", "md5_1", phashes=["0000000000000003"])
# vid-2: 距离=0 bits(完全匹配)
vid2 = self._make_existing_video("vid-2", "md5_2", phashes=["0000000000000001"])
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = [vid1, vid2]
fingerprint = VideoFingerprint(
md5="md5_new",
keyframe_phashes=["0000000000000001"],
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
# 返回第一个通过阈值的匹配(vid-1 距离=1 < 10
assert result["duplicate_of"] == "vid-1"
finally:
self._restore_repo(mod, orig)
def test_no_phashes_skips_video(self, deduplicator, mock_session):
"""已有视频无 phashes 时应被跳过。"""
existing = self._make_existing_video("vid-1", "md5_a", phashes=[])
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = [existing]
fingerprint = VideoFingerprint(
md5="md5_b",
keyframe_phashes=["abcdef01"],
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is None
finally:
self._restore_repo(mod, orig)
def test_phash_similarity_formula(self, deduplicator, mock_session):
"""验证相似度公式:similarity = 1.0 - (avg_distance / 64)。"""
# 使用已知距离的 phash 对
# "0000000000000000" vs "0000000000000001" → XOR = 1 → 1 bit → distance = 1
existing = self._make_existing_video("vid-1", "md5_a", phashes=["0000000000000000"])
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = [existing]
fingerprint = VideoFingerprint(
md5="md5_b",
keyframe_phashes=["0000000000000001"], # 1 bit different
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
# similarity = 1.0 - (1 / 64) = 0.984375
assert abs(result["similarity"] - (1.0 - 1.0 / 64)) < 1e-6
finally:
self._restore_repo(mod, orig)
def test_multiple_phashes_avg_distance(self, deduplicator, mock_session):
"""多帧 phash 使用平均最小距离。"""
# 已有视频有 2 帧 phash
existing = self._make_existing_video(
"vid-1", "md5_a",
phashes=["0000000000000000", "ffffffffffffffff"],
)
mock_repo = MagicMock()
mock_repo.list_by_project.return_value = [existing]
# 新视频有 1 帧 phash,与第一帧距离=0,与第二帧距离=64
# min_distance = 0, avg = 0 → 匹配
fingerprint = VideoFingerprint(
md5="md5_b",
keyframe_phashes=["0000000000000000"],
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
)
orig, mod = self._patch_repo(mock_repo)
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
assert result["similarity"] == 1.0 # avg_distance = 0
finally:
self._restore_repo(mod, orig)
+275
View File
@@ -0,0 +1,275 @@
"""查重域模型单元测试。
覆盖:
- DuplicationRecord.create() 工厂方法及验证
- DuplicationRecord 状态转换(mark_processing / mark_completed / mark_failed
- DuplicationRecord.can_retry() / reset_for_retry()
- DuplicateSegment.create() 工厂方法及验证
"""
from __future__ import annotations
import pytest
from packages.domain.duplication import DuplicateSegment, DuplicationRecord
class TestDuplicationRecordCreate:
"""DuplicationRecord.create() 工厂方法测试。"""
def test_create_success(self):
record = DuplicationRecord.create(
user_id="user-1",
filename="test.mp4",
file_size=1024,
storage_key="oss/key/test.mp4",
duration_seconds=30.0,
)
assert record.user_id == "user-1"
assert record.filename == "test.mp4"
assert record.file_size == 1024
assert record.storage_key == "oss/key/test.mp4"
assert record.duration_seconds == 30.0
assert record.status == "pending"
assert record.duplicate_rate is None
assert record.duplicate_count == 0
assert record.error_message == ""
assert record.segments == []
assert record.video_fingerprint is None
assert record.id # 自动生成 ID
def test_create_with_default_duration(self):
record = DuplicationRecord.create(
user_id="user-1",
filename="test.mp4",
file_size=1024,
storage_key="oss/key",
)
assert record.duration_seconds == 0.0
def test_create_empty_user_id_raises(self):
with pytest.raises(ValueError, match="user_id"):
DuplicationRecord.create(
user_id="",
filename="test.mp4",
file_size=1024,
storage_key="oss/key",
)
def test_create_whitespace_user_id_raises(self):
with pytest.raises(ValueError, match="user_id"):
DuplicationRecord.create(
user_id=" ",
filename="test.mp4",
file_size=1024,
storage_key="oss/key",
)
def test_create_empty_filename_raises(self):
with pytest.raises(ValueError, match="filename"):
DuplicationRecord.create(
user_id="user-1",
filename="",
file_size=1024,
storage_key="oss/key",
)
def test_create_zero_file_size_raises(self):
with pytest.raises(ValueError, match="file_size"):
DuplicationRecord.create(
user_id="user-1",
filename="test.mp4",
file_size=0,
storage_key="oss/key",
)
def test_create_negative_file_size_raises(self):
with pytest.raises(ValueError, match="file_size"):
DuplicationRecord.create(
user_id="user-1",
filename="test.mp4",
file_size=-100,
storage_key="oss/key",
)
class TestDuplicationRecordStateTransitions:
"""状态转换测试。"""
@pytest.fixture
def record(self):
return DuplicationRecord.create(
user_id="user-1",
filename="test.mp4",
file_size=1024,
storage_key="oss/key",
)
def test_mark_processing(self, record):
record.mark_processing()
assert record.status == "processing"
def test_mark_completed_success(self, record):
record.mark_processing()
segments = [
DuplicateSegment.create(
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.mark_completed(duplicate_rate=15.0, duplicate_count=1, segments=segments)
assert record.status == "completed"
assert record.duplicate_rate == 15.0
assert record.duplicate_count == 1
assert len(record.segments) == 1
def test_mark_completed_invalid_rate_raises(self, record):
record.mark_processing()
with pytest.raises(ValueError, match="duplicate_rate"):
record.mark_completed(duplicate_rate=101.0, duplicate_count=0, segments=[])
def test_mark_completed_negative_rate_raises(self, record):
record.mark_processing()
with pytest.raises(ValueError, match="duplicate_rate"):
record.mark_completed(duplicate_rate=-1.0, duplicate_count=0, segments=[])
def test_mark_failed(self, record):
record.mark_processing()
record.mark_failed("处理超时")
assert record.status == "failed"
assert record.error_message == "处理超时"
class TestDuplicationRecordRetry:
"""can_retry() 和 reset_for_retry() 测试。"""
@pytest.fixture
def record(self):
return DuplicationRecord.create(
user_id="user-1",
filename="test.mp4",
file_size=1024,
storage_key="oss/key",
)
def test_mark_failed_sets_status_and_error(self, record):
record.mark_processing()
record.mark_failed("处理失败")
assert record.status == "failed"
assert record.error_message == "处理失败"
def test_mark_failed_updates_timestamp(self, record):
old_updated = record.updated_at
record.mark_processing()
record.mark_failed("错误")
assert record.updated_at >= old_updated
def test_failed_record_preserves_result_fields(self, record):
"""mark_failed 不改变 duplicate_rate 等结果字段(由 use case 层重置)。"""
record.mark_processing()
record.mark_completed(duplicate_rate=10.0, duplicate_count=1, segments=[])
record.mark_failed("重试失败")
assert record.status == "failed"
assert record.error_message == "重试失败"
assert record.duplicate_rate == 10.0
class TestDuplicateSegmentCreate:
"""DuplicateSegment.create() 工厂方法测试。"""
def test_create_success(self):
seg = DuplicateSegment.create(
source_start=1.0,
source_end=5.0,
matched_video_id="vid-1",
matched_video_name="existing.mp4",
matched_start=2.0,
matched_end=6.0,
similarity=85.5,
)
assert seg.source_start == 1.0
assert seg.source_end == 5.0
assert seg.matched_video_id == "vid-1"
assert seg.matched_video_name == "existing.mp4"
assert seg.matched_start == 2.0
assert seg.matched_end == 6.0
assert seg.similarity == 85.5
assert seg.id # 自动生成 ID
def test_create_negative_source_start_raises(self):
with pytest.raises(ValueError, match="invalid source segment range"):
DuplicateSegment.create(
source_start=-1.0,
source_end=5.0,
matched_video_id="vid-1",
matched_video_name="v.mp4",
matched_start=0.0,
matched_end=5.0,
similarity=80.0,
)
def test_create_source_end_le_start_raises(self):
with pytest.raises(ValueError, match="invalid source segment range"):
DuplicateSegment.create(
source_start=5.0,
source_end=5.0,
matched_video_id="vid-1",
matched_video_name="v.mp4",
matched_start=0.0,
matched_end=5.0,
similarity=80.0,
)
def test_create_negative_matched_start_raises(self):
with pytest.raises(ValueError, match="invalid matched segment range"):
DuplicateSegment.create(
source_start=0.0,
source_end=5.0,
matched_video_id="vid-1",
matched_video_name="v.mp4",
matched_start=-1.0,
matched_end=5.0,
similarity=80.0,
)
def test_create_matched_end_le_start_raises(self):
with pytest.raises(ValueError, match="invalid matched segment range"):
DuplicateSegment.create(
source_start=0.0,
source_end=5.0,
matched_video_id="vid-1",
matched_video_name="v.mp4",
matched_start=2.0,
matched_end=1.0,
similarity=80.0,
)
def test_create_similarity_out_of_range_raises(self):
with pytest.raises(ValueError, match="similarity"):
DuplicateSegment.create(
source_start=0.0,
source_end=5.0,
matched_video_id="vid-1",
matched_video_name="v.mp4",
matched_start=0.0,
matched_end=5.0,
similarity=101.0,
)
def test_create_negative_similarity_raises(self):
with pytest.raises(ValueError, match="similarity"):
DuplicateSegment.create(
source_start=0.0,
source_end=5.0,
matched_video_id="vid-1",
matched_video_name="v.mp4",
matched_start=0.0,
matched_end=5.0,
similarity=-1.0,
)
+232
View File
@@ -0,0 +1,232 @@
"""查重应用层用例单元测试。
覆盖:
- UploadForDuplicationUseCase — 创建查重记录
- ListDuplicationRecordsUseCase — 列表查询(含分页)
- GetDuplicationDetailUseCase — 详情查询
- DeleteDuplicationRecordUseCase — 删除记录
- RetryDuplicationUseCase — 重试查重(含状态校验)
"""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from packages.application.duplication import (
ListDuplicationRecordsUseCase,
GetDuplicationDetailUseCase,
DeleteDuplicationRecordUseCase,
RetryDuplicationUseCase,
UploadForDuplicationCommand,
UploadForDuplicationUseCase,
)
from packages.domain.duplication import DuplicateSegment, DuplicationRecord
def _make_record(status="pending", **kwargs):
"""创建测试用 DuplicationRecord。"""
record = DuplicationRecord.create(
user_id=kwargs.get("user_id", "user-1"),
filename=kwargs.get("filename", "test.mp4"),
file_size=kwargs.get("file_size", 1024),
storage_key=kwargs.get("storage_key", "oss/key"),
duration_seconds=kwargs.get("duration", 30.0),
)
if status != "pending":
record.mark_processing()
if status == "completed":
record.mark_completed(duplicate_rate=15.0, duplicate_count=1, segments=[])
elif status == "failed":
record.mark_failed("处理失败")
return record
class TestUploadForDuplicationUseCase:
"""上传查重用例测试。"""
def test_execute_creates_and_persists_record(self):
mock_repo = MagicMock()
mock_repo.create.side_effect = lambda r: r
use_case = UploadForDuplicationUseCase(mock_repo)
command = UploadForDuplicationCommand(
user_id="user-1",
filename="video.mp4",
file_size=2048,
storage_key="oss/video.mp4",
duration_seconds=60.0,
)
result = use_case.execute(command)
assert result.user_id == "user-1"
assert result.filename == "video.mp4"
assert result.file_size == 2048
assert result.status == "pending"
mock_repo.create.assert_called_once()
def test_execute_with_default_duration(self):
mock_repo = MagicMock()
mock_repo.create.side_effect = lambda r: r
use_case = UploadForDuplicationUseCase(mock_repo)
command = UploadForDuplicationCommand(
user_id="user-1",
filename="video.mp4",
file_size=1024,
storage_key="oss/key",
)
result = use_case.execute(command)
assert result.duration_seconds == 0.0
def test_execute_invalid_user_id_raises(self):
mock_repo = MagicMock()
use_case = UploadForDuplicationUseCase(mock_repo)
command = UploadForDuplicationCommand(
user_id="",
filename="video.mp4",
file_size=1024,
storage_key="oss/key",
)
with pytest.raises(ValueError, match="user_id"):
use_case.execute(command)
mock_repo.create.assert_not_called()
class TestListDuplicationRecordsUseCase:
"""列表查询用例测试。"""
def test_execute_returns_records(self):
records = [_make_record(), _make_record(filename="b.mp4")]
mock_repo = MagicMock()
mock_repo.list_by_user.return_value = records
use_case = ListDuplicationRecordsUseCase(mock_repo)
result = use_case.execute("user-1")
assert len(result) == 2
mock_repo.list_by_user.assert_called_once_with("user-1", offset=0, limit=50)
def test_execute_with_pagination(self):
mock_repo = MagicMock()
mock_repo.list_by_user.return_value = []
use_case = ListDuplicationRecordsUseCase(mock_repo)
use_case.execute("user-1", offset=10, limit=20)
mock_repo.list_by_user.assert_called_once_with("user-1", offset=10, limit=20)
def test_execute_empty_user_id_raises(self):
mock_repo = MagicMock()
use_case = ListDuplicationRecordsUseCase(mock_repo)
with pytest.raises(ValueError):
use_case.execute("")
mock_repo.list_by_user.assert_not_called()
def test_execute_whitespace_user_id_raises(self):
mock_repo = MagicMock()
use_case = ListDuplicationRecordsUseCase(mock_repo)
with pytest.raises(ValueError):
use_case.execute(" ")
def test_execute_strips_user_id(self):
mock_repo = MagicMock()
mock_repo.list_by_user.return_value = []
use_case = ListDuplicationRecordsUseCase(mock_repo)
use_case.execute(" user-1 ")
mock_repo.list_by_user.assert_called_once_with("user-1", offset=0, limit=50)
class TestGetDuplicationDetailUseCase:
"""详情查询用例测试。"""
def test_execute_returns_record(self):
record = _make_record()
mock_repo = MagicMock()
mock_repo.get.return_value = record
use_case = GetDuplicationDetailUseCase(mock_repo)
result = use_case.execute(record.id)
assert result is record
mock_repo.get.assert_called_once_with(record.id)
def test_execute_returns_none_for_missing(self):
mock_repo = MagicMock()
mock_repo.get.return_value = None
use_case = GetDuplicationDetailUseCase(mock_repo)
result = use_case.execute("nonexistent")
assert result is None
class TestDeleteDuplicationRecordUseCase:
"""删除用例测试。"""
def test_execute_deletes_record(self):
mock_repo = MagicMock()
mock_repo.delete.return_value = True
use_case = DeleteDuplicationRecordUseCase(mock_repo)
result = use_case.execute("record-1")
assert result is True
mock_repo.delete.assert_called_once_with("record-1")
def test_execute_returns_false_for_missing(self):
mock_repo = MagicMock()
mock_repo.delete.return_value = False
use_case = DeleteDuplicationRecordUseCase(mock_repo)
result = use_case.execute("nonexistent")
assert result is False
class TestRetryDuplicationUseCase:
"""重试用例测试。"""
def test_execute_resets_failed_record(self):
record = _make_record(status="failed")
mock_repo = MagicMock()
mock_repo.get.return_value = record
mock_repo.update.side_effect = lambda r: r
use_case = RetryDuplicationUseCase(mock_repo)
result = use_case.execute(record.id)
assert result is not None
assert result.status == "pending"
assert result.error_message == ""
assert result.duplicate_rate is None
assert result.duplicate_count == 0
assert result.segments == []
mock_repo.update.assert_called_once()
def test_execute_returns_none_for_missing(self):
mock_repo = MagicMock()
mock_repo.get.return_value = None
use_case = RetryDuplicationUseCase(mock_repo)
result = use_case.execute("nonexistent")
assert result is None
mock_repo.update.assert_not_called()
def test_execute_calls_repo_get_and_update(self):
record = _make_record(status="failed")
mock_repo = MagicMock()
mock_repo.get.return_value = record
mock_repo.update.side_effect = lambda r: r
use_case = RetryDuplicationUseCase(mock_repo)
use_case.execute(record.id)
mock_repo.get.assert_called_once_with(record.id)
mock_repo.update.assert_called_once()