"""Asset InMemory Repository 实现""" from packages.domain import Asset class InMemoryAssetRepository: """素材 InMemory 仓储实现(用于同步用例和测试)""" def __init__(self): self._assets: dict[str, Asset] = {} def create(self, asset: Asset) -> Asset: self._assets[asset.id] = asset return asset def get(self, asset_id: str) -> Asset | None: return self._assets.get(asset_id) def list_by_project(self, project_id: str) -> list[Asset]: return [asset for asset in self._assets.values() if asset.project_id == project_id] def list_by_library(self, library_id: str) -> list[Asset]: return [asset for asset in self._assets.values() if asset.library_id == library_id] def find_by_library(self, library_id: str) -> list[Asset]: """Alias for list_by_library to match the port interface.""" return self.list_by_library(library_id) def find_by_library_and_file_type(self, library_id: str, file_type: str) -> list[Asset]: return [ asset for asset in self._assets.values() if asset.library_id == library_id and asset.mime_type and asset.mime_type.startswith(file_type) ] def update(self, asset: Asset) -> Asset: self._assets[asset.id] = asset return asset def delete(self, asset_id: str) -> bool: if asset_id in self._assets: del self._assets[asset_id] return True return False def batch_delete(self, asset_ids: list[str]) -> int: """批量删除素材(软删除,标记 status=deleted),返回实际影响数量。""" from datetime import datetime, timezone from packages.domain import AssetStatus count = 0 for aid in asset_ids: asset = self._assets.get(aid) if asset and asset.status != AssetStatus.DELETED: asset.status = AssetStatus.DELETED asset.updated_at = datetime.now(timezone.utc) count += 1 return count def batch_update_metadata(self, asset_ids: list[str], metadata_patch: dict[str, object]) -> int: """批量更新素材 metadata(合并 patch),返回实际影响数量。""" from datetime import datetime, timezone count = 0 for aid in asset_ids: asset = self._assets.get(aid) if asset: asset.metadata = {**asset.metadata, **metadata_patch} asset.updated_at = datetime.now(timezone.utc) count += 1 return count def batch_add_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int: """批量给素材添加标签(合并去重),返回实际影响数量。""" from datetime import datetime, timezone count = 0 for aid in asset_ids: asset = self._assets.get(aid) if asset: changed = False for tid in tag_ids: if tid not in asset.tag_ids: asset.tag_ids.append(tid) changed = True if changed: asset.updated_at = datetime.now(timezone.utc) count += 1 return count def batch_replace_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int: """批量替换素材标签(全量覆盖),返回实际影响数量。""" from datetime import datetime, timezone count = 0 for aid in asset_ids: asset = self._assets.get(aid) if asset: asset.tag_ids = list(tag_ids) asset.updated_at = datetime.now(timezone.utc) count += 1 return count def find_by_project( self, project_id: str, skip: int = 0, limit: int = 100, ) -> list[Asset]: items = [a for a in self._assets.values() if a.project_id == project_id] return items[skip : skip + limit] def find_by_id(self, asset_id: str) -> Asset | None: return self._assets.get(asset_id) def find_by_tag_ids( self, tag_ids: list[str], skip: int = 0, limit: int = 100, ) -> list[Asset]: """查找包含所有指定标签的素材。""" if not tag_ids: return [] tag_set = set(tag_ids) items = [a for a in self._assets.values() if tag_set.issubset(set(a.tag_ids))] return items[skip : skip + limit] def find_by_library_and_file_hash( self, library_id: str, file_hash: str, ) -> Asset | None: """按素材库 + 文件哈希查找已有素材(去重检测)。""" if not file_hash: return None for asset in self._assets.values(): if asset.library_id == library_id and asset.file_hash == file_hash: return asset return None