"""Application service for file management workflows.""" import mimetypes import uuid from pathlib import Path from fastapi import UploadFile from sqlmodel import Session from agenteval.storage.db import FileCategoryDB, FileRecordDB from agenteval.storage.file_repository import FileCategoryNode, FileCategoryRepository, FileRecordRepository from agenteval.storage.file_storage import ( FileStorage, FileTooLargeError, InvalidStorageNameError, StagedDeletion, ) class FileManagementError(Exception): """Base class for expected file-management errors.""" class FileValidationError(FileManagementError): """Raised when user input is invalid.""" class CategoryNotFoundError(FileManagementError): """Raised when a requested category does not exist.""" class FileRecordNotFoundError(FileManagementError): """Raised when a requested file record does not exist.""" class PhysicalFileNotFoundError(FileManagementError): """Raised when metadata exists but its physical file is missing.""" class UploadTooLargeError(FileManagementError): """Raised when an upload exceeds its configured size limit.""" class FileManagementService: """Coordinate database transactions and rollback-friendly disk operations.""" def __init__( self, session: Session, storage: FileStorage, allowed_extensions: set[str], max_upload_size_mb: int, ): self.session = session self.storage = storage self.allowed_extensions = {extension.lower() for extension in allowed_extensions} self.max_upload_size_mb = max_upload_size_mb self.categories = FileCategoryRepository(session) self.files = FileRecordRepository(session) @property def max_upload_bytes(self) -> int: return self.max_upload_size_mb * 1024 * 1024 def list_categories(self) -> list[FileCategoryNode]: return self.categories.get_tree() def create_category(self, name: str, parent_id: str | None = None) -> FileCategoryDB: normalized_name = name.strip() if not normalized_name: raise FileValidationError("分类名称不能为空") if parent_id and not self.categories.get(parent_id): raise FileValidationError("父分类不存在") try: category = self.categories.create(normalized_name, parent_id) self.session.commit() return category except Exception: self.session.rollback() raise def update_category(self, category_id: str, name: str) -> FileCategoryDB: normalized_name = name.strip() if not normalized_name: raise FileValidationError("分类名称不能为空") try: category = self.categories.update(category_id, normalized_name) if not category: raise CategoryNotFoundError self.session.commit() return category except Exception: self.session.rollback() raise def delete_category(self, category_id: str) -> None: if not self.categories.get(category_id): raise CategoryNotFoundError category_ids = self.categories.get_subtree_ids(category_id) records = self.files.list_by_category_ids(category_ids) staged = self._stage_records(records) try: self.categories.delete(category_id) self.session.commit() except Exception: self.session.rollback() self.storage.restore_deletions(staged) raise self.storage.purge_deletions(staged) def list_files(self, category_id: str | None = None) -> list[FileRecordDB]: if category_id and not self.categories.get(category_id): raise CategoryNotFoundError return self.files.list_all(category_id) async def upload(self, upload: UploadFile, category_id: str | None = None) -> FileRecordDB: filename = upload.filename or "" if not filename: raise FileValidationError("文件名为空") if category_id and not self.categories.get(category_id): raise FileValidationError("分类不存在") extension = Path(filename).suffix.lstrip(".").lower() if not extension: raise FileValidationError("文件没有扩展名") if extension not in self.allowed_extensions: allowed = ", ".join(sorted(self.allowed_extensions)) raise FileValidationError(f"不支持的文件类型 .{extension},允许的类型: {allowed}") storage_name = f"{uuid.uuid4()}.{extension}" try: pending = await self.storage.write_upload(upload, storage_name, self.max_upload_bytes) except FileTooLargeError as error: raise UploadTooLargeError from error final_path = None try: mime_type, _ = mimetypes.guess_type(filename) record = self.files.create( original_name=filename, storage_name=storage_name, file_size=pending.file_size, mime_type=mime_type or "application/octet-stream", file_ext=extension, category_id=category_id, ) final_path = self.storage.promote(pending) self.session.commit() return record except Exception: self.session.rollback() self.storage.discard(pending.temp_path) if final_path: self.storage.discard(final_path) raise def get_download(self, file_id: str) -> tuple[FileRecordDB, Path]: record = self.files.get(file_id) if not record: raise FileRecordNotFoundError try: file_path = self.storage.path_for(record.storage_name) except InvalidStorageNameError as error: raise PhysicalFileNotFoundError from error if not file_path.is_file(): raise PhysicalFileNotFoundError return record, file_path def delete_file(self, file_id: str) -> None: record = self.files.get(file_id) if not record: raise FileRecordNotFoundError staged = self._stage_records([record]) try: self.files.delete(file_id) self.session.commit() except Exception: self.session.rollback() self.storage.restore_deletions(staged) raise self.storage.purge_deletions(staged) def _stage_records(self, records: list[FileRecordDB]) -> list[StagedDeletion]: try: return self.storage.stage_deletions([record.storage_name for record in records]) except InvalidStorageNameError as error: raise PhysicalFileNotFoundError from error