Compare commits
13 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 1ff81dcd0a | |||
| db9ee89ffa | |||
| a0d4f6e111 | |||
| fac80b1f77 | |||
| 159a62f9a5 | |||
| 109d7afbc7 | |||
| 9d31818222 | |||
| af4dd31dd1 | |||
| 8ecf381a9d | |||
| 244691d335 | |||
| ee4fff42f0 | |||
| b0018e747b | |||
| cbca0c3584 |
@@ -0,0 +1,46 @@
|
||||
"""add video_fingerprint_chunks table for per-chunk fingerprint storage
|
||||
|
||||
Revision ID: 063_fingerprint_chunks
|
||||
Revises: 062_edit_plan_id
|
||||
Create Date: 2026-09-03
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "063_fingerprint_chunks"
|
||||
down_revision = "062_edit_plan_id"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"video_fingerprint_chunks",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("video_id", sa.String(36), nullable=False),
|
||||
sa.Column("project_id", sa.String(36), nullable=False),
|
||||
sa.Column("user_id", sa.String(36), nullable=False, server_default=""),
|
||||
sa.Column("start_time_ms", sa.Integer, nullable=False),
|
||||
sa.Column("end_time_ms", sa.Integer, nullable=False),
|
||||
sa.Column("phash_binary", sa.String(16), nullable=False),
|
||||
sa.Column("color_histogram", sa.JSON, nullable=False),
|
||||
sa.Column("frame_count", sa.Integer, nullable=False, server_default="1"),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime,
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
)
|
||||
op.create_index("ix_vfc_video_id", "video_fingerprint_chunks", ["video_id"])
|
||||
op.create_index("ix_vfc_project_id", "video_fingerprint_chunks", ["project_id"])
|
||||
op.create_index("ix_vfc_user_id", "video_fingerprint_chunks", ["user_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_vfc_user_id", table_name="video_fingerprint_chunks")
|
||||
op.drop_index("ix_vfc_project_id", table_name="video_fingerprint_chunks")
|
||||
op.drop_index("ix_vfc_video_id", table_name="video_fingerprint_chunks")
|
||||
op.drop_table("video_fingerprint_chunks")
|
||||
@@ -131,6 +131,7 @@ class PlanGeneratorService:
|
||||
editing_mode,
|
||||
random_selection=random_preview,
|
||||
asset_durations=asset_durations,
|
||||
user_id=created_by_user_id,
|
||||
)
|
||||
|
||||
# 5. 持久化所有 clips 并计算总时长
|
||||
@@ -218,6 +219,7 @@ class PlanGeneratorService:
|
||||
*,
|
||||
random_selection: bool = False,
|
||||
asset_durations: dict[str, float] | None = None,
|
||||
user_id: str = "",
|
||||
) -> None:
|
||||
"""按 editing_mode 将素材分配到 clips(就地修改,未持久化).
|
||||
|
||||
@@ -239,6 +241,14 @@ class PlanGeneratorService:
|
||||
asset_ids = list(asset_ids) # 复制避免修改调用方原列表
|
||||
random.shuffle(asset_ids)
|
||||
|
||||
# 查询已有视频的已用区间(跨视频避让)
|
||||
external_used_segments = None
|
||||
if user_id and self._clip_repo:
|
||||
try:
|
||||
external_used_segments = self._clip_repo.list_used_segments_by_user(user_id, limit_recent=50)
|
||||
except Exception:
|
||||
logger.warning("跨视频避让查询失败,回退到纯随机", exc_info=True)
|
||||
|
||||
distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
@@ -246,6 +256,7 @@ class PlanGeneratorService:
|
||||
random_selection=random_selection,
|
||||
asset_durations=asset_durations,
|
||||
asset_scene_points=asset_scene_points,
|
||||
external_used_segments=external_used_segments,
|
||||
)
|
||||
|
||||
def _fetch_asset_scene_points(self, asset_ids: List[str]) -> dict[str, list[float]]:
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
#!/usr/bin/env python3
|
||||
"""存量指纹重建脚本 — 为已有视频生成 video_fingerprint_chunks 分片数据。
|
||||
|
||||
功能:
|
||||
- 查询 generated_videos 中 video_fingerprint IS NOT NULL 但尚无分片数据的视频
|
||||
- 从 OSS 下载视频 → 用新的分片算法重新计算指纹 → 写入分片表
|
||||
- 支持 --dry-run(只打印不写入)和 --batch-size(默认 50)
|
||||
- 幂等:已存在分片数据的视频跳过
|
||||
|
||||
用法:
|
||||
# 预览(不写入)
|
||||
python rebuild_fingerprint_chunks.py --dry-run
|
||||
|
||||
# 执行重建
|
||||
python rebuild_fingerprint_chunks.py --batch-size 50
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
|
||||
# 确保可以 import worker_app 和 packages
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "..", "worker"))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", ".."))
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
)
|
||||
logger = logging.getLogger("rebuild_fingerprint_chunks")
|
||||
|
||||
|
||||
def find_videos_needing_rebuild(session, batch_size: int) -> list[dict]:
|
||||
"""查询需要重建分片指纹的视频。"""
|
||||
from sqlalchemy import and_
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel, VideoFingerprintChunkModel
|
||||
|
||||
# 有 video_fingerprint 的视频
|
||||
has_fingerprint = GeneratedVideoModel.video_fingerprint.isnot(None)
|
||||
has_fingerprint = and_(has_fingerprint, GeneratedVideoModel.video_fingerprint != "")
|
||||
|
||||
# 排除已有分片数据的视频
|
||||
subq = session.query(VideoFingerprintChunkModel.video_id).distinct().subquery()
|
||||
no_chunks = ~GeneratedVideoModel.id.in_(subq)
|
||||
|
||||
videos = (
|
||||
session.query(GeneratedVideoModel)
|
||||
.filter(and_(has_fingerprint, no_chunks))
|
||||
.order_by(GeneratedVideoModel.generated_at.desc())
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
return [
|
||||
{
|
||||
"id": v.id,
|
||||
"project_id": v.project_id,
|
||||
"user_id": v.user_id or "",
|
||||
"duration": v.duration,
|
||||
}
|
||||
for v in videos
|
||||
]
|
||||
|
||||
|
||||
def rebuild_one(video_info: dict, dry_run: bool = False) -> int:
|
||||
"""重建单个视频的分片数据。返回写入的 chunk 数量。"""
|
||||
from video_processing.dedup import VideoDeduplicator, _save_fingerprint_chunks
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import VideoFingerprintChunkModel
|
||||
from packages.shared.storage import get_storage_service
|
||||
|
||||
video_id = video_info["id"]
|
||||
project_id = video_info["project_id"]
|
||||
user_id = video_info["user_id"]
|
||||
|
||||
if dry_run:
|
||||
logger.info("[DRY-RUN] Would rebuild video %s (project=%s)", video_id, project_id)
|
||||
return 0
|
||||
|
||||
session = SessionLocal()
|
||||
temp_dir = tempfile.mkdtemp()
|
||||
|
||||
try:
|
||||
# 再次检查幂等性
|
||||
existing_count = (
|
||||
session.query(VideoFingerprintChunkModel).filter(VideoFingerprintChunkModel.video_id == video_id).count()
|
||||
)
|
||||
if existing_count > 0:
|
||||
logger.info("Video %s already has %d chunks, skipping", video_id, existing_count)
|
||||
return 0
|
||||
|
||||
# 下载视频
|
||||
storage_service = get_storage_service()
|
||||
local_path = os.path.join(temp_dir, f"{video_id}.mp4")
|
||||
storage_key = f"projects/{project_id}/generated/{video_id}/{video_id}.mp4"
|
||||
storage_service.download_file(storage_key, local_path)
|
||||
|
||||
# 重新计算指纹
|
||||
deduplicator = VideoDeduplicator()
|
||||
fingerprint = deduplicator.compute_fingerprint(local_path)
|
||||
|
||||
# 写入分片表
|
||||
_save_fingerprint_chunks(fingerprint, video_id, project_id, user_id, session)
|
||||
session.commit()
|
||||
|
||||
chunk_count = len(fingerprint.chunks)
|
||||
logger.info("Rebuilt %d chunks for video %s", chunk_count, video_id)
|
||||
return chunk_count
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to rebuild video %s: %s", video_id, e)
|
||||
session.rollback()
|
||||
return -1
|
||||
finally:
|
||||
session.close()
|
||||
import shutil
|
||||
|
||||
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="存量指纹重建脚本")
|
||||
parser.add_argument("--dry-run", action="store_true", help="只打印不写入")
|
||||
parser.add_argument("--batch-size", type=int, default=50, help="每批处理数量(默认 50)")
|
||||
parser.add_argument("--total-limit", type=int, default=0, help="总处理数量限制(0=不限制)")
|
||||
args = parser.parse_args()
|
||||
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
session = SessionLocal()
|
||||
|
||||
try:
|
||||
videos = find_videos_needing_rebuild(session, args.batch_size)
|
||||
logger.info("Found %d videos needing rebuild", len(videos))
|
||||
|
||||
if args.dry_run:
|
||||
for v in videos:
|
||||
logger.info("[DRY-RUN] Video %s | project=%s | duration=%.1fs", v["id"], v["project_id"], v["duration"])
|
||||
return
|
||||
|
||||
total_chunks = 0
|
||||
processed = 0
|
||||
failed = 0
|
||||
|
||||
for v in videos:
|
||||
if args.total_limit > 0 and processed >= args.total_limit:
|
||||
break
|
||||
|
||||
result = rebuild_one(v, dry_run=False)
|
||||
if result < 0:
|
||||
failed += 1
|
||||
else:
|
||||
total_chunks += result
|
||||
processed += 1
|
||||
|
||||
logger.info(
|
||||
"Rebuild complete: processed=%d, chunks=%d, failed=%d",
|
||||
processed,
|
||||
total_chunks,
|
||||
failed,
|
||||
)
|
||||
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -84,7 +84,7 @@ export const extractVideoVoice = async (
|
||||
|
||||
return new Promise((resolve, reject) => {
|
||||
const xhr = new XMLHttpRequest()
|
||||
xhr.open("POST", "/api/v1/tts/extract-video-voice")
|
||||
xhr.open("POST", "/api/v1/voices/extract-voice")
|
||||
|
||||
// 携带认证 token(从 localStorage 获取,与 apiClient 拦截器一致)
|
||||
const token = localStorage.getItem("access_token")
|
||||
|
||||
@@ -136,8 +136,9 @@ export const MaterialVoiceTab: React.FC<MaterialVoiceTabProps> = ({
|
||||
const material = mapAssetToMaterial(asset)
|
||||
// duration 优先取顶层(后端从 metadata 提取),兜底 metadata
|
||||
const cardDuration = asset.duration || material.duration || 0
|
||||
// AI 生成素材标识(metadata.source === "tts_job")
|
||||
const isAiMaterial = (asset.metadata as Record<string, unknown>)?.source === "tts_job"
|
||||
// AI 生成素材标识:兼容旧素材(无 source 字段但有 tts_job_id)
|
||||
const meta = asset.metadata as Record<string, unknown>
|
||||
const isAiMaterial = meta?.source === "tts_job" || !!meta?.tts_job_id
|
||||
const isPlaying = playingId === asset.id
|
||||
const isSelected = selectedIds.has(asset.id)
|
||||
// 播放中以 audio 真实时长为准,未播放显示卡片时长
|
||||
@@ -184,8 +185,10 @@ export const MaterialVoiceTab: React.FC<MaterialVoiceTabProps> = ({
|
||||
</div>
|
||||
|
||||
<div className="xx-voice-info vmat-info">
|
||||
<div className="xx-voice-name" title={asset.name}>
|
||||
{asset.name}
|
||||
<div className="xx-voice-name-row">
|
||||
<div className="xx-voice-name" title={asset.name}>
|
||||
{asset.name}
|
||||
</div>
|
||||
{isAiMaterial && <span className="vmat-ai-badge">AI</span>}
|
||||
</div>
|
||||
<div className="xx-voice-subtitle">
|
||||
|
||||
@@ -1,11 +1,16 @@
|
||||
"""Video deduplication module - compute fingerprints and detect duplicates."""
|
||||
"""Video deduplication module - compute fingerprints and detect duplicates.
|
||||
|
||||
Dynamic keyframe detection + sliding window temporal matching (Issue #1659).
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import statistics
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
@@ -15,10 +20,34 @@ from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
|
||||
from packages.adapters.sqlalchemy_impl.models import VideoFingerprintChunkModel
|
||||
from packages.shared.storage import get_storage_service
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 关键帧检测常量 ──────────────────────────────────────────────
|
||||
SCENE_CHANGE_THRESHOLD = 30 # 灰度差异阈值
|
||||
MIN_KEYFRAME_INTERVAL_SEC = 1.0 # 最小关键帧间隔(秒)
|
||||
MAX_KEYFRAMES = 30 # 最大关键帧数
|
||||
MIN_KEYFRAMES = 5 # 最小关键帧数
|
||||
LONG_VIDEO_SEGMENT_SEC = 30 # 长视频每段秒数
|
||||
LONG_VIDEO_DURATION_THRESHOLD_SEC = 180 # 3 分钟阈值
|
||||
MIN_FRAMES_PER_SEGMENT = 2 # 长视频每段最少帧数
|
||||
|
||||
# ── 滑动窗口匹配常量 ────────────────────────────────────────────
|
||||
SEGMENT_MATCH_THRESHOLD = 8 # 帧匹配汉明距离阈值
|
||||
MIN_CONSECUTIVE_MATCHES = 5 # 最少连续匹配帧数
|
||||
MAX_GAP = 2 # 允许的最大间隙帧数
|
||||
|
||||
# ── 融合判定常量 ────────────────────────────────────────────────
|
||||
PHASH_WEIGHT = 0.7 # pHash 权重
|
||||
HISTOGRAM_WEIGHT = 0.3 # 直方图权重
|
||||
MATCH_RATIO_THRESHOLD = 0.7 # 至少 70% 帧匹配
|
||||
DUPLICATE_THRESHOLD = 0.70 # 融合后相似度阈值
|
||||
|
||||
|
||||
# ── 感知哈希 & 颜色直方图工具函数 ────────────────────────────────
|
||||
|
||||
|
||||
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
|
||||
"""计算图像的感知哈希(pHash),基于 DCT(离散余弦变换)。
|
||||
@@ -80,6 +109,131 @@ def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]:
|
||||
return hist
|
||||
|
||||
|
||||
# ── 关键帧检测 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def detect_keyframe_timestamps(
|
||||
video_path: str,
|
||||
*,
|
||||
min_interval_sec: float = MIN_KEYFRAME_INTERVAL_SEC,
|
||||
max_frames: int = MAX_KEYFRAMES,
|
||||
min_frames: int = MIN_KEYFRAMES,
|
||||
) -> list[float]:
|
||||
"""检测视频中的场景切换点,返回关键帧时间戳列表(秒)。
|
||||
|
||||
算法:
|
||||
1. 降采样到 320x240,逐帧转灰度
|
||||
2. 计算相邻帧灰度差异(像素均值差)
|
||||
3. 差异 > SCENE_CHANGE_THRESHOLD(30) 标记为候选关键帧
|
||||
4. 相邻关键帧间隔 < min_interval_sec 的,保留差异更大的那个
|
||||
5. 数量裁剪到 [min_frames, max_frames]
|
||||
|
||||
对于长视频(>3分钟):
|
||||
- 每 30 秒一个分段
|
||||
- 每个分段至少选 2 个关键帧(如果分段内无场景切换,均匀取 2 帧)
|
||||
"""
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
if not cap.isOpened():
|
||||
raise RuntimeError(f"Cannot open video: {video_path}")
|
||||
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
duration = frame_count / fps if fps > 0 else 0
|
||||
|
||||
if duration <= 0:
|
||||
cap.release()
|
||||
return []
|
||||
|
||||
# 逐帧检测场景切换
|
||||
candidates: list[tuple[float, float]] = [] # (timestamp_sec, diff_score)
|
||||
prev_gray = None
|
||||
|
||||
while True:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
# 降采样 + 灰度
|
||||
small = cv2.resize(frame, (320, 240))
|
||||
gray = cv2.cvtColor(small, cv2.COLOR_BGR2GRAY).astype(np.float32)
|
||||
|
||||
if prev_gray is not None:
|
||||
diff = float(np.mean(np.abs(gray - prev_gray)))
|
||||
if diff > SCENE_CHANGE_THRESHOLD:
|
||||
pos_ms = cap.get(cv2.CAP_PROP_POS_MSEC)
|
||||
candidates.append((pos_ms / 1000.0, diff))
|
||||
|
||||
prev_gray = gray
|
||||
|
||||
cap.release()
|
||||
|
||||
# 按最小间隔过滤(保留差异更大的)
|
||||
filtered: list[tuple[float, float]] = []
|
||||
for ts, diff in sorted(candidates):
|
||||
if filtered and (ts - filtered[-1][0]) < min_interval_sec:
|
||||
if diff > filtered[-1][1]:
|
||||
filtered[-1] = (ts, diff)
|
||||
else:
|
||||
filtered.append((ts, diff))
|
||||
|
||||
keyframe_times = [ts for ts, _ in filtered]
|
||||
|
||||
# 数量不足 min_frames 时,在时间轴上均匀补充
|
||||
if len(keyframe_times) < min_frames:
|
||||
uniform = [duration * (i + 0.5) / min_frames for i in range(min_frames)]
|
||||
keyframe_times = sorted(set(uniform) | set(keyframe_times))
|
||||
# 如果合并后还不足 min_frames,直接用均匀分布
|
||||
if len(keyframe_times) < min_frames:
|
||||
keyframe_times = uniform
|
||||
|
||||
# 数量超过 max_frames 时,均匀采样
|
||||
if len(keyframe_times) > max_frames:
|
||||
step = len(keyframe_times) / max_frames
|
||||
keyframe_times = [keyframe_times[int(i * step)] for i in range(max_frames)]
|
||||
|
||||
# 长视频分段保底(>3分钟)
|
||||
if duration > LONG_VIDEO_DURATION_THRESHOLD_SEC:
|
||||
segment_count = int(duration / LONG_VIDEO_SEGMENT_SEC)
|
||||
for seg_idx in range(segment_count):
|
||||
seg_start = seg_idx * LONG_VIDEO_SEGMENT_SEC
|
||||
seg_end = min((seg_idx + 1) * LONG_VIDEO_SEGMENT_SEC, duration)
|
||||
seg_frames = [t for t in keyframe_times if seg_start <= t < seg_end]
|
||||
if len(seg_frames) < MIN_FRAMES_PER_SEGMENT:
|
||||
# 均匀补齐
|
||||
for i in range(MIN_FRAMES_PER_SEGMENT):
|
||||
t = seg_start + LONG_VIDEO_SEGMENT_SEC * (i + 0.5) / MIN_FRAMES_PER_SEGMENT
|
||||
if t not in keyframe_times and seg_start <= t < seg_end:
|
||||
keyframe_times.append(t)
|
||||
keyframe_times.sort()
|
||||
|
||||
return keyframe_times
|
||||
|
||||
|
||||
# ── 数据类 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class FingerprintChunk:
|
||||
"""单个分片指纹数据。"""
|
||||
|
||||
start_time_ms: int
|
||||
end_time_ms: int
|
||||
phash_binary: str
|
||||
color_histogram: list[float]
|
||||
frame_count: int = 1
|
||||
|
||||
|
||||
@dataclass
|
||||
class DuplicateSegment:
|
||||
"""一段重复片段的描述。"""
|
||||
|
||||
query_start_ms: int
|
||||
query_end_ms: int
|
||||
target_start_ms: int
|
||||
target_end_ms: int
|
||||
avg_distance: float # 该段内帧的平均汉明距离
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoFingerprint:
|
||||
"""Video fingerprint containing multiple similarity metrics."""
|
||||
@@ -89,6 +243,7 @@ class VideoFingerprint:
|
||||
color_histograms: list[list[float]]
|
||||
duration: float
|
||||
resolution: tuple[int, int]
|
||||
chunks: list[FingerprintChunk] = field(default_factory=list)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
# 注意:color_histograms 里的值可能是 np.float32(来自 cv2.normalize),
|
||||
@@ -101,8 +256,168 @@ class VideoFingerprint:
|
||||
"color_histograms": native_histograms,
|
||||
"duration": float(self.duration),
|
||||
"resolution": [int(self.resolution[0]), int(self.resolution[1])],
|
||||
"chunks": [
|
||||
{
|
||||
"start_time_ms": c.start_time_ms,
|
||||
"end_time_ms": c.end_time_ms,
|
||||
"phash_binary": c.phash_binary,
|
||||
"color_histogram": [float(v) for v in c.color_histogram],
|
||||
"frame_count": c.frame_count,
|
||||
}
|
||||
for c in self.chunks
|
||||
],
|
||||
}
|
||||
|
||||
def to_chunk_models(self, video_id: str, project_id: str, user_id: str = "") -> list[VideoFingerprintChunkModel]:
|
||||
"""将分片数据转为 SQLAlchemy Model 列表,用于批量写入 video_fingerprint_chunks 表。"""
|
||||
models = []
|
||||
for chunk in self.chunks:
|
||||
models.append(
|
||||
VideoFingerprintChunkModel(
|
||||
id=uuid4().hex,
|
||||
video_id=video_id,
|
||||
project_id=project_id,
|
||||
user_id=user_id,
|
||||
start_time_ms=chunk.start_time_ms,
|
||||
end_time_ms=chunk.end_time_ms,
|
||||
phash_binary=chunk.phash_binary,
|
||||
color_histogram=[float(v) for v in chunk.color_histogram],
|
||||
frame_count=chunk.frame_count,
|
||||
)
|
||||
)
|
||||
return models
|
||||
|
||||
|
||||
# ── 滑动窗口时序匹配 ────────────────────────────────────────────
|
||||
|
||||
|
||||
def find_duplicate_segments(
|
||||
query_chunks: list,
|
||||
target_chunks: list,
|
||||
*,
|
||||
match_threshold: int = SEGMENT_MATCH_THRESHOLD,
|
||||
min_consecutive: int = MIN_CONSECUTIVE_MATCHES,
|
||||
max_gap: int = MAX_GAP,
|
||||
) -> list[DuplicateSegment]:
|
||||
"""滑动窗口时序匹配:找出两组分片之间的重复片段。
|
||||
|
||||
算法:
|
||||
1. 对每个 query chunk,找到 target 中汉明距离最小的 chunk
|
||||
2. 距离 <= match_threshold 视为匹配
|
||||
3. 找连续匹配的 run(允许 max_gap 帧间隙)
|
||||
4. 连续匹配数 >= min_consecutive 的 run 报告为重复片段
|
||||
|
||||
Args:
|
||||
query_chunks: 查询视频的分片列表(FingerprintChunk 或 dict)
|
||||
target_chunks: 目标视频的分片列表
|
||||
match_threshold: 汉明距离匹配阈值
|
||||
min_consecutive: 最少连续匹配帧数
|
||||
max_gap: 允许的最大间隙帧数
|
||||
|
||||
Returns:
|
||||
DuplicateSegment 列表
|
||||
"""
|
||||
if not query_chunks or not target_chunks:
|
||||
return []
|
||||
|
||||
def _get_phash(chunk) -> str:
|
||||
if isinstance(chunk, dict):
|
||||
return chunk["phash_binary"]
|
||||
return chunk.phash_binary
|
||||
|
||||
def _get_start(chunk) -> int:
|
||||
if isinstance(chunk, dict):
|
||||
return chunk["start_time_ms"]
|
||||
return chunk.start_time_ms
|
||||
|
||||
def _get_end(chunk) -> int:
|
||||
if isinstance(chunk, dict):
|
||||
return chunk["end_time_ms"]
|
||||
return chunk.end_time_ms
|
||||
|
||||
# Step 1: 逐帧匹配
|
||||
frame_matches: list[tuple[bool, int, int]] = [] # (is_match, min_dist, best_target_idx)
|
||||
for qc in query_chunks:
|
||||
qc_phash = _get_phash(qc)
|
||||
best_dist = 64
|
||||
best_idx = 0
|
||||
for j, tc in enumerate(target_chunks):
|
||||
d = hamming_distance(qc_phash, _get_phash(tc))
|
||||
if d < best_dist:
|
||||
best_dist = d
|
||||
best_idx = j
|
||||
frame_matches.append((best_dist <= match_threshold, best_dist, best_idx))
|
||||
|
||||
# Step 2: 找连续匹配的 runs
|
||||
runs: list[tuple[int, int]] = [] # list of (start_idx, end_idx)
|
||||
run_start = None
|
||||
gap_count = 0
|
||||
|
||||
for i, (is_match, _dist, _idx) in enumerate(frame_matches):
|
||||
if is_match:
|
||||
if run_start is None:
|
||||
run_start = i
|
||||
gap_count = 0 # 重置间隙
|
||||
else:
|
||||
if run_start is not None:
|
||||
gap_count += 1
|
||||
if gap_count > max_gap:
|
||||
# 中断当前 run
|
||||
run_end = i - gap_count # 最后一个匹配帧的索引
|
||||
# 计算 run 内的实际匹配帧数(总跨度 - 间隙数)
|
||||
total_gaps = sum(1 for k in range(run_start, run_end + 1) if not frame_matches[k][0])
|
||||
matching_count = (run_end - run_start + 1) - total_gaps
|
||||
if matching_count >= min_consecutive:
|
||||
runs.append((run_start, run_end))
|
||||
run_start = None
|
||||
gap_count = 0
|
||||
|
||||
# 处理末尾 run
|
||||
if run_start is not None:
|
||||
last_idx = len(frame_matches) - 1
|
||||
# 回退找到最后一个匹配帧的位置(跳过尾部非匹配帧)
|
||||
while last_idx >= run_start and not frame_matches[last_idx][0]:
|
||||
last_idx -= 1
|
||||
if last_idx >= run_start:
|
||||
# 计算 run 内的总间隙数
|
||||
total_gaps = sum(1 for k in range(run_start, last_idx + 1) if not frame_matches[k][0])
|
||||
matching_count = (last_idx - run_start + 1) - total_gaps
|
||||
if matching_count >= min_consecutive:
|
||||
runs.append((run_start, last_idx))
|
||||
|
||||
# Step 3: 构建 DuplicateSegment
|
||||
segments: list[DuplicateSegment] = []
|
||||
for start, end in runs:
|
||||
query_start = _get_start(query_chunks[start])
|
||||
query_end = _get_end(query_chunks[end])
|
||||
|
||||
# 取目标范围(按最佳匹配的目标 chunk 时间范围)
|
||||
target_indices = [frame_matches[k][2] for k in range(start, end + 1) if frame_matches[k][0]]
|
||||
if target_indices:
|
||||
t_min = min(target_indices)
|
||||
t_max = max(target_indices)
|
||||
target_start = _get_start(target_chunks[t_min])
|
||||
target_end = _get_end(target_chunks[t_max])
|
||||
else:
|
||||
target_start = _get_start(target_chunks[0])
|
||||
target_end = _get_end(target_chunks[-1])
|
||||
|
||||
avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1)) / (end - start + 1)
|
||||
segments.append(
|
||||
DuplicateSegment(
|
||||
query_start_ms=query_start,
|
||||
query_end_ms=query_end,
|
||||
target_start_ms=target_start,
|
||||
target_end_ms=target_end,
|
||||
avg_distance=avg_dist,
|
||||
)
|
||||
)
|
||||
|
||||
return segments
|
||||
|
||||
|
||||
# ── VideoDeduplicator ───────────────────────────────────────────
|
||||
|
||||
|
||||
class VideoDeduplicator:
|
||||
"""Video deduplication using multiple fingerprint methods."""
|
||||
@@ -111,7 +426,12 @@ class VideoDeduplicator:
|
||||
HISTOGRAM_THRESHOLD = 0.85
|
||||
|
||||
def compute_fingerprint(self, video_path: str) -> VideoFingerprint:
|
||||
"""Compute video fingerprint using MD5, pHash, and color histogram."""
|
||||
"""Compute video fingerprint using dynamic keyframe detection.
|
||||
|
||||
使用 detect_keyframe_timestamps() 检测内容感知关键帧,
|
||||
在每个关键帧处取帧计算 pHash + color_histogram。
|
||||
同时保留 MD5 计算和分片数据结构。
|
||||
"""
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
if not cap.isOpened():
|
||||
raise RuntimeError(f"Cannot open video: {video_path}")
|
||||
@@ -122,43 +442,122 @@ class VideoDeduplicator:
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
md5_hash = hashlib.md5(usedforsecurity=False)
|
||||
keyframe_phashes = []
|
||||
color_histograms = []
|
||||
cap.release()
|
||||
|
||||
frame_interval = max(1, frame_count // 10)
|
||||
for i in range(0, frame_count, frame_interval):
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, i)
|
||||
# 1. 检测关键帧时间戳
|
||||
keyframe_times = detect_keyframe_timestamps(video_path)
|
||||
|
||||
if not keyframe_times:
|
||||
return VideoFingerprint(
|
||||
md5="",
|
||||
keyframe_phashes=[],
|
||||
color_histograms=[],
|
||||
duration=duration,
|
||||
resolution=(width, height),
|
||||
chunks=[],
|
||||
)
|
||||
|
||||
# 2. 打开视频,逐个关键帧取帧
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
md5_hash = hashlib.md5(usedforsecurity=False)
|
||||
chunks: list[FingerprintChunk] = []
|
||||
|
||||
for i, t_sec in enumerate(keyframe_times):
|
||||
seek_ms = t_sec * 1000
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, seek_ms)
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
continue
|
||||
|
||||
# MD5 计算
|
||||
_, buffer = cv2.imencode(".jpg", frame)
|
||||
md5_hash.update(buffer)
|
||||
|
||||
keyframe_phashes.append(compute_phash(frame))
|
||||
color_histograms.append(compute_color_histogram(frame))
|
||||
phash = compute_phash(frame)
|
||||
hist = compute_color_histogram(frame)
|
||||
|
||||
# 计算分片时间范围(从前一个关键帧到下一个关键帧的中点)
|
||||
prev_boundary = keyframe_times[i - 1] * 1000 if i > 0 else 0
|
||||
next_boundary = keyframe_times[i + 1] * 1000 if i < len(keyframe_times) - 1 else duration * 1000
|
||||
start_ms = int((prev_boundary + seek_ms) / 2)
|
||||
end_ms = int((seek_ms + next_boundary) / 2)
|
||||
|
||||
chunks.append(
|
||||
FingerprintChunk(
|
||||
start_time_ms=start_ms,
|
||||
end_time_ms=end_ms,
|
||||
phash_binary=phash,
|
||||
color_histogram=hist,
|
||||
frame_count=1,
|
||||
)
|
||||
)
|
||||
|
||||
cap.release()
|
||||
|
||||
# 向后兼容:聚合 keyframe_phashes / color_histograms
|
||||
keyframe_phashes = [c.phash_binary for c in chunks]
|
||||
color_histograms = [c.color_histogram for c in chunks]
|
||||
|
||||
return VideoFingerprint(
|
||||
md5=md5_hash.hexdigest(),
|
||||
keyframe_phashes=keyframe_phashes,
|
||||
color_histograms=color_histograms,
|
||||
duration=duration,
|
||||
resolution=(width, height),
|
||||
chunks=chunks,
|
||||
)
|
||||
|
||||
def _get_existing_chunks(self, video_id: str, session: Session) -> list[dict]:
|
||||
"""从 video_fingerprint_chunks 表读取分片数据。返回空列表表示无分片数据。"""
|
||||
rows = (
|
||||
session.query(VideoFingerprintChunkModel)
|
||||
.filter(VideoFingerprintChunkModel.video_id == video_id)
|
||||
.order_by(VideoFingerprintChunkModel.start_time_ms)
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"phash_binary": r.phash_binary,
|
||||
"color_histogram": r.color_histogram,
|
||||
"start_time_ms": r.start_time_ms,
|
||||
"end_time_ms": r.end_time_ms,
|
||||
}
|
||||
for r in rows
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _bhattacharyya_coefficient(hist_a: list[float], hist_b: list[float]) -> float:
|
||||
"""Bhattacharyya 系数:Σ √(a[i] * b[i]),范围 [0, 1],1=完全相同。"""
|
||||
min_len = min(len(hist_a), len(hist_b))
|
||||
a = hist_a[:min_len]
|
||||
b = hist_b[:min_len]
|
||||
return float(sum(np.sqrt(ai * bi) for ai, bi in zip(a, b, strict=False)))
|
||||
|
||||
@staticmethod
|
||||
def _compute_histogram_similarity(
|
||||
histograms_a: list[list[float]],
|
||||
histograms_b: list[list[float]],
|
||||
) -> float:
|
||||
"""对每组直方图,找到最佳匹配的 Bhattacharyya 系数,取平均。"""
|
||||
if not histograms_a or not histograms_b:
|
||||
return 0.0
|
||||
similarities = []
|
||||
for ha in histograms_a:
|
||||
best = 0.0
|
||||
for hb in histograms_b:
|
||||
bc = VideoDeduplicator._bhattacharyya_coefficient(ha, hb)
|
||||
best = max(best, bc)
|
||||
similarities.append(best)
|
||||
return sum(similarities) / len(similarities) if similarities else 0.0
|
||||
|
||||
def check_duplicate(self, fingerprint: VideoFingerprint, project_id: str, session: Session) -> Optional[dict]:
|
||||
"""检查视频是否与项目中已有视频重复。
|
||||
|
||||
判定逻辑(按优先级):
|
||||
1. MD5 精确匹配:完全一致则 similarity=1.0,立即返回
|
||||
2. pHash 相似度:计算新视频每帧 phash 与已有视频每帧 phash 的最小汉明距离,
|
||||
取所有帧的平均值 avg_distance。若 avg_distance < PHASH_THRESHOLD(10),
|
||||
则判定为重复,similarity = 1.0 - (avg_distance / 64)
|
||||
查重逻辑:
|
||||
1. MD5 精确匹配 → similarity=1.0
|
||||
2. pHash 中位数距离 + 帧匹配比例 + 直方图融合判定
|
||||
|
||||
注意:返回第一个通过阈值的匹配(非最优匹配)。
|
||||
判定为重复后,调用 find_duplicate_segments() 获取具体重复片段。
|
||||
|
||||
Args:
|
||||
fingerprint: 待检测视频的指纹
|
||||
@@ -166,7 +565,7 @@ class VideoDeduplicator:
|
||||
session: 数据库会话
|
||||
|
||||
Returns:
|
||||
重复信息字典(含 duplicate, duplicate_of, reason, similarity),
|
||||
重复信息字典(含 duplicate, duplicate_of, reason, similarity, duplicate_segments),
|
||||
或 None 表示未找到重复。
|
||||
"""
|
||||
video_repo = SQLAlchemyGeneratedVideoRepository(session)
|
||||
@@ -182,28 +581,77 @@ class VideoDeduplicator:
|
||||
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", [])
|
||||
# 优先从分片表读取已有视频的分片 phash
|
||||
existing_phashes = []
|
||||
chunk_data = self._get_existing_chunks(existing.id, session)
|
||||
if chunk_data:
|
||||
existing_phashes = [c["phash_binary"] for c in chunk_data]
|
||||
else:
|
||||
# 回退:从 JSON 字段读取(存量旧视频)
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
|
||||
if not existing_phashes:
|
||||
continue
|
||||
|
||||
# 计算每个新关键帧到已有关键帧的最小汉明距离,取平均
|
||||
# 计算每个新关键帧到已有关键帧的最小汉明距离
|
||||
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:
|
||||
# 帧匹配比例检查
|
||||
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
|
||||
match_ratio = matching_frames / len(min_distances) if min_distances else 0
|
||||
if match_ratio < 0.7:
|
||||
continue
|
||||
|
||||
phash_similarity = 1.0 - (avg_distance / 64)
|
||||
# 中位数距离
|
||||
median_distance = statistics.median(min_distances) if min_distances else 64
|
||||
if median_distance >= self.PHASH_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 直方图融合
|
||||
existing_histograms = []
|
||||
if chunk_data:
|
||||
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
|
||||
else:
|
||||
existing_histograms = ef.get("color_histograms", [])
|
||||
|
||||
phash_similarity = 1.0 - (median_distance / 64)
|
||||
hist_similarity = (
|
||||
self._compute_histogram_similarity(fingerprint.color_histograms, existing_histograms)
|
||||
if existing_histograms
|
||||
else 0.5
|
||||
)
|
||||
combined_score = 0.7 * phash_similarity + 0.3 * hist_similarity
|
||||
|
||||
# DUPLICATE_THRESHOLD from module level
|
||||
if combined_score < DUPLICATE_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 滑动窗口时序匹配:获取具体重复片段
|
||||
existing_chunk_objects = (
|
||||
chunk_data
|
||||
if chunk_data
|
||||
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
|
||||
)
|
||||
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
|
||||
|
||||
return {
|
||||
"duplicate": True,
|
||||
"duplicate_of": existing.id,
|
||||
"reason": "phash_similar",
|
||||
"similarity": phash_similarity,
|
||||
"reason": "phash_histogram_fusion",
|
||||
"similarity": combined_score,
|
||||
"duplicate_segments": [
|
||||
{
|
||||
"query_start_ms": s.query_start_ms,
|
||||
"query_end_ms": s.query_end_ms,
|
||||
"target_start_ms": s.target_start_ms,
|
||||
"target_end_ms": s.target_end_ms,
|
||||
"avg_distance": round(s.avg_distance, 2),
|
||||
}
|
||||
for s in segments
|
||||
],
|
||||
}
|
||||
|
||||
return None
|
||||
@@ -217,7 +665,8 @@ class VideoDeduplicator:
|
||||
) -> Optional[dict]:
|
||||
"""检查视频是否与同批次内其他视频重复。
|
||||
|
||||
逻辑与 check_duplicate 一致(MD5 + pHash),但搜索范围限定为同 batch_id 的视频。
|
||||
逻辑与 check_duplicate 一致(MD5 + pHash + 直方图融合 + 时序匹配),
|
||||
但搜索范围限定为同 batch_id 的视频。
|
||||
|
||||
Args:
|
||||
fingerprint: 待检测视频的指纹
|
||||
@@ -247,7 +696,14 @@ class VideoDeduplicator:
|
||||
"similarity": 1.0,
|
||||
}
|
||||
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
# 优先从分片表读取
|
||||
existing_phashes = []
|
||||
chunk_data = self._get_existing_chunks(existing.id, session)
|
||||
if chunk_data:
|
||||
existing_phashes = [c["phash_binary"] for c in chunk_data]
|
||||
else:
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
|
||||
if not existing_phashes:
|
||||
continue
|
||||
|
||||
@@ -255,59 +711,63 @@ class VideoDeduplicator:
|
||||
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:
|
||||
# 帧匹配比例检查
|
||||
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
|
||||
match_ratio = matching_frames / len(min_distances) if min_distances else 0
|
||||
if match_ratio < 0.7:
|
||||
continue
|
||||
|
||||
phash_similarity = 1.0 - (avg_distance / 64)
|
||||
median_distance = statistics.median(min_distances) if min_distances else 64
|
||||
if median_distance >= self.PHASH_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 直方图融合
|
||||
existing_histograms = []
|
||||
if chunk_data:
|
||||
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
|
||||
else:
|
||||
existing_histograms = ef.get("color_histograms", [])
|
||||
|
||||
phash_similarity = 1.0 - (median_distance / 64)
|
||||
hist_similarity = (
|
||||
self._compute_histogram_similarity(fingerprint.color_histograms, existing_histograms)
|
||||
if existing_histograms
|
||||
else 0.5
|
||||
)
|
||||
combined_score = 0.7 * phash_similarity + 0.3 * hist_similarity
|
||||
|
||||
# DUPLICATE_THRESHOLD from module level
|
||||
if combined_score < DUPLICATE_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 滑动窗口时序匹配
|
||||
existing_chunk_objects = (
|
||||
chunk_data
|
||||
if chunk_data
|
||||
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
|
||||
)
|
||||
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
|
||||
|
||||
return {
|
||||
"duplicate": True,
|
||||
"duplicate_of": existing.id,
|
||||
"reason": "batch_phash_similar",
|
||||
"similarity": phash_similarity,
|
||||
"reason": "batch_phash_histogram_fusion",
|
||||
"similarity": combined_score,
|
||||
"duplicate_segments": [
|
||||
{
|
||||
"query_start_ms": s.query_start_ms,
|
||||
"query_end_ms": s.query_end_ms,
|
||||
"target_start_ms": s.target_start_ms,
|
||||
"target_end_ms": s.target_end_ms,
|
||||
"avg_distance": round(s.avg_distance, 2),
|
||||
}
|
||||
for s in segments
|
||||
],
|
||||
}
|
||||
|
||||
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
|
||||
|
||||
def compute_duplicate_rate(
|
||||
self,
|
||||
fingerprint: VideoFingerprint,
|
||||
@@ -320,9 +780,9 @@ class VideoDeduplicator:
|
||||
"""计算当前视频与用户库内已有视频的最高相似度百分比。
|
||||
|
||||
优先按 user_id 全局比较(跨项目),user_id 为空时回退到项目级比较。
|
||||
遍历最近 200 个其他有指纹的视频,对每个计算相似度:
|
||||
遍历最近 200 个其他有指纹的视频,对每个计算融合相似度:
|
||||
- MD5 精确匹配 → 100%
|
||||
- pHash 相似度 → (1.0 - avg_distance / 64) * 100
|
||||
- pHash + 直方图融合 → 0.7 * phash_sim + 0.3 * hist_sim
|
||||
取最高值作为 duplicate_rate(0~100)。
|
||||
如果没有其他视频可比较,返回 0.0。
|
||||
|
||||
@@ -336,7 +796,6 @@ class VideoDeduplicator:
|
||||
Returns:
|
||||
duplicate_rate: 0~100 的浮点数
|
||||
"""
|
||||
# 限制查询最近 200 个视频,避免大库内存溢出
|
||||
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
|
||||
|
||||
# 优先按 user_id 全局比较(跨项目),否则回退到项目级
|
||||
@@ -351,7 +810,7 @@ class VideoDeduplicator:
|
||||
)
|
||||
logger.debug("compute_duplicate_rate: project-level fallback project_id=%s", project_id)
|
||||
|
||||
# 排除当前视频自身(记录可能已写入 DB,必须在查询层排除)
|
||||
# 排除当前视频自身
|
||||
if current_video_id:
|
||||
query = query.filter(GeneratedVideoModel.id != current_video_id)
|
||||
|
||||
@@ -372,8 +831,14 @@ class VideoDeduplicator:
|
||||
if fingerprint.md5 == ef.get("md5"):
|
||||
return 100.0
|
||||
|
||||
# pHash 相似度
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
# 优先从分片表读取
|
||||
existing_phashes = []
|
||||
chunk_data = self._get_existing_chunks(existing.id, session)
|
||||
if chunk_data:
|
||||
existing_phashes = [c["phash_binary"] for c in chunk_data]
|
||||
else:
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
|
||||
if not existing_phashes or not fingerprint.keyframe_phashes:
|
||||
continue
|
||||
|
||||
@@ -381,13 +846,59 @@ class VideoDeduplicator:
|
||||
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 64
|
||||
similarity = (1.0 - avg_distance / 64) * 100
|
||||
max_similarity = max(max_similarity, similarity)
|
||||
|
||||
# 帧匹配比例检查
|
||||
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
|
||||
match_ratio = matching_frames / len(min_distances) if min_distances else 0
|
||||
if match_ratio < 0.7:
|
||||
continue
|
||||
|
||||
median_distance = statistics.median(min_distances) if min_distances else 64
|
||||
|
||||
# 直方图融合
|
||||
existing_histograms = []
|
||||
if chunk_data:
|
||||
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
|
||||
else:
|
||||
existing_histograms = ef.get("color_histograms", [])
|
||||
|
||||
phash_similarity = (1.0 - median_distance / 64) * 100
|
||||
hist_similarity = (
|
||||
self._compute_histogram_similarity(fingerprint.color_histograms, existing_histograms) * 100
|
||||
if existing_histograms
|
||||
else 50.0
|
||||
)
|
||||
combined_score = 0.7 * phash_similarity + 0.3 * hist_similarity
|
||||
max_similarity = max(max_similarity, combined_score)
|
||||
|
||||
return round(max(max_similarity, 0.0), 2)
|
||||
|
||||
|
||||
def _save_fingerprint_chunks(
|
||||
fingerprint: VideoFingerprint,
|
||||
video_id: str,
|
||||
project_id: str,
|
||||
user_id: str,
|
||||
session: Session,
|
||||
) -> None:
|
||||
"""将指纹分片数据批量写入 video_fingerprint_chunks 表。幂等:已有数据时跳过。"""
|
||||
# 幂等检查:已有分片数据则跳过
|
||||
existing_count = (
|
||||
session.query(VideoFingerprintChunkModel).filter(VideoFingerprintChunkModel.video_id == video_id).count()
|
||||
)
|
||||
if existing_count > 0:
|
||||
logger.debug("Fingerprint chunks already exist for video %s (%d chunks), skipping", video_id, existing_count)
|
||||
return
|
||||
|
||||
if not fingerprint.chunks:
|
||||
logger.warning("No chunks in fingerprint for video %s, skipping chunk save", video_id)
|
||||
return
|
||||
|
||||
chunk_models = fingerprint.to_chunk_models(video_id, project_id, user_id)
|
||||
session.bulk_save_objects(chunk_models)
|
||||
logger.info("Saved %d fingerprint chunks for video %s", len(chunk_models), video_id)
|
||||
|
||||
|
||||
@celery_app.task(bind=True, max_retries=3, name="worker.check_duplicate")
|
||||
def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
|
||||
"""Celery task to check if generated video is a duplicate."""
|
||||
@@ -421,6 +932,10 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
|
||||
video.duplicate_of = None
|
||||
|
||||
video_repo.update(video)
|
||||
|
||||
# 写入分片表
|
||||
_save_fingerprint_chunks(fingerprint, generated_video_id, video.project_id, video.user_id, session)
|
||||
|
||||
session.commit()
|
||||
|
||||
logger.info(f"Duplicate check completed for video {generated_video_id}: is_duplicate={video.is_duplicate}")
|
||||
|
||||
@@ -100,6 +100,14 @@ def create_video_record_and_dedup(
|
||||
|
||||
generated_video.video_fingerprint = fingerprint.to_dict()
|
||||
|
||||
# 写入分片指纹表
|
||||
from video_processing.dedup import _save_fingerprint_chunks
|
||||
|
||||
try:
|
||||
_save_fingerprint_chunks(fingerprint, video_id, project_id, user_id, session)
|
||||
except Exception as chunk_err:
|
||||
logger.warning("Failed to save fingerprint chunks for %s: %s", video_id, chunk_err)
|
||||
|
||||
# (a) 历史成片查重
|
||||
duplicate_result = deduplicator.check_duplicate(fingerprint, project_id, session)
|
||||
|
||||
|
||||
@@ -131,3 +131,65 @@ class SQLAlchemyEditPlanClipRepository:
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
def list_used_segments_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
limit_recent: int = 50,
|
||||
) -> dict[str, list[tuple[float, float]]]:
|
||||
"""查询用户已有视频中已使用的素材区间(跨视频避让).
|
||||
|
||||
JOIN edit_plans 表,按 created_by_user_id 过滤,只查 status='completed'
|
||||
的 plan 下 status='rendered' 且 asset_id 非空的 clips。按 plan 的
|
||||
created_at DESC 取最近 limit_recent 个 plan。
|
||||
|
||||
Returns:
|
||||
{asset_id: [(start_time, start_time + duration), ...]}
|
||||
空结果返回空 dict。
|
||||
"""
|
||||
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
|
||||
|
||||
if not user_id:
|
||||
return {}
|
||||
|
||||
# 1. 查出最近 limit_recent 个已完成 plan 的 ID
|
||||
recent_plan_ids = [
|
||||
row[0]
|
||||
for row in self.session.query(EditPlanModel.id)
|
||||
.filter(
|
||||
EditPlanModel.created_by_user_id == user_id,
|
||||
EditPlanModel.status == "completed",
|
||||
)
|
||||
.order_by(EditPlanModel.created_at.desc())
|
||||
.limit(limit_recent)
|
||||
.all()
|
||||
]
|
||||
|
||||
if not recent_plan_ids:
|
||||
return {}
|
||||
|
||||
# 2. 查这些 plan 下已渲染、有素材的 clips
|
||||
clips = (
|
||||
self.session.query(
|
||||
EditPlanClipModel.asset_id,
|
||||
EditPlanClipModel.start_time,
|
||||
EditPlanClipModel.duration,
|
||||
)
|
||||
.filter(
|
||||
EditPlanClipModel.plan_id.in_(recent_plan_ids),
|
||||
EditPlanClipModel.status == "rendered",
|
||||
EditPlanClipModel.asset_id != "",
|
||||
EditPlanClipModel.asset_id.isnot(None),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 3. 聚合为 {asset_id: [(start, start+duration), ...]}
|
||||
result: dict[str, list[tuple[float, float]]] = {}
|
||||
for asset_id, start_time, duration in clips:
|
||||
if asset_id not in result:
|
||||
result[asset_id] = []
|
||||
result[asset_id].append((start_time or 0.0, (start_time or 0.0) + (duration or 0.0)))
|
||||
|
||||
return result
|
||||
|
||||
@@ -620,3 +620,20 @@ class CoverTemplateModel(Base):
|
||||
config = Column(JSON, nullable=False, default=dict)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class VideoFingerprintChunkModel(Base):
|
||||
"""分片视频指纹 — 每个视频按时间分片存储 pHash + color_histogram."""
|
||||
|
||||
__tablename__ = "video_fingerprint_chunks"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
video_id = Column(String(36), nullable=False, index=True)
|
||||
project_id = Column(String(36), nullable=False, index=True)
|
||||
user_id = Column(String(36), nullable=False, index=True, default="")
|
||||
start_time_ms = Column(Integer, nullable=False)
|
||||
end_time_ms = Column(Integer, nullable=False)
|
||||
phash_binary = Column(String(16), nullable=False)
|
||||
color_histogram = Column(JSON, nullable=False)
|
||||
frame_count = Column(Integer, nullable=False, default=1)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -169,6 +169,7 @@ def distribute_assets(
|
||||
random_selection: bool = False,
|
||||
asset_durations: dict[str, float] | None = None,
|
||||
asset_scene_points: dict[str, list[float]] | None = None,
|
||||
external_used_segments: dict[str, list[tuple[float, float]]] | None = None,
|
||||
) -> None:
|
||||
"""按 editing_mode 将素材分配到 clips(就地修改).
|
||||
|
||||
@@ -188,6 +189,7 @@ def distribute_assets(
|
||||
random_selection: 是否随机选择素材(用于预览生成)
|
||||
asset_durations: 素材 ID -> 时长(秒)映射,用于设置 start_time
|
||||
asset_scene_points: 素材 ID -> 场景切换点列表(metadata 缓存)
|
||||
external_used_segments: 跨视频已用区间(来自其他视频的 clips),注入到分配逻辑中避让
|
||||
"""
|
||||
if not asset_ids or not clips:
|
||||
return
|
||||
@@ -198,16 +200,16 @@ def distribute_assets(
|
||||
random.shuffle(asset_ids)
|
||||
|
||||
if editing_mode == EditingMode.ONE_TAKE.value:
|
||||
_distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points)
|
||||
_distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
|
||||
elif editing_mode == EditingMode.PIP.value:
|
||||
_distribute_pip(clips, asset_ids, asset_durations, asset_scene_points)
|
||||
_distribute_pip(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
|
||||
elif editing_mode == EditingMode.VOICE_OVER.value:
|
||||
_distribute_voice_over(clips, asset_ids, asset_durations, asset_scene_points)
|
||||
_distribute_voice_over(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
|
||||
elif editing_mode == EditingMode.VOICE_PIP.value:
|
||||
_distribute_voice_pip(clips, asset_ids, asset_durations, asset_scene_points)
|
||||
_distribute_voice_pip(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
|
||||
else:
|
||||
# 未知模式,退化为 one_take
|
||||
_distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points)
|
||||
_distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
|
||||
|
||||
|
||||
def _resolve_start_time(
|
||||
@@ -248,9 +250,12 @@ def _distribute_one_take(
|
||||
asset_ids: List[str],
|
||||
asset_durations: dict[str, float] | None = None,
|
||||
asset_scene_points: dict[str, list[float]] | None = None,
|
||||
external_used_segments: dict[str, list[tuple[float, float]]] | None = None,
|
||||
) -> None:
|
||||
"""ONE_TAKE: 素材按顺序依次分配给 main 类型 clips."""
|
||||
used_segments: dict[str, list[tuple[float, float]]] = {}
|
||||
used_segments: dict[str, list[tuple[float, float]]] = (
|
||||
{k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {}
|
||||
)
|
||||
main_clips = [c for c in clips if c.clip_type == ClipType.MAIN.value]
|
||||
for i, clip in enumerate(main_clips):
|
||||
if i < len(asset_ids):
|
||||
@@ -271,9 +276,12 @@ def _distribute_pip(
|
||||
asset_ids: List[str],
|
||||
asset_durations: dict[str, float] | None = None,
|
||||
asset_scene_points: dict[str, list[float]] | None = None,
|
||||
external_used_segments: dict[str, list[tuple[float, float]]] | None = None,
|
||||
) -> None:
|
||||
"""PIP: 第1个素材→main(全屏背景),其余→overlay clips."""
|
||||
used_segments: dict[str, list[tuple[float, float]]] = {}
|
||||
used_segments: dict[str, list[tuple[float, float]]] = (
|
||||
{k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {}
|
||||
)
|
||||
# 第1个素材 → main clip
|
||||
main_clips = [c for c in clips if c.clip_type == ClipType.MAIN.value]
|
||||
if main_clips and asset_ids:
|
||||
@@ -310,9 +318,12 @@ def _distribute_voice_over(
|
||||
asset_ids: List[str],
|
||||
asset_durations: dict[str, float] | None = None,
|
||||
asset_scene_points: dict[str, list[float]] | None = None,
|
||||
external_used_segments: dict[str, list[tuple[float, float]]] | None = None,
|
||||
) -> None:
|
||||
"""VOICE_OVER: 素材→main clips (B-roll)."""
|
||||
used_segments: dict[str, list[tuple[float, float]]] = {}
|
||||
used_segments: dict[str, list[tuple[float, float]]] = (
|
||||
{k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {}
|
||||
)
|
||||
main_clips = [c for c in clips if c.clip_type == ClipType.MAIN.value]
|
||||
for i, clip in enumerate(main_clips):
|
||||
if i < len(asset_ids):
|
||||
@@ -333,9 +344,12 @@ def _distribute_voice_pip(
|
||||
asset_ids: List[str],
|
||||
asset_durations: dict[str, float] | None = None,
|
||||
asset_scene_points: dict[str, list[float]] | None = None,
|
||||
external_used_segments: dict[str, list[tuple[float, float]]] | None = None,
|
||||
) -> None:
|
||||
"""VOICE_PIP: 第1个→background, 第2个→corner_voice, 其余→b_roll."""
|
||||
used_segments: dict[str, list[tuple[float, float]]] = {}
|
||||
used_segments: dict[str, list[tuple[float, float]]] = (
|
||||
{k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {}
|
||||
)
|
||||
bg_clips = [c for c in clips if c.clip_type == "background"]
|
||||
voice_clips = [c for c in clips if c.clip_type == "corner_voice"]
|
||||
broll_clips = [c for c in clips if c.clip_type == "b_roll"]
|
||||
|
||||
@@ -0,0 +1,334 @@
|
||||
"""Tests for Issue #1670 — 跨视频片段避让(生成前注入已用区间)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import (
|
||||
SQLAlchemyEditPlanClipRepository,
|
||||
)
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
from packages.domain.plan_generator_utils import (
|
||||
_distribute_one_take,
|
||||
distribute_assets,
|
||||
)
|
||||
|
||||
# ── Repository 层测试 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListUsedSegmentsByUser:
|
||||
"""测试 list_used_segments_by_user 方法."""
|
||||
|
||||
def _make_repo(self, session_mock):
|
||||
return SQLAlchemyEditPlanClipRepository(session_mock)
|
||||
|
||||
def test_empty_user_id_returns_empty_dict(self):
|
||||
"""空 user_id 直接返回空 dict,不查 DB."""
|
||||
session = MagicMock()
|
||||
repo = self._make_repo(session)
|
||||
result = repo.list_used_segments_by_user("")
|
||||
assert result == {}
|
||||
session.query.assert_not_called()
|
||||
|
||||
def test_no_completed_plans_returns_empty_dict(self):
|
||||
"""用户没有已完成的 plan 时返回空 dict."""
|
||||
session = MagicMock()
|
||||
# Mock plan query returns empty
|
||||
plan_query = MagicMock()
|
||||
plan_query.filter.return_value = plan_query
|
||||
plan_query.order_by.return_value = plan_query
|
||||
plan_query.limit.return_value = plan_query
|
||||
plan_query.all.return_value = []
|
||||
session.query.return_value = plan_query
|
||||
|
||||
repo = self._make_repo(session)
|
||||
result = repo.list_used_segments_by_user("user_123")
|
||||
assert result == {}
|
||||
|
||||
def test_aggregates_clips_from_multiple_plans(self):
|
||||
"""从多个已完成 plan 的 clips 聚合已用区间."""
|
||||
session = MagicMock()
|
||||
|
||||
# Mock plan query: 2 completed plans
|
||||
plan_query = MagicMock()
|
||||
plan_query.filter.return_value = plan_query
|
||||
plan_query.order_by.return_value = plan_query
|
||||
plan_query.limit.return_value = plan_query
|
||||
plan_query.all.return_value = [("plan_1",), ("plan_2",)]
|
||||
session.query.return_value = plan_query
|
||||
|
||||
# Mock clip query: clips from both plans
|
||||
clip_query = MagicMock()
|
||||
clip_query.filter.return_value = clip_query
|
||||
clip_query.all.return_value = [
|
||||
("asset_A", 0.0, 5.0), # plan_1, asset A: 0~5s
|
||||
("asset_A", 10.0, 3.0), # plan_1, asset A: 10~13s
|
||||
("asset_B", 2.0, 4.0), # plan_2, asset B: 2~6s
|
||||
]
|
||||
# Second session.query call is for clips
|
||||
session.query.side_effect = [plan_query, clip_query]
|
||||
|
||||
repo = self._make_repo(session)
|
||||
result = repo.list_used_segments_by_user("user_123")
|
||||
|
||||
assert "asset_A" in result
|
||||
assert len(result["asset_A"]) == 2
|
||||
assert result["asset_A"][0] == (0.0, 5.0)
|
||||
assert result["asset_A"][1] == (10.0, 13.0)
|
||||
assert "asset_B" in result
|
||||
assert result["asset_B"][0] == (2.0, 6.0)
|
||||
|
||||
def test_respects_limit_recent_parameter(self):
|
||||
"""limit_recent 参数限制查询的 plan 数量."""
|
||||
session = MagicMock()
|
||||
|
||||
plan_query = MagicMock()
|
||||
plan_query.filter.return_value = plan_query
|
||||
plan_query.order_by.return_value = plan_query
|
||||
plan_query.limit.return_value = plan_query
|
||||
plan_query.all.return_value = [("plan_1",)]
|
||||
session.query.return_value = plan_query
|
||||
|
||||
clip_query = MagicMock()
|
||||
clip_query.filter.return_value = clip_query
|
||||
clip_query.all.return_value = [("asset_X", 1.0, 2.0)]
|
||||
session.query.side_effect = [plan_query, clip_query]
|
||||
|
||||
repo = self._make_repo(session)
|
||||
result = repo.list_used_segments_by_user("user_123", limit_recent=10)
|
||||
|
||||
# Verify limit was called with the parameter
|
||||
plan_query.limit.assert_called_once_with(10)
|
||||
assert "asset_X" in result
|
||||
|
||||
|
||||
# ── Domain 层测试 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestDistributeAssetsWithExternalSegments:
|
||||
"""测试 distribute_assets 传入 external_used_segments 的行为."""
|
||||
|
||||
def _make_clips(self, count: int, duration: float = 3.0) -> list[EditPlanClip]:
|
||||
"""创建指定数量的 MAIN 类型 clips."""
|
||||
return [
|
||||
EditPlanClip(
|
||||
id=f"clip_{i}",
|
||||
plan_id="plan_1",
|
||||
clip_type="main",
|
||||
order=i,
|
||||
template_clip_config_id="",
|
||||
asset_id="",
|
||||
text_content="",
|
||||
start_time=0.0,
|
||||
duration=duration,
|
||||
status=EditPlanClipStatus.PENDING,
|
||||
)
|
||||
for i in range(count)
|
||||
]
|
||||
|
||||
def test_external_used_segments_none_backward_compatible(self):
|
||||
"""external_used_segments=None 时行为不变(向后兼容)."""
|
||||
clips = self._make_clips(3)
|
||||
asset_ids = ["asset_1", "asset_2", "asset_3"]
|
||||
asset_durations = {aid: 30.0 for aid in asset_ids}
|
||||
|
||||
# Should not raise
|
||||
distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
"one_take",
|
||||
asset_durations=asset_durations,
|
||||
external_used_segments=None,
|
||||
)
|
||||
|
||||
# All clips should have assets assigned
|
||||
for clip in clips:
|
||||
assert clip.asset_id != ""
|
||||
|
||||
def test_external_used_segments_avoids_existing_ranges(self):
|
||||
"""传入 external_used_segments 后,新分配的 start_time 避开已有区间."""
|
||||
clips = self._make_clips(2, duration=3.0)
|
||||
asset_ids = ["asset_1"]
|
||||
asset_durations = {"asset_1": 30.0}
|
||||
|
||||
# Pretend asset_1 0~10s is already used by another video
|
||||
external = {"asset_1": [(0.0, 10.0)]}
|
||||
|
||||
# Run multiple times to check that start_time always avoids 0~10s
|
||||
# (with some randomness, but the avoidance should be consistent)
|
||||
for _ in range(10):
|
||||
test_clips = self._make_clips(1, duration=3.0)
|
||||
distribute_assets(
|
||||
test_clips,
|
||||
asset_ids,
|
||||
"one_take",
|
||||
asset_durations=asset_durations,
|
||||
external_used_segments=external,
|
||||
)
|
||||
start = test_clips[0].start_time
|
||||
# Start time + duration (3s) should not overlap with 0~10
|
||||
# i.e., start >= 10.0 or start + 3 <= 0.0 (impossible since start >= 0)
|
||||
assert (
|
||||
start >= 10.0 or start + 3.0 <= 0.0 or start >= 10.0
|
||||
), f"start_time {start} overlaps with existing segment 0~10"
|
||||
|
||||
def test_external_used_segments_deep_copy(self):
|
||||
"""external_used_segments 会被深拷贝,不会修改外部数据."""
|
||||
external = {"asset_1": [(0.0, 5.0)]}
|
||||
original = {"asset_1": [(0.0, 5.0)]}
|
||||
|
||||
clips = self._make_clips(1, duration=2.0)
|
||||
asset_ids = ["asset_1"]
|
||||
asset_durations = {"asset_1": 20.0}
|
||||
|
||||
distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
"one_take",
|
||||
asset_durations=asset_durations,
|
||||
external_used_segments=external,
|
||||
)
|
||||
|
||||
# External dict should be unchanged
|
||||
assert external == original
|
||||
|
||||
def test_empty_external_used_segments_same_as_none(self):
|
||||
"""空 dict 的 external_used_segments 行为与 None 相同."""
|
||||
clips = self._make_clips(2, duration=3.0)
|
||||
asset_ids = ["asset_1", "asset_2"]
|
||||
asset_durations = {aid: 30.0 for aid in asset_ids}
|
||||
|
||||
# Should not raise and should assign assets normally
|
||||
distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
"one_take",
|
||||
asset_durations=asset_durations,
|
||||
external_used_segments={},
|
||||
)
|
||||
for clip in clips:
|
||||
assert clip.asset_id != ""
|
||||
|
||||
|
||||
# ── Service 层测试 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestServiceLayerIntegration:
|
||||
"""测试 _distribute_assets 在 service 层的查询逻辑."""
|
||||
|
||||
def _make_service(self, clip_repo_mock, asset_repo_mock=None):
|
||||
"""创建 PlanGeneratorService 并注入 mock repos."""
|
||||
|
||||
from apps.api.app.services.plan_generator_service import PlanGeneratorService
|
||||
|
||||
with (
|
||||
patch("apps.api.app.services.plan_generator_service.SQLAlchemyEditPlanRepository"),
|
||||
patch(
|
||||
"apps.api.app.services.plan_generator_service.SQLAlchemyEditPlanClipRepository",
|
||||
return_value=clip_repo_mock,
|
||||
),
|
||||
):
|
||||
db = MagicMock()
|
||||
svc = PlanGeneratorService(db, asset_repo=asset_repo_mock)
|
||||
svc._clip_repo = clip_repo_mock
|
||||
return svc
|
||||
|
||||
def _make_clip(self):
|
||||
return EditPlanClip(
|
||||
id="clip_1",
|
||||
plan_id="plan_1",
|
||||
clip_type="main",
|
||||
order=0,
|
||||
template_clip_config_id="",
|
||||
asset_id="",
|
||||
text_content="",
|
||||
start_time=0.0,
|
||||
duration=3.0,
|
||||
status=EditPlanClipStatus.PENDING,
|
||||
)
|
||||
|
||||
def test_query_called_with_user_id(self):
|
||||
"""有 user_id 时调用 list_used_segments_by_user."""
|
||||
clip_repo = MagicMock()
|
||||
clip_repo.list_used_segments_by_user.return_value = {"asset_A": [(0.0, 5.0)]}
|
||||
asset_repo = MagicMock()
|
||||
asset_repo.get.return_value = None # smart_match fallback
|
||||
|
||||
svc = self._make_service(clip_repo, asset_repo)
|
||||
clips = [self._make_clip()]
|
||||
|
||||
svc._distribute_assets(
|
||||
clips,
|
||||
["asset_A"],
|
||||
"one_take",
|
||||
asset_durations={"asset_A": 30.0},
|
||||
user_id="user_123",
|
||||
)
|
||||
|
||||
clip_repo.list_used_segments_by_user.assert_called_once_with("user_123", limit_recent=50)
|
||||
|
||||
def test_query_not_called_without_user_id(self):
|
||||
"""无 user_id 时不调用查询."""
|
||||
clip_repo = MagicMock()
|
||||
asset_repo = MagicMock()
|
||||
asset_repo.get.return_value = None
|
||||
|
||||
svc = self._make_service(clip_repo, asset_repo)
|
||||
clips = [self._make_clip()]
|
||||
|
||||
svc._distribute_assets(
|
||||
clips,
|
||||
["asset_A"],
|
||||
"one_take",
|
||||
asset_durations={"asset_A": 30.0},
|
||||
user_id="",
|
||||
)
|
||||
|
||||
clip_repo.list_used_segments_by_user.assert_not_called()
|
||||
|
||||
def test_query_failure_does_not_block_generation(self):
|
||||
"""查询失败时不阻塞生成,回退到纯随机."""
|
||||
clip_repo = MagicMock()
|
||||
clip_repo.list_used_segments_by_user.side_effect = Exception("DB error")
|
||||
asset_repo = MagicMock()
|
||||
asset_repo.get.return_value = None
|
||||
|
||||
svc = self._make_service(clip_repo, asset_repo)
|
||||
clips = [self._make_clip()]
|
||||
|
||||
# Should not raise
|
||||
svc._distribute_assets(
|
||||
clips,
|
||||
["asset_A"],
|
||||
"one_take",
|
||||
asset_durations={"asset_A": 30.0},
|
||||
user_id="user_123",
|
||||
)
|
||||
|
||||
# Clip should still get an asset assigned (fallback to random)
|
||||
assert clips[0].asset_id == "asset_A"
|
||||
|
||||
def test_preview_and_final_both_query(self):
|
||||
"""预览和正式生成都触发查询."""
|
||||
for random_selection in [True, False]:
|
||||
clip_repo = MagicMock()
|
||||
clip_repo.list_used_segments_by_user.return_value = {}
|
||||
asset_repo = MagicMock()
|
||||
asset_repo.get.return_value = None
|
||||
|
||||
svc = self._make_service(clip_repo, asset_repo)
|
||||
clips = [self._make_clip()]
|
||||
|
||||
svc._distribute_assets(
|
||||
clips,
|
||||
["asset_A"],
|
||||
"one_take",
|
||||
random_selection=random_selection,
|
||||
asset_durations={"asset_A": 30.0},
|
||||
user_id="user_123",
|
||||
)
|
||||
|
||||
clip_repo.list_used_segments_by_user.assert_called_once()
|
||||
@@ -285,8 +285,10 @@ class TestVideoDeduplicatorCheckDuplicate:
|
||||
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"
|
||||
assert result["similarity"] == pytest.approx(
|
||||
0.85, abs=0.01
|
||||
) # combined: 0.7*1.0 + 0.3*0.5 (no hist fallback)
|
||||
assert result["reason"] == "phash_histogram_fusion"
|
||||
finally:
|
||||
self._restore_repo(mod, orig)
|
||||
|
||||
@@ -425,8 +427,11 @@ class TestVideoDeduplicatorCheckDuplicate:
|
||||
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
|
||||
# 新算法: median_distance=1, phash_sim=1-1/64=0.984375
|
||||
# 无直方图 → hist_sim=0.5(fallback)
|
||||
# combined = 0.7*0.984375 + 0.3*0.5 = 0.839062
|
||||
expected_sim = 0.7 * (1.0 - 1.0 / 64) + 0.3 * 0.5
|
||||
assert abs(result["similarity"] - expected_sim) < 1e-6
|
||||
finally:
|
||||
self._restore_repo(mod, orig)
|
||||
|
||||
@@ -456,7 +461,9 @@ class TestVideoDeduplicatorCheckDuplicate:
|
||||
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
|
||||
# 新算法: median_distance=0, phash_sim=1.0, hist_sim=0.5(fallback)
|
||||
# combined = 0.7*1.0 + 0.3*0.5 = 0.85
|
||||
assert result["similarity"] == pytest.approx(0.85, abs=0.01)
|
||||
finally:
|
||||
self._restore_repo(mod, orig)
|
||||
|
||||
@@ -539,7 +546,7 @@ class TestVideoDeduplicatorCheckBatchDuplicate:
|
||||
result = deduplicator.check_batch_duplicate(fingerprint, "batch-1", "vid-self", mock_session)
|
||||
assert result is not None
|
||||
assert result["duplicate"] is True
|
||||
assert result["reason"] == "batch_phash_similar"
|
||||
assert result["reason"] == "batch_phash_histogram_fusion"
|
||||
finally:
|
||||
self._restore_repo(mod, orig)
|
||||
|
||||
|
||||
@@ -181,70 +181,68 @@ class TestVideoFingerprint:
|
||||
assert d["color_histograms"] == []
|
||||
|
||||
|
||||
class TestAverageHistogramSimilarity:
|
||||
"""_average_histogram_similarity 直方图相似度测试."""
|
||||
class TestBhattacharyyaCoefficient:
|
||||
"""_bhattacharyya_coefficient Bhattacharyya 系数测试."""
|
||||
|
||||
def test_identical_histograms(self):
|
||||
"""完全相同的直方图相似度为1.0."""
|
||||
hist = [[0.5, 0.5, 0.0], [0.3, 0.4, 0.3]]
|
||||
sim = VideoDeduplicator._average_histogram_similarity(hist, hist)
|
||||
assert sim == pytest.approx(1.0)
|
||||
"""完全相同的直方图系数为1.0."""
|
||||
hist = [0.5, 0.5, 0.0, 0.3]
|
||||
bc = VideoDeduplicator._bhattacharyya_coefficient(hist, hist)
|
||||
# Σ √(a[i]*a[i]) = Σ a[i] = 1.0 (normalized)
|
||||
assert bc == pytest.approx(sum(h for h in hist))
|
||||
|
||||
def test_empty_first_list(self):
|
||||
def test_zero_histograms(self):
|
||||
"""全零直方图系数为0."""
|
||||
bc = VideoDeduplicator._bhattacharyya_coefficient([0.0, 0.0], [0.0, 0.0])
|
||||
assert bc == 0.0
|
||||
|
||||
def test_orthogonal_histograms(self):
|
||||
"""正交直方图(无重叠)系数为0."""
|
||||
bc = VideoDeduplicator._bhattacharyya_coefficient([1.0, 0.0], [0.0, 1.0])
|
||||
assert bc == pytest.approx(0.0)
|
||||
|
||||
def test_different_lengths(self):
|
||||
"""不同长度直方图取最小长度对齐."""
|
||||
bc = VideoDeduplicator._bhattacharyya_coefficient([1.0, 1.0, 0.0, 0.0], [1.0, 1.0])
|
||||
# 对齐到前2维: √(1*1) + √(1*1) = 2.0
|
||||
assert bc == pytest.approx(2.0)
|
||||
|
||||
def test_known_value(self):
|
||||
"""已知值验证."""
|
||||
# [0.25, 0.25, 0.25, 0.25] vs [0.25, 0.25, 0.25, 0.25]
|
||||
# BC = 4 * √(0.25 * 0.25) = 4 * 0.25 = 1.0
|
||||
hist = [0.25, 0.25, 0.25, 0.25]
|
||||
bc = VideoDeduplicator._bhattacharyya_coefficient(hist, hist)
|
||||
assert bc == pytest.approx(1.0)
|
||||
|
||||
|
||||
class TestComputeHistogramSimilarity:
|
||||
"""_compute_histogram_similarity 多帧直方图相似度测试."""
|
||||
|
||||
def test_identical_histogram_groups(self):
|
||||
"""完全相同的两组直方图."""
|
||||
hist = [[0.5, 0.5], [0.3, 0.4]]
|
||||
sim = VideoDeduplicator._compute_histogram_similarity(hist, hist)
|
||||
# Each hist finds best match = itself
|
||||
assert sim > 0.0
|
||||
|
||||
def test_empty_first(self):
|
||||
"""第一组为空返回0."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity([], [[0.5, 0.5]])
|
||||
assert sim == 0.0
|
||||
assert VideoDeduplicator._compute_histogram_similarity([], [[0.5]]) == 0.0
|
||||
|
||||
def test_empty_second_list(self):
|
||||
def test_empty_second(self):
|
||||
"""第二组为空返回0."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity([[0.5, 0.5]], [])
|
||||
assert sim == 0.0
|
||||
assert VideoDeduplicator._compute_histogram_similarity([[0.5]], []) == 0.0
|
||||
|
||||
def test_both_empty(self):
|
||||
"""两组都为空返回0."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity([], [])
|
||||
assert sim == 0.0
|
||||
assert VideoDeduplicator._compute_histogram_similarity([], []) == 0.0
|
||||
|
||||
def test_orthogonal_histograms(self):
|
||||
"""正交直方图相似度为0."""
|
||||
# [1, 0] 和 [0, 1] 正交
|
||||
sim = VideoDeduplicator._average_histogram_similarity([[1.0, 0.0]], [[0.0, 1.0]])
|
||||
assert sim == pytest.approx(0.0)
|
||||
|
||||
def test_partial_similarity(self):
|
||||
"""部分相似."""
|
||||
# [1, 1] 和 [1, 0] 的余弦相似度 = 1/√2 ≈ 0.707
|
||||
sim = VideoDeduplicator._average_histogram_similarity([[1.0, 1.0]], [[1.0, 0.0]])
|
||||
assert sim == pytest.approx(1.0 / (2**0.5), rel=0.01)
|
||||
|
||||
def test_multiple_frames_best_match(self):
|
||||
def test_best_match_selection(self):
|
||||
"""多帧时取最佳匹配."""
|
||||
# 第一帧完全不同,第二帧完全相同 → 平均 best = (0 + 1) / 2 = 0.5
|
||||
sim = VideoDeduplicator._average_histogram_similarity(
|
||||
[[1.0, 0.0], [0.0, 1.0]],
|
||||
[[0.0, 1.0]], # 只有一帧,和第一帧0相似,和第二帧1相似
|
||||
)
|
||||
# 第一帧最佳匹配=0,第二帧最佳匹配=1,平均=0.5
|
||||
assert sim == pytest.approx(0.5)
|
||||
|
||||
def test_zero_norm_histogram_skipped(self):
|
||||
"""零范数直方图被跳过."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity([[0.0, 0.0]], [[1.0, 1.0]])
|
||||
# 第一组的零范数被跳过,similarities为空,返回0
|
||||
assert sim == 0.0
|
||||
|
||||
def test_different_length_histograms(self):
|
||||
"""不同长度的直方图取最小长度对齐."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity(
|
||||
[[1.0, 1.0, 0.0, 0.0]], # 4维
|
||||
[[1.0, 1.0]], # 2维
|
||||
)
|
||||
# 对齐到前2维,都是[1,1],相似度1.0
|
||||
# ha[0] 与 hb[0] 正交,与 hb[1] 完全相同
|
||||
a = [[1.0, 0.0]]
|
||||
b = [[0.0, 1.0], [1.0, 0.0]]
|
||||
sim = VideoDeduplicator._compute_histogram_similarity(a, b)
|
||||
# Best match for [1,0]: max(BC([1,0],[0,1]), BC([1,0],[1,0])) = max(0, 1) = 1
|
||||
assert sim == pytest.approx(1.0)
|
||||
|
||||
def test_similarity_in_zero_one_range(self):
|
||||
"""相似度在[0, 1]范围内."""
|
||||
hist_a = [np.random.rand(96).tolist() for _ in range(5)]
|
||||
hist_b = [np.random.rand(96).tolist() for _ in range(5)]
|
||||
sim = VideoDeduplicator._average_histogram_similarity(hist_a, hist_b)
|
||||
assert 0.0 <= sim <= 1.0
|
||||
|
||||
@@ -0,0 +1,532 @@
|
||||
"""Issue #1659: 动态抽帧 + 滑动窗口时序匹配 单元测试.
|
||||
|
||||
覆盖:
|
||||
- detect_keyframe_timestamps: 关键帧检测(mock cv2)
|
||||
- find_duplicate_segments: 滑动窗口时序匹配
|
||||
- DuplicateSegment 数据类
|
||||
- _bhattacharyya_coefficient / _compute_histogram_similarity
|
||||
- 帧匹配比例条件 (match_ratio < 0.7 → 跳过)
|
||||
- 中位数 vs 均值(抵抗异常值)
|
||||
- 向后兼容(无分片数据时不崩溃)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
def _mock_module(**attrs):
|
||||
"""Create a mock module with __spec__ to avoid AttributeError."""
|
||||
m = MagicMock()
|
||||
m.__spec__ = None
|
||||
for k, v in attrs.items():
|
||||
setattr(m, k, v)
|
||||
return m
|
||||
|
||||
|
||||
# ── Module-level setup: mock deps, import dedup, then restore sys.modules ──
|
||||
_SAVED_MODULES_KEYS = set(sys.modules.keys())
|
||||
_SAVED_MODULES_VALUES = {
|
||||
k: sys.modules.get(k)
|
||||
for k in [
|
||||
"cv2",
|
||||
"celery",
|
||||
"sqlalchemy",
|
||||
"sqlalchemy.orm",
|
||||
"sqlalchemy.engine",
|
||||
"sqlalchemy.ext",
|
||||
"sqlalchemy.ext.declarative",
|
||||
"worker_app.db",
|
||||
"worker_app.celery_app",
|
||||
"worker_app.core.config",
|
||||
"packages.adapters.sqlalchemy_impl.session",
|
||||
"packages.adapters.sqlalchemy_impl.generated_video_repository",
|
||||
"packages.adapters.sqlalchemy_impl.models",
|
||||
"packages.shared.config",
|
||||
"packages.shared.storage",
|
||||
]
|
||||
}
|
||||
|
||||
sys.modules["cv2"] = _mock_module()
|
||||
|
||||
_mock_celery = MagicMock()
|
||||
_mock_celery.Task = MagicMock
|
||||
_mock_celery.Celery = MagicMock
|
||||
_mock_celery.__spec__ = None
|
||||
sys.modules["celery"] = _mock_celery
|
||||
|
||||
_mock_sqla = MagicMock()
|
||||
_mock_sqla.__path__ = []
|
||||
_mock_sqla.__spec__ = None
|
||||
sys.modules["sqlalchemy"] = _mock_sqla
|
||||
|
||||
_mock_sqla_orm = MagicMock()
|
||||
_mock_sqla_orm.__path__ = []
|
||||
_mock_sqla_orm.__spec__ = None
|
||||
_mock_sqla_orm.Session = MagicMock
|
||||
sys.modules["sqlalchemy.orm"] = _mock_sqla_orm
|
||||
sys.modules["sqlalchemy.engine"] = _mock_module()
|
||||
sys.modules["sqlalchemy.ext"] = _mock_module()
|
||||
sys.modules["sqlalchemy.ext.declarative"] = _mock_module()
|
||||
|
||||
sys.modules["worker_app.db"] = _mock_module(SessionLocal=MagicMock())
|
||||
sys.modules["worker_app.celery_app"] = _mock_module(celery_app=MagicMock())
|
||||
sys.modules["worker_app.core.config"] = _mock_module(get_settings=MagicMock(return_value=MagicMock()))
|
||||
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.session"] = _mock_module(
|
||||
Base=MagicMock(),
|
||||
build_engine=MagicMock(),
|
||||
build_session_factory=MagicMock(),
|
||||
ensure_database_exists=MagicMock(),
|
||||
initialize_database=MagicMock(),
|
||||
)
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.generated_video_repository"] = _mock_module(
|
||||
SQLAlchemyGeneratedVideoRepository=MagicMock
|
||||
)
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.models"] = _mock_module(
|
||||
VideoFingerprintChunkModel=MagicMock,
|
||||
GeneratedVideoModel=MagicMock,
|
||||
)
|
||||
sys.modules["packages.shared.config"] = _mock_module(get_shared_settings=MagicMock(return_value=MagicMock()))
|
||||
sys.modules["packages.shared.storage"] = _mock_module()
|
||||
|
||||
# Save a reference to the dedup module for use in tests (after sys.modules restore)
|
||||
import video_processing.dedup as _dedup_mod
|
||||
from video_processing.dedup import ( # noqa: E402
|
||||
DUPLICATE_THRESHOLD,
|
||||
HISTOGRAM_WEIGHT,
|
||||
LONG_VIDEO_DURATION_THRESHOLD_SEC,
|
||||
MATCH_RATIO_THRESHOLD,
|
||||
MAX_GAP,
|
||||
MAX_KEYFRAMES,
|
||||
MIN_CONSECUTIVE_MATCHES,
|
||||
MIN_KEYFRAME_INTERVAL_SEC,
|
||||
MIN_KEYFRAMES,
|
||||
PHASH_WEIGHT,
|
||||
SCENE_CHANGE_THRESHOLD,
|
||||
SEGMENT_MATCH_THRESHOLD,
|
||||
DuplicateSegment,
|
||||
FingerprintChunk,
|
||||
VideoDeduplicator,
|
||||
VideoFingerprint,
|
||||
detect_keyframe_timestamps,
|
||||
find_duplicate_segments,
|
||||
hamming_distance,
|
||||
)
|
||||
|
||||
# ── Restore sys.modules immediately after import ──
|
||||
for _key in list(sys.modules.keys()):
|
||||
if _key not in _SAVED_MODULES_KEYS:
|
||||
del sys.modules[_key]
|
||||
for _key, _value in _SAVED_MODULES_VALUES.items():
|
||||
if _value is not None:
|
||||
sys.modules[_key] = _value
|
||||
elif _key in sys.modules:
|
||||
del sys.modules[_key]
|
||||
del _SAVED_MODULES_KEYS, _SAVED_MODULES_VALUES, _key, _value
|
||||
|
||||
|
||||
# ── Helper ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _make_chunk(start_ms: int, end_ms: int, phash: str, hist: list[float] | None = None) -> FingerprintChunk:
|
||||
"""创建测试用 FingerprintChunk."""
|
||||
return FingerprintChunk(
|
||||
start_time_ms=start_ms,
|
||||
end_time_ms=end_ms,
|
||||
phash_binary=phash,
|
||||
color_histogram=hist or [0.1] * 96,
|
||||
frame_count=1,
|
||||
)
|
||||
|
||||
|
||||
# ── TestDuplicateSegment ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestDuplicateSegment:
|
||||
"""DuplicateSegment 数据类测试."""
|
||||
|
||||
def test_creation(self):
|
||||
"""正常创建."""
|
||||
seg = DuplicateSegment(
|
||||
query_start_ms=1000,
|
||||
query_end_ms=5000,
|
||||
target_start_ms=2000,
|
||||
target_end_ms=6000,
|
||||
avg_distance=3.5,
|
||||
)
|
||||
assert seg.query_start_ms == 1000
|
||||
assert seg.avg_distance == 3.5
|
||||
|
||||
def test_fields(self):
|
||||
"""所有字段可访问."""
|
||||
seg = DuplicateSegment(0, 1000, 500, 1500, 2.0)
|
||||
assert seg.query_end_ms == 1000
|
||||
assert seg.target_start_ms == 500
|
||||
assert seg.target_end_ms == 1500
|
||||
|
||||
|
||||
# ── TestDetectKeyframeTimestamps ────────────────────────────────
|
||||
|
||||
|
||||
class TestDetectKeyframeTimestamps:
|
||||
"""detect_keyframe_timestamps 关键帧检测测试.
|
||||
|
||||
由于 cv2 在单元测试环境中是 mock,这里只测试边界条件。
|
||||
完整的视频处理测试在集成测试中进行。
|
||||
"""
|
||||
|
||||
def test_cannot_open_video_raises(self):
|
||||
"""无法打开视频时抛出 RuntimeError."""
|
||||
cv2_mock = _dedup_mod.cv2
|
||||
mock_cap = MagicMock()
|
||||
mock_cap.isOpened.return_value = False
|
||||
cv2_mock.VideoCapture.return_value = mock_cap
|
||||
|
||||
import pytest
|
||||
|
||||
with pytest.raises(RuntimeError, match="Cannot open video"):
|
||||
detect_keyframe_timestamps("/fake/path.mp4")
|
||||
|
||||
def test_zero_duration_returns_empty(self):
|
||||
"""视频时长为 0 时返回空列表."""
|
||||
cv2_mock = _dedup_mod.cv2
|
||||
mock_cap = MagicMock()
|
||||
mock_cap.isOpened.return_value = True
|
||||
# cv2.CAP_PROP_FPS etc. are Mock objects; configure get() to return 0 for frame_count
|
||||
mock_cap.get.return_value = 0
|
||||
mock_cap.read.return_value = (False, None)
|
||||
cv2_mock.VideoCapture.return_value = mock_cap
|
||||
|
||||
result = detect_keyframe_timestamps("/fake/zero.mp4")
|
||||
assert result == []
|
||||
|
||||
def test_function_signature(self):
|
||||
"""验证函数签名和默认参数."""
|
||||
import inspect
|
||||
|
||||
sig = inspect.signature(detect_keyframe_timestamps)
|
||||
params = sig.parameters
|
||||
assert "video_path" in params
|
||||
assert "min_interval_sec" in params
|
||||
assert "max_frames" in params
|
||||
assert "min_frames" in params
|
||||
# 默认值
|
||||
assert params["min_interval_sec"].default == 1.0
|
||||
assert params["max_frames"].default == 30
|
||||
assert params["min_frames"].default == 5
|
||||
|
||||
|
||||
# ── TestFindDuplicateSegments ───────────────────────────────────
|
||||
|
||||
|
||||
class TestFindDuplicateSegments:
|
||||
"""find_duplicate_segments 滑动窗口时序匹配测试."""
|
||||
|
||||
def test_identical_chunks_full_match(self):
|
||||
"""两组完全相同的 chunks → 整段匹配."""
|
||||
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, "aaaaaaaaaaaaaaaa") for i in range(10)]
|
||||
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, "aaaaaaaaaaaaaaaa") for i in range(10)]
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
assert len(segments) >= 1
|
||||
# 应该覆盖大部分范围
|
||||
total_query_range = segments[-1].query_end_ms - segments[0].query_start_ms
|
||||
assert total_query_range > 5000 # 至少覆盖 5 秒
|
||||
|
||||
def test_completely_different_chunks(self):
|
||||
"""两组完全不同的 chunks → 空列表."""
|
||||
# 距离都 > 阈值
|
||||
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, "0000000000000000") for i in range(10)]
|
||||
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, "ffffffffffffffff") for i in range(10)]
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
assert segments == []
|
||||
|
||||
def test_partial_overlap(self):
|
||||
"""部分重叠 → 只返回重叠段."""
|
||||
# 前 5 帧相同,后 5 帧不同
|
||||
same_hash = "aaaaaaaaaaaaaaaa"
|
||||
diff_hash_a = "0000000000000000"
|
||||
diff_hash_b = "ffffffffffffffff"
|
||||
|
||||
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(5)] + [
|
||||
_make_chunk(i * 1000, (i + 1) * 1000, diff_hash_a) for i in range(5, 10)
|
||||
]
|
||||
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(5)] + [
|
||||
_make_chunk(i * 1000, (i + 1) * 1000, diff_hash_b) for i in range(5, 10)
|
||||
]
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
# 应该只有前 5 帧的匹配段
|
||||
if segments:
|
||||
assert segments[0].query_end_ms <= 5000
|
||||
|
||||
def test_min_consecutive_not_met(self):
|
||||
"""连续 4 帧匹配(< min_consecutive=5)→ 不报重复.
|
||||
|
||||
注意:使用不同的 hash 对,确保后半部分帧距离 > 阈值。
|
||||
"""
|
||||
same_hash = "aaaaaaaaaaaaaaaa"
|
||||
# 4 帧匹配,后面 6 帧各自不同(在 query 和 target 中使用不同 hash)
|
||||
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
|
||||
_make_chunk(i * 1000, (i + 1) * 1000, "bbbbbbbbbbbbbbbb") for i in range(4, 10)
|
||||
]
|
||||
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
|
||||
_make_chunk(i * 1000, (i + 1) * 1000, "cccccccccccccccc") for i in range(4, 10)
|
||||
]
|
||||
|
||||
# hamming("bbbb...", "cccc...") should be > 8 (SEGMENT_MATCH_THRESHOLD)
|
||||
# b=1011, c=1100 → 4 bits differ per hex digit × 16 digits = 64 bits total? No...
|
||||
# Actually: hamming_distance("bbbbbbbbbbbbbbbb", "cccccccccccccccc")
|
||||
# b=0xb=1011, c=0xc=1100 → XOR=0111=0x7 → 3 bits per digit × 16 = 48
|
||||
# That's > 8 so won't match
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
# 只有 4 帧匹配(< min_consecutive=5),所以不报告
|
||||
assert segments == []
|
||||
|
||||
def test_max_gap_behavior(self):
|
||||
"""5 帧匹配 + 1 帧间隙 + 3 帧匹配 → 验证 max_gap 行为.
|
||||
|
||||
关键:间隙帧必须在 query 和 target 中使用不同 hash,使其真正不匹配。
|
||||
"""
|
||||
match_hash = "aaaaaaaaaaaaaaaa"
|
||||
gap_hash_a = "bbbbbbbbbbbbbbbb" # query 端
|
||||
gap_hash_b = "cccccccccccccccc" # target 端(与 query 端距离 > 8)
|
||||
tail_hash_a = "dddddddddddddddd"
|
||||
tail_hash_b = "eeeeeeeeeeeeeeee"
|
||||
|
||||
# 5 帧匹配, 1 帧间隙, 3 帧匹配, 5 帧不匹配
|
||||
hashes_a = [match_hash] * 5 + [gap_hash_a] + [match_hash] * 3 + [tail_hash_a] * 5
|
||||
hashes_b = [match_hash] * 5 + [gap_hash_b] + [match_hash] * 3 + [tail_hash_b] * 5
|
||||
|
||||
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, h) for i, h in enumerate(hashes_a)]
|
||||
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, h) for i, h in enumerate(hashes_b)]
|
||||
|
||||
# max_gap=2, 所以 1 帧间隙会被合并
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b, max_gap=2)
|
||||
# 5 match + 1 gap + 3 match = run of 9(间隙被桥接)
|
||||
assert len(segments) == 1
|
||||
# run 覆盖 indices 0-8(5 match + 1 gap + 3 match),但 gap 帧不计入 match
|
||||
# query_start = chunks_a[0].start = 0
|
||||
# query_end = chunks_a[8].end = 9000
|
||||
assert segments[0].query_start_ms == 0
|
||||
assert segments[0].query_end_ms == 9000
|
||||
|
||||
def test_max_gap_exceeded(self):
|
||||
"""间隙超过 max_gap → 分成两段."""
|
||||
match_hash = "aaaaaaaaaaaaaaaa"
|
||||
gap_hash_a = "bbbbbbbbbbbbbbbb"
|
||||
gap_hash_b = "cccccccccccccccc"
|
||||
tail_hash_a = "dddddddddddddddd"
|
||||
tail_hash_b = "eeeeeeeeeeeeeeee"
|
||||
|
||||
# 5 帧匹配, 3 帧间隙 (> max_gap=2), 5 帧匹配, 5 帧不匹配
|
||||
hashes_a = [match_hash] * 5 + [gap_hash_a] * 3 + [match_hash] * 5 + [tail_hash_a] * 5
|
||||
hashes_b = [match_hash] * 5 + [gap_hash_b] * 3 + [match_hash] * 5 + [tail_hash_b] * 5
|
||||
|
||||
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, h) for i, h in enumerate(hashes_a)]
|
||||
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, h) for i, h in enumerate(hashes_b)]
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b, max_gap=2)
|
||||
# 3 帧间隙 > max_gap=2 → 分成两段(每段 5 帧匹配)
|
||||
assert len(segments) == 2
|
||||
|
||||
def test_empty_chunks(self):
|
||||
"""空 chunks 返回空列表."""
|
||||
assert find_duplicate_segments([], [_make_chunk(0, 1000, "aa")]) == []
|
||||
assert find_duplicate_segments([_make_chunk(0, 1000, "aa")], []) == []
|
||||
assert find_duplicate_segments([], []) == []
|
||||
|
||||
def test_dict_chunks_compatibility(self):
|
||||
"""dict 格式的 chunks 也能正常工作."""
|
||||
chunks_a = [
|
||||
{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": i * 1000, "end_time_ms": (i + 1) * 1000}
|
||||
for i in range(10)
|
||||
]
|
||||
chunks_b = [
|
||||
{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": i * 1000, "end_time_ms": (i + 1) * 1000}
|
||||
for i in range(10)
|
||||
]
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
assert len(segments) >= 1
|
||||
|
||||
def test_segment_time_ranges(self):
|
||||
"""返回的 segment 时间范围正确.
|
||||
|
||||
每个 query chunk 匹配到 target 中对应的 chunk(相同 hash),
|
||||
确保 target 时间范围正确映射。
|
||||
"""
|
||||
|
||||
# 给每个 chunk 唯一的 hash(但保证 query[i] == target[i])
|
||||
def _unique_hash(i: int) -> str:
|
||||
return format(i, "016x")
|
||||
|
||||
chunks_a = [_make_chunk(i * 2000, (i + 1) * 2000, _unique_hash(i)) for i in range(7)]
|
||||
chunks_b = [_make_chunk(i * 2000, (i + 1) * 2000, _unique_hash(i)) for i in range(7)]
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
assert len(segments) >= 1
|
||||
seg = segments[0]
|
||||
assert seg.query_start_ms == 0
|
||||
assert seg.query_end_ms == 14000
|
||||
# target 应该映射到正确的范围
|
||||
assert seg.target_start_ms == 0
|
||||
assert seg.target_end_ms == 14000
|
||||
assert seg.avg_distance == 0.0 # 完全相同
|
||||
|
||||
|
||||
# ── TestMedianVsMean ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestMedianVsMean:
|
||||
"""中位数 vs 均值:验证中位数抵抗异常值."""
|
||||
|
||||
def test_median_resists_outlier(self):
|
||||
"""距离 [3,3,3,3,30]:均值=8.4,中位数=3.
|
||||
中位数 < PHASH_THRESHOLD(10),均值也 < 10。
|
||||
但更极端的:[3,3,3,3,60]:均值=14.4,中位数=3.
|
||||
"""
|
||||
import statistics
|
||||
|
||||
distances = [3, 3, 3, 3, 60]
|
||||
assert statistics.median(distances) == 3
|
||||
assert sum(distances) / len(distances) == 14.4
|
||||
# 中位数 < 10 → 通过阈值
|
||||
assert statistics.median(distances) < 10
|
||||
|
||||
|
||||
# ── TestMatchRatioCondition ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestMatchRatioCondition:
|
||||
"""帧匹配比例条件测试."""
|
||||
|
||||
def test_ratio_below_threshold_skips(self):
|
||||
"""10 帧中只有 5 帧距离 < 10 → match_ratio=0.5 < 0.7 → 跳过."""
|
||||
distances = [3, 5, 7, 8, 9, 15, 20, 25, 30, 40]
|
||||
threshold = 10
|
||||
matching = sum(1 for d in distances if d < threshold)
|
||||
ratio = matching / len(distances)
|
||||
assert ratio == 0.5
|
||||
assert ratio < 0.7 # 应该被跳过
|
||||
|
||||
def test_ratio_above_threshold_passes(self):
|
||||
"""10 帧中 8 帧距离 < 10 → match_ratio=0.8 >= 0.7 → 通过."""
|
||||
distances = [3, 5, 7, 8, 9, 3, 5, 7, 20, 30]
|
||||
threshold = 10
|
||||
matching = sum(1 for d in distances if d < threshold)
|
||||
ratio = matching / len(distances)
|
||||
assert ratio == 0.8
|
||||
assert ratio >= 0.7 # 应该通过
|
||||
|
||||
|
||||
# ── TestBhattacharyyaFusion ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestBhattacharyyaFusion:
|
||||
"""直方图融合逻辑测试."""
|
||||
|
||||
def test_high_phash_high_hist_is_duplicate(self):
|
||||
"""pHash 高相似 + 直方图高相似 → combined_score 高."""
|
||||
phash_similarity = 0.95 # median_distance ≈ 3
|
||||
hist_similarity = 0.90
|
||||
combined = 0.7 * phash_similarity + 0.3 * hist_similarity
|
||||
assert combined > 0.70 # DUPLICATE_THRESHOLD
|
||||
|
||||
def test_high_phash_low_hist_maybe_not(self):
|
||||
"""pHash 高相似 + 直方图低相似 → combined_score 取决于权重."""
|
||||
phash_similarity = 0.85 # median_distance ≈ 10
|
||||
hist_similarity = 0.10
|
||||
combined = 0.7 * phash_similarity + 0.3 * hist_similarity
|
||||
# 0.7 * 0.85 + 0.3 * 0.10 = 0.595 + 0.03 = 0.625 < 0.70
|
||||
assert combined < 0.70
|
||||
|
||||
def test_no_histogram_fallback(self):
|
||||
"""无直方图数据时 hist_similarity 回退到 0.5."""
|
||||
phash_similarity = 0.90
|
||||
hist_similarity = 0.5 # fallback
|
||||
combined = 0.7 * phash_similarity + 0.3 * hist_similarity
|
||||
# 0.7 * 0.90 + 0.3 * 0.5 = 0.63 + 0.15 = 0.78 > 0.70
|
||||
assert combined > 0.70
|
||||
|
||||
|
||||
# ── TestBackwardCompatibility ───────────────────────────────────
|
||||
|
||||
|
||||
class TestBackwardCompatibility:
|
||||
"""向后兼容测试."""
|
||||
|
||||
def test_no_chunks_no_crash(self):
|
||||
"""已有视频无分片数据 → find_duplicate_segments 返回空列表."""
|
||||
# 模拟:fingerprint 有 chunks,但 existing 只有 JSON phashes
|
||||
query_chunks = [_make_chunk(i * 1000, (i + 1) * 1000, "aaaaaaaaaaaaaaaa") for i in range(10)]
|
||||
# 没有 start_time_ms/end_time_ms 的简化 dict
|
||||
target_as_dicts = [{"phash_binary": "aaaaaaaaaaaaaaaa"} for _ in range(10)]
|
||||
|
||||
# find_duplicate_segments 需要 start_time_ms/end_time_ms
|
||||
# 在没有的情况下应该不崩溃(用默认值)
|
||||
# 实际上我们的实现用 _get_start/_get_end 访问,缺 key 会 KeyError
|
||||
# 所以 check_duplicate 传入时会补上默认值
|
||||
target_with_defaults = [
|
||||
{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": 0, "end_time_ms": 0} for _ in range(10)
|
||||
]
|
||||
segments = find_duplicate_segments(query_chunks, target_with_defaults)
|
||||
# 不会崩溃
|
||||
assert isinstance(segments, list)
|
||||
|
||||
def test_few_chunks_no_crash(self):
|
||||
"""少量 chunk 不崩溃."""
|
||||
chunks_a = [_make_chunk(0, 5000, "aaaaaaaaaaaaaaaa")]
|
||||
chunks_b = [{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": 0, "end_time_ms": 5000}]
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
# 1 帧 < min_consecutive=5,不会报重复
|
||||
assert segments == []
|
||||
|
||||
|
||||
# ── TestConstants ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestConstants:
|
||||
"""常量值验证 — 使用已在模块顶部导入的常量,避免重新 import."""
|
||||
|
||||
def test_segment_match_threshold(self):
|
||||
# 从已导入的 find_duplicate_segments 默认参数间接验证
|
||||
assert SEGMENT_MATCH_THRESHOLD == 8
|
||||
|
||||
def test_min_consecutive_matches(self):
|
||||
assert MIN_CONSECUTIVE_MATCHES == 5
|
||||
|
||||
def test_max_gap(self):
|
||||
assert MAX_GAP == 2
|
||||
|
||||
def test_scene_change_threshold(self):
|
||||
assert SCENE_CHANGE_THRESHOLD == 30
|
||||
|
||||
def test_min_keyframe_interval(self):
|
||||
assert MIN_KEYFRAME_INTERVAL_SEC == 1.0
|
||||
|
||||
def test_max_keyframes(self):
|
||||
assert MAX_KEYFRAMES == 30
|
||||
|
||||
def test_min_keyframes(self):
|
||||
assert MIN_KEYFRAMES == 5
|
||||
|
||||
def test_long_video_threshold(self):
|
||||
assert LONG_VIDEO_DURATION_THRESHOLD_SEC == 180
|
||||
|
||||
def test_duplicate_threshold(self):
|
||||
assert DUPLICATE_THRESHOLD == 0.70
|
||||
|
||||
def test_phash_weight(self):
|
||||
assert PHASH_WEIGHT == 0.7
|
||||
|
||||
def test_histogram_weight(self):
|
||||
assert HISTOGRAM_WEIGHT == 0.3
|
||||
|
||||
def test_match_ratio_threshold(self):
|
||||
assert MATCH_RATIO_THRESHOLD == 0.7
|
||||
@@ -122,7 +122,7 @@ class TestComputeDuplicateRate:
|
||||
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
|
||||
|
||||
# hamming distance = 2, similarity = (1 - 2/64) * 100 = 96.875
|
||||
assert rate == pytest.approx(96.88, abs=0.1)
|
||||
assert rate == pytest.approx(82.81, abs=0.1) # 新算法: 0.7*(1-2/64)*100 + 0.3*50
|
||||
|
||||
def test_excludes_self_video(self):
|
||||
from video_processing.dedup import VideoDeduplicator
|
||||
@@ -186,7 +186,7 @@ class TestComputeDuplicateRate:
|
||||
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
|
||||
|
||||
# max similarity: e2 distance=1, (1-1/64)*100 = 98.4375
|
||||
assert rate == pytest.approx(98.44, abs=0.1)
|
||||
assert rate == pytest.approx(83.91, abs=0.1) # 新算法: 0.7*(1-1/64)*100 + 0.3*50
|
||||
|
||||
def test_user_id_scope_cross_project(self):
|
||||
"""传 user_id 时应跨项目查询,而非仅当前项目."""
|
||||
|
||||
@@ -0,0 +1,282 @@
|
||||
"""分片指纹存储单元测试 — Issue #1657.
|
||||
|
||||
覆盖:
|
||||
- 分片策略:60秒视频 → 30片,120秒视频 → 24片
|
||||
- VideoFingerprint.to_chunk_models() 输出正确
|
||||
- _save_fingerprint_chunks 幂等性(已有数据跳过)
|
||||
- to_dict() 向后兼容
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
def _mock_module(**attrs):
|
||||
"""Create a mock module with __spec__ to avoid AttributeError."""
|
||||
m = MagicMock()
|
||||
m.__spec__ = None
|
||||
for k, v in attrs.items():
|
||||
setattr(m, k, v)
|
||||
return m
|
||||
|
||||
|
||||
# ── Module-level setup: mock deps, import dedup, then restore sys.modules ──
|
||||
_SAVED_MODULES_KEYS = set(sys.modules.keys())
|
||||
_SAVED_MODULES_VALUES = {
|
||||
k: sys.modules.get(k)
|
||||
for k in [
|
||||
"cv2",
|
||||
"celery",
|
||||
"sqlalchemy",
|
||||
"sqlalchemy.orm",
|
||||
"sqlalchemy.engine",
|
||||
"sqlalchemy.ext",
|
||||
"sqlalchemy.ext.declarative",
|
||||
"worker_app.db",
|
||||
"worker_app.celery_app",
|
||||
"worker_app.core.config",
|
||||
"packages.adapters.sqlalchemy_impl.session",
|
||||
"packages.adapters.sqlalchemy_impl.generated_video_repository",
|
||||
"packages.adapters.sqlalchemy_impl.models",
|
||||
"packages.shared.config",
|
||||
"packages.shared.storage",
|
||||
]
|
||||
}
|
||||
|
||||
# Set up mocks
|
||||
sys.modules["cv2"] = _mock_module()
|
||||
|
||||
_mock_celery = MagicMock()
|
||||
_mock_celery.Task = MagicMock
|
||||
_mock_celery.Celery = MagicMock
|
||||
_mock_celery.__spec__ = None
|
||||
sys.modules["celery"] = _mock_celery
|
||||
|
||||
_mock_sqla = MagicMock()
|
||||
_mock_sqla.__path__ = []
|
||||
_mock_sqla.__spec__ = None
|
||||
sys.modules["sqlalchemy"] = _mock_sqla
|
||||
|
||||
_mock_sqla_orm = MagicMock()
|
||||
_mock_sqla_orm.__path__ = []
|
||||
_mock_sqla_orm.__spec__ = None
|
||||
_mock_sqla_orm.Session = MagicMock
|
||||
sys.modules["sqlalchemy.orm"] = _mock_sqla_orm
|
||||
sys.modules["sqlalchemy.engine"] = _mock_module()
|
||||
sys.modules["sqlalchemy.ext"] = _mock_module()
|
||||
sys.modules["sqlalchemy.ext.declarative"] = _mock_module()
|
||||
|
||||
sys.modules["worker_app.db"] = _mock_module(SessionLocal=MagicMock())
|
||||
sys.modules["worker_app.celery_app"] = _mock_module(celery_app=MagicMock())
|
||||
sys.modules["worker_app.core.config"] = _mock_module(get_settings=MagicMock(return_value=MagicMock()))
|
||||
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.session"] = _mock_module(
|
||||
Base=MagicMock(),
|
||||
build_engine=MagicMock(),
|
||||
build_session_factory=MagicMock(),
|
||||
ensure_database_exists=MagicMock(),
|
||||
initialize_database=MagicMock(),
|
||||
)
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.generated_video_repository"] = _mock_module()
|
||||
|
||||
|
||||
# Mock VideoFingerprintChunkModel with class-level column attributes
|
||||
class _FakeChunkModel:
|
||||
video_id = MagicMock()
|
||||
project_id = MagicMock()
|
||||
user_id = MagicMock()
|
||||
start_time_ms = MagicMock()
|
||||
end_time_ms = MagicMock()
|
||||
phash_binary = MagicMock()
|
||||
color_histogram = MagicMock()
|
||||
frame_count = MagicMock()
|
||||
created_at = MagicMock()
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
for k, v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.models"] = _mock_module(
|
||||
VideoFingerprintChunkModel=_FakeChunkModel,
|
||||
)
|
||||
sys.modules["packages.shared.config"] = _mock_module(get_shared_settings=MagicMock(return_value=MagicMock()))
|
||||
sys.modules["packages.shared.storage"] = _mock_module()
|
||||
|
||||
# Import dedup while mocks are active
|
||||
from video_processing.dedup import ( # noqa: E402
|
||||
FingerprintChunk,
|
||||
VideoFingerprint,
|
||||
_save_fingerprint_chunks,
|
||||
)
|
||||
|
||||
# ── Restore sys.modules immediately after import ──
|
||||
for _key in list(sys.modules.keys()):
|
||||
if _key not in _SAVED_MODULES_KEYS:
|
||||
del sys.modules[_key]
|
||||
for _key, _value in _SAVED_MODULES_VALUES.items():
|
||||
if _value is not None:
|
||||
sys.modules[_key] = _value
|
||||
elif _key in sys.modules:
|
||||
del sys.modules[_key]
|
||||
del _SAVED_MODULES_KEYS, _SAVED_MODULES_VALUES, _key, _value
|
||||
|
||||
|
||||
class TestVideoFingerprintToChunkModels:
|
||||
"""测试 VideoFingerprint.to_chunk_models() 输出。"""
|
||||
|
||||
def test_to_chunk_models_output(self):
|
||||
"""to_chunk_models 返回正确的 Model 列表。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc123",
|
||||
keyframe_phashes=["a1b2", "c3d4"],
|
||||
color_histograms=[[0.1] * 96, [0.2] * 96],
|
||||
duration=10.0,
|
||||
resolution=(1920, 1080),
|
||||
chunks=[
|
||||
FingerprintChunk(start_time_ms=0, end_time_ms=2000, phash_binary="a1b2", color_histogram=[0.1] * 96),
|
||||
FingerprintChunk(start_time_ms=2000, end_time_ms=4000, phash_binary="c3d4", color_histogram=[0.2] * 96),
|
||||
],
|
||||
)
|
||||
|
||||
models = fp.to_chunk_models(video_id="v1", project_id="p1", user_id="u1")
|
||||
|
||||
assert len(models) == 2
|
||||
assert models[0].video_id == "v1"
|
||||
assert models[0].project_id == "p1"
|
||||
assert models[0].user_id == "u1"
|
||||
assert models[0].start_time_ms == 0
|
||||
assert models[0].end_time_ms == 2000
|
||||
assert models[0].phash_binary == "a1b2"
|
||||
assert models[1].start_time_ms == 2000
|
||||
assert models[1].end_time_ms == 4000
|
||||
assert models[1].phash_binary == "c3d4"
|
||||
|
||||
def test_to_chunk_models_empty_chunks(self):
|
||||
"""空 chunks 列表返回空 Model 列表。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=[],
|
||||
color_histograms=[],
|
||||
duration=0,
|
||||
resolution=(0, 0),
|
||||
chunks=[],
|
||||
)
|
||||
|
||||
models = fp.to_chunk_models(video_id="v1", project_id="p1")
|
||||
assert models == []
|
||||
|
||||
|
||||
class TestSaveFingerprintChunksIdempotent:
|
||||
"""测试 _save_fingerprint_chunks 幂等性。"""
|
||||
|
||||
def test_save_skips_existing(self):
|
||||
"""已有分片数据时跳过写入。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=["a1b2"],
|
||||
color_histograms=[[0.1] * 96],
|
||||
duration=5.0,
|
||||
resolution=(1920, 1080),
|
||||
chunks=[
|
||||
FingerprintChunk(start_time_ms=0, end_time_ms=2000, phash_binary="a1b2", color_histogram=[0.1] * 96),
|
||||
],
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
# Mock: 已有 1 条分片数据
|
||||
session.query.return_value.filter.return_value.count.return_value = 1
|
||||
|
||||
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
|
||||
|
||||
# bulk_save_objects 不应被调用
|
||||
session.bulk_save_objects.assert_not_called()
|
||||
|
||||
def test_save_writes_new(self):
|
||||
"""无分片数据时写入。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=["a1b2"],
|
||||
color_histograms=[[0.1] * 96],
|
||||
duration=5.0,
|
||||
resolution=(1920, 1080),
|
||||
chunks=[
|
||||
FingerprintChunk(start_time_ms=0, end_time_ms=2000, phash_binary="a1b2", color_histogram=[0.1] * 96),
|
||||
],
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
# Mock: 无分片数据
|
||||
session.query.return_value.filter.return_value.count.return_value = 0
|
||||
|
||||
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
|
||||
|
||||
# bulk_save_objects 应被调用一次
|
||||
session.bulk_save_objects.assert_called_once()
|
||||
saved_models = session.bulk_save_objects.call_args[0][0]
|
||||
assert len(saved_models) == 1
|
||||
assert saved_models[0].video_id == "v1"
|
||||
assert saved_models[0].phash_binary == "a1b2"
|
||||
|
||||
def test_save_skips_no_chunks(self):
|
||||
"""指纹无 chunks 时跳过。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=[],
|
||||
color_histograms=[],
|
||||
duration=0,
|
||||
resolution=(0, 0),
|
||||
chunks=[],
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
session.query.return_value.filter.return_value.count.return_value = 0
|
||||
|
||||
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
|
||||
|
||||
# bulk_save_objects 不应被调用
|
||||
session.bulk_save_objects.assert_not_called()
|
||||
|
||||
|
||||
class TestFingerprintToDictBackwardCompat:
|
||||
"""测试 to_dict() 向后兼容性。"""
|
||||
|
||||
def test_to_dict_includes_chunks(self):
|
||||
"""to_dict() 包含 chunks 字段。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc123",
|
||||
keyframe_phashes=["a1b2"],
|
||||
color_histograms=[[0.1] * 96],
|
||||
duration=5.0,
|
||||
resolution=(1920, 1080),
|
||||
chunks=[
|
||||
FingerprintChunk(start_time_ms=0, end_time_ms=2000, phash_binary="a1b2", color_histogram=[0.1] * 96),
|
||||
],
|
||||
)
|
||||
|
||||
d = fp.to_dict()
|
||||
|
||||
assert "chunks" in d
|
||||
assert len(d["chunks"]) == 1
|
||||
assert d["chunks"][0]["start_time_ms"] == 0
|
||||
assert d["chunks"][0]["end_time_ms"] == 2000
|
||||
assert d["chunks"][0]["phash_binary"] == "a1b2"
|
||||
|
||||
def test_to_dict_preserves_legacy_fields(self):
|
||||
"""to_dict() 保留 keyframe_phashes 和 color_histograms 字段。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=["a1b2", "c3d4"],
|
||||
color_histograms=[[0.1] * 96, [0.2] * 96],
|
||||
duration=10.0,
|
||||
resolution=(1920, 1080),
|
||||
)
|
||||
|
||||
d = fp.to_dict()
|
||||
|
||||
assert "keyframe_phashes" in d
|
||||
assert "color_histograms" in d
|
||||
assert len(d["keyframe_phashes"]) == 2
|
||||
assert len(d["color_histograms"]) == 2
|
||||
Reference in New Issue
Block a user