"""查重记录 SQLAlchemy 仓库实现。""" from __future__ import annotations import json from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import DuplicationRecordModel, DuplicationSegmentModel from packages.domain.duplication import DuplicateSegment, DuplicationRecord class SQLAlchemyDuplicationRecordRepository: def __init__(self, session: Session): self.session = session def create(self, record: DuplicationRecord) -> DuplicationRecord: model = DuplicationRecordModel( id=record.id, user_id=record.user_id, filename=record.filename, file_size=record.file_size, storage_key=record.storage_key, duration_seconds=record.duration_seconds, status=record.status, duplicate_rate=record.duplicate_rate, duplicate_count=record.duplicate_count, video_fingerprint=json.dumps(record.video_fingerprint) if record.video_fingerprint else None, error_message=record.error_message, created_at=record.created_at, updated_at=record.updated_at, ) self.session.add(model) self.session.commit() return record def get(self, record_id: str) -> DuplicationRecord | None: model = self.session.query(DuplicationRecordModel).filter(DuplicationRecordModel.id == record_id).first() if model is None: return None return self._to_domain(model) def list_by_user(self, user_id: str, *, offset: int = 0, limit: int = 50) -> list[DuplicationRecord]: models = ( self.session.query(DuplicationRecordModel) .filter(DuplicationRecordModel.user_id == user_id) .order_by(DuplicationRecordModel.created_at.desc()) .offset(offset) .limit(limit) .all() ) return [self._to_domain(m) for m in models] def update(self, record: DuplicationRecord) -> DuplicationRecord: model = self.session.query(DuplicationRecordModel).filter(DuplicationRecordModel.id == record.id).first() if model is None: return record model.status = record.status model.duplicate_rate = record.duplicate_rate model.duplicate_count = record.duplicate_count model.video_fingerprint = json.dumps(record.video_fingerprint) if record.video_fingerprint else None model.error_message = record.error_message model.updated_at = record.updated_at # 更新 segments:先删后建 self.session.query(DuplicationSegmentModel).filter(DuplicationSegmentModel.record_id == record.id).delete() for seg in record.segments: seg_model = DuplicationSegmentModel( id=seg.id, record_id=record.id, source_start=seg.source_start, source_end=seg.source_end, matched_video_id=seg.matched_video_id, matched_video_name=seg.matched_video_name, matched_start=seg.matched_start, matched_end=seg.matched_end, similarity=seg.similarity, ) self.session.add(seg_model) self.session.commit() return record def delete(self, record_id: str) -> bool: """ 删除查重记录及其关联片段。 注意:必须先删片段再删记录,防止进程崩溃时产生孤儿片段数据。 """ # 先删关联片段,再删主记录(安全顺序) self.session.query(DuplicationSegmentModel).filter(DuplicationSegmentModel.record_id == record_id).delete() count = self.session.query(DuplicationRecordModel).filter(DuplicationRecordModel.id == record_id).delete() self.session.commit() return count > 0 def _to_domain(self, model: DuplicationRecordModel) -> DuplicationRecord: segment_models = ( self.session.query(DuplicationSegmentModel).filter(DuplicationSegmentModel.record_id == model.id).all() ) segments = [ DuplicateSegment( id=s.id, source_start=s.source_start, source_end=s.source_end, matched_video_id=s.matched_video_id, matched_video_name=s.matched_video_name, matched_start=s.matched_start, matched_end=s.matched_end, similarity=s.similarity, ) for s in segment_models ] fp_raw = getattr(model, "video_fingerprint", None) return DuplicationRecord( id=model.id, user_id=model.user_id, filename=model.filename, file_size=int(model.file_size or 0), storage_key=model.storage_key, duration_seconds=model.duration_seconds, status=model.status, duplicate_rate=model.duplicate_rate, duplicate_count=int(model.duplicate_count or 0), video_fingerprint=json.loads(fp_raw) if fp_raw else None, error_message=getattr(model, "error_message", ""), segments=segments, created_at=model.created_at, updated_at=model.updated_at, )