Compare commits

..

1 Commits

Author SHA1 Message Date
xiaoxia d0c368df41 fix: correct extract-voice API path from /tts to /voices prefix
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 3s
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 / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Validate - Style (pull_request) Successful in 1m56s
AI Code Review / AI Code Review (pull_request) Successful in 2m0s
CI/CD Pipeline / PR Build API Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m4s
CI/CD Pipeline / PR Build Worker 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
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m20s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 3m4s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m6s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Successful in 1m31s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 5m9s
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
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 2m20s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Has been cancelled
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 1m47s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 28m27s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 5s
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
Backend route is mounted under /voices, not /tts
2026-09-03 20:40:48 +08:00
32 changed files with 395 additions and 3029 deletions
@@ -1,46 +0,0 @@
"""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")
@@ -1,25 +0,0 @@
"""add match_count and visual_similarity to generated_videos
Revision ID: 064_match_count_visual_sim
Revises: 063_fingerprint_chunks
Create Date: 2026-09-03
"""
import sqlalchemy as sa
from alembic import op
revision = "064_match_count_visual_sim"
down_revision = "063_fingerprint_chunks"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column("generated_videos", sa.Column("match_count", sa.Integer(), nullable=True, server_default="0"))
op.add_column("generated_videos", sa.Column("visual_similarity", sa.Float(), nullable=True, server_default="0.0"))
def downgrade() -> None:
op.drop_column("generated_videos", "visual_similarity")
op.drop_column("generated_videos", "match_count")
@@ -131,7 +131,6 @@ class PlanGeneratorService:
editing_mode,
random_selection=random_preview,
asset_durations=asset_durations,
user_id=created_by_user_id,
)
# 5. 持久化所有 clips 并计算总时长
@@ -219,7 +218,6 @@ class PlanGeneratorService:
*,
random_selection: bool = False,
asset_durations: dict[str, float] | None = None,
user_id: str = "",
) -> None:
"""按 editing_mode 将素材分配到 clips(就地修改,未持久化).
@@ -241,14 +239,6 @@ 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,
@@ -256,7 +246,6 @@ 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]]:
@@ -1,174 +0,0 @@
#!/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()
-4
View File
@@ -20,10 +20,6 @@ export interface DuplicationRecord {
duplicate_rate?: number
/** 重复片段数 */
duplicate_count?: number
/** 视觉相似度(0-100),#1660 新增 */
visual_similarity?: number
/** 匹配帧数,#1660 新增 */
match_count?: number
/** 创建时间 */
created_at: string
/** 更新时间 */
-8
View File
@@ -23,10 +23,6 @@ export interface ProductItem {
project_name?: string
/** 查重率(百分比) */
duplicate_rate?: number
/** 视觉相似度(0-100),#1660 新增 */
visual_similarity?: number
/** 匹配帧数,#1660 新增 */
match_count?: number
created_at?: string
updated_at?: string
}
@@ -76,8 +72,4 @@ export interface VideoItem {
download_url: string
generated_at: string
duplicate_rate?: number
/** 视觉相似度(0-100),#1660 新增 */
visual_similarity?: number
/** 匹配帧数,#1660 新增 */
match_count?: number
}
-2
View File
@@ -30,7 +30,5 @@ export function mapVideoToProductItem(video: VideoItem): ProductItem {
created_at: video.generated_at,
updated_at: video.generated_at,
duplicate_rate: video.duplicate_rate,
visual_similarity: video.visual_similarity,
match_count: video.match_count,
}
}
@@ -98,7 +98,7 @@ const DuplicationDetail: React.FC = () => {
<div className="dup-detail-grid">
<RiskCard riskLevel={riskLevel} similarityPercent={similarityPercent} />
<InfoCard detail={detail} />
<SegmentsSection segments={detail.segments} totalDuration={detail.duration_seconds} />
<SegmentsSection segments={detail.segments} />
</div>
</div>
)
@@ -1,7 +1,7 @@
import React from "react"
import { Button, Tag, Tooltip } from "@/components/ui"
import type { DuplicationRecord } from "@/api/duplication"
import { STATUS_CONFIG, RISK_TAG_VARIANT, RISK_LABELS } from "../constants"
import { STATUS_CONFIG } from "../constants"
import { getRiskLevel, formatSize, formatDuration } from "../utils"
interface ResultCardProps {
@@ -54,9 +54,6 @@ const ResultCard: React.FC<ResultCardProps> = ({ record, onView, onDelete, onRet
/>
</div>
<span className={`dup-score-value ${riskLevel}`}>{rateValue.toFixed(1)}%</span>
<Tag variant={RISK_TAG_VARIANT[riskLevel]} className="dup-score-risk-tag">
{RISK_LABELS[riskLevel]}
</Tag>
</>
) : record.status === "failed" ? (
<Tooltip title="重新查重">
@@ -2,81 +2,34 @@ import React from "react"
import { Tag } from "@/components/ui"
import type { DuplicateSegment } from "@/api/duplication"
import { SegmentCard } from "./SegmentCard"
import { formatTime } from "../utils"
interface SegmentsSectionProps {
segments?: DuplicateSegment[]
/** 视频总时长(秒),用于渲染时间轴 */
totalDuration?: number
}
/** 片段相似度 → 风险等级(时间轴配色用) */
const getSegmentRisk = (similarity: number): "low" | "medium" | "high" => {
if (similarity >= 90) return "high"
if (similarity >= 70) return "medium"
return "low"
}
/**
* 重复片段列表区域(含时间轴可视化)
* 重复片段列表区域
*/
export const SegmentsSection: React.FC<SegmentsSectionProps> = ({
segments = [],
totalDuration,
}) => {
const showTimeline = segments.length > 0 && totalDuration !== undefined && totalDuration > 0
export const SegmentsSection: React.FC<SegmentsSectionProps> = ({ segments = [] }) => (
<div className="dup-checks-section">
<h3>
🔍
<Tag variant="primary" style={{ marginLeft: 8 }}>
{segments.length}
</Tag>
</h3>
return (
<div className="dup-checks-section">
<h3>
🔍
<Tag variant="primary" style={{ marginLeft: 8 }}>
{segments.length}
</Tag>
</h3>
{showTimeline && (
<div className="dup-timeline">
<div className="dup-timeline-bar">
{segments.map((seg, i) => {
const left = (seg.source_start / totalDuration) * 100
const width = Math.max(
((seg.source_end - seg.source_start) / totalDuration) * 100,
0.5,
)
const segRisk = getSegmentRisk(seg.similarity)
return (
<div
key={seg.id ?? i}
className={`dup-timeline-segment ${segRisk}`}
style={{
left: `${Math.min(left, 100)}%`,
width: `${Math.min(width, 100 - Math.min(left, 100))}%`,
}}
title={`${formatTime(seg.source_start)} - ${formatTime(seg.source_end)} · 相似度 ${seg.similarity.toFixed(0)}% · ${seg.matched_video_name}`}
/>
)
})}
</div>
<div className="dup-timeline-labels">
<span>0s</span>
<span>{formatTime(totalDuration ?? 0)}</span>
</div>
</div>
)}
{segments.length > 0 ? (
<div className="dup-checks-list">
{segments.map((segment, index) => (
<SegmentCard key={segment.id} segment={segment} index={index} />
))}
</div>
) : (
<div className="dup-results-empty" style={{ padding: "32px 0" }}>
<div className="dup-results-empty-icon">🎉</div>
<p></p>
</div>
)}
</div>
)
}
{segments.length > 0 ? (
<div className="dup-checks-list">
{segments.map((segment, index) => (
<SegmentCard key={segment.id} segment={segment} index={index} />
))}
</div>
) : (
<div className="dup-results-empty" style={{ padding: "32px 0" }}>
<div className="dup-results-empty-icon">🎉</div>
<p></p>
</div>
)}
</div>
)
@@ -831,61 +831,3 @@
font-size: 16px;
}
}
/* ============================================================
查重率风险标签(列表卡片)
============================================================ */
.dup-score-risk-tag {
flex-shrink: 0;
margin-left: 2px;
}
/* ============================================================
重复片段时间轴可视化(#1662)
============================================================ */
.dup-timeline {
margin: 16px 0;
padding: 0 8px;
}
.dup-timeline-bar {
position: relative;
height: 24px;
background: var(--bg-secondary, #f1f5f9);
border-radius: 4px;
overflow: hidden;
}
.dup-timeline-segment {
position: absolute;
top: 2px;
height: 20px;
border-radius: 3px;
opacity: 0.8;
cursor: pointer;
transition: opacity 0.2s;
}
.dup-timeline-segment:hover {
opacity: 1;
}
.dup-timeline-segment.low {
background: #22c55e;
}
.dup-timeline-segment.medium {
background: #f59e0b;
}
.dup-timeline-segment.high {
background: #ef4444;
}
.dup-timeline-labels {
display: flex;
justify-content: space-between;
font-size: 12px;
color: var(--text-secondary);
margin-top: 4px;
}
+3 -3
View File
@@ -1,9 +1,9 @@
/** 根据查重率获取风险等级 */
export const getRiskLevel = (rate?: number): "low" | "medium" | "high" => {
if (rate === undefined) return "low"
if (rate < 15) return "low" // <15% 绿色(安全)
if (rate <= 30) return "medium" // 15-30% 黄色(注意)
return "high" // >30% 红色(危险)
if (rate <= 10) return "low"
if (rate <= 30) return "medium"
return "high"
}
/** 格式化时间(秒 → mm:ss */
@@ -2,7 +2,6 @@ import React from "react"
import type { ProductItem } from "../../../api/products"
import { STATUS_MAP } from "../constants"
import { formatDuration, formatFileSize, formatDate } from "../detailUtils"
import { getRiskLevel } from "../../duplication/utils"
interface ProductInfoPanelProps {
product: ProductItem
@@ -45,26 +44,12 @@ export const ProductInfoPanel: React.FC<ProductInfoPanelProps> = ({ product }) =
</div>
<div className="xx-detail-meta-item">
<span className="xx-detail-meta-label"></span>
<span
className={`xx-detail-meta-value dup-risk-text dup-risk-${getRiskLevel(product.duplicate_rate)}`}
>
<span className="xx-detail-meta-value">
{(product.duplicate_rate ?? 0) > 0
? `${(product.duplicate_rate ?? 0).toFixed(1)}%`
: "-"}
</span>
</div>
{product.visual_similarity != null && (
<div className="xx-detail-meta-item">
<span className="xx-detail-meta-label"></span>
<span className="xx-detail-meta-value">{product.visual_similarity.toFixed(1)}%</span>
</div>
)}
{product.match_count != null && (
<div className="xx-detail-meta-item">
<span className="xx-detail-meta-label"></span>
<span className="xx-detail-meta-value">{product.match_count}</span>
</div>
)}
<div className="xx-detail-meta-item">
<span className="xx-detail-meta-label"></span>
<span className="xx-detail-meta-value">{formatDate(product.created_at ?? "")}</span>
-16
View File
@@ -1076,19 +1076,3 @@
gap: var(--space-sm);
}
}
/* 查重率风险颜色(#1662) */
.xx-detail-meta-value.dup-risk-low {
color: var(--success-color, #22c55e);
font-weight: 600;
}
.xx-detail-meta-value.dup-risk-medium {
color: var(--warning-color, #f59e0b);
font-weight: 600;
}
.xx-detail-meta-value.dup-risk-high {
color: var(--error-color, #ef4444);
font-weight: 600;
}
@@ -136,9 +136,8 @@ export const MaterialVoiceTab: React.FC<MaterialVoiceTabProps> = ({
const material = mapAssetToMaterial(asset)
// duration 优先取顶层(后端从 metadata 提取),兜底 metadata
const cardDuration = asset.duration || material.duration || 0
// AI 生成素材标识:兼容旧素材(无 source 字段但有 tts_job_id
const meta = asset.metadata as Record<string, unknown>
const isAiMaterial = meta?.source === "tts_job" || !!meta?.tts_job_id
// AI 生成素材标识metadata.source === "tts_job"
const isAiMaterial = (asset.metadata as Record<string, unknown>)?.source === "tts_job"
const isPlaying = playingId === asset.id
const isSelected = selectedIds.has(asset.id)
// 播放中以 audio 真实时长为准,未播放显示卡片时长
@@ -185,10 +184,8 @@ export const MaterialVoiceTab: React.FC<MaterialVoiceTabProps> = ({
</div>
<div className="xx-voice-info vmat-info">
<div className="xx-voice-name-row">
<div className="xx-voice-name" title={asset.name}>
{asset.name}
</div>
<div className="xx-voice-name" title={asset.name}>
{asset.name}
{isAiMaterial && <span className="vmat-ai-badge">AI</span>}
</div>
<div className="xx-voice-subtitle">
@@ -1,29 +0,0 @@
import { describe, it, expect } from "vitest"
import { getRiskLevel } from "@/pages/duplication/utils"
describe("getRiskLevel (#1662 阈值 <15 / 15-30 / >30)", () => {
it("undefined 返回 low(兼容无数据)", () => {
expect(getRiskLevel(undefined)).toBe("low")
})
it("<15% 为低风险", () => {
expect(getRiskLevel(0)).toBe("low")
expect(getRiskLevel(10)).toBe("low")
expect(getRiskLevel(14.9)).toBe("low")
})
it("15% 边界为中风险", () => {
expect(getRiskLevel(15)).toBe("medium")
})
it("15-30% 为中风险", () => {
expect(getRiskLevel(20)).toBe("medium")
expect(getRiskLevel(30)).toBe("medium")
})
it(">30% 为高风险", () => {
expect(getRiskLevel(30.1)).toBe("high")
expect(getRiskLevel(80)).toBe("high")
expect(getRiskLevel(100)).toBe("high")
})
})
File diff suppressed because it is too large Load Diff
+6 -31
View File
@@ -100,24 +100,8 @@ 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) 历史成片查重(跨项目全局 + 时长预过滤)
duration_sec = fingerprint.duration / 1000 if fingerprint.duration else 0
duplicate_result = deduplicator.check_duplicate(
fingerprint,
project_id,
session,
scope="user",
user_id=user_id,
duration_sec=duration_sec,
)
# (a) 历史成片查重
duplicate_result = deduplicator.check_duplicate(fingerprint, project_id, session)
# (b) 批次内查重(仅当有 batch_id 时)
if not duplicate_result and batch_id:
@@ -137,26 +121,17 @@ def create_video_record_and_dedup(
generated_video.is_duplicate = False
generated_video.duplicate_of = None
# 计算重复率百分比(项目全局
# 计算重复率百分比(项目内所有已有视频对比取最高相似度
try:
rate_result = deduplicator.compute_duplicate_rate(
dup_rate = deduplicator.compute_duplicate_rate(
fingerprint,
project_id,
video_id,
session,
scope="user",
user_id=user_id,
)
generated_video.duplicate_rate = rate_result["duplicate_rate"]
generated_video.match_count = rate_result["match_count"]
generated_video.visual_similarity = rate_result["visual_similarity"]
logger.info(
"Duplicate rate for %s: %.2f%% (visual_sim=%.3f, matches=%d)",
video_id,
rate_result["duplicate_rate"],
rate_result["visual_similarity"],
rate_result["match_count"],
)
generated_video.duplicate_rate = dup_rate
logger.info("Duplicate rate for %s: %.2f%%", video_id, dup_rate)
except Exception as rate_err:
logger.warning("Failed to compute duplicate_rate for %s: %s", video_id, rate_err)
generated_video.duplicate_rate = None
@@ -131,65 +131,3 @@ 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
@@ -31,8 +31,6 @@ class SQLAlchemyGeneratedVideoRepository:
is_duplicate=video.is_duplicate,
duplicate_of=video.duplicate_of,
duplicate_rate=video.duplicate_rate,
match_count=getattr(video, "match_count", 0),
visual_similarity=getattr(video, "visual_similarity", 0.0),
generated_at=video.generated_at,
created_at=video.created_at,
)
@@ -64,8 +62,6 @@ class SQLAlchemyGeneratedVideoRepository:
is_duplicate=getattr(model, "is_duplicate", False),
duplicate_of=getattr(model, "duplicate_of", None),
duplicate_rate=getattr(model, "duplicate_rate", None),
match_count=getattr(model, "match_count", 0) or 0,
visual_similarity=getattr(model, "visual_similarity", 0.0) or 0.0,
generated_at=model.generated_at,
created_at=model.created_at,
)
@@ -81,8 +77,6 @@ class SQLAlchemyGeneratedVideoRepository:
model.is_duplicate = video.is_duplicate
model.duplicate_of = video.duplicate_of
model.duplicate_rate = video.duplicate_rate
model.match_count = getattr(video, "match_count", 0)
model.visual_similarity = getattr(video, "visual_similarity", 0.0)
self.session.add(model)
self.session.commit()
return video
@@ -91,24 +85,6 @@ class SQLAlchemyGeneratedVideoRepository:
models = self.session.query(GeneratedVideoModel).filter(GeneratedVideoModel.project_id == project_id).all()
return [self._to_domain(model) for model in models]
def list_by_user(self, user_id: str, *, duration_min: float = 0, duration_max: float = 0) -> list[GeneratedVideo]:
"""按 user_id 查询用户所有项目的视频(跨项目查重)。
Args:
user_id: 用户 ID
duration_min: 时长下限0 表示不限
duration_max: 时长上限0 表示不限
"""
query = self.session.query(GeneratedVideoModel).filter(
GeneratedVideoModel.user_id == user_id,
)
if duration_min > 0:
query = query.filter(GeneratedVideoModel.duration >= duration_min)
if duration_max > 0:
query = query.filter(GeneratedVideoModel.duration <= duration_max)
models = query.all()
return [self._to_domain(model) for model in models]
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]:
models = (
self.session.query(GeneratedVideoModel)
@@ -232,8 +208,6 @@ class SQLAlchemyGeneratedVideoRepository:
is_duplicate=getattr(model, "is_duplicate", False),
duplicate_of=getattr(model, "duplicate_of", None),
duplicate_rate=getattr(model, "duplicate_rate", None),
match_count=getattr(model, "match_count", 0) or 0,
visual_similarity=getattr(model, "visual_similarity", 0.0) or 0.0,
generated_at=model.generated_at,
created_at=model.created_at,
)
@@ -340,8 +340,6 @@ class GeneratedVideoModel(Base):
is_duplicate = Column(Boolean, nullable=False, default=False)
duplicate_of = Column(String(36), nullable=True)
duplicate_rate = Column(Float, nullable=True)
match_count = Column(Integer, nullable=True, default=0)
visual_similarity = Column(Float, nullable=True, default=0.0)
class TitleLibraryModel(Base):
@@ -622,20 +620,3 @@ 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))
-2
View File
@@ -27,8 +27,6 @@ class GeneratedVideo:
is_duplicate: bool = False
duplicate_of: str | None = None
duplicate_rate: float | None = None
match_count: int = 0
visual_similarity: float = 0.0
generated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
+9 -23
View File
@@ -169,7 +169,6 @@ 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(就地修改).
@@ -189,7 +188,6 @@ 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
@@ -200,16 +198,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, external_used_segments)
_distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points)
elif editing_mode == EditingMode.PIP.value:
_distribute_pip(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
_distribute_pip(clips, asset_ids, asset_durations, asset_scene_points)
elif editing_mode == EditingMode.VOICE_OVER.value:
_distribute_voice_over(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
_distribute_voice_over(clips, asset_ids, asset_durations, asset_scene_points)
elif editing_mode == EditingMode.VOICE_PIP.value:
_distribute_voice_pip(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
_distribute_voice_pip(clips, asset_ids, asset_durations, asset_scene_points)
else:
# 未知模式,退化为 one_take
_distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points, external_used_segments)
_distribute_one_take(clips, asset_ids, asset_durations, asset_scene_points)
def _resolve_start_time(
@@ -250,12 +248,9 @@ 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]]] = (
{k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {}
)
used_segments: dict[str, list[tuple[float, float]]] = {}
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):
@@ -276,12 +271,9 @@ 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]]] = (
{k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {}
)
used_segments: dict[str, list[tuple[float, float]]] = {}
# 第1个素材 → main clip
main_clips = [c for c in clips if c.clip_type == ClipType.MAIN.value]
if main_clips and asset_ids:
@@ -318,12 +310,9 @@ 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]]] = (
{k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {}
)
used_segments: dict[str, list[tuple[float, float]]] = {}
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):
@@ -344,12 +333,9 @@ 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]]] = (
{k: list(v) for k, v in external_used_segments.items()} if external_used_segments else {}
)
used_segments: dict[str, list[tuple[float, float]]] = {}
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"]
-334
View File
@@ -1,334 +0,0 @@
"""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()
+6 -13
View File
@@ -285,10 +285,8 @@ class TestVideoDeduplicatorCheckDuplicate:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
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"
assert result["similarity"] == 1.0 # distance=0 → 1.0
assert result["reason"] == "phash_similar"
finally:
self._restore_repo(mod, orig)
@@ -427,11 +425,8 @@ class TestVideoDeduplicatorCheckDuplicate:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
# 新算法: 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
# similarity = 1.0 - (1 / 64) = 0.984375
assert abs(result["similarity"] - (1.0 - 1.0 / 64)) < 1e-6
finally:
self._restore_repo(mod, orig)
@@ -461,9 +456,7 @@ class TestVideoDeduplicatorCheckDuplicate:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
assert result["duplicate"] is True
# 新算法: 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)
assert result["similarity"] == 1.0 # avg_distance = 0
finally:
self._restore_repo(mod, orig)
@@ -546,7 +539,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_histogram_fusion"
assert result["reason"] == "batch_phash_similar"
finally:
self._restore_repo(mod, orig)
+3 -15
View File
@@ -43,11 +43,7 @@ class TestDedupHelpersUserIdPassthrough:
mock_deduplicator = MagicMock()
mock_deduplicator.compute_fingerprint.return_value = mock_fingerprint
mock_deduplicator.check_duplicate.return_value = None
mock_deduplicator.compute_duplicate_rate.return_value = {
"duplicate_rate": 42.5,
"visual_similarity": 0.7,
"match_count": 2,
}
mock_deduplicator.compute_duplicate_rate.return_value = 42.5
with (
patch(
@@ -89,11 +85,7 @@ class TestDedupHelpersUserIdPassthrough:
mock_deduplicator = MagicMock()
mock_deduplicator.compute_fingerprint.return_value = mock_fingerprint
mock_deduplicator.check_duplicate.return_value = None
mock_deduplicator.compute_duplicate_rate.return_value = {
"duplicate_rate": 0.0,
"visual_similarity": 0.0,
"match_count": 0,
}
mock_deduplicator.compute_duplicate_rate.return_value = 0.0
with (
patch(
@@ -132,11 +124,7 @@ class TestDedupHelpersUserIdPassthrough:
mock_deduplicator = MagicMock()
mock_deduplicator.compute_fingerprint.return_value = mock_fingerprint
mock_deduplicator.check_duplicate.return_value = None
mock_deduplicator.compute_duplicate_rate.return_value = {
"duplicate_rate": 78.5,
"visual_similarity": 0.85,
"match_count": 3,
}
mock_deduplicator.compute_duplicate_rate.return_value = 78.5
with (
patch(
+55 -53
View File
@@ -181,68 +181,70 @@ class TestVideoFingerprint:
assert d["color_histograms"] == []
class TestBhattacharyyaCoefficient:
"""_bhattacharyya_coefficient Bhattacharyya 系数测试."""
class TestAverageHistogramSimilarity:
"""_average_histogram_similarity 直方图相似度测试."""
def test_identical_histograms(self):
"""完全相同的直方图系数为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))
"""完全相同的直方图相似度为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)
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):
def test_empty_first_list(self):
"""第一组为空返回0."""
assert VideoDeduplicator._compute_histogram_similarity([], [[0.5]]) == 0.0
sim = VideoDeduplicator._average_histogram_similarity([], [[0.5, 0.5]])
assert sim == 0.0
def test_empty_second(self):
def test_empty_second_list(self):
"""第二组为空返回0."""
assert VideoDeduplicator._compute_histogram_similarity([[0.5]], []) == 0.0
sim = VideoDeduplicator._average_histogram_similarity([[0.5, 0.5]], [])
assert sim == 0.0
def test_both_empty(self):
"""两组都为空返回0."""
assert VideoDeduplicator._compute_histogram_similarity([], []) == 0.0
sim = VideoDeduplicator._average_histogram_similarity([], [])
assert sim == 0.0
def test_best_match_selection(self):
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):
"""多帧时取最佳匹配."""
# 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
# 第一帧完全不同,第二帧完全相同 → 平均 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
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
-532
View File
@@ -1,532 +0,0 @@
"""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
+171 -74
View File
@@ -19,14 +19,14 @@ sys.path.insert(0, str(ROOT / "apps" / "worker"))
class TestComputeDuplicateRate:
"""Test VideoDeduplicator.compute_duplicate_rate."""
def _make_fingerprint(self, md5="abc123", phashes=None, duration_ms=10000):
def _make_fingerprint(self, md5="abc123", phashes=None):
from video_processing.dedup import VideoFingerprint
return VideoFingerprint(
md5=md5,
keyframe_phashes=phashes or ["ff00ff00ff00ff00"],
color_histograms=[],
duration=duration_ms,
duration=10.0,
resolution=(1920, 1080),
)
@@ -56,116 +56,184 @@ class TestComputeDuplicateRate:
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = []
query_mock = MagicMock()
query_mock.filter.return_value = query_mock
query_mock.order_by.return_value.limit.return_value.all.return_value = []
session.query.return_value = query_mock
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
assert rate["duplicate_rate"] == 0.0
assert rate["match_count"] == 0
assert isinstance(rate, dict)
assert rate == 0.0
def test_md5_match_returns_100(self):
from video_processing.dedup import VideoDeduplicator
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
deduplicator = VideoDeduplicator()
fingerprint = self._make_fingerprint(md5="exact_md5")
fingerprint = self._make_fingerprint(md5="exact_match_md5")
session = MagicMock()
existing = self._make_existing_video("vid2", {"md5": "exact_md5", "keyframe_phashes": ["aa"]})
existing = self._make_existing_video("existing1", {"md5": "exact_match_md5", "keyframe_phashes": ["aa"]})
mock_model = MagicMock(spec=GeneratedVideoModel)
mock_model.id = existing.id
mock_model.project_id = existing.project_id
mock_model.video_fingerprint = existing.video_fingerprint
mock_model.generated_at = "2026-01-01"
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = [existing]
mock_repo._to_domain.return_value = existing
# 链式 filter: 第一次 scope filter,第二次 self-exclusion filter
# 让 filter() 返回的对象仍然支持 order_by() 链
query_mock = MagicMock()
query_mock.filter.return_value = query_mock # filter → filter chainable
query_mock.order_by.return_value.limit.return_value.all.return_value = [mock_model]
session.query.return_value = query_mock
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
assert rate["duplicate_rate"] == 100.0
assert rate["match_count"] == 1
assert rate == 100.0
def test_phash_similarity_computed(self):
from video_processing.dedup import VideoDeduplicator
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
deduplicator = VideoDeduplicator()
# Two very similar phashes
fingerprint = self._make_fingerprint(
md5="new",
phashes=["ff00ff00ff00ff00", "ff00ff00ff00ff01"],
)
fingerprint = self._make_fingerprint(md5="different_md5", phashes=["ff00ff00ff00ff00"])
session = MagicMock()
existing = self._make_existing_video(
"vid2",
{"md5": "other", "keyframe_phashes": ["ff00ff00ff00ff00", "ff00ff00ff00ff02"]},
"existing1",
{"md5": "other_md5", "keyframe_phashes": ["ff00ff00ff00ff03"]},
)
mock_model = MagicMock(spec=GeneratedVideoModel)
mock_model.id = existing.id
mock_model.project_id = existing.project_id
mock_model.video_fingerprint = existing.video_fingerprint
mock_model.generated_at = "2026-01-01"
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = [existing]
mock_repo._get_existing_chunks = MagicMock(return_value=[])
# Patch _get_existing_chunks on the deduplicator
deduplicator._get_existing_chunks = MagicMock(return_value=[])
mock_repo._to_domain.return_value = existing
query_mock = MagicMock()
query_mock.filter.return_value = query_mock
query_mock.order_by.return_value.limit.return_value.all.return_value = [mock_model]
session.query.return_value = query_mock
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
# With identical phashes, frame_match_rate should be high
assert rate["duplicate_rate"] >= 0.0
assert isinstance(rate, dict)
assert "visual_similarity" in rate
# hamming distance = 2, similarity = (1 - 2/64) * 100 = 96.875
assert rate == pytest.approx(96.88, abs=0.1)
def test_excludes_self_video(self):
from video_processing.dedup import VideoDeduplicator
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
deduplicator = VideoDeduplicator()
fingerprint = self._make_fingerprint(md5="same_md5")
session = MagicMock()
self_video = self._make_existing_video("vid1", {"md5": "same_md5", "keyframe_phashes": ["aa"]})
mock_model = MagicMock(spec=GeneratedVideoModel)
mock_model.id = self_video.id
mock_model.project_id = self_video.project_id
mock_model.video_fingerprint = self_video.video_fingerprint
mock_model.generated_at = "2026-01-01"
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo._to_domain.return_value = self_video
query_mock = MagicMock()
query_mock.filter.return_value = query_mock
query_mock.order_by.return_value.limit.return_value.all.return_value = [mock_model]
session.query.return_value = query_mock
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
assert rate == 0.0
def test_takes_max_similarity(self):
from video_processing.dedup import VideoDeduplicator
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
deduplicator = VideoDeduplicator()
fingerprint = self._make_fingerprint(
md5="new",
phashes=["aa00aa00aa00aa00"],
)
fingerprint = self._make_fingerprint(md5="new_md5", phashes=["ff00ff00ff00ff00"])
session = MagicMock()
# Two existing videos with different phashes
existing1 = self._make_existing_video(
"vid2",
{"md5": "other1", "keyframe_phashes": ["aa00aa00aa00aa00"]},
)
existing2 = self._make_existing_video(
"vid3",
{"md5": "other2", "keyframe_phashes": ["ff00ff00ff00ff00"]},
)
existing1 = self._make_existing_video("e1", {"md5": "md5_1", "keyframe_phashes": ["ff00ff00ff00ff0f"]})
existing2 = self._make_existing_video("e2", {"md5": "md5_2", "keyframe_phashes": ["ff00ff00ff00ff01"]})
mock_model1 = MagicMock(spec=GeneratedVideoModel)
mock_model1.id = existing1.id
mock_model1.project_id = existing1.project_id
mock_model1.video_fingerprint = existing1.video_fingerprint
mock_model1.generated_at = "2026-01-02"
mock_model2 = MagicMock(spec=GeneratedVideoModel)
mock_model2.id = existing2.id
mock_model2.project_id = existing2.project_id
mock_model2.video_fingerprint = existing2.video_fingerprint
mock_model2.generated_at = "2026-01-01"
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = [existing1, existing2]
deduplicator._get_existing_chunks = MagicMock(return_value=[])
mock_repo._to_domain.side_effect = [existing1, existing2]
query_mock = MagicMock()
query_mock.filter.return_value = query_mock
query_mock.order_by.return_value.limit.return_value.all.return_value = [
mock_model1,
mock_model2,
]
session.query.return_value = query_mock
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
# Should take the max across all videos
assert rate["duplicate_rate"] >= 0.0
assert isinstance(rate["duplicate_rate"], float)
# max similarity: e2 distance=1, (1-1/64)*100 = 98.4375
assert rate == pytest.approx(98.44, abs=0.1)
def test_user_id_scope_cross_project(self):
"""传 user_id 时应跨项目查询,而非仅当前项目."""
from video_processing.dedup import VideoDeduplicator
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
deduplicator = VideoDeduplicator()
fingerprint = self._make_fingerprint(md5="exact_md5_x")
fingerprint = self._make_fingerprint(md5="cross_proj_md5")
session = MagicMock()
existing = self._make_existing_video("vid2", {"md5": "exact_md5_x", "keyframe_phashes": ["aa"]})
# 模拟一个不同项目但同一用户的视频
existing = self._make_existing_video(
"existing_other_proj", {"md5": "cross_proj_md5", "keyframe_phashes": ["aa"]}
)
existing.project_id = "proj2" # 不同项目
existing.user_id = "user1"
mock_model = MagicMock(spec=GeneratedVideoModel)
mock_model.id = existing.id
mock_model.project_id = existing.project_id
mock_model.user_id = existing.user_id
mock_model.video_fingerprint = existing.video_fingerprint
mock_model.generated_at = "2026-01-01"
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_user.return_value = [existing]
mock_repo._to_domain.return_value = existing
query_mock = MagicMock()
query_mock.filter.return_value = query_mock
query_mock.order_by.return_value.limit.return_value.all.return_value = [mock_model]
session.query.return_value = query_mock
rate = deduplicator.compute_duplicate_rate(
fingerprint,
"proj1",
"vid1",
session,
scope="user",
user_id="user1",
)
# Should use list_by_user and find the match
mock_repo.list_by_user.assert_called_once_with("user1")
assert rate["duplicate_rate"] == 100.0
# 应通过 user_id 过滤,且匹配到跨项目视频
assert rate == 100.0
def test_return_dict_structure(self):
"""compute_duplicate_rate returns dict with three fields."""
def test_user_id_empty_falls_back_to_project(self):
"""user_id 为空时应回退到 project_id 过滤."""
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
@@ -174,29 +242,58 @@ class TestComputeDuplicateRate:
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = []
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
query_mock = MagicMock()
query_mock.filter.return_value = query_mock
query_mock.order_by.return_value.limit.return_value.all.return_value = []
session.query.return_value = query_mock
assert isinstance(rate, dict)
assert "duplicate_rate" in rate
assert "visual_similarity" in rate
assert "match_count" in rate
assert isinstance(rate["duplicate_rate"], float)
assert isinstance(rate["visual_similarity"], float)
assert isinstance(rate["match_count"], int)
rate = deduplicator.compute_duplicate_rate(
fingerprint,
"proj1",
"vid1",
session,
user_id="",
)
def test_backward_compat_no_scope(self):
"""Not passing scope defaults to project-level."""
from video_processing.dedup import VideoDeduplicator
assert rate == 0.0
# 验证使用的是 project_id 过滤(回退路径)
# 通过检查 filter 被调用时的参数来间接验证
deduplicator = VideoDeduplicator()
fingerprint = self._make_fingerprint()
session = MagicMock()
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = []
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
class TestDuplicateRateAPI:
"""Test that duplicate_rate is returned in API responses."""
mock_repo.list_by_project.assert_called_once_with("proj1")
assert rate["duplicate_rate"] == 0.0
def test_video_item_response_has_duplicate_rate(self):
from app.schemas.video_center import VideoItemResponse
resp = VideoItemResponse(
id="v1",
project_id="p1",
generation_task_id="t1",
name="test.mp4",
file_url="https://example.com/test.mp4",
file_size=1000,
duration=10.0,
width=1920,
height=1080,
fps=25.0,
duplicate_rate=75.5,
)
assert resp.duplicate_rate == 75.5
def test_video_item_response_duplicate_rate_default_none(self):
from app.schemas.video_center import VideoItemResponse
resp = VideoItemResponse(
id="v1",
project_id="p1",
generation_task_id="t1",
name="test.mp4",
file_url="https://example.com/test.mp4",
file_size=1000,
duration=10.0,
width=1920,
height=1080,
fps=25.0,
)
assert resp.duplicate_rate is None
-367
View File
@@ -1,367 +0,0 @@
"""Tests for Issue #1660 — 查重率百分比计算 + 跨项目查重."""
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
sys.modules.setdefault("cv2", MagicMock())
ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(ROOT / "apps" / "api"))
sys.path.insert(0, str(ROOT / "packages"))
sys.path.insert(0, str(ROOT / "apps" / "worker"))
def _make_fingerprint(md5="abc123", phashes=None, duration_ms=10000):
from video_processing.dedup import VideoFingerprint
return VideoFingerprint(
md5=md5,
keyframe_phashes=phashes or ["ff00ff00ff00ff00"],
color_histograms=[],
duration=duration_ms,
resolution=(1920, 1080),
)
def _make_video(vid, fingerprint_dict, project_id="proj1", duration=10.0):
from packages.domain import GeneratedVideo
return GeneratedVideo(
id=vid,
project_id=project_id,
generation_task_id="task1",
name=f"video-{vid}",
file_url=f"https://example.com/{vid}.mp4",
file_size=1000,
duration=duration,
width=1920,
height=1080,
fps=25.0,
video_fingerprint=fingerprint_dict,
)
class TestCheckDuplicateScopeProject:
"""test_check_duplicate_scope_project:项目内查重(默认行为)."""
def test_default_scope_queries_by_project(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint(md5="unique_md5")
session = MagicMock()
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = []
result = deduplicator.check_duplicate(fingerprint, "proj1", session)
mock_repo.list_by_project.assert_called_once_with("proj1")
assert result is None
def test_project_scope_finds_duplicate(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint(md5="same_md5")
session = MagicMock()
existing = _make_video("vid2", {"md5": "same_md5", "keyframe_phashes": ["aa"]})
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = [existing]
result = deduplicator.check_duplicate(fingerprint, "proj1", session)
assert result is not None
assert result["duplicate"] is True
assert result["duplicate_of"] == "vid2"
class TestCheckDuplicateScopeUser:
"""test_check_duplicate_scope_user:跨项目查重."""
def test_user_scope_queries_by_user(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint(md5="unique_md5")
session = MagicMock()
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_user.return_value = []
result = deduplicator.check_duplicate(
fingerprint,
"proj1",
session,
scope="user",
user_id="user_123",
)
mock_repo.list_by_user.assert_called_once()
assert result is None
def test_user_scope_finds_cross_project_duplicate(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint(md5="cross_proj_md5")
session = MagicMock()
# Existing video from a different project
existing = _make_video("vid_other", {"md5": "cross_proj_md5"}, project_id="proj_other")
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_user.return_value = [existing]
result = deduplicator.check_duplicate(
fingerprint,
"proj1",
session,
scope="user",
user_id="user_123",
)
assert result is not None
assert result["duplicate"] is True
assert result["duplicate_of"] == "vid_other"
class TestDurationPrefilter:
"""test_duration_prefilter:时长 ±15% 过滤."""
def test_duration_prefilter_passes_correct_range(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint(duration_ms=30000) # 30s video
session = MagicMock()
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_user.return_value = []
deduplicator.check_duplicate(
fingerprint,
"proj1",
session,
scope="user",
user_id="user1",
duration_sec=30.0,
)
# Should pass duration_min=25.5, duration_max=34.5 (30 ± 15%)
call_args = mock_repo.list_by_user.call_args
assert call_args[1]["duration_min"] == pytest.approx(25.5, abs=0.1)
assert call_args[1]["duration_max"] == pytest.approx(34.5, abs=0.1)
def test_no_duration_prefilter_when_zero(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint()
session = MagicMock()
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_user.return_value = []
deduplicator.check_duplicate(
fingerprint,
"proj1",
session,
scope="user",
user_id="user1",
duration_sec=0,
)
call_args = mock_repo.list_by_user.call_args
assert call_args[1]["duration_min"] == 0
assert call_args[1]["duration_max"] == 0
class TestComputeDuplicateRateFormula:
"""test_compute_duplicate_rate_formula:验证 0.4 * frame_match_rate + 0.6 * temporal_coverage_rate."""
def test_formula_with_matching_frames(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
# 10 frames, all identical to existing → frame_match_rate = 1.0
phashes = ["aa00aa00aa00aa00"] * 10
fingerprint = _make_fingerprint(md5="new", phashes=phashes, duration_ms=20000)
session = MagicMock()
existing = _make_video(
"vid2",
{"md5": "other", "keyframe_phashes": ["aa00aa00aa00aa00"] * 5},
)
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = [existing]
deduplicator._get_existing_chunks = MagicMock(return_value=[])
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
# frame_match_rate=1.0, temporal_coverage depends on segments
# duplicate_rate = (1.0 * 0.4 + temporal_coverage * 0.6) * 100
assert rate["duplicate_rate"] >= 40.0 # At minimum, frame_match contributes 40%
def test_no_match_returns_zero(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
# Completely different phashes
fingerprint = _make_fingerprint(md5="new", phashes=["ff00ff00ff00ff00"])
session = MagicMock()
existing = _make_video(
"vid2",
{"md5": "other", "keyframe_phashes": ["00ff00ff00ff00ff"]},
)
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = [existing]
deduplicator._get_existing_chunks = MagicMock(return_value=[])
rate = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
# Very different phashes, match_ratio < 0.3 → skipped
assert rate["duplicate_rate"] == 0.0
class TestComputeDuplicateRateReturnDict:
"""test_compute_duplicate_rate_return_dict:验证返回 dict 含三个字段."""
def test_return_structure(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint()
session = MagicMock()
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = []
result = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
assert isinstance(result, dict)
assert set(result.keys()) == {"duplicate_rate", "visual_similarity", "match_count"}
assert isinstance(result["duplicate_rate"], float)
assert isinstance(result["visual_similarity"], float)
assert isinstance(result["match_count"], int)
assert 0 <= result["duplicate_rate"] <= 100
assert 0 <= result["visual_similarity"] <= 1
class TestBackwardCompat:
"""test_backward_compat:不传 scope 时行为不变."""
def test_default_scope_is_project(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint()
session = MagicMock()
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = []
# Call without scope parameter
result = deduplicator.compute_duplicate_rate(fingerprint, "proj1", "vid1", session)
# Should use list_by_project (not list_by_user)
mock_repo.list_by_project.assert_called_once_with("proj1")
mock_repo.list_by_user.assert_not_called()
assert result["duplicate_rate"] == 0.0
def test_check_duplicate_default_scope_backward_compat(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
fingerprint = _make_fingerprint()
session = MagicMock()
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
mock_repo = MockRepo.return_value
mock_repo.list_by_project.return_value = []
result = deduplicator.check_duplicate(fingerprint, "proj1", session)
mock_repo.list_by_project.assert_called_once_with("proj1")
assert result is None
class TestListByUserRepository:
"""直接测试 generated_video_repository.list_by_user() 的真实实现,覆盖 diff 代码行。"""
def _make_repo(self):
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
from packages.adapters.sqlalchemy_impl.models import Base, GeneratedVideoModel
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine)
session = Session()
repo = SQLAlchemyGeneratedVideoRepository(session)
return repo, session
def _insert_video(self, session, video_id, user_id, project_id, duration, **kw):
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
row = GeneratedVideoModel(
id=video_id,
user_id=user_id,
project_id=project_id,
generation_task_id=f"task-{video_id[:8]}",
name=f"video-{video_id[:8]}.mp4",
file_url=f"https://example.com/{video_id}.mp4",
file_size=1024,
duration=duration,
width=1280,
height=720,
fps=25.0,
status="completed",
)
session.add(row)
session.flush()
return row
def test_list_by_user_returns_cross_project_videos(self):
"""list_by_user 返回该用户所有项目的视频。"""
repo, session = self._make_repo()
self._insert_video(session, "v1", "user-a", "proj-1", 30.0)
self._insert_video(session, "v2", "user-a", "proj-2", 45.0)
self._insert_video(session, "v3", "user-b", "proj-1", 20.0)
results = repo.list_by_user("user-a")
assert len(results) == 2
ids = {r.id for r in results}
assert ids == {"v1", "v2"}
session.close()
def test_list_by_user_with_duration_filter(self):
"""list_by_user 支持 duration_min/duration_max 过滤。"""
repo, session = self._make_repo()
self._insert_video(session, "v1", "user-a", "proj-1", 10.0)
self._insert_video(session, "v2", "user-a", "proj-1", 30.0)
self._insert_video(session, "v3", "user-a", "proj-1", 60.0)
results = repo.list_by_user("user-a", duration_min=20.0, duration_max=50.0)
assert len(results) == 1
assert results[0].id == "v2"
session.close()
def test_list_by_user_empty_result(self):
"""list_by_user 无匹配时返回空列表。"""
repo, session = self._make_repo()
self._insert_video(session, "v1", "user-a", "proj-1", 30.0)
results = repo.list_by_user("user-nonexistent")
assert results == []
session.close()
-282
View File
@@ -1,282 +0,0 @@
"""分片指纹存储单元测试 — Issue #1657.
覆盖
- 分片策略60秒视频 30120秒视频 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
@@ -359,11 +359,6 @@ class TestThumbnailInDedupHelpers:
mock_dedup.compute_fingerprint.return_value = MagicMock(to_dict=lambda: {})
mock_dedup.check_duplicate.return_value = None
mock_dedup.check_batch_duplicate.return_value = None
mock_dedup.compute_duplicate_rate.return_value = {
"duplicate_rate": 0.0,
"visual_similarity": 0.0,
"match_count": 0,
}
result = create_video_record_and_dedup(
generation_task_id="task-thumb-reuse",
@@ -406,11 +401,6 @@ class TestThumbnailInDedupHelpers:
mock_dedup.compute_fingerprint.return_value = MagicMock(to_dict=lambda: {})
mock_dedup.check_duplicate.return_value = None
mock_dedup.check_batch_duplicate.return_value = None
mock_dedup.compute_duplicate_rate.return_value = {
"duplicate_rate": 0.0,
"visual_similarity": 0.0,
"match_count": 0,
}
result = create_video_record_and_dedup(
generation_task_id="task-thumb-gen",
@@ -453,11 +443,6 @@ class TestThumbnailInDedupHelpers:
mock_dedup.compute_fingerprint.return_value = MagicMock(to_dict=lambda: {})
mock_dedup.check_duplicate.return_value = None
mock_dedup.check_batch_duplicate.return_value = None
mock_dedup.compute_duplicate_rate.return_value = {
"duplicate_rate": 0.0,
"visual_similarity": 0.0,
"match_count": 0,
}
result = create_video_record_and_dedup(
generation_task_id="task-thumb-fail",