Files
xiaoxia-saas/tests/unit/test_dedup_v2.py
T
CI Bot a0d4f6e111
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 22s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 10s
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Successful in 1m20s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m27s
CI/CD Pipeline / Validate - Style (pull_request) Failing after 1m33s
AI Code Review / AI Code Review (pull_request) Failing after 2m11s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m19s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m43s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 4m35s
CI/CD Pipeline / Validate - Security (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
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 2m29s
style: auto-format with black + isort + prettier [skip ci-format-check]
2026-09-03 15:04:14 +00:00

533 lines
20 KiB
Python
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.
"""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-85 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