Files
xiaoxia-saas/apps/worker/video_processing/dedup.py
T
xiaoxia a40de0338e
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 15s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2m54s
AI Code Review / AI Code Review (pull_request) Failing after 4m9s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 3m6s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 3m39s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 3m29s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 2m46s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 7m14s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 3m33s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 6m42s
CI/CD Pipeline / Validate - Code Quality (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Failing after 384h2m26s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 384h2m30s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 384h2m34s
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Failing after 384h3m35s
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Failing after 384h3m37s
CI/CD Pipeline / PR Build Web Image (pull_request) Failing after 384h4m7s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 384h4m19s
CI/CD Pipeline / Frontend Lint (pull_request) Failing after 384h4m21s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 384h5m50s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 384h5m54s
CI/CD Pipeline / Check push changed paths (pull_request) Failing after 384h9m16s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 384h36m39s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Failing after 384h37m50s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 384h40m3s
feat: 成片 duplicate_rate 百分比 + 素材高频使用自动排除
任务一:生成视频补 duplicate_rate 百分比
- GeneratedVideoModel 加 duplicate_rate 列 (Float, nullable)
- GeneratedVideo 领域实体 + repository 映射同步更新
- Alembic migration 059
- VideoDeduplicator.compute_duplicate_rate:遍历项目内已有视频,
  MD5精确=100%,pHash用 (1-avg_distance/64)*100,取最高值
- dedup_helpers 在指纹计算后调用并存入 duplicate_rate
- VideoItemResponse schema 加 duplicate_rate 字段
- _to_video_response 透传 duplicate_rate

任务二:素材级自动排除高频使用素材
- asset_segment_tracker.get_asset_recent_use_counts:
  统计每个素材在最近 N 个不同 plan_id 中的使用次数
- smart_match_assets 路由加高频排除层(阈值 3 次 / 最近 5 个视频)
- 排除后不够 limit 时自动放宽,宁多不少
- 异常时降级跳过排除,不影响正常选素材

测试:12 个新用例覆盖 duplicate_rate 计算 + API 响应 + 高频统计
2026-08-31 15:14:34 +08:00

420 lines
15 KiB
Python
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Video deduplication module - compute fingerprints and detect duplicates."""
import hashlib
import logging
import os
import tempfile
from dataclasses import dataclass
from typing import Optional
import cv2
import numpy as np
from celery import Task
from sqlalchemy.orm import Session
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.shared.storage import get_storage_service
logger = logging.getLogger(__name__)
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
"""计算图像的感知哈希(pHash),基于 DCT(离散余弦变换)。
算法步骤:
1. 将图像缩放到 hash_size*4 × hash_size*4(默认 32×32)
2. 转为灰度图,应用 2D DCT 提取频率分量
3. 取左上角 hash_size×hash_size 的低频分量(默认 8×8 = 64 bit)
4. 排除 DC 分量([0,0] 位置),计算中位数
5. 每个分量与中位数比较,生成二值 hash
Args:
image: BGR 格式的 numpy 图像数组
hash_size: 哈希边长,默认 8(生成 64-bit hash)
Returns:
十六进制字符串表示的感知哈希
"""
# Resize to 32x32 for DCT
resized = cv2.resize(image, (hash_size * 4, hash_size * 4))
gray = cv2.cvtColor(resized, cv2.COLOR_BGR2GRAY).astype(np.float32)
# Apply 2D DCT
dct = cv2.dct(gray)
# Take top-left 8x8 low-frequency components
dct_low = dct[:hash_size, :hash_size]
# Compute median (excluding DC component at [0,0])
dct_low[0, 0] = 0
median = np.median(dct_low)
# Generate hash based on comparison with median
diff = (dct_low > median).astype(int)
hash_str = "".join(str(b) for row in diff for b in row)
return hex(int(hash_str, 2))[2:]
def hamming_distance(hash1: str, hash2: str) -> int:
"""计算两个十六进制哈希之间的汉明距离(不同 bit 位数)。
使用 XOR 异或 + bit 计数:bin(h1 ^ h2).count("1")。
例如:hamming_distance("00", "ff") = 8(8 个 bit 全不同)。
Args:
hash1: 十六进制字符串
hash2: 十六进制字符串
Returns:
不同 bit 的数量
"""
h1, h2 = int(hash1, 16), int(hash2, 16)
return bin(h1 ^ h2).count("1")
def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]:
"""Compute color histogram for an image."""
hist = []
for i in range(3):
h = cv2.calcHist([image], [i], None, [bins], [0, 256])
h = cv2.normalize(h, h).flatten()
hist.extend(h)
return hist
@dataclass
class VideoFingerprint:
"""Video fingerprint containing multiple similarity metrics."""
md5: str
keyframe_phashes: list[str]
color_histograms: list[list[float]]
duration: float
resolution: tuple[int, int]
def to_dict(self) -> dict:
# 注意:color_histograms 里的值可能是 np.float32(来自 cv2.normalize),
# 直接存进 dict 后 SQLAlchemy JSON 序列化会报 "float32 is not JSON serializable"。
# 这里统一转成 Python 原生 float。
native_histograms = [[float(v) for v in hist] for hist in self.color_histograms]
return {
"md5": self.md5,
"keyframe_phashes": self.keyframe_phashes,
"color_histograms": native_histograms,
"duration": float(self.duration),
"resolution": [int(self.resolution[0]), int(self.resolution[1])],
}
class VideoDeduplicator:
"""Video deduplication using multiple fingerprint methods."""
PHASH_THRESHOLD = 10
HISTOGRAM_THRESHOLD = 0.85
def compute_fingerprint(self, video_path: str) -> VideoFingerprint:
"""Compute video fingerprint using MD5, pHash, and color histogram."""
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
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 = []
frame_interval = max(1, frame_count // 10)
for i in range(0, frame_count, frame_interval):
cap.set(cv2.CAP_PROP_POS_FRAMES, i)
ret, frame = cap.read()
if not ret:
continue
_, buffer = cv2.imencode(".jpg", frame)
md5_hash.update(buffer)
keyframe_phashes.append(compute_phash(frame))
color_histograms.append(compute_color_histogram(frame))
cap.release()
return VideoFingerprint(
md5=md5_hash.hexdigest(),
keyframe_phashes=keyframe_phashes,
color_histograms=color_histograms,
duration=duration,
resolution=(width, height),
)
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)
注意:返回第一个通过阈值的匹配(非最优匹配)。
Args:
fingerprint: 待检测视频的指纹
project_id: 项目 ID,仅在同一项目内搜索
session: 数据库会话
Returns:
重复信息字典(含 duplicate, duplicate_of, reason, similarity),
或 None 表示未找到重复。
"""
video_repo = SQLAlchemyGeneratedVideoRepository(session)
existing_videos = video_repo.list_by_project(project_id)
for existing in existing_videos:
if not existing.video_fingerprint:
continue
ef = existing.video_fingerprint
# 精确匹配:MD5 完全一致
if fingerprint.md5 == ef.get("md5"):
return {"duplicate": True, "duplicate_of": existing.id, "reason": "exact_md5_match", "similarity": 1.0}
# 感知哈希相似度
existing_phashes = ef.get("keyframe_phashes", [])
if 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:
continue
phash_similarity = 1.0 - (avg_distance / 64)
return {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "phash_similar",
"similarity": phash_similarity,
}
return None
def check_batch_duplicate(
self,
fingerprint: VideoFingerprint,
batch_id: str,
current_video_id: str,
session: Session,
) -> Optional[dict]:
"""检查视频是否与同批次内其他视频重复。
逻辑与 check_duplicate 一致(MD5 + pHash),但搜索范围限定为同 batch_id 的视频。
Args:
fingerprint: 待检测视频的指纹
batch_id: 批次 ID
current_video_id: 当前视频 ID(排除自身)
session: 数据库会话
Returns:
重复信息字典,或 None 表示未找到重复
"""
video_repo = SQLAlchemyGeneratedVideoRepository(session)
batch_videos = video_repo.list_by_batch(batch_id)
for existing in batch_videos:
if existing.id == current_video_id:
continue
if not existing.video_fingerprint:
continue
ef = existing.video_fingerprint
if fingerprint.md5 == ef.get("md5"):
return {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "batch_exact_md5_match",
"similarity": 1.0,
}
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:
continue
phash_similarity = 1.0 - (avg_distance / 64)
return {
"duplicate": True,
"duplicate_of": existing.id,
"reason": "batch_phash_similar",
"similarity": phash_similarity,
}
return None
@staticmethod
def _average_histogram_similarity(histograms_a: list[list[float]], histograms_b: list[list[float]]) -> float:
"""
计算两组颜色直方图之间的平均余弦相似度。
对每组直方图对取最小长度对齐,计算余弦相似度后取平均。
Args:
histograms_a: 第一组直方图(每帧一个 list)
histograms_b: 第二组直方图
Returns:
平均余弦相似度,范围 [0, 1]
"""
if not histograms_a or not histograms_b:
return 0.0
similarities = []
for ha in histograms_a:
best = 0.0
vec_a = np.array(ha, dtype=np.float64)
norm_a = np.linalg.norm(vec_a)
if norm_a == 0:
continue
for hb in histograms_b:
vec_b = np.array(hb, dtype=np.float64)
# 对齐长度
min_len = min(len(vec_a), len(vec_b))
va, vb = vec_a[:min_len], vec_b[:min_len]
norm_b = np.linalg.norm(vb)
if norm_b == 0:
continue
sim = float(np.dot(va, vb) / (norm_a * norm_b))
best = max(best, sim)
similarities.append(best)
return sum(similarities) / len(similarities) if similarities else 0.0
def compute_duplicate_rate(
self,
fingerprint: VideoFingerprint,
project_id: str,
current_video_id: str | None,
session: Session,
) -> float:
"""计算当前视频与项目内已有视频的最高相似度百分比。
遍历项目内所有其他有指纹的视频,对每个计算相似度:
- MD5 精确匹配 → 100%
- pHash 相似度 → (1.0 - avg_distance / 64) * 100
取最高值作为 duplicate_rate(0~100)。
如果没有其他视频可比较,返回 0.0。
Args:
fingerprint: 当前视频的指纹
project_id: 项目 ID
current_video_id: 当前视频 ID(排除自身,可为 None)
session: 数据库会话
Returns:
duplicate_rate: 0~100 的浮点数
"""
video_repo = SQLAlchemyGeneratedVideoRepository(session)
existing_videos = video_repo.list_by_project(project_id)
max_similarity = 0.0
for existing in existing_videos:
if current_video_id and existing.id == current_video_id:
continue
if not existing.video_fingerprint:
continue
ef = existing.video_fingerprint
# MD5 精确匹配 → 100%
if fingerprint.md5 == ef.get("md5"):
return 100.0
# pHash 相似度
existing_phashes = ef.get("keyframe_phashes", [])
if not existing_phashes or not fingerprint.keyframe_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 64
similarity = (1.0 - avg_distance / 64) * 100
max_similarity = max(max_similarity, similarity)
return round(max(max_similarity, 0.0), 2)
@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."""
session = SessionLocal()
temp_dir = tempfile.mkdtemp()
try:
video_repo = SQLAlchemyGeneratedVideoRepository(session)
storage_service = get_storage_service()
deduplicator = VideoDeduplicator()
video = video_repo.get(generated_video_id)
if video is None:
raise ValueError(f"Generated video {generated_video_id} not found")
local_path = os.path.join(temp_dir, f"{generated_video_id}.mp4")
storage_service.download_file(
f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path
)
fingerprint = deduplicator.compute_fingerprint(local_path)
duplicate_result = deduplicator.check_duplicate(fingerprint, video.project_id, session)
video.video_fingerprint = fingerprint.to_dict()
if duplicate_result:
video.is_duplicate = True
video.duplicate_of = duplicate_result["duplicate_of"]
else:
video.is_duplicate = False
video.duplicate_of = None
video_repo.update(video)
session.commit()
logger.info(f"Duplicate check completed for video {generated_video_id}: is_duplicate={video.is_duplicate}")
return {
"ok": True,
"video_id": generated_video_id,
"is_duplicate": video.is_duplicate,
"duplicate_of": video.duplicate_of,
"fingerprint": fingerprint.to_dict(),
}
except Exception as e:
logger.error(f"Duplicate check failed for {generated_video_id}: {str(e)}")
session.rollback()
raise self.retry(exc=e, countdown=60) from e
finally:
session.close()
import shutil
shutil.rmtree(temp_dir, ignore_errors=True)