Add transactional file storage workflows, typed API contracts, recursive category handling, frontend component separation, and Files API coverage.
157 lines
5.2 KiB
Python
157 lines
5.2 KiB
Python
"""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
|