"""Database repositories for file categories and uploaded file records.""" from dataclasses import dataclass, field from typing import Optional from sqlmodel import Session, col, select from agenteval.storage.db import FileCategoryDB, FileRecordDB, get_session, new_uuid, utc_now @dataclass class FileCategoryNode: """Framework-independent category tree node.""" id: str name: str parent_id: Optional[str] children: list["FileCategoryNode"] = field(default_factory=list) class FileCategoryRepository: """Persistence operations for the file category tree.""" def __init__(self, session: Optional[Session] = None): self.session = session or get_session() def list_all(self) -> list[FileCategoryDB]: statement = select(FileCategoryDB).order_by(FileCategoryDB.created_at.asc()) return list(self.session.exec(statement).all()) def get(self, category_id: str) -> Optional[FileCategoryDB]: return self.session.get(FileCategoryDB, category_id) def get_tree(self) -> list[FileCategoryNode]: """Build the complete category tree with one database query.""" categories = self.list_all() node_map = { category.id: FileCategoryNode( id=category.id, name=category.name, parent_id=category.parent_id, ) for category in categories } roots: list[FileCategoryNode] = [] for category in categories: node = node_map[category.id] parent = node_map.get(category.parent_id) if category.parent_id else None if parent: parent.children.append(node) else: roots.append(node) return roots def get_subtree_ids(self, category_id: str) -> list[str]: """Return the category and all descendant IDs with one database query.""" categories = self.list_all() children_by_parent: dict[str, list[str]] = {} for category in categories: if category.parent_id: children_by_parent.setdefault(category.parent_id, []).append(category.id) result: list[str] = [] pending = [category_id] while pending: current = pending.pop() result.append(current) pending.extend(children_by_parent.get(current, [])) return result def create(self, name: str, parent_id: Optional[str] = None) -> FileCategoryDB: category = FileCategoryDB( id=new_uuid(), name=name, parent_id=parent_id, created_at=utc_now(), updated_at=utc_now(), ) self.session.add(category) self.session.flush() return category def update(self, category_id: str, name: str) -> Optional[FileCategoryDB]: category = self.session.get(FileCategoryDB, category_id) if not category: return None category.name = name category.updated_at = utc_now() self.session.add(category) self.session.flush() return category def delete(self, category_id: str) -> bool: category = self.session.get(FileCategoryDB, category_id) if not category: return False self.session.delete(category) self.session.flush() return True class FileRecordRepository: """Persistence operations for uploaded file metadata.""" def __init__(self, session: Optional[Session] = None): self.session = session or get_session() def list_all(self, category_id: Optional[str] = None) -> list[FileRecordDB]: statement = select(FileRecordDB) if category_id: category_ids = FileCategoryRepository(self.session).get_subtree_ids(category_id) statement = statement.where(col(FileRecordDB.category_id).in_(category_ids)) statement = statement.order_by(FileRecordDB.created_at.desc()) return list(self.session.exec(statement).all()) def list_by_category_ids(self, category_ids: list[str]) -> list[FileRecordDB]: if not category_ids: return [] statement = select(FileRecordDB).where(col(FileRecordDB.category_id).in_(category_ids)) return list(self.session.exec(statement).all()) def get(self, file_id: str) -> Optional[FileRecordDB]: return self.session.get(FileRecordDB, file_id) def create( self, original_name: str, storage_name: str, file_size: int, mime_type: str, file_ext: str, category_id: Optional[str] = None, ) -> FileRecordDB: record = FileRecordDB( id=new_uuid(), original_name=original_name, storage_name=storage_name, category_id=category_id, file_size=file_size, mime_type=mime_type, file_ext=file_ext, created_at=utc_now(), ) self.session.add(record) self.session.flush() return record def delete(self, file_id: str) -> bool: record = self.session.get(FileRecordDB, file_id) if not record: return False self.session.delete(record) self.session.flush() return True