fix(#1743): 批量变体独立选片——完整重跑单视频选片+批次20%重叠避让+查重超阈重渲+封面独立 #1745

Merged
auto-approve-bot merged 2 commits from fix/batch-variant-independent-plans-1743 into develop 2026-09-06 18:07:25 +08:00
16 changed files with 1893 additions and 204 deletions
+48 -13
View File
@@ -420,42 +420,77 @@ def create_preview_generation_task(
logger.error("[预览生成] 创建失败: %s", e, exc_info=True)
raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e
# ── 克隆独立变体 plan:N 个预览全部克隆(预览不污染源 plan)──
# 源 plan 不存在(无编辑历史)时各任务走自身随机选片流程,不克隆。
# ── 独立变体 plan#1743)──
# count=1:克隆源 plan(预览不污染源 plan,仅起点重算),行为与旧版一致;
# count>1:变体 0 保留源 plan,变体 1..N-1 用 reselect_plan_for_variant 完整
# 重跑单视频选片(素材洗牌+镜头洗牌+起点随机+跨变体避让+批次 20% 重叠重选),
# 所见即所得——预览变体差异即正式成片差异。
source_plan_id = created_tasks[0].source_edit_plan_id if created_tasks else ""
if source_plan_id:
if source_plan_id and count == 1:
# 单预览:克隆一份(原逻辑)
try:
from app.services.edit_plan_service import EditPlanService
_plan_svc = EditPlanService(db)
for variant_index in range(count):
variant_plan = _plan_svc.clone_plan_for_variant(
source_plan_id,
created_by_user_id=user_id,
name_suffix="预览变体",
)
variant_plan_ids.append(variant_plan.id)
except Exception as e:
logger.error("[预览生成] 克隆预览 plan 异常: %s", e, exc_info=True)
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览计划创建失败")
raise HTTPException(
status_code=500,
detail="创建预览任务失败:无法生成独立剪辑计划,请重试",
) from e
elif source_plan_id and count > 1:
try:
from app.services.edit_plan_service import EditPlanService
_plan_svc = EditPlanService(db)
# 变体 0 直接用源 plan;变体 1..N-1 独立选片
variant_plan_ids.append(source_plan_id)
batch_asset_pool = list(dict.fromkeys(request.asset_ids or []))
for variant_index in range(1, count):
last_err: Exception | None = None
variant_plan = None
for _attempt in range(2): # 1 次重试,抗 DB 瞬时抖动
try:
variant_plan = _plan_svc.clone_plan_for_variant(
variant_plan = _plan_svc.reselect_plan_for_variant(
source_plan_id,
batch_asset_pool,
created_by_user_id=user_id,
name_suffix=f"预览变体{variant_index + 1}" if count > 1 else "预览变体",
name_suffix=f"预览变体{variant_index + 1}",
)
break
except Exception as clone_err: # noqa: PERF203
last_err = clone_err
except ValueError as ve:
logger.warning("[预览生成] 变体独立选片失败(素材不足): %s", ve)
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体选片失败")
raise HTTPException(
status_code=400,
detail=f"批量预览第 {variant_index + 1} 个视频无法独立选片:{ve}"
"请增加素材库中的视频素材后重试。",
) from ve
except Exception as reselection_err: # noqa: PERF203
last_err = reselection_err
logger.warning(
"[预览生成] 克隆变体 plan 失败(尝试%d/2): variant=%d error=%s",
"[预览生成] 变体独立选片失败(尝试%d/2): variant=%d error=%s",
_attempt + 1,
variant_index,
clone_err,
reselection_err,
exc_info=True,
)
if variant_plan is None:
logger.error(
"[预览生成] 克隆预览变体 plan 重试仍失败: variant=%d source=%s",
"[预览生成] 变体独立选片重试仍失败: variant=%d source=%s",
variant_index,
source_plan_id,
exc_info=last_err,
)
# 标记已创建任务失败
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体计划创建失败")
raise HTTPException(
@@ -466,7 +501,7 @@ def create_preview_generation_task(
except HTTPException:
raise
except Exception as e:
logger.error("[预览生成] 克隆变体 plan 异常: %s", e, exc_info=True)
logger.error("[预览生成] 变体 plan 生成异常: %s", e, exc_info=True)
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体计划创建失败")
raise HTTPException(
+76 -19
View File
@@ -432,39 +432,83 @@ def create_generation_task(
logger.info("画中画已下线,strategy_id %s → one_take", effective_strategy_id)
effective_strategy_id = "one_take"
# 批量生成时每个任务关联独立克隆 plan(片段起点重算),
# 禁止 N 条任务共用同一 source_edit_plan_id 导致片段一模一样。
# 在创建任何任务【之前】预克隆全部变体:克隆失败直接中断(此时无脏数据),
# 绝不静默退回共用源 plan(否则批量视频内容重复,违反去重诉求)。
# 批量生成(count>1):每个变体必须走与单视频完全相同的独立选片流程(#1743)。
# - 变体 0 保留源 plan(保留用户编辑结果);
# - 变体 1..N-1 用 reselect_plan_for_variant 完整重跑选片(素材洗牌 + 镜头洗牌
# + 起点随机 + 跨变体区间避让 + 批次 20% 重叠重选),而非"克隆只改起点";
# - count>1 但没有源 plan(前端未传 source_edit_plan_id 且无模板 plan)时,
# 不允许 N 个任务兜底共用同一 plan,直接 4xx 中断(宁可不生成,也不出同源成片)。
# 在创建任何任务【之前】预生成全部变体 plan:失败直接中断(此时无脏数据)。
variant_plan_ids: list[str] = []
if count > 1 and request.source_edit_plan_id:
if count > 1:
from app.services.edit_plan_service import EditPlanService
_plan_svc = EditPlanService(db)
# 解析批量源 plan:优先前端传入;否则按 template_id + user 查最新(与单任务兜底同源)
batch_source_plan_id = request.source_edit_plan_id
if not batch_source_plan_id and request.template_id:
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
_latest = (
db.query(EditPlanModel)
.filter(
EditPlanModel.template_id == request.template_id,
EditPlanModel.created_by_user_id == user_id,
)
.order_by(EditPlanModel.created_at.desc())
.first()
)
if _latest:
batch_source_plan_id = _latest.id
except Exception:
logger.warning("[生成任务] 批量源 plan 解析失败", exc_info=True)
if not batch_source_plan_id:
# 无任何可用源 plan:批量变体无从选片,明确报错,严禁静默共用/同源
logger.error("[生成任务] 批量 count=%d 但无可编辑计划(无 source_edit_plan_id/template plan", count)
raise HTTPException(
status_code=400,
detail="批量生成需要先完成预览生成(缺少剪辑计划)。请先生成预览后再批量创建。",
)
# 批次素材池:请求显式素材 + 库自动匹配素材(resolved_asset_ids
batch_asset_pool = list(dict.fromkeys(resolved_asset_ids or []))
for task_index in range(1, count):
variant = None
last_err: Exception | None = None
for _attempt in range(2): # 1 次重试,抗 DB 瞬时抖动
try:
variant = _plan_svc.clone_plan_for_variant(
request.source_edit_plan_id,
variant = _plan_svc.reselect_plan_for_variant(
batch_source_plan_id,
batch_asset_pool,
created_by_user_id=user_id,
name_suffix=f"批量{task_index + 1}",
)
break
except Exception as clone_err: # noqa: PERF203
last_err = clone_err
except ValueError as ve:
# 素材不足等可预期错误:不重试,直接中断并给出明确提示
logger.warning("[生成任务] 变体独立选片失败(素材不足): %s", ve)
raise HTTPException(
status_code=400,
detail=f"批量生成第 {task_index + 1} 个视频无法独立选片:{ve}"
"请增加素材库中的视频素材后重试。",
) from ve
except Exception as reselection_err: # noqa: PERF203
last_err = reselection_err
logger.warning(
"[生成任务] 克隆变体 plan 失败(尝试%d/2): source=%s error=%s",
"[生成任务] 变体独立选片失败(尝试%d/2): source=%s error=%s",
_attempt + 1,
request.source_edit_plan_id,
clone_err,
batch_source_plan_id,
reselection_err,
exc_info=True,
)
if variant is None:
logger.error(
"[生成任务] 克隆变体 plan 重试仍失败,中断批量创建: source=%s",
request.source_edit_plan_id,
"[生成任务] 变体独立选片重试仍失败,中断批量创建: source=%s",
batch_source_plan_id,
exc_info=last_err,
)
raise HTTPException(
@@ -475,10 +519,11 @@ def create_generation_task(
try:
for task_index in range(count):
# 第 1 条复用源 plan(保留用户编辑结果);其余使用预克隆的独立变体 plan。
# 无源 plansource_edit_plan_id 为空)时无可克隆对象,variant_plan_ids
# 为空列表:各任务走自身随机选片流程,不做索引访问(防 IndexError)
effective_plan_id = request.source_edit_plan_id
# 变体 0 复用源 plan(保留用户编辑结果);变体 1..N-1 用预生成的独立选片 plan。
# count>1 时上方已保证存在源 plan 且变体 plan 数量 == count-1。
effective_plan_id = (
request.source_edit_plan_id or batch_source_plan_id if count > 1 else request.source_edit_plan_id
)
if task_index > 0 and variant_plan_ids:
effective_plan_id = variant_plan_ids[task_index - 1]
@@ -523,7 +568,19 @@ def create_generation_task(
try:
# 兜底关联编辑计划:前端未传 source_edit_plan_id 时,
# 通过 template_id + user_id 在 DB 层直接查找最新的 plan。
# 必须在 enqueue 之前执行,避免 worker 读取时 source_edit_plan_id 为空(竞态条件)
# 必须在 enqueue 之前执行,避免 worker 读取时 source_edit_plan_id 为空(竞态条件)
# #1743:批量(count>1)场景严禁兜底共用——变体 plan 已在上方预生成,
# 走到这里还缺 plan 说明预生成漏配,直接报错中断,不允许 N 任务关联同一 plan。
if not task.source_edit_plan_id and count > 1:
logger.error(
"[生成任务] 批量任务缺少独立 plan(禁止共用兜底): task_index=%d task_id=%s",
task_index,
task.id,
)
raise HTTPException(
status_code=500,
detail="创建批量任务失败:变体剪辑计划缺失,请重新预览后再批量生成。",
)
if not task.source_edit_plan_id and request.template_id:
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
+135
View File
@@ -409,8 +409,13 @@ class EditPlanService:
clip_type=clip_item.get("clip_type", "main"),
order=order,
asset_id=clip_item.get("asset_id", ""),
text_content=clip_item.get("text_content", ""),
start_time=clip_item.get("start_time", 0.0),
duration=clip_item.get("duration", 0.0),
transition_effect=clip_item.get("transition_effect", "cut"),
transition_duration=clip_item.get("transition_duration", 0.0),
playback_speed=clip_item.get("playback_speed", 1.0),
config=clip_item.get("config") or None,
)
model = EditPlanClipModel(
id=clip.id,
@@ -459,6 +464,136 @@ class EditPlanService:
logger.exception("事务性替换片段失败: plan_id=%s", plan_id)
raise
def reselect_plan_for_variant(
self,
source_plan_id: str,
candidate_asset_ids: list[str],
*,
created_by_user_id: str = "",
name_suffix: str = "变体",
rng=None,
) -> EditPlan:
"""为批量变体生成独立 plan:完整重跑单视频选片流程(#1743)。
与 clone_plan_for_variant(只重算起点、素材/顺序不变)不同,本方法:
- 源 plan 片段骨架(clip_type/order/duration/文案/转场)保留;
- 素材池 shuffle 随机分配 + main 片段顺序洗牌;
- 起点走场景镜头洗牌/随机起点/历史区间避让(与单视频同一入口);
- 批次内同素材区间重叠 >20% 自动重选起点;
- 新片段区间 record_used_segments 写回素材 metadata(跨变体/跨任务避让)。
Args:
source_plan_id: 源 plan(任务 0 / 预览源)。
candidate_asset_ids: 素材池(源 plan 素材 ∪ 批次素材)。
created_by_user_id: 新 plan 归属用户。
name_suffix: plan 名后缀。
rng: 可选随机数(测试注入种子)。
Raises:
ValueError: 源 plan 不存在/无片段、素材池为空或时长全未知。
"""
from packages.adapters.sqlalchemy_impl.models import AssetModel
from packages.domain.plan_generator_utils import extract_scene_points_from_metadata
from packages.domain.variant_plan_selector import reselect_clips_for_variant
source = self.get_plan_or_raise(source_plan_id)
# 分页读取源 plan 全部片段
clips: List[EditPlanClip] = []
skip, page = 0, 500
while True:
batch = self._clip_repo.list_by_plan(source_plan_id, skip=skip, limit=page)
if not batch:
break
clips.extend(batch)
if len(batch) < page:
break
skip += page
if not clips:
raise ValueError(f"源 plan 无片段,无法生成变体: {source_plan_id}")
source_clips_data = [
{
"order": c.order if c.order is not None else i,
"asset_id": c.asset_id,
"start_time": float(c.start_time or 0.0),
"duration": float(c.duration or 0.0),
"clip_type": c.clip_type,
"playback_speed": float(c.playback_speed or 1.0),
"transition_effect": c.transition_effect,
"transition_duration": float(c.transition_duration or 0.0),
"text_content": c.text_content or "",
"config": c.config or {},
}
for i, c in enumerate(clips)
]
db = self._clip_repo.session
# 素材池 = 源 plan 素材 ∪ 调用方传入素材(去重保序)
pool_ids: list[str] = []
seen = set()
for aid in [c.asset_id for c in clips if c.asset_id] + list(candidate_asset_ids or []):
if aid and aid not in seen:
seen.add(aid)
pool_ids.append(aid)
# 时长 + 场景点
durations: dict[str, float] = {}
scene_points: dict[str, list[float]] = {}
if pool_ids:
for m in db.query(AssetModel).filter(AssetModel.id.in_(pool_ids)).all():
durations[m.id] = float(getattr(m, "duration", 0.0) or 0.0)
pts = extract_scene_points_from_metadata(getattr(m, "metadata", None))
if pts:
scene_points[m.id] = pts
historical = get_used_segments(db, pool_ids)
# 创建新 plan(复制模板归属与 config
new_plan = self.create_plan(
template_id=source.template_id,
name=f"{source.name or '剪辑计划'} · {name_suffix}",
config=dict(source.config or {}),
total_duration=source.total_duration,
project_id=source.project_id or "",
created_by_user_id=created_by_user_id or (source.created_by_user_id or ""),
)
# 批次内区间:以源 plan(变体 0)片段为初始避让对象
batch_segments: dict[str, list[tuple[float, float]]] = {}
for c in clips:
if c.asset_id and float(c.duration or 0) > 0:
st = float(c.start_time or 0.0)
batch_segments.setdefault(c.asset_id, []).append((st, st + float(c.duration)))
clips_data = reselect_clips_for_variant(
source_clips_data,
pool_ids,
asset_durations=durations,
asset_scene_points=scene_points,
historical_used_segments=historical,
batch_segments=batch_segments,
rng=rng,
)
# 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit
for item in clips_data:
aid = item.get("asset_id", "")
if aid:
st = float(item.get("start_time", 0.0))
record_used_segments(db, aid, st, st + float(item.get("duration", 0.0)), new_plan.id)
self.replace_all_clips_transactional(new_plan.id, clips_data)
logger.info(
"变体独立选片完成: source=%s new=%s clips=%d assets=%d",
source_plan_id,
new_plan.id,
len(clips_data),
len(pool_ids),
)
return new_plan
def clone_plan_for_variant(
self,
source_plan_id: str,
+14 -4
View File
@@ -34,7 +34,9 @@ def create_video_record_and_dedup(
fps: float = 25.0,
name: str = "",
thumbnail_url: str = "",
) -> int:
) -> dict:
"""Returns: {"video_count": int, "is_duplicate": bool, "batch_similarity": float|None,
"duplicate_of": str|None} —— batch_similarity 为批次内最高相似度(无批次查重时 None)。"""
"""创建 GeneratedVideo 记录,计算指纹并执行查重(历史 + 批次)。
采用两阶段持久化:先计算所有指纹/查重数据(内存),
@@ -76,6 +78,7 @@ def create_video_record_and_dedup(
# ── Phase 2: 计算指纹 & 查重(全部在内存) ────────────────
deduplicator = VideoDeduplicator()
fingerprint = None
batch_similarity: float | None = None
try:
fingerprint = deduplicator.compute_fingerprint(video_path)
@@ -105,9 +108,11 @@ def create_video_record_and_dedup(
)
# (b) 批次内查重(仅当有 batch_id 时)
batch_similarity: float | None = None
if not duplicate_result and batch_id:
duplicate_result = deduplicator.check_batch_duplicate(fingerprint, batch_id, video_id, session)
if duplicate_result:
batch_similarity = float(duplicate_result.get("similarity", 0.0))
if duplicate_result:
generated_video.is_duplicate = True
generated_video.duplicate_of = duplicate_result["duplicate_of"]
@@ -161,7 +166,12 @@ def create_video_record_and_dedup(
generated_video.is_duplicate,
generated_video.duplicate_rate,
)
return 1
return {
"video_count": 1,
"is_duplicate": bool(generated_video.is_duplicate),
"batch_similarity": batch_similarity,
"duplicate_of": generated_video.duplicate_of,
}
except Exception as e:
logger.error(
"Failed to create video record / dedup for task %s: %s",
@@ -169,4 +179,4 @@ def create_video_record_and_dedup(
e,
)
session.rollback()
return 0
return {"video_count": 0, "is_duplicate": False, "batch_similarity": None, "duplicate_of": None}
+190 -55
View File
@@ -383,35 +383,61 @@ def _load_task_info(task_id: str) -> dict | None:
session.close()
def _upload_and_record(
# ── #1743 批量变体重渲/封面判定(纯函数,便于单测) ──────────────────────
BATCH_RENDER_SIMILARITY_LIMIT = 0.20
"""批次内成片查重相似度阈值:超过则重选独立 plan 重渲一次(20%)。"""
def should_rerender_for_batch_dedup(*, batch_id: str, render_attempt: int, batch_similarity) -> bool:
"""批次内查重后判定是否需要重选 plan 重渲。
条件(全部满足才重渲):批次任务、首版(attempt==0)、查重率已得出、相似度 > 20%
非批次任务 / 已是重渲版 / 查重率缺失 / 相似度达标 → 不重渲。
"""
if not batch_id:
return False
if render_attempt >= 1:
return False
if batch_similarity is None:
return False
return float(batch_similarity) > BATCH_RENDER_SIMILARITY_LIMIT
def pick_batch_cover_index(task_id: str, candidate_count: int) -> int:
"""批次变体封面帧选取:按 task_id md5 稳定哈希分散到候选帧。
同任务重试结果稳定;批次内不同 task_id 哈希后分散,避免 N 个变体都抽 frame_0
导致封面雷同。非批次调用方应直接取 0(主流程按 batch_id 区分)。
"""
if candidate_count <= 1:
return 0
import hashlib
return int(hashlib.md5(task_id.encode()).hexdigest(), 16) % candidate_count
def _upload_rendered_video(
task_id: str,
output_path: Path,
project_id: str,
batch_id: str,
editing_mode,
user_id: str = "",
video_name: str = "",
thumbnail_url: str = "",
) -> tuple[str, float, int, int]:
"""上传 OSS、创建视频记录并查重。
*,
attempt: int = 0,
) -> tuple[str, str]:
"""上传成片到 OSS(不落库)。attempt>0 时文件名带轮次后缀,避免覆盖首版。
Returns:
(file_url, duration, file_size, video_count)
Returns: (file_url, storage_key)
"""
# project_id 可能为空(模板编辑器草稿不属于任何项目),过滤空段避免 OSS key 出现 //
path_parts = [p for p in ("generated", "projects", project_id, "tasks", task_id, output_path.name) if p]
suffix = f"_v{attempt}" if attempt > 0 else ""
stem = output_path.stem
name = f"{stem}{suffix}{output_path.suffix or '.mp4'}"
path_parts = [p for p in ("generated", "projects", project_id, "tasks", task_id, name) if p]
storage_key = "/".join(path_parts)
file_size = output_path.stat().st_size
# 上传 OSS
logger.info("[task_id=%s] [OSS上传] 开始上传: size=%d", task_id, file_size)
upload_start = time.monotonic()
logger.info("[task_id=%s] [OSS上传] 开始上传(attempt=%d): size=%d", task_id, attempt, output_path.stat().st_size)
file_url = upload_to_oss(output_path, storage_key)
upload_elapsed = time.monotonic() - upload_start
if not file_url:
raise RuntimeError(f"OSS 上传失败: task_id={task_id}, storage_key={storage_key}")
# 校验 URL 可达性(P0-2: 私有 bucket 用预签名 + object_exists 降级)
verify_url = get_signed_download_url(file_url, expires_seconds=300) or file_url
if not _verify_url_accessible(verify_url):
from video_processing.oss_helpers import normalize_storage_key, oss_bucket
@@ -420,25 +446,58 @@ def _upload_and_record(
key = normalize_storage_key(file_url)
if not (bucket and bucket.object_exists(key)):
raise RuntimeError(
f"OSS 上传后 URL 不可访问且 object_exists 失败: file_url={file_url}, " f"storage_key={storage_key}"
f"OSS 上传后 URL 不可访问且 object_exists 失败: file_url={file_url}, storage_key={storage_key}"
)
logger.info(
"URL 校验失败但 object_exists 确认文件存在,视为上传成功: storage_key=%s",
key,
)
logger.info("URL 校验失败但 object_exists 确认文件存在,视为上传成功: storage_key=%s", key)
return file_url, storage_key
logger.info(
"[task_id=%s] [OSS上传] 成功: 耗时=%.1fs, file_url=%s",
task_id,
upload_elapsed,
file_url,
)
# 创建 GeneratedVideo 记录 + 查重
duration = probe_duration(output_path)
def _reselect_plan_for_batch_retry(task_id: str, plan_id: str, task_info: dict) -> str | None:
"""批次内查重超阈值后,为当前任务重新独立选片生成新 plan(#1743 自动重渲)。
复用 API 侧同一套 EditPlanService.reselect_plan_for_variantpackages 层
variant_plan_selector 纯核心),素材池来自任务 asset_ids + 源 plan 素材。
成功返回新 plan_id;失败返回 None(调用方放弃重渲,保留首版)。
"""
try:
from app.services.edit_plan_service import EditPlanService
db = SessionLocal()
try:
svc = EditPlanService(db)
asset_pool = list(task_info.get("task_asset_ids") or [])
new_plan = svc.reselect_plan_for_variant(
plan_id,
asset_pool,
created_by_user_id=task_info.get("user_id", ""),
name_suffix="重渲变体",
)
return new_plan.id
finally:
db.close()
except Exception:
logger.warning("[task_id=%s] 批次重渲前重选 plan 失败,放弃重渲", task_id, exc_info=True)
return None
def _record_video_and_dedup(
*,
task_id: str,
project_id: str,
batch_id: str,
editing_mode,
user_id: str,
file_url: str,
file_size: int,
video_path: str,
video_name: str = "",
thumbnail_url: str = "",
) -> dict:
"""成片落库 + 指纹查重(含批次内)。返回查重信息 dict。"""
duration = probe_duration(Path(video_path))
dedup_session = SessionLocal()
try:
video_count = create_video_record_and_dedup(
result = create_video_record_and_dedup(
generation_task_id=task_id,
project_id=project_id,
user_id=user_id,
@@ -446,7 +505,7 @@ def _upload_and_record(
file_url=file_url,
file_size=file_size,
duration=duration,
video_path=str(output_path),
video_path=video_path,
mode=editing_mode.value,
session=dedup_session,
name=video_name,
@@ -454,11 +513,12 @@ def _upload_and_record(
)
finally:
dedup_session.close()
return file_url, duration, file_size, video_count or 1
result["duration"] = duration
return result
# ── Celery Task ──────────────────────────────────────────────────────────────
# ── Celery Task ──────────────────────────────────────────────────────────────
def _sync_task_config_to_plan(source_edit_plan_id: str, task_info: dict, db) -> str | None:
@@ -729,19 +789,29 @@ def generate_video(self, task_id: str) -> dict:
gen_task.append_log("渲染模式", "从草稿数据渲染(与预览一致)")
_flush_logs(task_id, gen_task)
output_path, render_duration, cover_candidates, voiceover_tmp_path, render_temp_dir, thumbnail_url = (
_render_from_edit_plan(
# ── 渲染→上传→查重→(批次超阈值则重选 plan 重渲一次)循环(#1743)──
current_plan_id = source_edit_plan_id
file_url = ""
duration = 0.0
file_size = 0
video_count = 1
file_size_final = 0
for render_attempt in range(2): # 首版 + 最多 1 次重渲
(
output_path,
render_duration,
cover_candidates,
voiceover_tmp_path,
render_temp_dir,
thumbnail_url,
) = _render_from_edit_plan(
task_id=task_id,
source_edit_plan_id=source_edit_plan_id,
source_edit_plan_id=current_plan_id,
task_info=task_info,
)
)
# 从这里开始,render_temp_dir 已赋值,必须确保异常时也能清理
try:
if gen_task:
gen_task.append_log("渲染", f"渲染完成, 时长={render_duration:.1f}s")
gen_task.append_log("渲染", f"渲染完成(第{render_attempt + 1}版), 时长={render_duration:.1f}s")
_flush_logs(task_id, gen_task)
_update_task_progress(task_id, 80, "渲染完成")
# ── 3.5 随机边缘裁剪降重(#1664) ──────────────────────────
@@ -751,7 +821,7 @@ def generate_video(self, task_id: str) -> dict:
cropped_path = random_edge_crop(output_path)
if cropped_path != output_path:
output_path = cropped_path
if gen_task:
if gen_task and render_attempt == 0:
gen_task.append_log("边缘裁剪", "已应用随机 2-5% 边缘裁剪降重")
_flush_logs(task_id, gen_task)
logger.info("[task_id=%s] 随机边缘裁剪完成: %s", task_id, output_path)
@@ -762,45 +832,110 @@ def generate_video(self, task_id: str) -> dict:
crop_err,
exc_info=True,
)
if gen_task:
gen_task.append_log("边缘裁剪", f"裁剪失败,使用原始视频: {crop_err}")
_flush_logs(task_id, gen_task)
# ── 4. 上传 OSS + 查重记录 ───────────────────────────────
# ── 4. 上传 OSS(不落库) ───────────────────────────────
_update_task_progress(task_id, 85, "开始上传")
file_url, duration, file_size, video_count = _upload_and_record(
file_url, _storage_key = _upload_rendered_video(
task_id=task_id,
output_path=output_path,
project_id=project_id,
attempt=render_attempt,
)
file_size = output_path.stat().st_size
# ── 4.5 落库 + 查重(批次任务检查批次内相似度) ───────────
dedup_info = _record_video_and_dedup(
task_id=task_id,
project_id=project_id,
batch_id=batch_id,
editing_mode=editing_mode,
user_id=user_id,
file_url=file_url,
file_size=file_size,
video_path=str(output_path),
video_name=task_info.get("video_title", ""),
thumbnail_url=thumbnail_url,
)
duration = dedup_info.get("duration", render_duration)
video_count = dedup_info.get("video_count", 1)
batch_sim = dedup_info.get("batch_similarity")
if gen_task:
gen_task.append_log(
"OSS上传",
f"上传成功, 大小={file_size}",
f"{render_attempt + 1}上传成功, 大小={file_size}"
+ (f", 批次相似度={batch_sim:.0%}" if batch_sim is not None else ""),
file_size=file_size,
file_url=file_url,
)
_flush_logs(task_id, gen_task)
_update_task_progress(task_id, 95, "上传完成")
finally:
# 清理渲染临时目录(无论后续步骤成功与否都清理)
# 非批次 / 相似度达标 / 已是最后一次 → 结束循环
if not should_rerender_for_batch_dedup(
batch_id=batch_id,
render_attempt=render_attempt,
batch_similarity=batch_sim,
):
file_size_final = file_size
break
# 批次内相似度过高:重选独立 plan 后重渲一次
logger.warning(
"[task_id=%s] 批次内查重相似度 %.2f 超阈值 %.2f,重选 plan 重渲",
task_id,
batch_sim,
BATCH_RENDER_SIMILARITY_LIMIT,
)
if gen_task:
gen_task.append_log("批次查重", f"与批次内成片相似度过高({batch_sim:.0%}),重新选片渲染")
_flush_logs(task_id, gen_task)
new_plan_id = _reselect_plan_for_batch_retry(task_id, current_plan_id, task_info)
if not new_plan_id:
logger.warning("[task_id=%s] 重选 plan 失败,保留首版", task_id)
file_size_final = file_size
break
# 回写任务关联的 plan(重渲版以新 plan 渲染)
try:
_ps = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
_pr = SQLAlchemyGenerationTaskRepository(_ps)
_gt = _pr.get(task_id)
if _gt:
_gt.source_edit_plan_id = new_plan_id
_pr.update(_gt)
finally:
_ps.close()
except Exception:
logger.warning("[task_id=%s] 回写重渲 plan_id 失败", task_id, exc_info=True)
current_plan_id = new_plan_id
# 清理本轮临时目录,下一轮重新渲染
if render_temp_dir:
import shutil
shutil.rmtree(render_temp_dir, ignore_errors=True)
logger.info("[task_id=%s] 渲染临时目录已清理: %s", task_id, render_temp_dir)
render_temp_dir = None
file_size = file_size_final or file_size
_update_task_progress(task_id, 95, "上传完成")
# 渲染结束后清理临时目录(重渲循环内每轮已清理,此处兜底最后一轮)
if render_temp_dir:
import shutil
shutil.rmtree(render_temp_dir, ignore_errors=True)
logger.info("[task_id=%s] 渲染临时目录已清理: %s", task_id, render_temp_dir)
# ── 4.5 封面帧持久化 ────────────────────────────────────────────
try:
if cover_candidates:
first = cover_candidates[0]
# #1743:批量变体封面差异化——候选帧按 task_id 稳定哈希分散选取
# (同任务重试稳定,批次内不同任务落在不同帧位),非批次取首帧。
_cover_idx = pick_batch_cover_index(task_id, len(cover_candidates)) if batch_id else 0
first = cover_candidates[_cover_idx]
cover_frame_url = first.get("image_url") or first.get("url") or ""
if cover_frame_url:
_cover_session = SessionLocal()
+250
View File
@@ -0,0 +1,250 @@
"""批量变体独立选片核心(#1743)。
总原则:多视频 = 单视频逻辑 × N。批量正式生成/批量预览时,变体 1..N-1
不再"克隆源 plan 只重算起点"(那会导致同批素材、同顺序、同速度,成片同源),
而是**完整重跑单视频的选片流程**
1. 源 plan 片段骨架(clip_type/order/duration/text/transition)保持不变 —— 保留
模板结构与用户编辑结果;
2. 素材池上做完整随机重选:
- 素材组合随机(shuffle 素材池 + smart_match 评分噪声由调用方排序决定);
- main 片段之间随机洗牌顺序(片段顺序显著不同);
- 起点走场景镜头洗牌 + 随机起点 + 历史已用区间避让
pick_scene_aware_start / _calc_random_start_time,与单视频同一入口);
- 跨变体/跨任务避让:get_used_segments 读取素材 metadata 持久化的已用区间,
record_used_segments 随新片段写回(同事务),N 个变体串行选片时天然互相避让;
3. 批次内片段重叠检查:重选后与"本批次已选定片段"对比,同一 asset 时间区间
重叠占比 > 阈值(默认 20%)则该片段重选起点,最多重试若干次。
本模块只产出 clips_datadict 列表,供 EditPlanService.replace_all_clips_transactional
落库),不碰 DB 事务边界;素材时长/场景点/已用区间由调用方注入,便于单测。
"""
from __future__ import annotations
import logging
import random
from packages.domain.plan_generator_utils import _resolve_start_time
logger = logging.getLogger(__name__)
# ── 阈值常量 ────────────────────────────────────────────────────────────────
BATCH_CLIP_OVERLAP_LIMIT = 0.20
"""批次内同一素材片段时间区间重叠占比上限(20%)。超过则重选起点。"""
VARIANT_RESELECT_MAX_ATTEMPTS = 6
"""单片段重叠避让/起点重选的最大尝试次数。"""
MAIN_CLIP_TYPES = {"main"}
"""参与素材洗牌重选的片段类型(intro/outro/overlay 等固定角色片段保持源 plan)。"""
def _clip_overlap_ratio(
asset_id: str,
start: float,
duration: float,
batch_segments: dict[str, list[tuple[float, float]]],
) -> float:
"""计算新区间 [start, start+duration) 与批次内同素材已选区间的重叠占比。
返回重叠总时长 / 片段时长。
"""
if not asset_id or duration <= 0:
return 0.0
end = start + duration
overlap = 0.0
for seg_start, seg_end in batch_segments.get(asset_id, []):
ov = max(0.0, min(end, seg_end) - max(start, seg_start))
overlap += ov
return min(1.0, overlap / duration)
def reselect_clips_for_variant(
source_clips: list[dict],
candidate_asset_ids: list[str],
*,
asset_durations: dict[str, float],
asset_scene_points: dict[str, list[float]] | None = None,
historical_used_segments: dict[str, list[tuple[float, float]]] | None = None,
batch_segments: dict[str, list[tuple[float, float]]] | None = None,
rng: random.Random | None = None,
) -> list[dict]:
"""为一个变体基于源片段骨架重新独立选片。
Args:
source_clips: 源 plan 片段(dict 列表,每项至少含
order/asset_id/start_time/duration/clip_type,可含
playback_speed/transition_effect/transition_duration/text_content)。
candidate_asset_ids: 素材池(源 plan 素材 ∪ 批次任务素材),将被 shuffle
后随机分配给 main 片段。
asset_durations: {asset_id: 时长秒},起点避让/区间计算必需。
asset_scene_points: {asset_id: 场景切换点},有则走镜头洗牌选起点。
historical_used_segments: 素材 metadata 中持久化的历史已用区间
(跨任务/跨变体避让),函数内会就地追加本变体选中的区间。
batch_segments: 本批次已选片段区间(变体间避让 + 20% 重叠检查),
函数内会就地追加本变体选中的区间。
rng: 可选随机数生成器(测试可注入固定种子)。
Returns:
clips_data: 与源片段等长、order 对齐的新片段 dict 列表。
Raises:
ValueError: 源片段为空 / 素材池为空 / 素材时长全为 0(无法差异化选片)。
"""
rng = rng or random.Random()
if not source_clips:
raise ValueError("源 plan 无片段,无法为变体重新选片")
if not candidate_asset_ids:
raise ValueError("素材池为空,无法为变体独立选片(不允许退回同源成片)")
# 仅保留时长可知(>0)的素材;时长未知无法做区间避让/重叠计算
usable_assets = [a for a in dict.fromkeys(candidate_asset_ids) if asset_durations.get(a, 0.0) > 0]
if not usable_assets:
raise ValueError("素材池时长全部未知(0),无法为变体独立选片")
# 历史已用区间:复制一份,本变体选中的区间就地追加(随 clip record 持久化由调用方负责)
used_segments: dict[str, list[tuple[float, float]]] = (
{k: list(v) for k, v in (historical_used_segments or {}).items()} if historical_used_segments else {}
)
batch_segments = batch_segments if batch_segments is not None else {}
# 按 order 排序源片段,保持骨架顺序
ordered = sorted(source_clips, key=lambda c: c.get("order", 0))
# ── 1. 素材洗牌:素材池 shuffle(组合随机) ─────────────────────────────
shuffled_pool = list(usable_assets)
rng.shuffle(shuffled_pool)
# ── 2. main 片段之间洗牌顺序(顺序随机) ────────────────────────────────
main_indexes = [i for i, c in enumerate(ordered) if c.get("clip_type", "main") in MAIN_CLIP_TYPES]
rng.shuffle(main_indexes)
result: list[dict | None] = [None] * len(ordered)
pool_cursor = 0
for idx in main_indexes:
src = ordered[idx]
dur = float(src.get("duration", 0.0) or 0.0)
if dur <= 0:
# 异常片段:原样保留
result[idx] = _base_clip_data(
src, asset_id=src.get("asset_id", ""), start=float(src.get("start_time", 0.0))
)
continue
# 轮询取洗牌后素材(素材数 < 片段数时循环复用,但组合/顺序已随机)
asset_id = shuffled_pool[pool_cursor % len(shuffled_pool)]
pool_cursor += 1
total = asset_durations.get(asset_id, 0.0)
eff_dur = min(dur, total) if total > 0 else dur
# ── 3. 起点重选(镜头洗牌/随机起点/历史避让)+ 批次重叠避让 ──────────
start = _pick_start_with_overlap_avoid(
asset_id=asset_id,
clip_duration=eff_dur,
asset_durations=asset_durations,
used_segments=used_segments,
asset_scene_points=asset_scene_points,
batch_segments=batch_segments,
rng=rng,
)
interval = (start, start + eff_dur)
used_segments.setdefault(asset_id, []).append(interval)
batch_segments.setdefault(asset_id, []).append(interval)
result[idx] = _base_clip_data(src, asset_id=asset_id, start=start)
# ── 4. 非 main 片段(intro/outro/overlay 等固定角色):保留源素材,仅重算起点 ──
for idx, c in enumerate(ordered):
if result[idx] is not None:
continue
src = c
aid = src.get("asset_id", "")
dur = float(src.get("duration", 0.0) or 0.0)
start = float(src.get("start_time", 0.0))
total = asset_durations.get(aid, 0.0)
if aid and dur > 0 and total > 0:
eff_dur = min(dur, total)
new_start = _pick_start_with_overlap_avoid(
asset_id=aid,
clip_duration=eff_dur,
asset_durations=asset_durations,
used_segments=used_segments,
asset_scene_points=asset_scene_points,
batch_segments=batch_segments,
rng=rng,
)
start = new_start
interval = (start, start + eff_dur)
used_segments.setdefault(aid, []).append(interval)
batch_segments.setdefault(aid, []).append(interval)
result[idx] = _base_clip_data(src, asset_id=aid, start=start)
return [c for c in result if c is not None]
def _pick_start_with_overlap_avoid(
*,
asset_id: str,
clip_duration: float,
asset_durations: dict[str, float],
used_segments: dict[str, list[tuple[float, float]]],
asset_scene_points: dict[str, list[float]] | None,
batch_segments: dict[str, list[tuple[float, float]]],
rng: random.Random,
) -> float:
"""选起点:优先单视频同一入口(镜头洗牌/随机/历史避让),再叠加批次 20% 重叠避让。
批次内重叠超阈值时在素材可用范围内随机抖动重选,最多 VARIANT_RESELECT_MAX_ATTEMPTS 次;
仍超阈值则返回最后一次结果(素材极少时的尽力而为,不阻塞生成)。
"""
total = asset_durations.get(asset_id, 0.0)
max_start = max(0.0, total - clip_duration)
candidate = _resolve_start_time(
asset_id,
clip_duration,
asset_durations,
used_segments,
asset_scene_points,
)
if candidate is None:
candidate = rng.uniform(0.0, max_start) if max_start > 0 else 0.0
best_start = candidate
best_ratio = _clip_overlap_ratio(asset_id, candidate, clip_duration, batch_segments)
if best_ratio <= BATCH_CLIP_OVERLAP_LIMIT:
return candidate
# 重叠超阈值:在可用范围内随机重试
for _ in range(VARIANT_RESELECT_MAX_ATTEMPTS):
alt = rng.uniform(0.0, max_start) if max_start > 0 else 0.0
ratio = _clip_overlap_ratio(asset_id, alt, clip_duration, batch_segments)
if ratio < best_ratio:
best_start, best_ratio = alt, ratio
if ratio <= BATCH_CLIP_OVERLAP_LIMIT:
return alt
logger.info(
"变体选片批次重叠避让达上限,采用最优起点: asset=%s overlap_ratio=%.2f",
asset_id,
best_ratio,
)
return best_start
def _base_clip_data(src: dict, *, asset_id: str, start: float) -> dict:
"""从源片段构造落库 dict(保留骨架/转场/文案/速度,替换素材与起点)。"""
return {
"order": src.get("order", 0),
"asset_id": asset_id,
"start_time": round(float(start), 3),
"duration": float(src.get("duration", 0.0) or 0.0),
"clip_type": src.get("clip_type", "main"),
"playback_speed": float(src.get("playback_speed", 1.0) or 1.0),
"transition_effect": src.get("transition_effect", "cut"),
"transition_duration": float(src.get("transition_duration", 0.0) or 0.0),
"text_content": src.get("text_content", ""),
"config": src.get("config") or {},
}
+33 -5
View File
@@ -347,8 +347,9 @@ class TestCreateGenerationTask:
assert mock_celery.send_task.call_args[0][0] == "worker.generate_video"
@patch("app.core.task_enqueue.celery_app")
def test_create_batch_tasks(self, mock_celery, client):
"""批量创建多个生成任务。"""
def test_create_batch_tasks_without_plan_rejected(self, mock_celery, client):
"""#1743:批量 count=3 但无剪辑计划(未传 source_edit_plan_id 且模板兜底无 plan
→ 400 中断,严禁 N 任务兜底共用同一 plan 产出同源成片。"""
mock_celery.send_task = MagicMock()
resp = client.post(
@@ -361,17 +362,44 @@ class TestCreateGenerationTask:
"count": 3,
},
)
assert resp.status_code == 200
assert resp.status_code == 400
assert "预览" in resp.json()["detail"] or "剪辑计划" in resp.json()["detail"]
mock_celery.send_task.assert_not_called()
@patch("app.core.task_enqueue.celery_app")
def test_create_batch_tasks_with_independent_variant_plans(self, mock_celery, client):
"""#1743:批量 count=3 且有源 plan → 变体 0 用源 plan,变体 1/2 各自 reselect
独立选片,3 个任务关联 3 个不同 plan。"""
mock_celery.send_task = MagicMock()
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = [
MagicMock(id="variant-plan-1"),
MagicMock(id="variant-plan-2"),
]
resp = client.post(
"/api/v1/generation/tasks",
json={
"project_id": "proj-1",
"asset_library_id": "lib-1",
"strategy_id": "strategy-default",
"voice_library_id": "voice-lib-1",
"source_edit_plan_id": "source-plan-1",
"count": 3,
},
)
assert resp.status_code == 200, resp.text
data = resp.json()
assert len(data["items"]) == 3
assert data["total"] == 3
# 验证所有任务都有不同的 ID
task_ids = [t["id"] for t in data["items"]]
assert len(set(task_ids)) == 3
# 同一批次应有相同的 batch_id
# 同一批次 batch_id 相同
batch_ids = [t["batch_id"] for t in data["items"] if t["batch_id"]]
assert len(batch_ids) == 3
assert len(set(batch_ids)) == 1
# 变体 1/2 各自独立选片
assert MockPlanSvc.return_value.reselect_plan_for_variant.call_count == 2
def test_create_task_project_not_found(self, client):
"""项目不存在返回 404。"""
+216 -35
View File
@@ -189,30 +189,32 @@ class TestBatchPreviewRoute:
for i, item in enumerate(resp.items):
assert item.variant_index == i
def test_preview_count_3_clones_three_variant_plans(self):
"""有源 plan 时N=3 克隆 3 个独立变体 plan(预览全部克隆,不用源 plan)"""
def test_preview_count_3_reselects_independent_variant_plans(self):
"""#1743有源 plan 时 N=3,变体0保留源 plan,变体1/2 各自独立选片(reselect)。"""
from app.api.routes.generation_preview import create_preview_generation_task
tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)]
repo = _repo_mock()
cloned_plan_ids = ["clone_1", "clone_2", "clone_3"]
reselect_plan_ids = ["reselect_1", "reselect_2"]
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.side_effect = tasks
with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
clone_results = [MagicMock(id=pid) for pid in cloned_plan_ids]
MockPlanSvc.return_value.clone_plan_for_variant.side_effect = clone_results
reselect_results = [MagicMock(id=pid) for pid in reselect_plan_ids]
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = reselect_results
create_preview_generation_task(
_make_preview_request(preview_count=3),
authenticated_user=_make_user(),
generation_task_repository=repo,
db=MagicMock(),
)
# 克隆被调用 3 次
assert MockPlanSvc.return_value.clone_plan_for_variant.call_count == 3
# 每个任务关联到不同的克隆 plan
for i, task in enumerate(tasks):
assert task.source_edit_plan_id == cloned_plan_ids[i]
# 变体 1..N-1 各独立选片一次(共 2 次);count>1 不再走 clone
assert MockPlanSvc.return_value.reselect_plan_for_variant.call_count == 2
MockPlanSvc.return_value.clone_plan_for_variant.assert_not_called()
# 变体0保留源 plan;变体1/2 关联各自独立选出的 plan
assert tasks[0].source_edit_plan_id == "source_plan"
assert tasks[1].source_edit_plan_id == "reselect_1"
assert tasks[2].source_edit_plan_id == "reselect_2"
def test_preview_variant_titles_injected_per_variant(self):
"""titles[] 按变体注入 title_config.text"""
@@ -335,8 +337,8 @@ class TestBatchPreviewRoute:
)
assert exc.value.status_code == 429
def test_preview_clone_failure_marks_all_failed(self):
"""克隆变体 plan 失败 → 已创建任务全部标记 failed 并 500"""
def test_preview_reselect_failure_marks_all_failed(self):
"""#1743:变体独立选片(reselect)重试仍失败 → 已创建任务全部标记 failed 并 500"""
from app.api.routes.generation_preview import create_preview_generation_task
from fastapi import HTTPException
@@ -345,7 +347,7 @@ class TestBatchPreviewRoute:
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.side_effect = tasks
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.clone_plan_for_variant.side_effect = RuntimeError("db down")
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = RuntimeError("db down")
with pytest.raises(HTTPException) as exc:
create_preview_generation_task(
_make_preview_request(preview_count=3),
@@ -357,6 +359,73 @@ class TestBatchPreviewRoute:
# 所有已创建任务都被标记 failed
assert all(t.status == GenerationTaskStatus.FAILED for t in tasks)
def test_preview_reselect_value_error_returns_400(self):
"""#1743:预览 count>1 reselect 素材不足(ValueError)→ 400,已建任务标 failed。"""
from app.api.routes.generation_preview import create_preview_generation_task
from fastapi import HTTPException
tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)]
repo = _repo_mock()
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.side_effect = tasks
with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = ValueError("素材池为空")
with pytest.raises(HTTPException) as exc:
create_preview_generation_task(
_make_preview_request(preview_count=3),
authenticated_user=_make_user(),
generation_task_repository=repo,
db=MagicMock(),
)
assert exc.value.status_code == 400
assert "无法独立选片" in exc.value.detail
assert all(t.status == GenerationTaskStatus.FAILED for t in tasks)
def test_preview_reselect_retry_exhausted_returns_500(self):
"""#1743:预览 count>1 reselect 连续失败(非 ValueError)→ 500,已建任务标 failed。"""
from app.api.routes.generation_preview import create_preview_generation_task
from fastapi import HTTPException
tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)]
repo = _repo_mock()
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.side_effect = tasks
with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = RuntimeError("db down")
with pytest.raises(HTTPException) as exc:
create_preview_generation_task(
_make_preview_request(preview_count=3),
authenticated_user=_make_user(),
generation_task_repository=repo,
db=MagicMock(),
)
assert exc.value.status_code == 500
assert all(t.status == GenerationTaskStatus.FAILED for t in tasks)
def test_preview_count1_with_source_plan_clones(self):
"""#1743 零回归:预览 count=1 且有源 plan 仍走 clone(不 reselect)。"""
from app.api.routes.generation_preview import create_preview_generation_task
task = _make_task(task_id="task_1", source_plan_id="source_plan")
repo = _repo_mock()
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.return_value = task
with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.clone_plan_for_variant.return_value = MagicMock(id="clone_1")
resp = create_preview_generation_task(
_make_preview_request(preview_count=1),
authenticated_user=_make_user(),
generation_task_repository=repo,
db=MagicMock(),
)
assert resp.total == 1
MockPlanSvc.return_value.clone_plan_for_variant.assert_called_once()
MockPlanSvc.return_value.reselect_plan_for_variant.assert_not_called()
assert task.source_edit_plan_id == "clone_1"
# ════════════════════════════════════════════════════════════════════════════
# 批量正式生成:变体配置注入
@@ -366,7 +435,7 @@ class TestBatchPreviewRoute:
class TestBatchGenerationVariantConfig:
"""POST /tasks count=N 时变体独立配置。"""
def _call_create_tasks(self, request, repo=None):
def _call_create_tasks(self, request, repo=None, db_latest_plan=None):
from app.api.routes.generation_tasks import create_generation_task
repo = repo or MagicMock()
@@ -379,9 +448,9 @@ class TestBatchGenerationVariantConfig:
asset_repo = MagicMock()
asset_repo.find_by_id.return_value = None
# db.query().filter()...first() 返回 None:不走兜底关联编辑计划
# db.query().filter()...first()db_latest_plan 非空时模拟模板兜底查到最新 plan
db = MagicMock()
db.query.return_value.filter.return_value.order_by.return_value.first.return_value = None
db.query.return_value.filter.return_value.order_by.return_value.first.return_value = db_latest_plan
return create_generation_task(
request,
@@ -407,20 +476,29 @@ class TestBatchGenerationVariantConfig:
t.title_config = cmd.title_config
t.voice_library_id = cmd.voice_library_id
t.cover_url = cmd.cover_url
t.source_edit_plan_id = cmd.source_edit_plan_id # #1743usecase 落库关联 plan
return t
MockUC.return_value.execute.side_effect = _execute
with patch.object(routes, "safe_enqueue_generation_task", return_value=True):
req = CreateGenerationTaskRequest(
template_id="tpl_1",
asset_ids=["a1"],
count=3,
title_config={"font": "宋体"},
titles=["成片标题1", "成片标题2", "成片标题3"],
voice_library_ids=["v1", "v2", "v3"],
cover_urls=["http://c1", "http://c2", "http://c3"],
)
resp = self._call_create_tasks(req)
# #1743count>1 必须有源 plan,变体 1..N-1 走 reselect 独立选片
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = [
MagicMock(id="reselect_1"),
MagicMock(id="reselect_2"),
]
req = CreateGenerationTaskRequest(
template_id="tpl_1",
asset_ids=["a1"],
count=3,
source_edit_plan_id="source_plan",
title_config={"font": "宋体"},
titles=["成片标题1", "成片标题2", "成片标题3"],
voice_library_ids=["v1", "v2", "v3"],
cover_urls=["http://c1", "http://c2", "http://c3"],
)
resp = self._call_create_tasks(req)
assert MockPlanSvc.return_value.reselect_plan_for_variant.call_count == 2
assert resp.total == 3
assert [c.title_config["text"] for c in captured] == ["成片标题1", "成片标题2", "成片标题3"]
assert [c.voice_library_id for c in captured] == ["v1", "v2", "v3"]
@@ -467,21 +545,124 @@ class TestBatchGenerationVariantConfig:
def _execute(cmd):
captured.append(cmd)
return tasks[len(captured) - 1]
t = tasks[len(captured) - 1]
t.source_edit_plan_id = cmd.source_edit_plan_id # #1743usecase 落库关联 plan
return t
MockUC.return_value.execute.side_effect = _execute
with patch.object(routes, "safe_enqueue_generation_task", return_value=True):
req = CreateGenerationTaskRequest(
template_id="tpl_1",
asset_ids=["a1"],
count=3,
voice_library_ids=["shared_voice"],
cover_urls=["http://shared"],
)
self._call_create_tasks(req)
# #1743count>1 必须有源 plan,变体 1..N-1 走 reselect 独立选片
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = [
MagicMock(id="reselect_1"),
MagicMock(id="reselect_2"),
]
req = CreateGenerationTaskRequest(
template_id="tpl_1",
asset_ids=["a1"],
count=3,
source_edit_plan_id="source_plan",
voice_library_ids=["shared_voice"],
cover_urls=["http://shared"],
)
self._call_create_tasks(req)
assert MockPlanSvc.return_value.reselect_plan_for_variant.call_count == 2
assert all(c.voice_library_id == "shared_voice" for c in captured)
assert all(c.cover_url == "http://shared" for c in captured)
def test_count3_template_fallback_plan_used(self):
"""#1743:未传 source_edit_plan_id 时,模板兜底查到最新 plan 即作为批量源 plan。"""
from app.api.routes import generation_tasks as routes
tasks = [_make_task(task_id=f"gen_fb_{i}") for i in range(3)]
def _execute(cmd):
t = tasks[len([c for c in getattr(_execute, "caps", [])])]
t.source_edit_plan_id = cmd.source_edit_plan_id
_execute.caps.append(cmd)
return t
_execute.caps = []
latest = MagicMock(id="fallback_plan_id")
req = CreateGenerationTaskRequest(template_id="tpl_1", asset_ids=["a1"], count=3)
with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.side_effect = _execute
with patch.object(routes, "safe_enqueue_generation_task", return_value=True):
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = [
MagicMock(id="reselect_1"),
MagicMock(id="reselect_2"),
]
self._call_create_tasks(req, db_latest_plan=latest)
# 兜底 plan 被用作源;变体0关联兜底 plan,变体1/2关联 reselect plan
assert _execute.caps[0].source_edit_plan_id == "fallback_plan_id"
assert _execute.caps[1].source_edit_plan_id == "reselect_1"
assert _execute.caps[2].source_edit_plan_id == "reselect_2"
def test_count3_reselect_value_error_returns_400(self):
"""#1743reselect 素材不足(ValueError)→ 400 明确报错,零任务入队。"""
from app.api.routes import generation_tasks as routes
from fastapi import HTTPException
enqueue = MagicMock(return_value=True)
req = CreateGenerationTaskRequest(
template_id="tpl_1", asset_ids=["a1"], count=3, source_edit_plan_id="source_plan"
)
with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.side_effect = lambda cmd: _make_task(task_id="should_not_run")
with patch.object(routes, "safe_enqueue_generation_task", enqueue):
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = ValueError("素材池为空")
with pytest.raises(HTTPException) as exc:
self._call_create_tasks(req)
assert exc.value.status_code == 400
assert "无法独立选片" in exc.value.detail
enqueue.assert_not_called()
def test_count3_reselect_retry_exhausted_returns_500(self):
"""#1743reselect 连续 2 次都非 ValueError 失败 → 500,零任务入队。"""
from app.api.routes import generation_tasks as routes
from fastapi import HTTPException
enqueue = MagicMock(return_value=True)
req = CreateGenerationTaskRequest(
template_id="tpl_1", asset_ids=["a1"], count=3, source_edit_plan_id="source_plan"
)
with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.side_effect = lambda cmd: _make_task(task_id="should_not_run")
with patch.object(routes, "safe_enqueue_generation_task", enqueue):
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = RuntimeError("db down")
with pytest.raises(HTTPException) as exc:
self._call_create_tasks(req)
assert exc.value.status_code == 500
enqueue.assert_not_called()
def test_count3_missing_plan_after_prebuild_raises_500(self):
"""#1743 兜底守卫:任务落库时 plan 丢失(usecase 未透传)→ 500,严禁静默共用。"""
from app.api.routes import generation_tasks as routes
from fastapi import HTTPException
req = CreateGenerationTaskRequest(
template_id="tpl_1", asset_ids=["a1"], count=3, source_edit_plan_id="source_plan"
)
with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC:
# usecase 返回的任务 source_edit_plan_id 为空(模拟落库丢 plan
MockUC.return_value.execute.side_effect = lambda cmd: _make_task(task_id="lost_plan")
with patch.object(routes, "safe_enqueue_generation_task", return_value=True):
with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
MockPlanSvc.return_value.reselect_plan_for_variant.side_effect = [
MagicMock(id="reselect_1"),
MagicMock(id="reselect_2"),
]
with pytest.raises(HTTPException) as exc:
self._call_create_tasks(req)
assert exc.value.status_code == 500
assert "变体剪辑计划缺失" in exc.value.detail
class TestVariantValueHelper:
"""_variant_value 取值逻辑。"""
+244
View File
@@ -0,0 +1,244 @@
"""#1743 dedup_helpers 批次查重 + worker 重选 plan 重试函数测试。
覆盖:
- 批次任务且无历史重复时走 check_batch_duplicatebatch_similarity 透传到返回 dict
- 批次任务历史已重复 → 不再做批次查重,is_duplicate=True / duplicate_of 透传
- 非批次任务 batch_similarity 恒为 None
- _reselect_plan_for_batch_retry:成功返回新 plan_id;异常返回 None(保留首版)
"""
from __future__ import annotations
import os
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
for sub in ("apps/worker", "apps/api", "packages", ""):
p = str(REPO_ROOT / sub) if sub else str(REPO_ROOT)
if p not in sys.path:
sys.path.insert(0, p)
from sqlalchemy import create_engine # noqa: E402
from sqlalchemy.orm import sessionmaker # noqa: E402
os.environ.setdefault("DATABASE_URL", "sqlite:///:memory:")
from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
@pytest.fixture
def session():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(engine)
s = sessionmaker(bind=engine)()
try:
yield s
finally:
s.close()
def _dedup_kwargs(batch_id):
return dict(
generation_task_id="task-batch-1",
project_id="proj-1",
user_id="user-1",
batch_id=batch_id,
file_url="https://oss.example.com/v.mp4",
file_size=1024,
duration=10.0,
video_path="/tmp/fake.mp4",
mode="edit_plan",
width=1280,
height=720,
fps=25.0,
)
class TestBatchDedupInHelpers:
"""dedup_helpers 内部 lazy import video_processing.dedup(依赖 cv2),
复用 test_generated_video_creation_logic 的 mock 模块注册模式。"""
@classmethod
def setup_class(cls):
import sys
if "cv2" not in sys.modules:
sys.modules["cv2"] = MagicMock()
# 保存 setup 前的真实模块引用,teardown 原样还原(沙箱无 cv2 时 reimport 会失败,
# 若直接 pop 掉 MagicMock 会让后续测试(如 test_dedup_1702)拿到残缺模块)
import video_processing
cls._orig_dedup = sys.modules.get("video_processing.dedup")
cls._orig_dedup_attr = getattr(video_processing, "dedup", None)
cls._orig_thumb = sys.modules.get("video_processing.thumbnail_generator")
cls._orig_thumb_attr = getattr(video_processing, "thumbnail_generator", None)
mock_dedup = MagicMock()
mock_dedup.VideoDeduplicator = MagicMock()
sys.modules["video_processing.dedup"] = mock_dedup
mock_thumb = MagicMock()
mock_thumb.extract_first_frame = MagicMock()
sys.modules["video_processing.thumbnail_generator"] = mock_thumb
video_processing.dedup = mock_dedup
video_processing.thumbnail_generator = mock_thumb
@classmethod
def teardown_class(cls):
import sys
import video_processing
# 还原 setup 前状态:原本有真模块→放回;原本没有→移除 mock
if cls._orig_dedup is not None:
sys.modules["video_processing.dedup"] = cls._orig_dedup
else:
sys.modules.pop("video_processing.dedup", None)
if cls._orig_dedup_attr is not None:
video_processing.dedup = cls._orig_dedup_attr
elif hasattr(video_processing, "dedup"):
delattr(video_processing, "dedup")
if cls._orig_thumb is not None:
sys.modules["video_processing.thumbnail_generator"] = cls._orig_thumb
else:
sys.modules.pop("video_processing.thumbnail_generator", None)
if cls._orig_thumb_attr is not None:
video_processing.thumbnail_generator = cls._orig_thumb_attr
elif hasattr(video_processing, "thumbnail_generator"):
delattr(video_processing, "thumbnail_generator")
def test_batch_similarity_returned_when_batch_duplicate_found(self, session):
"""批次任务 + 无历史重复 + 批次查重命中 → 返回 batch_similarity 与 is_duplicate。"""
from video_processing.dedup_helpers import create_video_record_and_dedup
with patch("video_processing.dedup.VideoDeduplicator") as mock_cls:
dedup = mock_cls.return_value
dedup.compute_fingerprint.return_value = MagicMock(to_dict=lambda: {})
dedup.check_duplicate.return_value = None # 历史无重复
dedup.check_batch_duplicate.return_value = {
"duplicate_of": "video-existing",
"reason": "batch_similar",
"similarity": 0.601,
}
dedup.compute_duplicate_rate.return_value = {
"duplicate_rate": 0.0,
"visual_similarity": 0.0,
"match_count": 0,
}
result = create_video_record_and_dedup(session=session, **_dedup_kwargs("batch-abc"))
assert result["video_count"] == 1
assert result["batch_similarity"] == pytest.approx(0.601)
assert result["is_duplicate"] is True
assert result["duplicate_of"] == "video-existing"
# 批次查重确实被调用(历史查重为 None 才走批次)
dedup.check_batch_duplicate.assert_called_once()
def test_batch_check_skipped_when_historical_duplicate(self, session):
"""历史查重已命中 → 不再批次查重,batch_similarity 为 None。"""
from video_processing.dedup_helpers import create_video_record_and_dedup
with patch("video_processing.dedup.VideoDeduplicator") as mock_cls:
dedup = mock_cls.return_value
dedup.compute_fingerprint.return_value = MagicMock(to_dict=lambda: {})
dedup.check_duplicate.return_value = {
"duplicate_of": "video-old",
"reason": "global",
"similarity": 0.85,
}
dedup.compute_duplicate_rate.return_value = {
"duplicate_rate": 85.0,
"visual_similarity": 0.85,
"match_count": 3,
}
result = create_video_record_and_dedup(session=session, **_dedup_kwargs("batch-abc"))
assert result["is_duplicate"] is True
assert result["duplicate_of"] == "video-old"
assert result["batch_similarity"] is None
dedup.check_batch_duplicate.assert_not_called()
def test_non_batch_never_runs_batch_check(self, session):
"""非批次任务(batch_id 为空)→ 不调用批次查重,batch_similarity None。"""
from video_processing.dedup_helpers import create_video_record_and_dedup
with patch("video_processing.dedup.VideoDeduplicator") as mock_cls:
dedup = mock_cls.return_value
dedup.compute_fingerprint.return_value = MagicMock(to_dict=lambda: {})
dedup.check_duplicate.return_value = None
dedup.check_batch_duplicate.return_value = {"similarity": 0.99}
dedup.compute_duplicate_rate.return_value = {
"duplicate_rate": 0.0,
"visual_similarity": 0.0,
"match_count": 0,
}
result = create_video_record_and_dedup(session=session, **_dedup_kwargs(""))
assert result["video_count"] == 1
assert result["batch_similarity"] is None
assert result["is_duplicate"] is False
dedup.check_batch_duplicate.assert_not_called()
def test_batch_no_duplicate_returns_none_similarity(self, session):
"""批次任务但批次查重也未命中 → batch_similarity None(驱动不重渲)。"""
from video_processing.dedup_helpers import create_video_record_and_dedup
with patch("video_processing.dedup.VideoDeduplicator") as mock_cls:
dedup = mock_cls.return_value
dedup.compute_fingerprint.return_value = MagicMock(to_dict=lambda: {})
dedup.check_duplicate.return_value = None
dedup.check_batch_duplicate.return_value = None
dedup.compute_duplicate_rate.return_value = {
"duplicate_rate": 0.0,
"visual_similarity": 0.0,
"match_count": 0,
}
result = create_video_record_and_dedup(session=session, **_dedup_kwargs("batch-xyz"))
assert result["batch_similarity"] is None
assert result["is_duplicate"] is False
class TestReselectPlanForBatchRetry:
def test_success_returns_new_plan_id(self):
"""重选成功 → 返回新 plan_id。"""
from worker_app.tasks.generation import _reselect_plan_for_batch_retry
fake_svc = MagicMock()
fake_svc.reselect_plan_for_variant.return_value = MagicMock(id="new-plan-999")
with patch("app.services.edit_plan_service.EditPlanService", return_value=fake_svc):
with patch("worker_app.tasks.generation.SessionLocal") as mock_session_local:
mock_session_local.return_value = MagicMock()
new_id = _reselect_plan_for_batch_retry(
"task-1",
"old-plan",
{"task_asset_ids": ["a1", "a2"], "user_id": "u-1"},
)
assert new_id == "new-plan-999"
fake_svc.reselect_plan_for_variant.assert_called_once()
args, kwargs = fake_svc.reselect_plan_for_variant.call_args
assert args[0] == "old-plan"
assert args[1] == ["a1", "a2"]
assert kwargs["created_by_user_id"] == "u-1"
def test_failure_returns_none_and_keeps_first_version(self):
"""重选抛异常 → 返回 None(调用方放弃重渲、保留首版),不抛出。"""
from worker_app.tasks.generation import _reselect_plan_for_batch_retry
fake_svc = MagicMock()
fake_svc.reselect_plan_for_variant.side_effect = RuntimeError("db down")
with patch("app.services.edit_plan_service.EditPlanService", return_value=fake_svc):
with patch("worker_app.tasks.generation.SessionLocal") as mock_session_local:
mock_session_local.return_value = MagicMock()
result = _reselect_plan_for_batch_retry("task-1", "old-plan", {"task_asset_ids": [], "user_id": "u-1"})
assert result is None
+108 -64
View File
@@ -1,7 +1,9 @@
"""AI Review 回归:批量生成 count>1 但 source_edit_plan_id 为空时不应 IndexError
"""#1743 批量生成无可用 plan 守卫
变体 plan 预克隆仅在 source_edit_plan_id 非空时执行;无源 plan 时
variant_plan_ids 为空,循环中禁止索引访问,各任务走自身随机选片流程。
新规则(P0 降重):count>1 批量生成时必须存在源 plan(前端传入或按模板兜底
解析到最新 plan),为每个变体独立选片;**无任何可用 plan 时直接 4xx 中断、
不创建任务**,严禁 N 个任务兜底共用同一 plan 产出同源成片。
N=1 单视频不受影响(无 plan 时走原有单任务流程)。
"""
import sys
@@ -9,6 +11,9 @@ from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from fastapi import HTTPException
REPO_ROOT = Path(__file__).resolve().parents[2]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
@@ -18,78 +23,117 @@ def _make_user():
return SimpleNamespace(user=SimpleNamespace(id="user-1"))
def _make_request(count):
def _make_request(count, **overrides):
from app.schemas.generation_task import CreateGenerationTaskRequest
return CreateGenerationTaskRequest(
fields = dict(
project_id="proj-1",
asset_library_id="lib-1",
strategy_id="one_take",
asset_ids=["a1"],
count=count,
source_edit_plan_id="", # 关键:无源 plan(空字符串为假值)
source_edit_plan_id="",
)
fields.update(overrides)
return CreateGenerationTaskRequest(**fields)
class TestBatchNoSourcePlanNoIndexError:
def test_count3_without_source_plan_creates_three_tasks(self):
"""count=3 且无 source_edit_plan_id:不克隆、不 IndexError、创建 3 个任务。"""
def _common_patches(latest_plan=None):
"""构造通用 patch 上下文(repo/usecase/enqueue 等)。latest_plan 为模板兜底 plan 或 None。"""
repo = MagicMock()
repo.count_pending_by_user.return_value = 0
repo.count_pending_total.return_value = 0
repo.create.side_effect = lambda t: t
repo.update.side_effect = lambda t: t
created = []
def _fake_execute(cmd):
task = MagicMock()
task.id = f"task-{len(created) + 1}"
task.source_edit_plan_id = cmd.source_edit_plan_id
task.status = "pending"
task.progress = 0.0
task.strategy_id = "one_take"
task.error_message = ""
task.cover_url = ""
task.title_config = {}
task.created_at = None
task.batch_id = "batch-1"
created.append(task)
return task
db = MagicMock()
# 模板兜底查最新 plan:返回 latest_planNone 表示查不到)
db.query.return_value.filter.return_value.order_by.return_value.first.return_value = latest_plan
mock_uc = patch("app.api.routes.generation_tasks.CreateGenerationTaskUseCase")
other_patches = [
patch("app.api.routes.generation_tasks.safe_enqueue_generation_task", return_value=True),
patch("app.api.routes.generation_tasks._writeback_edit_plan_config"),
patch(
"app.api.routes.generation_tasks._resolve_project_and_library",
return_value=("proj-1", ""),
),
]
return repo, db, created, mock_uc, other_patches, _fake_execute
class TestBatchNoSourcePlanGuard:
def test_count3_without_any_plan_rejects_4xx_and_creates_nothing(self):
"""count=3 且无源 plan、模板兜底也查不到 → 400 中断,零任务创建(严禁同源成片)。"""
from app.api.routes.generation_tasks import create_generation_task
repo = MagicMock()
repo.count_pending_by_user.return_value = 0
repo.count_pending_total.return_value = 0
repo.create.side_effect = lambda t: t
repo.update.side_effect = lambda t: t
repo, db, created, mock_uc, other_patches, fake_exec = _common_patches(latest_plan=None)
MockUC = mock_uc.start()
MockUC.return_value.execute.side_effect = fake_exec
for p in other_patches:
p.start()
all_patches = [mock_uc] + other_patches
try:
with pytest.raises(HTTPException) as exc_info:
create_generation_task(
_make_request(3),
authenticated_user=_make_user(),
generation_task_repository=repo,
project_repository=MagicMock(),
asset_repository=MagicMock(),
asset_library_repository=MagicMock(),
db=db,
)
assert exc_info.value.status_code == 400
finally:
for p in reversed(all_patches):
p.stop()
assert len(created) == 0, "无 plan 批量必须零任务创建"
created = []
def test_count1_without_plan_still_works(self):
"""N=1 单视频无 plan 不触发批量守卫(向后兼容,不 4xx)。"""
from app.api.routes.generation_tasks import create_generation_task
def _fake_execute(cmd):
task = MagicMock()
task.id = f"task-{len(created) + 1}"
task.source_edit_plan_id = cmd.source_edit_plan_id
task.status = "pending"
task.progress = 0.0
task.strategy_id = "one_take"
task.error_message = ""
task.cover_url = None
task.title_config = {}
task.created_at = None
task.batch_id = "batch-1"
created.append(task)
return task
with patch("app.api.routes.generation_tasks.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.side_effect = _fake_execute
with patch(
"app.api.routes.generation_tasks.safe_enqueue_generation_task",
return_value=True,
):
with patch("app.api.routes.generation_tasks._writeback_edit_plan_config"):
with patch(
"app.api.routes.generation_tasks._resolve_project_and_library",
return_value=("proj-1", ""),
):
# 核心断言:不得抛 IndexError(变体 plan 索引守卫)。
# 响应序列化字段与本回归无关,ValidationError 可接受,
# 但 IndexError 必须不出现。
try:
create_generation_task(
_make_request(3),
authenticated_user=_make_user(),
generation_task_repository=repo,
project_repository=MagicMock(),
asset_repository=MagicMock(),
asset_library_repository=MagicMock(),
db=MagicMock(),
)
except IndexError as exc: # pragma: no cover - 不应发生
pytest.fail(f"无源 plan 批量生成触发 IndexError: {exc}")
except Exception:
# 响应序列化等其他异常与本次守卫无关,忽略
pass
# 3 个任务全部创建(未因 IndexError 中断)
assert len(created) == 3
# 无源 plan 时所有任务 source_edit_plan_id 均为空
assert all(not t.source_edit_plan_id for t in created)
repo, db, created, mock_uc, other_patches, fake_exec = _common_patches(latest_plan=None)
MockUC = mock_uc.start()
MockUC.return_value.execute.side_effect = fake_exec
for p in other_patches:
p.start()
all_patches = [mock_uc] + other_patches
try:
create_generation_task(
_make_request(1),
authenticated_user=_make_user(),
generation_task_repository=repo,
project_repository=MagicMock(),
asset_repository=MagicMock(),
asset_library_repository=MagicMock(),
db=db,
)
except HTTPException as e:
assert e.status_code != 400, f"N=1 不应被批量守卫拦截: {e.detail}"
except Exception:
# MagicMock 任务对象下游响应序列化可能抛 ValidationError 等,与批量守卫无关;
# 任务已在 usecase.execute 中创建,下方断言 created==1 即证明守卫未拦截。
pass
finally:
for p in reversed(all_patches):
p.stop()
assert len(created) == 1
@@ -0,0 +1,80 @@
"""#1743 worker 批量重渲判定 + 封面哈希选帧纯函数测试。
- should_rerender_for_batch_dedup:批次首版查重率 >20% 才重渲(非批次/重渲版/无查重率不重渲)
- pick_batch_cover_indextask_id md5 稳定哈希分散候选帧(同任务稳定、跨任务分散)
"""
from __future__ import annotations
import hashlib
import sys
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
for sub in ("apps/worker", "apps/api", "packages", ""):
p = str(REPO_ROOT / sub) if sub else str(REPO_ROOT)
if p not in sys.path:
sys.path.insert(0, p)
from worker_app.tasks.generation import ( # noqa: E402
BATCH_RENDER_SIMILARITY_LIMIT,
pick_batch_cover_index,
should_rerender_for_batch_dedup,
)
class TestShouldRerenderForBatchDedup:
def test_non_batch_never_rerenders(self):
"""非批次任务(batch_id 为空)任何查重率都不重渲。"""
assert should_rerender_for_batch_dedup(batch_id="", render_attempt=0, batch_similarity=0.99) is False
assert should_rerender_for_batch_dedup(batch_id="", render_attempt=0, batch_similarity=None) is False
def test_batch_first_version_over_threshold_rerenders(self):
"""批次首版(attempt=0)查重率 >20% → 重渲。"""
assert should_rerender_for_batch_dedup(batch_id="batch-1", render_attempt=0, batch_similarity=0.601) is True
assert should_rerender_for_batch_dedup(batch_id="batch-1", render_attempt=0, batch_similarity=0.21) is True
def test_batch_first_version_under_threshold_no_rerender(self):
"""批次首版查重率 ≤20% → 不重渲。"""
assert should_rerender_for_batch_dedup(batch_id="batch-1", render_attempt=0, batch_similarity=0.20) is False
assert should_rerender_for_batch_dedup(batch_id="batch-1", render_attempt=0, batch_similarity=0.0) is False
assert should_rerender_for_batch_dedup(batch_id="batch-1", render_attempt=0, batch_similarity=0.05) is False
def test_rerendered_version_never_rerenders_again(self):
"""重渲版(attempt=1)即使仍超阈值也不再重渲(最多重渲一次)。"""
assert should_rerender_for_batch_dedup(batch_id="batch-1", render_attempt=1, batch_similarity=0.99) is False
def test_missing_similarity_no_rerender(self):
"""查重率缺失(None,非批次查重路径)→ 不重渲。"""
assert should_rerender_for_batch_dedup(batch_id="batch-1", render_attempt=0, batch_similarity=None) is False
def test_threshold_constant_is_20_percent(self):
assert BATCH_RENDER_SIMILARITY_LIMIT == 0.20
class TestPickBatchCoverIndex:
def test_stable_for_same_task(self):
"""同一 task_id 多次调用结果稳定(重试封面不变)。"""
first = pick_batch_cover_index("task-abc", 3)
for _ in range(5):
assert pick_batch_cover_index("task-abc", 3) == first
def test_zero_or_one_candidate_returns_zero(self):
assert pick_batch_cover_index("task-x", 0) == 0
assert pick_batch_cover_index("task-x", 1) == 0
def test_index_within_range(self):
for i in range(20):
idx = pick_batch_cover_index(f"task-{i}", 3)
assert 0 <= idx < 3
def test_matches_md5_formula(self):
"""与主流程公式一致:md5(task_id) % 候选数。"""
task_id = "task-formula-check"
expected = int(hashlib.md5(task_id.encode()).hexdigest(), 16) % 3
assert pick_batch_cover_index(task_id, 3) == expected
def test_batch_tasks_spread_across_frames(self):
"""30 个批次任务在 3 个候选帧上分散(不能全部落在 frame_0——旧 bug 回归守卫)。"""
indexes = {pick_batch_cover_index(f"batch-task-{i}", 3) for i in range(30)}
assert len(indexes) >= 2, f"封面帧应分散到多个帧位,实际全部落在: {indexes}"
+14 -4
View File
@@ -85,7 +85,10 @@ class TestTwoPhaseCommit:
session=session,
)
assert result == 1
# #1743:返回值由 int 改为 dict(含批次查重信息)
assert isinstance(result, dict)
assert result["video_count"] == 1
assert result["batch_similarity"] is None # 非批次任务无批次相似度
# create() should be called exactly once with the complete video object
mock_repo.create.assert_called_once()
created_video = mock_repo.create.call_args[0][0]
@@ -122,7 +125,10 @@ class TestTwoPhaseCommit:
session=session,
)
assert result == 1
# #1743:返回值由 int 改为 dict(含批次查重信息)
assert isinstance(result, dict)
assert result["video_count"] == 1
assert result["batch_similarity"] is None # 非批次任务无批次相似度
mock_repo.create.assert_called_once()
created_video = mock_repo.create.call_args[0][0]
assert created_video.duplicate_rate is None
@@ -161,7 +167,10 @@ class TestTwoPhaseCommit:
session=session,
)
assert result == 1
# #1743:返回值由 int 改为 dict(含批次查重信息)
assert isinstance(result, dict)
assert result["video_count"] == 1
assert result["batch_similarity"] is None # 非批次任务无批次相似度
mock_repo.create.assert_called_once()
created_video = mock_repo.create.call_args[0][0]
# Fingerprint should be set
@@ -229,6 +238,7 @@ class TestTwoPhaseCommit:
session=session,
)
assert result == 0
# #1743:失败路径返回 dictvideo_count=0
assert result["video_count"] == 0
session.commit.assert_not_called()
session.rollback.assert_called_once()
@@ -381,7 +381,8 @@ class TestThumbnailInDedupHelpers:
thumbnail_url=pre_thumb_url,
)
assert result == 1
# #1743:返回 dict(非批次 batch_similarity=None
assert result["video_count"] == 1
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
@@ -427,7 +428,8 @@ class TestThumbnailInDedupHelpers:
fps=25.0,
)
assert result == 1
# #1743:返回 dict(非批次 batch_similarity=None
assert result["video_count"] == 1
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
@@ -474,7 +476,7 @@ class TestThumbnailInDedupHelpers:
fps=25.0,
)
assert result == 1 # 不阻断
assert result["video_count"] == 1 # 不阻断#1743 dict 返回)
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
@@ -0,0 +1,206 @@
"""#1743 EditPlanService.reselect_plan_for_variant 服务层测试。
与 clone_plan_for_variant(只重算起点、素材/顺序不变)不同,reselect 完整重跑
单视频选片:素材池 shuffle + main 片段顺序洗牌 + 起点重选 + 批次 20% 重叠避让。
覆盖:
- 独立 plan:新 plan_id 与源不同、命名带后缀、模板/config 复制
- 新片段经 replace_all_clips_transactional 落库,素材/起点与源 plan 存在差异
- 源 plan 区间作为批次避让初始对象;record_used_segments 随新片段写回
- 源 plan 无片段 → ValueError(不创建同源变体)
"""
from __future__ import annotations
import os
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
REPO_ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(REPO_ROOT / "apps" / "api"))
sys.path.insert(0, str(Path(__file__).resolve().parent)) # tests/unit
from test_edit_plan_service import StubEditPlanClipRepository, StubEditPlanRepository, _make_service # noqa: E402
from packages.domain.edit_plan_clip import EditPlanClip # noqa: E402
@pytest.fixture
def svc_with_source():
"""带源 plan4 个 main 片段,素材 a1/a2/a3/a4+ 2 个批次候选素材的 service。"""
svc = _make_service()
svc._clip_repo.session = MagicMock()
source = svc.create_plan(template_id="tpl-001", name="9/6-草稿", total_duration=20.0, config={"title": "源配置"})
for i, aid in enumerate(["a1", "a2", "a3", "a4"]):
clip = EditPlanClip.create(
plan_id=source.id,
clip_type="main",
order=i,
asset_id=aid,
start_time=float(i * 5),
duration=5.0,
text_content=f"文案{i}",
)
svc._clip_repo.create(clip)
return svc, source
def _patch_deps(svc, durations):
"""统一 patch reselect 的 DB/素材/历史区间依赖。返回 (patches, mock_replace)。"""
asset_models = []
for aid, dur in durations.items():
m = MagicMock(id=aid)
m.duration = dur
m.metadata = None # extract_scene_points_from_metadata(None) → 无场景点
asset_models.append(m)
mock_replace = patch.object(svc, "replace_all_clips_transactional", return_value=4)
patches = [
patch("app.services.edit_plan_service.get_used_segments", return_value={}),
patch("app.services.edit_plan_service.record_used_segments", return_value=None),
patch("packages.domain.plan_generator_utils.extract_scene_points_from_metadata", return_value=[]),
patch("packages.adapters.sqlalchemy_impl.models.AssetModel", create=True),
mock_replace,
]
started = []
for p in patches:
started.append(p.start())
# started[-1] 是 replace_all 的 MagicMock
mock_replace_obj = started[-1]
# db.query(AssetModel).filter(...).all() → 带 duration 的素材 mock
svc._clip_repo.session.query.return_value.filter.return_value.all.return_value = asset_models
return patches, mock_replace_obj
class TestReselectPlanForVariant:
def test_creates_independent_plan_with_different_clips(self, svc_with_source):
"""reselect 产出新 plan(id/名称不同),片段素材或起点与源 plan 存在差异。"""
import random
svc, source = svc_with_source
durations = {"a1": 300.0, "a2": 300.0, "a3": 300.0, "a4": 300.0, "a5": 300.0, "a6": 300.0}
patches, mock_replace = _patch_deps(svc, durations)
try:
new_plan = svc.reselect_plan_for_variant(
source.id,
["a5", "a6"], # 批次素材并入素材池
created_by_user_id="u-1",
name_suffix="批量2",
rng=random.Random(42),
)
finally:
for p in reversed(patches):
p.stop()
# 新 plan 独立、归属/模板/config 复制
assert new_plan.id != source.id
assert "批量2" in new_plan.name
assert new_plan.template_id == "tpl-001"
assert new_plan.config == {"title": "源配置"}
assert new_plan.created_by_user_id == "u-1"
# 落库片段数 == 源片段数,且 order 对齐
clips_data = mock_replace.call_args.args[1]
assert len(clips_data) == 4
assert [c["order"] for c in clips_data] == [0, 1, 2, 3]
# 与源 plan 对比:素材序列或起点必须存在差异(降重核心——不是克隆)
source_pairs = [(c.asset_id, round(float(c.start_time), 2)) for c in svc._clip_repo.list_by_plan(source.id)]
new_pairs = [(c["asset_id"], round(float(c["start_time"]), 2)) for c in clips_data]
assert new_pairs != source_pairs, f"reselect 片段应与源 plan 不同,实际相同: {new_pairs}"
# 素材全部来自素材池(源 a1-a4 ∪ 批次 a5-a6
for c in clips_data:
assert c["asset_id"] in durations
def test_record_used_segments_called_per_clip(self, svc_with_source):
"""每个新片段区间调用 record_used_segments 写回(跨变体/跨任务避让链路)。"""
import random
svc, source = svc_with_source
durations = {"a1": 300.0, "a2": 300.0, "a3": 300.0, "a4": 300.0, "a5": 300.0}
asset_models = []
for aid in durations:
m = MagicMock(id=aid)
m.duration = durations[aid]
m.metadata = None
asset_models.append(m)
svc._clip_repo.session.query.return_value.filter.return_value.all.return_value = asset_models
started = [
patch("app.services.edit_plan_service.get_used_segments", return_value={}).start(),
patch("packages.domain.plan_generator_utils.extract_scene_points_from_metadata", return_value=[]).start(),
patch("packages.adapters.sqlalchemy_impl.models.AssetModel", create=True).start(),
patch.object(svc, "replace_all_clips_transactional", return_value=4).start(),
]
mock_record = patch("app.services.edit_plan_service.record_used_segments", return_value=None).start()
try:
svc.reselect_plan_for_variant(
source.id, ["a5"], created_by_user_id="u-1", name_suffix="批量2", rng=random.Random(5)
)
finally:
patch.stopall()
assert mock_record.call_count == 4, "4 个片段应各写一次 used_segment"
for call in mock_record.call_args_list:
args = call.args
assert args[1] in durations, f"asset_id {args[1]} 不在素材池" # asset_id
assert args[3] > args[2], "区间 end 应大于 start" # end > start
assert args[4], "new plan_id 应非空"
def test_source_clips_seed_batch_avoidance(self, svc_with_source):
"""源 plan 片段区间进入批次避让集:与源完全同区间的起点重叠率应超限被避开。"""
import random
svc, source = svc_with_source
# 素材池只有源素材(极端小池),时长充足
durations = {"a1": 600.0, "a2": 600.0, "a3": 600.0, "a4": 600.0}
patches, mock_replace = _patch_deps(svc, durations)
try:
new_plan = svc.reselect_plan_for_variant(
source.id, [], created_by_user_id="u-1", name_suffix="批量2", rng=random.Random(99)
)
finally:
for p in reversed(patches):
p.stop()
clips_data = mock_replace.call_args.args[1]
source_clips = svc._clip_repo.list_by_plan(source.id)
source_by_asset = {}
for c in source_clips:
source_by_asset.setdefault(c.asset_id, []).append(
(float(c.start_time), float(c.start_time) + float(c.duration))
)
# 同素材新片段与源区间的重叠占比均 ≤20%
from packages.domain.variant_plan_selector import _clip_overlap_ratio
for c in clips_data:
ratio = _clip_overlap_ratio(c["asset_id"], float(c["start_time"]), float(c["duration"]), source_by_asset)
assert (
ratio <= 0.20 + 1e-6
), f"变体片段与源 plan 同素材区间重叠超限: asset={c['asset_id']} ratio={ratio:.2%}"
assert new_plan.id != source.id
def test_source_plan_without_clips_raises(self):
"""源 plan 无片段 → ValueError(明确报错,不产出同源变体)。"""
svc = _make_service()
svc._clip_repo.session = MagicMock()
empty = svc.create_plan(template_id="tpl-x", name="空计划")
with pytest.raises(ValueError, match="源 plan 无片段"):
svc.reselect_plan_for_variant(empty.id, ["a1"], created_by_user_id="u-1", name_suffix="变体")
def test_missing_source_plan_raises(self):
"""源 plan 不存在 → get_plan_or_raise 抛错。"""
svc = _make_service()
svc._clip_repo.session = MagicMock()
with pytest.raises((ValueError, KeyError, LookupError)): # get_plan_or_raise 抛错
svc.reselect_plan_for_variant("not-exist-plan", ["a1"], created_by_user_id="u-1", name_suffix="变体")
+3 -2
View File
@@ -51,6 +51,7 @@ def test_generation_cleans_up_temp_dir():
assert "rmtree(render_temp_dir" in source, "generation.py 应清理 render_temp_dir"
# 验证清理发生在上传之后(通过查找顺序)
upload_pos = source.find("_upload_and_record")
# #1743_upload_and_record 拆分为 _upload_rendered_video(仅 OSS 上传)
upload_pos = source.find("_upload_rendered_video")
cleanup_pos = source.find("rmtree(render_temp_dir")
assert upload_pos > 0 and cleanup_pos > upload_pos, "清理临时目录应在 _upload_and_record 之后执行"
assert upload_pos > 0 and cleanup_pos > upload_pos, "清理临时目录应在 _upload_rendered_video 之后执行"
@@ -0,0 +1,271 @@
"""#1743 批量变体独立选片纯核心测试(packages/domain/variant_plan_selector.py)。
覆盖:
- 固定种子下 N 次独立选片:素材组合/片段顺序/起点显著不同(降重核心)
- 批次内同素材区间重叠 >20% 触发避让重选
- 异常输入:空源片段/空素材池/时长全 0 → ValueError(严禁退回同源)
- 非 main 片段(intro/outro/overlay)保留源骨架素材,仅重算起点
- 跨变体 batch_segments 就地累加(串行选片天然避让)
- 文案/转场/速度等骨架字段透传
"""
from __future__ import annotations
import random
import sys
from pathlib import Path
from unittest.mock import patch
import pytest
REPO_ROOT = Path(__file__).resolve().parents[2]
for sub in ("packages", ""):
p = str(REPO_ROOT / sub) if sub else str(REPO_ROOT)
if p not in sys.path:
sys.path.insert(0, p)
from packages.domain import variant_plan_selector as vps # noqa: E402
from packages.domain.variant_plan_selector import ( # noqa: E402
BATCH_CLIP_OVERLAP_LIMIT,
_clip_overlap_ratio,
reselect_clips_for_variant,
)
def _source_clips(assets=("a1", "a2", "a3"), dur=5.0, with_non_main=False):
"""构造源片段骨架:main 片段若干,可选 intro/outro 固定角色片段。"""
clips = []
order = 0
if with_non_main:
clips.append(
{
"order": order,
"asset_id": "intro_asset",
"start_time": 0.0,
"duration": 3.0,
"clip_type": "intro",
"playback_speed": 1.0,
"transition_effect": "fade",
"transition_duration": 0.5,
"text_content": "片头",
"config": {"role": "intro"},
}
)
order += 1
for i, aid in enumerate(assets):
clips.append(
{
"order": order,
"asset_id": aid,
"start_time": float(i * 10),
"duration": dur,
"clip_type": "main",
"playback_speed": 1.2,
"transition_effect": "cut",
"transition_duration": 0.0,
"text_content": f"文案{i}",
"config": {},
}
)
order += 1
if with_non_main:
clips.append(
{
"order": order,
"asset_id": "outro_asset",
"start_time": 0.0,
"duration": 2.0,
"clip_type": "outro",
"playback_speed": 1.0,
"transition_effect": "fade",
"transition_duration": 0.5,
"text_content": "片尾",
"config": {"role": "outro"},
}
)
return clips
def _durations(asset_ids, total=120.0, extra=None):
d = {a: total for a in asset_ids}
if extra:
d.update(extra)
return d
class TestReselectClipsValidation:
def test_empty_source_clips_raises(self):
with pytest.raises(ValueError, match="源 plan 无片段"):
reselect_clips_for_variant(
[],
["a1"],
asset_durations={"a1": 60.0},
rng=random.Random(1),
)
def test_empty_asset_pool_raises(self):
with pytest.raises(ValueError, match="素材池为空"):
reselect_clips_for_variant(
_source_clips(),
[],
asset_durations={},
rng=random.Random(1),
)
def test_all_zero_duration_raises(self):
with pytest.raises(ValueError, match="时长全部未知"):
reselect_clips_for_variant(
_source_clips(),
["a1", "a2"],
asset_durations={"a1": 0.0, "a2": 0.0},
rng=random.Random(1),
)
class TestReselectClipsDifferentiation:
def test_three_variants_differ_in_assets_order_and_starts(self):
"""核心验收:固定种子连续 3 次独立选片,素材组合/顺序/起点显著不同。"""
source = _source_clips(assets=("a1", "a2", "a3", "a4"))
pool = ["a1", "a2", "a3", "a4", "a5", "a6"]
durations = _durations(pool, total=300.0)
batch_segments: dict = {}
variants = []
for seed in range(3):
clips = reselect_clips_for_variant(
source,
pool,
asset_durations=durations,
batch_segments=batch_segments, # 串行调用:上一变体区间参与避让
rng=random.Random(100 + seed),
)
variants.append(clips)
main_sequences = []
for clips in variants:
main = [c for c in clips if c["clip_type"] == "main"]
main_sequences.append([(c["asset_id"], round(c["start_time"], 2)) for c in main])
# 1) 每个变体片段数与源骨架一致
for clips in variants:
main_clips = [c for c in clips if c["clip_type"] == "main"]
assert len(main_clips) == 4
# 2) 三个变体的素材序列不全相同(素材洗牌 + main 顺序洗牌生效)
seq_sets = {tuple(a for a, _ in seq) for seq in main_sequences}
assert len(seq_sets) >= 2, f"变体素材序列应存在差异,实际全部相同: {seq_sets}"
# 3) 起点组合不全相同(起点重选生效)
start_sets = {tuple(s for _, s in seq) for seq in main_sequences}
assert len(start_sets) >= 2, f"变体起点组合应存在差异,实际全部相同: {start_sets}"
# 4) batch_segments 跨变体累加(串行避让链路存在)
total_segments = sum(len(v) for v in batch_segments.values())
assert total_segments >= 12, f"3 变体 × 4 main 片段应累加 >=12 区间,实际 {total_segments}"
def test_skeleton_fields_preserved(self):
"""文案/转场/速度等骨架字段随片段透传(只换素材与起点)。"""
source = _source_clips(assets=("a1", "a2"), with_non_main=True)
pool = ["a1", "a2", "a3"]
durations = _durations(pool, total=120.0, extra={"intro_asset": 30.0, "outro_asset": 30.0})
clips = reselect_clips_for_variant(
source,
pool,
asset_durations=durations,
rng=random.Random(7),
)
by_order = {c["order"]: c for c in clips}
# intro/outro 骨架字段保留
intro = next(c for c in clips if c["clip_type"] == "intro")
outro = next(c for c in clips if c["clip_type"] == "outro")
assert intro["asset_id"] == "intro_asset"
assert intro["text_content"] == "片头"
assert intro["transition_effect"] == "fade"
assert intro["transition_duration"] == 0.5
assert outro["asset_id"] == "outro_asset"
assert outro["text_content"] == "片尾"
# main 片段文案/速度随骨架 order 保留
for c in clips:
if c["clip_type"] == "main":
assert c["playback_speed"] == 1.2
assert c["text_content"].startswith("文案")
class TestBatchOverlapAvoidance:
def test_overlap_ratio_calculation(self):
segs = {"a1": [(10.0, 20.0)]} # 已占 10s 区间
# 新区间 [10,20) 完全重叠 → 1.0
assert _clip_overlap_ratio("a1", 10.0, 10.0, segs) == pytest.approx(1.0)
# 新区间 [20,30) 零重叠 → 0.0
assert _clip_overlap_ratio("a1", 20.0, 10.0, segs) == pytest.approx(0.0)
# 新区间 [15,25) 重叠 5s / 10s → 0.5
assert _clip_overlap_ratio("a1", 15.0, 10.0, segs) == pytest.approx(0.5)
# 空 asset / 零时长 → 0
assert _clip_overlap_ratio("", 0.0, 10.0, segs) == 0.0
assert _clip_overlap_ratio("a1", 10.0, 0.0, segs) == 0.0
def test_over_limit_triggers_reselect_to_non_overlapping(self):
"""批次已占满素材前段时,避让重选应把起点挪到重叠 ≤20% 的位置。"""
source = _source_clips(assets=("a1",), dur=10.0)
durations = {"a1": 120.0}
# 批次已选区间:a1 [0, 100) 几乎占满前段
batch_segments = {"a1": [(0.0, 100.0)]}
# _resolve_start_time 第一次返回高重叠起点(2.0),之后 rng 抖动应找到低重叠位置
call_count = {"n": 0}
def fake_resolve(asset_id, clip_duration, asset_durations, used_segments, scene_points=None, on_exhausted=None):
call_count["n"] += 1
return 2.0 if call_count["n"] == 1 else None # 后续回退 rng.uniform
with patch.object(vps, "_resolve_start_time", side_effect=fake_resolve):
# rng.uniform 返回 105.0(与 [0,100) 零重叠);rng 是 random.Random 实例,
# 需 patch 类方法 uniform 才能生效
with patch.object(random.Random, "uniform", return_value=105.0):
clips = reselect_clips_for_variant(
source,
["a1"],
asset_durations=durations,
batch_segments=batch_segments,
rng=random.Random(3),
)
main = [c for c in clips if c["clip_type"] == "main"]
assert len(main) == 1
ratio = _clip_overlap_ratio("a1", main[0]["start_time"], 10.0, {"a1": [(0.0, 100.0)]})
assert ratio <= BATCH_CLIP_OVERLAP_LIMIT, f"避让后重叠应 ≤20%,实际 {ratio:.2%}"
assert main[0]["start_time"] == pytest.approx(105.0, abs=0.01)
def test_first_variant_segments_become_avoidance_target(self):
"""变体 0 选定区间后,变体 1 选同素材时批次区间生效(不与源区间完全重合)。"""
source = _source_clips(assets=("a1", "a2"), dur=8.0)
pool = ["a1", "a2", "a3"]
durations = _durations(pool, total=600.0)
batch_segments: dict = {}
v1 = reselect_clips_for_variant(
source, pool, asset_durations=durations, batch_segments=batch_segments, rng=random.Random(11)
)
v1_main = [(c["asset_id"], c["start_time"], c["duration"]) for c in v1 if c["clip_type"] == "main"]
# 快照 v1 之后的批次避让集(v2 调用会就地追加 v2 自身区间,断言必须用调用前快照)
import copy
batch_snapshot = copy.deepcopy(batch_segments)
v2 = reselect_clips_for_variant(
source, pool, asset_durations=durations, batch_segments=batch_segments, rng=random.Random(12)
)
# v2 与 v1(快照)同素材片段的区间重叠占比均 ≤20%
for c in v2:
if c["clip_type"] != "main" or not c["asset_id"]:
continue
ratio = _clip_overlap_ratio(c["asset_id"], c["start_time"], c["duration"], batch_snapshot)
assert (
ratio <= BATCH_CLIP_OVERLAP_LIMIT + 1e-6
), f"变体间片段重叠超限: asset={c['asset_id']} start={c['start_time']} ratio={ratio:.2%}"
# v1 片段确实进入了批次避让集
assert any(a in batch_segments for a, _, _ in v1_main)