"""原子片段加载与兜底 — #1970 智能剪辑流程重构 P1. 选片前从 ``asset_atom_clips`` 表加载素材池的原子片段;老素材/切片任务尚未 完成/切片失败导致某些素材没有片段时,按需求兜底:内存中按 3-6 秒临时均匀 切片(不存库,片段标记 is_fallback=True)。 本模块对 repository 做鸭子类型约束(只需 find_by_asset / find_candidates_for_selection 和 asset_repo.get),方便 API 侧(SQLAlchemy)与 worker 侧复用,也便于单测注入内存假实现。 """ from __future__ import annotations import logging from packages.domain.asset_atom_clip import AssetAtomClip from packages.domain.atom_clip_service import compute_fallback_clips logger = logging.getLogger(__name__) # 兜底均匀切片步长(秒),落在 3~6s 区间中段 FALLBACK_CLIP_SECONDS = 4.5 def load_atom_clips_for_assets( asset_ids: list[str], *, atom_clip_repo, asset_repo=None, ) -> dict[str, list[AssetAtomClip]]: """加载素材池的原子片段(缺失素材走内存兜底). Args: asset_ids: 候选素材 ID(去重保序)。 atom_clip_repo: AssetAtomClipRepository 实现(需有 ``find_candidates_for_selection`` 或 ``find_by_asset``)。 asset_repo: 可选,素材仓储(需有 ``get``),用于读取时长兜底切片。 为 None 时,没有原子片段的素材直接跳过(不兜底)。 Returns: {asset_id: [AssetAtomClip, ...]},仅包含至少有一个片段的素材, 片段按 clip_index 排序。 """ result: dict[str, list[AssetAtomClip]] = {} unique_ids = list(dict.fromkeys(asset_ids)) if not unique_ids: return result # 1. 批量查询已生成的原子片段 persisted: dict[str, list[AssetAtomClip]] = {} try: if hasattr(atom_clip_repo, "find_candidates_for_selection"): clips = atom_clip_repo.find_candidates_for_selection(unique_ids, limit=0) else: clips = [] for asset_id in unique_ids: clips.extend(atom_clip_repo.find_by_asset(asset_id)) for clip in clips: persisted.setdefault(clip.asset_id, []).append(clip) except Exception: logger.warning("加载 atom_clips 失败,全部走内存兜底", exc_info=True) persisted = {} for asset_id in unique_ids: clips = persisted.get(asset_id) if clips: clips.sort(key=lambda c: c.clip_index) result[asset_id] = clips continue # 2. 兜底:内存均匀切片(不存库) if asset_repo is None: continue duration = _safe_asset_duration(asset_repo, asset_id) if duration <= 0: continue result[asset_id] = compute_fallback_clips( asset_id, duration, clip_seconds=FALLBACK_CLIP_SECONDS, ) return result def flatten_candidates( clips_by_asset: dict[str, list[AssetAtomClip]], ) -> list[AssetAtomClip]: """把 {asset_id: [clips]} 摊平为候选片段列表(素材顺序内片段有序)。""" flat: list[AssetAtomClip] = [] for clips in clips_by_asset.values(): flat.extend(clips) return flat def _safe_asset_duration(asset_repo, asset_id: str) -> float: """安全读取素材时长,任何异常返回 0。""" try: asset = asset_repo.get(asset_id) if asset is None: return 0.0 return float(getattr(asset, "duration", 0.0) or 0.0) except Exception: logger.warning("读取素材时长失败: asset_id=%s", asset_id, exc_info=True) return 0.0