From d7514f4e651c0758b114f43268de94897230570e Mon Sep 17 00:00:00 2001 From: sinohqb Date: Fri, 17 Jul 2026 17:41:19 +0800 Subject: [PATCH] refactor(files): harden storage and split frontend Add transactional file storage workflows, typed API contracts, recursive category handling, frontend component separation, and Files API coverage. --- backend/agenteval/services/__init__.py | 1 + backend/agenteval/services/files.py | 191 ++++++ backend/agenteval/storage/file_repository.py | 182 +++-- backend/agenteval/storage/file_storage.py | 94 +++ backend/agenteval/web/file_schemas.py | 69 ++ backend/agenteval/web/routers/files.py | 291 ++++---- frontend/web/src/api.ts | 12 +- .../src/components/files/FileCategoryTree.tsx | 127 ++++ .../web/src/components/files/FileTable.tsx | 182 +++++ .../src/components/files/FileUploadModal.tsx | 121 ++++ frontend/web/src/hooks/useFiles.ts | 136 ++++ frontend/web/src/index.css | 193 +++++- frontend/web/src/pages/Files.tsx | 621 ++++-------------- frontend/web/src/utils/fileFormat.ts | 31 + frontend/web/src/utils/fileTree.ts | 32 + scripts/deploy-t480.sh | 1 + tests/integration/test_files_api.py | 198 ++++++ tests/unit/test_file_repository.py | 48 +- tests/unit/test_file_service.py | 92 +++ 19 files changed, 1804 insertions(+), 818 deletions(-) create mode 100644 backend/agenteval/services/__init__.py create mode 100644 backend/agenteval/services/files.py create mode 100644 backend/agenteval/storage/file_storage.py create mode 100644 backend/agenteval/web/file_schemas.py create mode 100644 frontend/web/src/components/files/FileCategoryTree.tsx create mode 100644 frontend/web/src/components/files/FileTable.tsx create mode 100644 frontend/web/src/components/files/FileUploadModal.tsx create mode 100644 frontend/web/src/hooks/useFiles.ts create mode 100644 frontend/web/src/utils/fileFormat.ts create mode 100644 frontend/web/src/utils/fileTree.ts create mode 100644 tests/integration/test_files_api.py create mode 100644 tests/unit/test_file_service.py diff --git a/backend/agenteval/services/__init__.py b/backend/agenteval/services/__init__.py new file mode 100644 index 0000000..a69ee20 --- /dev/null +++ b/backend/agenteval/services/__init__.py @@ -0,0 +1 @@ +"""Application services.""" diff --git a/backend/agenteval/services/files.py b/backend/agenteval/services/files.py new file mode 100644 index 0000000..a8338ed --- /dev/null +++ b/backend/agenteval/services/files.py @@ -0,0 +1,191 @@ +"""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 diff --git a/backend/agenteval/storage/file_repository.py b/backend/agenteval/storage/file_repository.py index 8fea3e4..94b627d 100644 --- a/backend/agenteval/storage/file_repository.py +++ b/backend/agenteval/storage/file_repository.py @@ -1,22 +1,25 @@ -"""Repository layer for file management (categories + records).""" +"""Database repositories for file categories and uploaded file records.""" -import os +from dataclasses import dataclass, field from typing import Optional -from sqlmodel import Session, select +from sqlmodel import Session, col, select -from agenteval.storage.db import ( - FILES_DIR, - FileCategoryDB, - FileRecordDB, - get_session, - new_uuid, - utc_now, -) +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: - """Repository for file categories (tree structure).""" + """Persistence operations for the file category tree.""" def __init__(self, session: Optional[Session] = None): self.session = session or get_session() @@ -28,87 +31,78 @@ class FileCategoryRepository: def get(self, category_id: str) -> Optional[FileCategoryDB]: return self.session.get(FileCategoryDB, category_id) - def get_tree(self) -> list[dict]: - """Return categories as a nested tree structure for frontend Tree component.""" - all_cats = self.list_all() - cat_map: dict[str, dict] = {} - roots: list[dict] = [] + 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 cat in all_cats: - node = { - "key": cat.id, - "title": cat.name, - "parent_id": cat.parent_id, - "children": [], - } - cat_map[cat.id] = node - - for cat in all_cats: - node = cat_map[cat.id] - if cat.parent_id and cat.parent_id in cat_map: - cat_map[cat.parent_id]["children"].append(node) + 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: - db = FileCategoryDB( + category = FileCategoryDB( id=new_uuid(), name=name, parent_id=parent_id, created_at=utc_now(), updated_at=utc_now(), ) - self.session.add(db) - self.session.commit() - self.session.refresh(db) - return db + self.session.add(category) + self.session.flush() + return category def update(self, category_id: str, name: str) -> Optional[FileCategoryDB]: - existing = self.session.get(FileCategoryDB, category_id) - if not existing: + category = self.session.get(FileCategoryDB, category_id) + if not category: return None - existing.name = name - existing.updated_at = utc_now() - self.session.add(existing) - self.session.commit() - self.session.refresh(existing) - return existing + category.name = name + category.updated_at = utc_now() + self.session.add(category) + self.session.flush() + return category def delete(self, category_id: str) -> bool: - """Delete a category and cascade-delete children + files. - - Physical files are cleaned up via FileRecordRepository.delete(). - We must explicitly delete files first to trigger physical cleanup, - because the DB cascade only removes rows. - """ - existing = self.session.get(FileCategoryDB, category_id) - if not existing: + category = self.session.get(FileCategoryDB, category_id) + if not category: return False - - # Collect all file IDs to clean up physical files - file_ids = self._collect_file_ids(existing) - - # Delete physical files - file_repo = FileRecordRepository(self.session) - for fid in file_ids: - file_repo._remove_physical_file(fid) - - self.session.delete(existing) - self.session.commit() + self.session.delete(category) + self.session.flush() return True - def _collect_file_ids(self, category: FileCategoryDB) -> list[str]: - """Recursively collect all file IDs under a category and its children.""" - file_ids = [f.id for f in category.files] - for child in category.children: - file_ids.extend(self._collect_file_ids(child)) - return file_ids - class FileRecordRepository: - """Repository for uploaded file records.""" + """Persistence operations for uploaded file metadata.""" def __init__(self, session: Optional[Session] = None): self.session = session or get_session() @@ -116,29 +110,16 @@ class FileRecordRepository: def list_all(self, category_id: Optional[str] = None) -> list[FileRecordDB]: statement = select(FileRecordDB) if category_id: - # Also include files in subcategories - cat_repo = FileCategoryRepository(self.session) - cat_ids = self._get_descendant_ids(category_id, cat_repo) - cat_ids.append(category_id) - from sqlmodel import col - - statement = statement.where(col(FileRecordDB.category_id).in_(cat_ids)) - else: - # Only filter when category_id is explicitly provided; - # None means "all files" (no filter). - pass + 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 _get_descendant_ids(self, parent_id: str, cat_repo: FileCategoryRepository) -> list[str]: - """Recursively collect IDs of all descendant categories.""" - ids: list[str] = [] - all_cats = cat_repo.list_all() - children = [c for c in all_cats if c.parent_id == parent_id] - for child in children: - ids.append(child.id) - ids.extend(self._get_descendant_ids(child.id, cat_repo)) - return ids + 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) @@ -152,7 +133,7 @@ class FileRecordRepository: file_ext: str, category_id: Optional[str] = None, ) -> FileRecordDB: - db = FileRecordDB( + record = FileRecordDB( id=new_uuid(), original_name=original_name, storage_name=storage_name, @@ -162,25 +143,14 @@ class FileRecordRepository: file_ext=file_ext, created_at=utc_now(), ) - self.session.add(db) - self.session.commit() - self.session.refresh(db) - return db + 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._remove_physical_file(file_id) self.session.delete(record) - self.session.commit() + self.session.flush() return True - - def _remove_physical_file(self, file_id: str) -> None: - """Remove the physical file from disk if it exists.""" - record = self.session.get(FileRecordDB, file_id) - if not record: - return - file_path = FILES_DIR / record.storage_name - if file_path.exists(): - os.remove(file_path) diff --git a/backend/agenteval/storage/file_storage.py b/backend/agenteval/storage/file_storage.py new file mode 100644 index 0000000..6075178 --- /dev/null +++ b/backend/agenteval/storage/file_storage.py @@ -0,0 +1,94 @@ +"""Physical storage operations for uploaded files.""" + +import os +import uuid +from dataclasses import dataclass +from pathlib import Path + +from fastapi import UploadFile + +UPLOAD_CHUNK_SIZE = 1024 * 1024 + + +class FileTooLargeError(Exception): + """Raised when a streamed upload exceeds its configured limit.""" + + +class InvalidStorageNameError(Exception): + """Raised when a storage name could escape the configured root.""" + + +@dataclass(frozen=True) +class PendingUpload: + temp_path: Path + storage_name: str + file_size: int + + +@dataclass(frozen=True) +class StagedDeletion: + original_path: Path + staged_path: Path + + +class FileStorage: + """Store files below one root and provide rollback-friendly operations.""" + + def __init__(self, root: Path): + self.root = root + self.root.mkdir(parents=True, exist_ok=True) + + def path_for(self, storage_name: str) -> Path: + if not storage_name or Path(storage_name).name != storage_name: + raise InvalidStorageNameError(storage_name) + return self.root / storage_name + + async def write_upload(self, upload: UploadFile, storage_name: str, max_bytes: int) -> PendingUpload: + self.path_for(storage_name) + temp_path = self.root / f".upload-{uuid.uuid4()}.tmp" + file_size = 0 + try: + with temp_path.open("xb") as output: + while chunk := await upload.read(UPLOAD_CHUNK_SIZE): + file_size += len(chunk) + if file_size > max_bytes: + raise FileTooLargeError + output.write(chunk) + except Exception: + temp_path.unlink(missing_ok=True) + raise + finally: + await upload.close() + return PendingUpload(temp_path=temp_path, storage_name=storage_name, file_size=file_size) + + def promote(self, pending: PendingUpload) -> Path: + final_path = self.path_for(pending.storage_name) + os.replace(pending.temp_path, final_path) + return final_path + + def discard(self, path: Path) -> None: + path.unlink(missing_ok=True) + + def stage_deletions(self, storage_names: list[str]) -> list[StagedDeletion]: + staged: list[StagedDeletion] = [] + try: + for storage_name in storage_names: + original_path = self.path_for(storage_name) + if not original_path.exists(): + continue + staged_path = self.root / f".delete-{uuid.uuid4()}.tmp" + os.replace(original_path, staged_path) + staged.append(StagedDeletion(original_path=original_path, staged_path=staged_path)) + except Exception: + self.restore_deletions(staged) + raise + return staged + + def restore_deletions(self, staged: list[StagedDeletion]) -> None: + for item in reversed(staged): + if item.staged_path.exists(): + os.replace(item.staged_path, item.original_path) + + def purge_deletions(self, staged: list[StagedDeletion]) -> None: + for item in staged: + item.staged_path.unlink(missing_ok=True) diff --git a/backend/agenteval/web/file_schemas.py b/backend/agenteval/web/file_schemas.py new file mode 100644 index 0000000..028ca9b --- /dev/null +++ b/backend/agenteval/web/file_schemas.py @@ -0,0 +1,69 @@ +"""HTTP schemas for the file management API.""" + +from datetime import datetime + +from pydantic import BaseModel, ConfigDict, Field, field_serializer + +from agenteval.storage.db import FileCategoryDB, FileRecordDB, iso_utc + + +class CategoryCreate(BaseModel): + name: str + parent_id: str | None = None + + +class CategoryUpdate(BaseModel): + name: str + + +class FileCategoryResponse(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: str + name: str + parent_id: str | None + children: list["FileCategoryResponse"] = Field(default_factory=list) + + @classmethod + def from_record(cls, category: FileCategoryDB) -> "FileCategoryResponse": + return cls( + id=category.id or "", + name=category.name, + parent_id=category.parent_id, + children=[], + ) + + +class FileRecordResponse(BaseModel): + id: str + original_name: str + file_size: int + mime_type: str + file_ext: str + category_id: str | None + created_at: datetime | None + + @field_serializer("created_at") + def serialize_created_at(self, value: datetime | None) -> str | None: + return iso_utc(value) + + @classmethod + def from_record(cls, record: FileRecordDB) -> "FileRecordResponse": + return cls( + id=record.id or "", + original_name=record.original_name, + file_size=record.file_size, + mime_type=record.mime_type, + file_ext=record.file_ext, + category_id=record.category_id, + created_at=record.created_at, + ) + + +class FileUploadConfigResponse(BaseModel): + allowed_extensions: list[str] + max_upload_size_mb: int + + +class OkResponse(BaseModel): + ok: bool = True diff --git a/backend/agenteval/web/routers/files.py b/backend/agenteval/web/routers/files.py index af0cdbe..523110f 100644 --- a/backend/agenteval/web/routers/files.py +++ b/backend/agenteval/web/routers/files.py @@ -1,203 +1,148 @@ -"""API routes for file management (categories + file upload/download).""" +"""API routes for file management.""" -import mimetypes -import uuid -from pathlib import Path +from typing import NoReturn from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile from fastapi.responses import FileResponse -from pydantic import BaseModel from sqlmodel import Session from agenteval.config import get_settings -from agenteval.storage.db import FILES_DIR, iso_utc -from agenteval.storage.file_repository import FileCategoryRepository, FileRecordRepository +from agenteval.services.files import ( + CategoryNotFoundError, + FileManagementError, + FileManagementService, + FileRecordNotFoundError, + FileValidationError, + PhysicalFileNotFoundError, + UploadTooLargeError, +) +from agenteval.storage.db import FILES_DIR +from agenteval.storage.file_storage import FileStorage from agenteval.web.deps import get_db +from agenteval.web.file_schemas import ( + CategoryCreate, + CategoryUpdate, + FileCategoryResponse, + FileRecordResponse, + FileUploadConfigResponse, + OkResponse, +) router = APIRouter() -# ── helpers ──────────────────────────────────────────────────────────── - -def _get_allowed_extensions() -> set[str]: - """Return the set of allowed lowercase extensions from settings.""" +def _allowed_extensions() -> set[str]: raw = get_settings().allowed_extensions - return {ext.strip().lower() for ext in raw.split(",") if ext.strip()} + return {extension.strip().lower() for extension in raw.split(",") if extension.strip()} -def _get_max_bytes() -> int: - return get_settings().max_upload_size_mb * 1024 * 1024 - - -def _validate_extension(filename: str) -> str: - """Validate the file extension and return the lowercase extension.""" - ext = Path(filename).suffix.lstrip(".").lower() - if not ext: - raise HTTPException(status_code=400, detail="文件没有扩展名") - allowed = _get_allowed_extensions() - if ext not in allowed: - raise HTTPException( - status_code=400, - detail=f"不支持的文件类型 .{ext},允许的类型: {', '.join(sorted(allowed))}", - ) - return ext - - -# ── API models ───────────────────────────────────────────────────────── - - -class CategoryCreate(BaseModel): - name: str - parent_id: str | None = None - - -class CategoryUpdate(BaseModel): - name: str - - -# ── Category endpoints ───────────────────────────────────────────────── - - -@router.get("/categories") -def list_categories(session: Session = Depends(get_db)) -> list[dict]: - """List categories as a nested tree.""" - return FileCategoryRepository(session).get_tree() - - -@router.post("/categories") -def create_category(body: CategoryCreate, session: Session = Depends(get_db)) -> dict: - if not body.name.strip(): - raise HTTPException(status_code=400, detail="分类名称不能为空") - cat = FileCategoryRepository(session).create( - name=body.name.strip(), - parent_id=body.parent_id, +def get_file_service(session: Session = Depends(get_db)) -> FileManagementService: + settings = get_settings() + return FileManagementService( + session=session, + storage=FileStorage(FILES_DIR), + allowed_extensions=_allowed_extensions(), + max_upload_size_mb=settings.max_upload_size_mb, ) - return { - "key": cat.id, - "title": cat.name, - "parent_id": cat.parent_id, - "children": [], - } -@router.put("/categories/{category_id}") -def update_category(category_id: str, body: CategoryUpdate, session: Session = Depends(get_db)) -> dict: - if not body.name.strip(): - raise HTTPException(status_code=400, detail="分类名称不能为空") - cat = FileCategoryRepository(session).update(category_id, body.name.strip()) - if not cat: - raise HTTPException(status_code=404, detail="分类不存在") - return {"ok": True, "name": cat.name} - - -@router.delete("/categories/{category_id}") -def delete_category(category_id: str, session: Session = Depends(get_db)) -> dict: - if not FileCategoryRepository(session).delete(category_id): - raise HTTPException(status_code=404, detail="分类不存在") - return {"ok": True} - - -# ── File endpoints ───────────────────────────────────────────────────── - - -@router.get("") -def list_files(category_id: str | None = None, session: Session = Depends(get_db)) -> list[dict]: - """List file records, optionally filtered by category.""" - records = FileRecordRepository(session).list_all(category_id=category_id) - return [ - { - "id": r.id, - "original_name": r.original_name, - "file_size": r.file_size, - "mime_type": r.mime_type, - "file_ext": r.file_ext, - "category_id": r.category_id, - "created_at": iso_utc(r.created_at), - } - for r in records - ] - - -@router.post("/upload") -async def upload_file( - file: UploadFile = File(...), - category_id: str | None = Form(default=None), - session: Session = Depends(get_db), -) -> dict: - """Upload a single file.""" - if not file.filename: - raise HTTPException(status_code=400, detail="文件名为空") - - ext = _validate_extension(file.filename) - - # Read content and validate size - content = await file.read() - max_bytes = _get_max_bytes() - if len(content) > max_bytes: +def _raise_http_error(error: FileManagementError) -> NoReturn: + if isinstance(error, UploadTooLargeError): raise HTTPException( status_code=413, detail=f"文件大小超过限制 ({get_settings().max_upload_size_mb}MB)", - ) + ) from error + if isinstance(error, FileValidationError): + raise HTTPException(status_code=400, detail=str(error)) from error + if isinstance(error, CategoryNotFoundError): + raise HTTPException(status_code=404, detail="分类不存在") from error + if isinstance(error, FileRecordNotFoundError): + raise HTTPException(status_code=404, detail="文件不存在") from error + if isinstance(error, PhysicalFileNotFoundError): + raise HTTPException(status_code=404, detail="物理文件不存在") from error + raise HTTPException(status_code=500, detail="文件操作失败") from error - # Generate storage name - storage_name = f"{uuid.uuid4()}.{ext}" - # Guess MIME type - mime_type, _ = mimetypes.guess_type(file.filename) - if not mime_type: - mime_type = "application/octet-stream" - - # Validate category exists if provided - if category_id: - cat = FileCategoryRepository(session).get(category_id) - if not cat: - raise HTTPException(status_code=400, detail="分类不存在") - - # Write physical file - file_path = FILES_DIR / storage_name - file_path.write_bytes(content) - - # Create DB record - record = FileRecordRepository(session).create( - original_name=file.filename, - storage_name=storage_name, - file_size=len(content), - mime_type=mime_type, - file_ext=ext, - category_id=category_id if category_id else None, +@router.get("/config", response_model=FileUploadConfigResponse) +def get_upload_config() -> FileUploadConfigResponse: + return FileUploadConfigResponse( + allowed_extensions=sorted(_allowed_extensions()), + max_upload_size_mb=get_settings().max_upload_size_mb, ) - return { - "id": record.id, - "original_name": record.original_name, - "file_size": record.file_size, - "mime_type": record.mime_type, - "file_ext": record.file_ext, - "category_id": record.category_id, - "created_at": iso_utc(record.created_at), - } + +@router.get("/categories", response_model=list[FileCategoryResponse]) +def list_categories(service: FileManagementService = Depends(get_file_service)): + return service.list_categories() + + +@router.post("/categories", response_model=FileCategoryResponse) +def create_category(body: CategoryCreate, service: FileManagementService = Depends(get_file_service)): + try: + return FileCategoryResponse.from_record(service.create_category(body.name, body.parent_id)) + except FileManagementError as error: + _raise_http_error(error) + + +@router.put("/categories/{category_id}", response_model=FileCategoryResponse) +def update_category( + category_id: str, + body: CategoryUpdate, + service: FileManagementService = Depends(get_file_service), +): + try: + return FileCategoryResponse.from_record(service.update_category(category_id, body.name)) + except FileManagementError as error: + _raise_http_error(error) + + +@router.delete("/categories/{category_id}", response_model=OkResponse) +def delete_category(category_id: str, service: FileManagementService = Depends(get_file_service)): + try: + service.delete_category(category_id) + return OkResponse() + except FileManagementError as error: + _raise_http_error(error) + + +@router.get("", response_model=list[FileRecordResponse]) +def list_files(category_id: str | None = None, service: FileManagementService = Depends(get_file_service)): + try: + return [FileRecordResponse.from_record(record) for record in service.list_files(category_id)] + except FileManagementError as error: + _raise_http_error(error) + + +@router.post("/upload", response_model=FileRecordResponse) +async def upload_file( + file: UploadFile = File(...), + category_id: str | None = Form(default=None), + service: FileManagementService = Depends(get_file_service), +): + try: + record = await service.upload(file, category_id) + return FileRecordResponse.from_record(record) + except FileManagementError as error: + _raise_http_error(error) @router.get("/{file_id}/download") -def download_file(file_id: str, session: Session = Depends(get_db)): - """Download a file by its record ID.""" - record = FileRecordRepository(session).get(file_id) - if not record: - raise HTTPException(status_code=404, detail="文件不存在") - - file_path = FILES_DIR / record.storage_name - if not file_path.exists(): - raise HTTPException(status_code=404, detail="物理文件不存在") - - return FileResponse( - path=str(file_path), - filename=record.original_name, - media_type=record.mime_type or "application/octet-stream", - ) +def download_file(file_id: str, service: FileManagementService = Depends(get_file_service)): + try: + record, file_path = service.get_download(file_id) + return FileResponse( + path=str(file_path), + filename=record.original_name, + media_type=record.mime_type or "application/octet-stream", + ) + except FileManagementError as error: + _raise_http_error(error) -@router.delete("/{file_id}") -def delete_file(file_id: str, session: Session = Depends(get_db)) -> dict: - if not FileRecordRepository(session).delete(file_id): - raise HTTPException(status_code=404, detail="文件不存在") - return {"ok": True} +@router.delete("/{file_id}", response_model=OkResponse) +def delete_file(file_id: str, service: FileManagementService = Depends(get_file_service)): + try: + service.delete_file(file_id) + return OkResponse() + except FileManagementError as error: + _raise_http_error(error) diff --git a/frontend/web/src/api.ts b/frontend/web/src/api.ts index 3e2004d..c943940 100644 --- a/frontend/web/src/api.ts +++ b/frontend/web/src/api.ts @@ -153,8 +153,8 @@ export const statsApi = { // ── File Management ────────────────────────────────────────────── export interface FileCategory { - key: string - title: string + id: string + name: string parent_id: string | null children: FileCategory[] } @@ -169,8 +169,14 @@ export interface FileRecord { created_at: string | null } +export interface FileUploadConfig { + allowed_extensions: string[] + max_upload_size_mb: number +} + export const filesApi = { // Categories + getConfig: () => api.get('/files/config'), listCategories: () => api.get('/files/categories'), createCategory: (data: { name: string; parent_id?: string | null }) => api.post('/files/categories', data), @@ -191,6 +197,6 @@ export const filesApi = { timeout: 120000, }) }, - downloadUrl: (id: string) => `/api/files/${id}/download`, + download: (id: string) => api.get(`/files/${id}/download`, { responseType: 'blob' }), delete: (id: string) => api.delete(`/files/${id}`), } diff --git a/frontend/web/src/components/files/FileCategoryTree.tsx b/frontend/web/src/components/files/FileCategoryTree.tsx new file mode 100644 index 0000000..2f2812b --- /dev/null +++ b/frontend/web/src/components/files/FileCategoryTree.tsx @@ -0,0 +1,127 @@ +import { + DeleteOutlined, + EditOutlined, + FileOutlined, + FolderAddOutlined, + FolderOutlined, + PlusOutlined, +} from '@ant-design/icons' +import { Button, Popconfirm, Tooltip, Tree } from 'antd' +import type { DataNode } from 'antd/es/tree' +import type { FileCategory } from '../../api' +import { colors, fontSizes } from '../../tokens' + +const ALL_FILES_KEY = '__all__' + +interface FileCategoryTreeProps { + categories: FileCategory[] + selectedCategoryId: string | null + onSelect: (categoryId: string | null) => void + onCreate: (parentId?: string) => void + onRename: (category: FileCategory) => void + onDelete: (categoryId: string) => Promise +} + +export default function FileCategoryTree({ + categories, + selectedCategoryId, + onSelect, + onCreate, + onRename, + onDelete, +}: FileCategoryTreeProps) { + const mapCategory = (category: FileCategory): DataNode => ({ + key: category.id, + icon: , + title: ( +
+ {category.name} + + +
+ ), + children: category.children.map(mapCategory), + }) + + const treeData: DataNode[] = [{ + key: ALL_FILES_KEY, + icon: , + title: ( + + + 全部文件 + + ), + children: categories.map(mapCategory), + }] + + return ( + + ) +} diff --git a/frontend/web/src/components/files/FileTable.tsx b/frontend/web/src/components/files/FileTable.tsx new file mode 100644 index 0000000..7f27c71 --- /dev/null +++ b/frontend/web/src/components/files/FileTable.tsx @@ -0,0 +1,182 @@ +import { + DeleteOutlined, + DownloadOutlined, + FileOutlined, + FolderOutlined, + InboxOutlined, + ReloadOutlined, + UploadOutlined, +} from '@ant-design/icons' +import { Button, Popconfirm, Space, Table, Tag, Tooltip } from 'antd' +import type { ColumnsType } from 'antd/es/table' +import { useMemo } from 'react' +import type { FileRecord } from '../../api' +import { colors, fontSizes } from '../../tokens' +import { formatDateTime } from '../../utils/date' +import { extensionColor, formatFileSize } from '../../utils/fileFormat' + +interface FileTableProps { + files: FileRecord[] + loading: boolean + categoryName: string + onRefresh: () => void + onOpenUpload: () => void + onDownload: (record: FileRecord) => Promise + onDelete: (id: string) => Promise +} + +export default function FileTable({ + files, + loading, + categoryName, + onRefresh, + onOpenUpload, + onDownload, + onDelete, +}: FileTableProps) { + const totalSize = useMemo(() => files.reduce((sum, file) => sum + file.file_size, 0), [files]) + const extensionCounts = useMemo(() => files.reduce>((counts, file) => { + counts[file.file_ext] = (counts[file.file_ext] || 0) + 1 + return counts + }, {}), [files]) + + const columns: ColumnsType = [ + { + title: '文件名', + dataIndex: 'original_name', + key: 'name', + ellipsis: true, + render: (name: string, record) => ( + + + {name} + + ), + }, + { + title: '大小', + dataIndex: 'file_size', + key: 'size', + width: 100, + responsive: ['sm'], + render: (size: number) => ( + + {formatFileSize(size)} + + ), + }, + { + title: '类型', + dataIndex: 'file_ext', + key: 'type', + width: 90, + render: (extension: string) => ( + + {extension.toUpperCase()} + + ), + }, + { + title: '上传时间', + dataIndex: 'created_at', + key: 'created_at', + width: 170, + responsive: ['md'], + render: (time: string | null) => ( + + {time ? formatDateTime(time) : '-'} + + ), + }, + { + title: '操作', + key: 'action', + width: 100, + fixed: 'right', + render: (_, record) => ( + + + + + +
+ + dataSource={files} + columns={columns} + rowKey="id" + loading={loading} + scroll={{ x: 620 }} + pagination={files.length > 10 ? { + pageSize: 15, + showSizeChanger: true, + pageSizeOptions: ['15', '30', '50'], + showTotal: (total) => `共 ${total} 个文件`, + } : false} + locale={{ + emptyText: ( +
+ + 暂无文件 +
+ ), + }} + /> +
+ + ) +} diff --git a/frontend/web/src/components/files/FileUploadModal.tsx b/frontend/web/src/components/files/FileUploadModal.tsx new file mode 100644 index 0000000..d437dc3 --- /dev/null +++ b/frontend/web/src/components/files/FileUploadModal.tsx @@ -0,0 +1,121 @@ +import { InboxOutlined } from '@ant-design/icons' +import { message, Modal, Select, Upload } from 'antd' +import type { UploadFile, UploadProps } from 'antd' +import { useEffect, useMemo, useState } from 'react' +import type { FileCategory, FileUploadConfig } from '../../api' +import { flattenCategoryOptions } from '../../utils/fileTree' + +const { Dragger } = Upload + +interface UploadResult { + succeeded: number + failed: number +} + +interface FileUploadModalProps { + open: boolean + categories: FileCategory[] + config: FileUploadConfig + initialCategoryId: string | null + onClose: () => void + onUpload: (files: File[], categoryId: string | null) => Promise +} + +export default function FileUploadModal({ + open, + categories, + config, + initialCategoryId, + onClose, + onUpload, +}: FileUploadModalProps) { + const [targetCategoryId, setTargetCategoryId] = useState() + const [fileList, setFileList] = useState([]) + const [uploading, setUploading] = useState(false) + + useEffect(() => { + if (open) { + setTargetCategoryId(initialCategoryId ?? undefined) + setFileList([]) + } + }, [initialCategoryId, open]) + + const options = useMemo(() => flattenCategoryOptions(categories), [categories]) + const accept = config.allowed_extensions.map((extension) => `.${extension}`).join(',') + + const uploadProps: UploadProps = { + multiple: true, + accept: accept || undefined, + fileList, + customRequest: () => undefined, + beforeUpload: (file) => { + const maxBytes = config.max_upload_size_mb * 1024 * 1024 + if (file.size > maxBytes) { + message.error(`${file.name} 超过 ${config.max_upload_size_mb}MB 限制`) + return Upload.LIST_IGNORE + } + return true + }, + onChange: ({ fileList: nextFileList }) => setFileList(nextFileList), + onRemove: (file) => { + setFileList((current) => current.filter((item) => item.uid !== file.uid)) + return true + }, + } + + const handleUpload = async () => { + const files = fileList.flatMap((item) => item.originFileObj ? [item.originFileObj] : []) + if (files.length === 0) return + + setUploading(true) + try { + const result = await onUpload(files, targetCategoryId ?? null) + if (result.succeeded > 0 && result.failed === 0) { + message.success(`成功上传 ${result.succeeded} 个文件`) + } else if (result.succeeded > 0) { + message.warning(`成功 ${result.succeeded} 个,失败 ${result.failed} 个`) + } else { + message.error('上传失败,请检查文件格式和大小') + } + onClose() + } finally { + setUploading(false) + } + } + + const formats = config.allowed_extensions.length > 0 + ? config.allowed_extensions.map((extension) => `.${extension}`).join(' ') + : '正在读取支持格式' + + return ( + void handleUpload()} + okText="开始上传" + confirmLoading={uploading} + okButtonProps={{ disabled: fileList.length === 0 }} + width={520} + destroyOnHidden + > +
+ 目标分类: + setUploadCategory(v)} - options={categories.flatMap((c) => [ - { value: c.key, label: c.title }, - ...c.children.map((ch) => ({ value: ch.key, label: ` └ ${ch.title}` })), - ])} - /> -
- -

-

点击或拖拽文件到此区域上传

-

- 支持 .txt .md .json .yaml .csv .xml .xlsx .png .jpg .gif .svg .zip .py .js .ts 格式,单文件最大 50MB -

-
-
+ categories={categories} + config={config} + initialCategoryId={selectedCategoryId} + onClose={() => setUploadOpen(false)} + onUpload={uploadFiles} + /> - {/* Category Dialog */} setCatDialogOpen(false)} - okText={catDialogMode === 'create' ? '创建' : '保存'} - okButtonProps={{ disabled: !catDialogName.trim() }} + title={categoryDialog.mode === 'create' ? '新建分类' : '重命名分类'} + open={categoryDialog.open} + onOk={() => void saveCategory()} + onCancel={() => setCategoryDialog(CLOSED_CATEGORY_DIALOG)} + okText={categoryDialog.mode === 'create' ? '创建' : '保存'} + confirmLoading={categorySaving} + okButtonProps={{ disabled: !categoryDialog.name.trim() }} width={360} + destroyOnHidden > setCatDialogName(e.target.value)} - onPressEnter={handleCatDialogOk} + value={categoryDialog.name} + onChange={(event) => setCategoryDialog((current) => ({ + ...current, + name: event.target.value, + }))} + onPressEnter={() => void saveCategory()} /> - + ) -} \ No newline at end of file +} diff --git a/frontend/web/src/utils/fileFormat.ts b/frontend/web/src/utils/fileFormat.ts new file mode 100644 index 0000000..b944239 --- /dev/null +++ b/frontend/web/src/utils/fileFormat.ts @@ -0,0 +1,31 @@ +export function formatFileSize(bytes: number): string { + if (bytes === 0) return '0 B' + const units = ['B', 'KB', 'MB', 'GB'] + const index = Math.min(Math.floor(Math.log(bytes) / Math.log(1024)), units.length - 1) + return `${(bytes / Math.pow(1024, index)).toFixed(index === 0 ? 0 : 1)} ${units[index]}` +} + +const extensionColors: Record = { + txt: '#666666', + md: '#1677ff', + json: '#389e0d', + yaml: '#08979c', + yml: '#08979c', + csv: '#d46b08', + xml: '#531dab', + xlsx: '#237804', + xls: '#237804', + png: '#c41d7f', + jpg: '#c41d7f', + jpeg: '#c41d7f', + gif: '#c41d7f', + svg: '#c41d7f', + zip: '#cf1322', + py: '#1d39c4', + js: '#ad8b00', + ts: '#0958d9', +} + +export function extensionColor(extension: string): string { + return extensionColors[extension.toLowerCase()] || '#595959' +} diff --git a/frontend/web/src/utils/fileTree.ts b/frontend/web/src/utils/fileTree.ts new file mode 100644 index 0000000..c62e512 --- /dev/null +++ b/frontend/web/src/utils/fileTree.ts @@ -0,0 +1,32 @@ +import type { FileCategory } from '../api' + +export interface CategoryOption { + value: string + label: string +} + +export function findCategory(categories: FileCategory[], id: string): FileCategory | undefined { + for (const category of categories) { + if (category.id === id) return category + const nested = findCategory(category.children, id) + if (nested) return nested + } + return undefined +} + +export function categoryContains(category: FileCategory, id: string): boolean { + return category.id === id || category.children.some((child) => categoryContains(child, id)) +} + +export function flattenCategoryOptions( + categories: FileCategory[], + depth = 0, +): CategoryOption[] { + return categories.flatMap((category) => [ + { + value: category.id, + label: `${' '.repeat(depth)}${depth > 0 ? '└ ' : ''}${category.name}`, + }, + ...flattenCategoryOptions(category.children, depth + 1), + ]) +} diff --git a/scripts/deploy-t480.sh b/scripts/deploy-t480.sh index 258badd..2d6b0fe 100755 --- a/scripts/deploy-t480.sh +++ b/scripts/deploy-t480.sh @@ -85,6 +85,7 @@ run rsync -az --delete \ --exclude='__pycache__' \ --exclude='.pytest_cache' \ --exclude='.ruff_cache' \ + --exclude='AGENTS.md' \ --exclude='data' \ --exclude='.env' \ --exclude='config/config.json' \ diff --git a/tests/integration/test_files_api.py b/tests/integration/test_files_api.py new file mode 100644 index 0000000..9e6ea4b --- /dev/null +++ b/tests/integration/test_files_api.py @@ -0,0 +1,198 @@ +"""Integration tests for the file management API.""" + +from pathlib import Path + +import pytest +from agenteval.config import get_settings +from agenteval.web.app import app +from agenteval.web.deps import get_db +from fastapi.testclient import TestClient +from sqlmodel import Session, SQLModel, create_engine + + +@pytest.fixture() +def files_client(tmp_path: Path, monkeypatch): + from agenteval.storage.db import ( # noqa: F401 + EvalResultDB, + EvalRunDB, + EvalTargetDB, + FileCategoryDB, + FileRecordDB, + ScenarioDB, + TurnDB, + ) + from agenteval.web.routers import files as files_module + + engine = create_engine( + f"sqlite:///{tmp_path / 'files_api.db'}", + connect_args={"check_same_thread": False}, + ) + SQLModel.metadata.create_all(engine) + session = Session(engine) + files_dir = tmp_path / "files" + files_dir.mkdir() + + monkeypatch.setattr(files_module, "FILES_DIR", files_dir) + settings = get_settings() + monkeypatch.setattr(settings, "allowed_extensions", "txt,json") + monkeypatch.setattr(settings, "max_upload_size_mb", 1) + + def override_get_db(): + yield session + + app.dependency_overrides[get_db] = override_get_db + with TestClient(app) as client: + yield client, session, files_dir + + app.dependency_overrides.clear() + session.close() + engine.dispose() + + +def _create_category(client: TestClient, name: str, parent_id: str | None = None) -> dict: + response = client.post( + "/api/files/categories", + json={"name": name, "parent_id": parent_id}, + ) + assert response.status_code == 200 + return response.json() + + +def test_upload_config_uses_server_settings(files_client): + client, _, _ = files_client + + response = client.get("/api/files/config") + + assert response.status_code == 200 + assert response.json() == { + "allowed_extensions": ["json", "txt"], + "max_upload_size_mb": 1, + } + + +def test_category_crud_returns_domain_fields_and_deep_tree(files_client): + client, _, _ = files_client + root = _create_category(client, "根分类") + child = _create_category(client, "二级", root["id"]) + leaf = _create_category(client, "三级", child["id"]) + + tree = client.get("/api/files/categories").json() + + assert tree[0]["id"] == root["id"] + assert tree[0]["name"] == "根分类" + assert tree[0]["children"][0]["id"] == child["id"] + assert tree[0]["children"][0]["children"][0]["id"] == leaf["id"] + + updated = client.put( + f"/api/files/categories/{leaf['id']}", + json={"name": "三级分类"}, + ) + assert updated.status_code == 200 + assert updated.json()["name"] == "三级分类" + + +def test_create_category_rejects_missing_parent(files_client): + client, _, _ = files_client + + response = client.post( + "/api/files/categories", + json={"name": "孤立分类", "parent_id": "missing"}, + ) + + assert response.status_code == 400 + assert response.json()["detail"] == "父分类不存在" + + +def test_list_files_includes_all_descendant_categories(files_client): + client, _, _ = files_client + root = _create_category(client, "根分类") + child = _create_category(client, "二级", root["id"]) + leaf = _create_category(client, "三级", child["id"]) + + upload = client.post( + "/api/files/upload", + data={"category_id": leaf["id"]}, + files={"file": ("deep.txt", b"deep content", "text/plain")}, + ) + assert upload.status_code == 200 + + response = client.get("/api/files", params={"category_id": root["id"]}) + + assert response.status_code == 200 + assert [item["original_name"] for item in response.json()] == ["deep.txt"] + + +def test_upload_download_and_delete_file(files_client): + client, _, files_dir = files_client + + uploaded = client.post( + "/api/files/upload", + files={"file": ("notes.txt", b"hello", "text/plain")}, + ) + + assert uploaded.status_code == 200 + record = uploaded.json() + assert record["original_name"] == "notes.txt" + assert record["file_size"] == 5 + assert len(list(files_dir.iterdir())) == 1 + + downloaded = client.get(f"/api/files/{record['id']}/download") + assert downloaded.status_code == 200 + assert downloaded.content == b"hello" + + deleted = client.delete(f"/api/files/{record['id']}") + assert deleted.status_code == 200 + assert list(files_dir.iterdir()) == [] + assert client.get("/api/files").json() == [] + + +def test_upload_rejects_invalid_category_and_extension(files_client): + client, _, files_dir = files_client + + invalid_category = client.post( + "/api/files/upload", + data={"category_id": "missing"}, + files={"file": ("notes.txt", b"hello", "text/plain")}, + ) + invalid_extension = client.post( + "/api/files/upload", + files={"file": ("notes.exe", b"hello", "application/octet-stream")}, + ) + + assert invalid_category.status_code == 400 + assert invalid_category.json()["detail"] == "分类不存在" + assert invalid_extension.status_code == 400 + assert list(files_dir.iterdir()) == [] + + +def test_upload_rejects_oversized_file_without_disk_artifact(files_client): + client, _, files_dir = files_client + + response = client.post( + "/api/files/upload", + files={"file": ("large.txt", b"x" * (1024 * 1024 + 1), "text/plain")}, + ) + + assert response.status_code == 413 + assert list(files_dir.iterdir()) == [] + assert client.get("/api/files").json() == [] + + +def test_delete_category_removes_descendant_files(files_client): + client, _, files_dir = files_client + root = _create_category(client, "根分类") + child = _create_category(client, "子分类", root["id"]) + uploaded = client.post( + "/api/files/upload", + data={"category_id": child["id"]}, + files={"file": ("child.txt", b"content", "text/plain")}, + ) + assert uploaded.status_code == 200 + assert len(list(files_dir.iterdir())) == 1 + + response = client.delete(f"/api/files/categories/{root['id']}") + + assert response.status_code == 200 + assert client.get("/api/files/categories").json() == [] + assert client.get("/api/files").json() == [] + assert list(files_dir.iterdir()) == [] diff --git a/tests/unit/test_file_repository.py b/tests/unit/test_file_repository.py index 2da7a23..750fa18 100644 --- a/tests/unit/test_file_repository.py +++ b/tests/unit/test_file_repository.py @@ -1,19 +1,23 @@ """Unit tests for FileCategoryRepository and FileRecordRepository.""" -import uuid from pathlib import Path import pytest -from sqlmodel import Session, SQLModel, create_engine - from agenteval.storage.file_repository import FileCategoryRepository, FileRecordRepository +from sqlmodel import Session, SQLModel, create_engine @pytest.fixture() def file_db_session(tmp_path: Path): """Isolated SQLite session with file management tables.""" from agenteval.storage.db import ( # noqa: F401 - EvalResultDB, EvalRunDB, EvalTargetDB, FileCategoryDB, FileRecordDB, ScenarioDB, TurnDB, + EvalResultDB, + EvalRunDB, + EvalTargetDB, + FileCategoryDB, + FileRecordDB, + ScenarioDB, + TurnDB, ) engine = create_engine( f"sqlite:///{tmp_path / 'file_test.db'}", @@ -28,19 +32,6 @@ def file_db_session(tmp_path: Path): engine.dispose() -@pytest.fixture() -def fake_files_dir(tmp_path: Path, monkeypatch): - """Redirect FILES_DIR to a temp directory so tests don't touch data/files.""" - import agenteval.storage.db as db_mod - import agenteval.storage.file_repository as repo_mod - - fake_dir = tmp_path / "files" - fake_dir.mkdir() - monkeypatch.setattr(db_mod, "FILES_DIR", fake_dir) - monkeypatch.setattr(repo_mod, "FILES_DIR", fake_dir) - return fake_dir - - # ── FileCategoryRepository ──────────────────────────────────────────────── def test_create_root_category(file_db_session): @@ -73,8 +64,8 @@ def test_get_tree_structure(file_db_session): repo.create("child2", parent_id=root.id) tree = repo.get_tree() assert len(tree) == 1 - assert tree[0]["key"] == root.id - assert len(tree[0]["children"]) == 2 + assert tree[0].id == root.id + assert len(tree[0].children) == 2 def test_get_tree_flat(file_db_session): @@ -83,7 +74,7 @@ def test_get_tree_flat(file_db_session): repo.create("B") tree = repo.get_tree() assert len(tree) == 2 - assert all(len(node["children"]) == 0 for node in tree) + assert all(len(node.children) == 0 for node in tree) def test_update_category_name(file_db_session): @@ -116,7 +107,7 @@ def test_delete_nonexistent_returns_false(file_db_session): assert repo.delete("ghost-id") is False -def test_delete_cascades_to_children(file_db_session, fake_files_dir): +def test_delete_cascades_to_children(file_db_session): cat_repo = FileCategoryRepository(file_db_session) file_repo = FileRecordRepository(file_db_session) @@ -203,16 +194,13 @@ def test_delete_file_record(file_db_session): assert repo.get(rec.id) is None -def test_delete_removes_physical_file(file_db_session, fake_files_dir): - repo = FileRecordRepository(file_db_session) - storage_name = f"{uuid.uuid4()}.txt" - physical = fake_files_dir / storage_name - physical.write_text("content") - assert physical.exists() +def test_get_subtree_ids(file_db_session): + repo = FileCategoryRepository(file_db_session) + root = repo.create("root") + child = repo.create("child", parent_id=root.id) + leaf = repo.create("leaf", parent_id=child.id) - rec = repo.create("orig.txt", storage_name, 7, "text/plain", "txt") - repo.delete(rec.id) - assert not physical.exists() + assert set(repo.get_subtree_ids(root.id)) == {root.id, child.id, leaf.id} def test_delete_nonexistent_file_returns_false(file_db_session): diff --git a/tests/unit/test_file_service.py b/tests/unit/test_file_service.py new file mode 100644 index 0000000..10b2348 --- /dev/null +++ b/tests/unit/test_file_service.py @@ -0,0 +1,92 @@ +"""Tests for file workflow transaction compensation and path safety.""" + +from io import BytesIO +from pathlib import Path + +import pytest +from agenteval.services.files import FileManagementService +from agenteval.storage.file_repository import FileRecordRepository +from agenteval.storage.file_storage import FileStorage, InvalidStorageNameError +from sqlmodel import Session, SQLModel, create_engine +from starlette.datastructures import Headers, UploadFile + + +@pytest.fixture() +def service_context(tmp_path: Path): + from agenteval.storage.db import ( # noqa: F401 + EvalResultDB, + EvalRunDB, + EvalTargetDB, + FileCategoryDB, + FileRecordDB, + ScenarioDB, + TurnDB, + ) + + engine = create_engine( + f"sqlite:///{tmp_path / 'file_service.db'}", + connect_args={"check_same_thread": False}, + ) + SQLModel.metadata.create_all(engine) + session = Session(engine) + storage = FileStorage(tmp_path / "files") + service = FileManagementService(session, storage, {"txt"}, 1) + yield service, session, storage + session.close() + engine.dispose() + + +def _upload(content: bytes = b"content") -> UploadFile: + return UploadFile( + filename="test.txt", + file=BytesIO(content), + headers=Headers({"content-type": "text/plain"}), + ) + + +async def test_upload_commit_failure_removes_physical_file(service_context, monkeypatch): + service, session, storage = service_context + + def fail_commit(): + raise RuntimeError("commit failed") + + monkeypatch.setattr(session, "commit", fail_commit) + + with pytest.raises(RuntimeError, match="commit failed"): + await service.upload(_upload()) + + assert list(storage.root.iterdir()) == [] + assert FileRecordRepository(session).list_all() == [] + + +def test_delete_commit_failure_restores_physical_file(service_context, monkeypatch): + service, session, storage = service_context + record = FileRecordRepository(session).create( + original_name="test.txt", + storage_name="stored.txt", + file_size=7, + mime_type="text/plain", + file_ext="txt", + ) + session.commit() + physical_file = storage.path_for("stored.txt") + physical_file.write_bytes(b"content") + + def fail_commit(): + raise RuntimeError("commit failed") + + monkeypatch.setattr(session, "commit", fail_commit) + + with pytest.raises(RuntimeError, match="commit failed"): + service.delete_file(record.id) + + assert physical_file.read_bytes() == b"content" + assert FileRecordRepository(session).get(record.id) is not None + + +@pytest.mark.parametrize("storage_name", ["../escape.txt", "folder/file.txt", ""]) +def test_storage_rejects_names_outside_root(service_context, storage_name): + _, _, storage = service_context + + with pytest.raises(InvalidStorageNameError): + storage.path_for(storage_name)