"""素材原子片段仓储 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 _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, 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, )