"""素材原子片段仓储 SQLAlchemy 实现。""" from __future__ import annotations from datetime import UTC, datetime from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import AssetAtomClipModel from packages.domain.asset_atom_clip import AssetAtomClip class SQLAlchemyAssetAtomClipRepository: def __init__(self, session: Session): self.session = session def create(self, clip: AssetAtomClip) -> AssetAtomClip: model = self._to_model(clip) self.session.add(model) self.session.flush() self.session.commit() return clip def batch_create(self, clips: list[AssetAtomClip]) -> list[AssetAtomClip]: if not clips: return [] models = [self._to_model(c) for c in clips] self.session.add_all(models) self.session.flush() self.session.commit() return clips def find_by_asset(self, asset_id: str) -> list[AssetAtomClip]: models = ( self.session.query(AssetAtomClipModel) .filter(AssetAtomClipModel.asset_id == asset_id) .order_by(AssetAtomClipModel.clip_index.asc()) .all() ) return [self._to_domain(m) for m in models] def find_by_id(self, clip_id: str) -> AssetAtomClip | None: model = self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.id == clip_id).first() if model is None: return None return self._to_domain(model) def find_by_ids(self, clip_ids: list[str]) -> list[AssetAtomClip]: if not clip_ids: return [] models = self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.id.in_(clip_ids)).all() return [self._to_domain(m) for m in models] def delete_by_asset(self, asset_id: str) -> int: count = ( self.session.query(AssetAtomClipModel) .filter(AssetAtomClipModel.asset_id == asset_id) .delete(synchronize_session=False) ) self.session.commit() return count def count_by_asset(self, asset_id: str) -> int: return self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.asset_id == asset_id).count() def find_candidates_for_selection( self, asset_ids: list[str], *, min_duration: float | None = None, max_duration: float | None = None, limit: int = 100, ) -> list[AssetAtomClip]: """按筛选条件查找候选原子片段,按时长排序。用于选片逻辑。""" query = self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.asset_id.in_(asset_ids)) if min_duration is not None: query = query.filter(AssetAtomClipModel.duration >= min_duration) if max_duration is not None: query = query.filter(AssetAtomClipModel.duration <= max_duration) query = query.order_by(AssetAtomClipModel.clip_index.asc()) if limit > 0: query = query.limit(limit) models = query.all() return [self._to_domain(m) for m in models] def update_ai_tags(self, clip_id: str, ai_tags: dict) -> bool: """更新指定片段的 ai_tags 字段.""" count = ( self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.id == clip_id).update({"ai_tags": ai_tags}) ) self.session.commit() return count > 0 def find_untagged(self, limit: int = 100, include_downgraded: bool = False) -> list[AssetAtomClip]: """查找未完成 AI 打标的片段,用于回填. 默认仅匹配 ai_tags IS NULL;include_downgraded=True 时额外包含 只有 inherited_tags 的降级记录(视觉 API 失败时写入,无 has_text 字段), 供强制回填(#1970 force backfill)使用。 """ query = self.session.query(AssetAtomClipModel) if include_downgraded: # as_string() → JSON/JSONB ->> 取值;NULL 记录或缺 has_text 键 # (降级记录)均为 NULL,has_text 为 true/false 的完整记录被排除 query = query.filter(AssetAtomClipModel.ai_tags["has_text"].as_string().is_(None)) else: query = query.filter(AssetAtomClipModel.ai_tags.is_(None)) models = query.order_by(AssetAtomClipModel.created_at.asc()).limit(limit).all() return [self._to_domain(m) for m in models] def _to_model(self, clip: AssetAtomClip) -> AssetAtomClipModel: return AssetAtomClipModel( id=clip.id, asset_id=clip.asset_id, start_time=clip.start_time, end_time=clip.end_time, duration=clip.duration, clip_index=clip.clip_index, tags=clip.tags, ai_tags=clip.ai_tags, scene_change_at=clip.scene_change_at, is_fallback=clip.is_fallback, created_at=clip.created_at or datetime.now(UTC), ) def _to_domain(self, model: AssetAtomClipModel) -> AssetAtomClip: return AssetAtomClip( id=model.id, asset_id=model.asset_id, start_time=model.start_time, end_time=model.end_time, duration=model.duration, clip_index=model.clip_index, tags=model.tags or [], scene_change_at=model.scene_change_at, is_fallback=model.is_fallback, created_at=model.created_at, )