import json from datetime import datetime, timezone from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import AssetModel, AssetTagModel from packages.domain import Asset, AssetStatus, ClassificationStatus class SQLAlchemyAssetRepository: def __init__(self, session: Session): self.session = session def find_by_library( self, library_id: str, skip: int = 0, limit: int = 100, ) -> list[Asset]: models = ( self.session.query(AssetModel) .filter(AssetModel.asset_library_id == library_id) .offset(skip) .limit(limit) .all() ) return [self._to_domain(model) for model in models] def find_by_project( self, project_id: str, skip: int = 0, limit: int = 100, ) -> list[Asset]: models = ( self.session.query(AssetModel).filter(AssetModel.project_id == project_id).offset(skip).limit(limit).all() ) return [self._to_domain(model) for model in models] def find_by_library_and_file_type( self, library_id: str, file_type: str, skip: int = 0, limit: int = 100, ) -> list[Asset]: models = ( self.session.query(AssetModel) .filter(AssetModel.asset_library_id == library_id, AssetModel.file_type == file_type) .offset(skip) .limit(limit) .all() ) return [self._to_domain(model) for model in models] def find_by_id(self, asset_id: str) -> Asset | None: model = self.session.query(AssetModel).filter(AssetModel.id == asset_id).first() if model is None: return None return self._to_domain(model) def get(self, asset_id: str) -> Asset | None: return self.find_by_id(asset_id) def create(self, asset: Asset) -> Asset: now = datetime.now(timezone.utc) model = AssetModel( id=asset.id, project_id=asset.project_id, asset_library_id=asset.library_id, name=asset.name, file_type=(asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type), file_size=asset.file_size, file_url=asset.storage_key, thumbnail_url=asset.thumbnail_url, duration=asset.duration, width=asset.width, height=asset.height, fps=asset.fps, codec=asset.codec, status=asset.status.value, classification_status=asset.classification_status.value, classification_result=(json.dumps(asset.metadata) if asset.metadata else None), quality_score=asset.quality_score, uploaded_by_user_id=asset.uploaded_by_user_id or "system", file_hash=asset.file_hash or None, created_at=asset.created_at, updated_at=now, ) self.session.add(model) self.session.flush() self._sync_asset_tags(asset.id, asset.tag_ids) self.session.commit() return asset def update(self, asset: Asset) -> Asset: model = self.session.query(AssetModel).filter(AssetModel.id == asset.id).first() if model is None: raise ValueError(f"Asset {asset.id} not found") model.name = asset.name model.file_size = asset.file_size model.file_url = asset.storage_key model.thumbnail_url = asset.thumbnail_url model.duration = asset.duration model.width = asset.width model.height = asset.height model.fps = asset.fps model.codec = asset.codec model.status = asset.status.value model.classification_status = asset.classification_status.value model.classification_result = json.dumps(asset.metadata) if asset.metadata else None model.quality_score = asset.quality_score model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id model.file_hash = asset.file_hash or model.file_hash model.updated_at = datetime.now(timezone.utc) self.session.flush() self._sync_asset_tags(asset.id, asset.tag_ids) self.session.commit() return asset def delete(self, asset_id: str) -> bool: model = self.session.query(AssetModel).filter(AssetModel.id == asset_id).first() if model: self.session.delete(model) self.session.commit() return True return False def batch_delete(self, asset_ids: list[str]) -> int: """批量删除素材,返回实际删除数量。""" if not asset_ids: return 0 count = self.session.query(AssetModel).filter(AssetModel.id.in_(asset_ids)).delete(synchronize_session=False) self.session.commit() return count def count_by_project(self, project_id: str) -> int: return self.session.query(AssetModel).filter(AssetModel.project_id == project_id).count() def count_by_project_ids(self, project_ids: list[str]) -> int: if not project_ids: return 0 return self.session.query(AssetModel).filter(AssetModel.project_id.in_(project_ids)).count() def sum_storage_by_project_ids(self, project_ids: list[str]) -> int: if not project_ids: return 0 from sqlalchemy import func result = ( self.session.query(func.coalesce(func.sum(AssetModel.file_size), 0)) .filter(AssetModel.project_id.in_(project_ids)) .scalar() ) return int(result or 0) def search_candidates( self, project_id: str, *, file_type: str | None = None, min_quality_score: float | None = None, min_duration: float | None = None, max_duration: float | None = None, classification_category: str | None = None, tags: list[str] | None = None, status: str | None = None, limit: int = 50, ) -> list[Asset]: """按筛选条件搜索候选素材,按质量分降序排列。""" query = self.session.query(AssetModel).filter( AssetModel.project_id == project_id, ) if file_type is not None: query = query.filter(AssetModel.file_type == file_type) if min_quality_score is not None: query = query.filter(AssetModel.quality_score >= min_quality_score) if min_duration is not None: query = query.filter(AssetModel.duration >= min_duration) if max_duration is not None: query = query.filter(AssetModel.duration <= max_duration) if status is not None: query = query.filter(AssetModel.status == status) if classification_category is not None: # classification_result 是 JSON Text,用 LIKE 匹配 category 字段 query = query.filter(AssetModel.classification_result.like(f'%"{classification_category}"%')) query = query.order_by(AssetModel.quality_score.desc().nullslast()) if limit > 0: query = query.limit(limit) models = query.all() candidates = [self._to_domain(m) for m in models] # 内存中过滤 tags(tags 存在 metadata 中) if tags: tag_set = set(tags) candidates = [a for a in candidates if tag_set.issubset(set(a.metadata.get("tags", [])))] return candidates def _to_domain(self, model: AssetModel) -> Asset: metadata = {} if model.classification_result: try: metadata = json.loads(model.classification_result) except Exception: metadata = {} mime_type = model.file_type if "/" not in mime_type: mime_type = { "video": "video/mp4", "audio": "audio/mpeg", "image": "image/jpeg", }.get(mime_type, mime_type) # 查询关联的 tag_ids tag_ids = [ row.tag_id for row in self.session.query(AssetTagModel.tag_id).filter(AssetTagModel.asset_id == model.id).all() ] return Asset( id=model.id, project_id=model.project_id, library_id=model.asset_library_id, name=model.name, storage_key=model.file_url, mime_type=mime_type, file_size=int(model.file_size or 0), thumbnail_url=model.thumbnail_url, duration=model.duration, width=int(model.width) if model.width is not None else None, height=int(model.height) if model.height is not None else None, fps=model.fps, codec=model.codec, status=AssetStatus(model.status), classification_status=ClassificationStatus(model.classification_status), quality_score=model.quality_score, uploaded_by_user_id=model.uploaded_by_user_id, file_hash=model.file_hash or "", metadata=metadata, tag_ids=tag_ids, created_at=model.created_at, updated_at=model.updated_at, ) def _sync_asset_tags(self, asset_id: str, tag_ids: list[str]) -> None: """同步素材-标签关联表(全量替换)。""" self.session.query(AssetTagModel).filter(AssetTagModel.asset_id == asset_id).delete(synchronize_session=False) for tag_id in tag_ids: self.session.add(AssetTagModel(asset_id=asset_id, tag_id=tag_id)) def find_by_tag_ids( self, tag_ids: list[str], skip: int = 0, limit: int = 100, ) -> list[Asset]: """查找包含所有指定标签的素材。""" if not tag_ids: return [] from sqlalchemy import func # 找出同时拥有所有指定 tag_id 的 asset_id tag_set = set(tag_ids) asset_ids = ( self.session.query(AssetTagModel.asset_id) .filter(AssetTagModel.tag_id.in_(tag_set)) .group_by(AssetTagModel.asset_id) .having(func.count(AssetTagModel.tag_id) == len(tag_set)) .all() ) ids = [row[0] for row in asset_ids] if not ids: return [] models = self.session.query(AssetModel).filter(AssetModel.id.in_(ids)).offset(skip).limit(limit).all() return [self._to_domain(m) for m in models] def find_by_library_and_file_hash( self, library_id: str, file_hash: str, ) -> Asset | None: """按素材库 + 文件哈希查找已有素材(去重检测)。""" if not file_hash: return None model = ( self.session.query(AssetModel) .filter( AssetModel.asset_library_id == library_id, AssetModel.file_hash == file_hash, ) .first() ) if model is None: return None return self._to_domain(model)