diff --git a/app/controllers/file_controller.py b/app/controllers/file_controller.py index 433aed6..a326ef5 100644 --- a/app/controllers/file_controller.py +++ b/app/controllers/file_controller.py @@ -10,7 +10,12 @@ from sqlalchemy.orm import Session from ..database import get_db from ..dao.uploaded_file_dao import UploadedFileDAO -from ..schemas.file import FileListResponse, FileUploadResponse, UploadedFileOut +from ..schemas.file import ( + FileListResponse, + FileUploadResponse, + SftpRegisterRequest, + UploadedFileOut, +) from ..services.upload_service import UploadService router = APIRouter(prefix="/api/files", tags=["files"]) @@ -49,6 +54,42 @@ def list_files( return FileListResponse(total=total, items=items) +@router.get( + "/exists", + response_model=UploadedFileOut, + summary="按 sha256 查询是否已上传", + description="命中返回 200 + 元数据;未命中返回 404。供客户端在上传前去重。", +) +def file_exists( + sha256: str, + service: UploadService = Depends(_service), +) -> UploadedFileOut: + out = service.find_by_sha256(sha256) + if out is None: + raise HTTPException(404, "未找到匹配的 sha256") + return out + + +@router.post( + "/register-sftp", + response_model=FileUploadResponse, + summary="登记一个已通过 SFTP 落盘的文件", + description=( + "客户端先把文件 SFTP 到 ``incoming/``,再用本接口登记入库。" + "服务端会计算 sha256(已存在则去重)、把文件原子改名到 ``YYYY/MM/.``、写 DB 行。" + ), +) +def register_sftp( + body: SftpRegisterRequest, + service: UploadService = Depends(_service), +) -> FileUploadResponse: + return service.register_sftp( + filename=body.filename, + original_filename=body.original_filename, + uploaded_by=body.uploaded_by, + ) + + @router.get("/{file_id}", response_model=UploadedFileOut, summary="查询单个文件元数据") def get_file(file_id: int, service: UploadService = Depends(_service)) -> UploadedFileOut: out = service.get_out(file_id) diff --git a/app/dao/uploaded_file_dao.py b/app/dao/uploaded_file_dao.py index e8227c6..7638917 100644 --- a/app/dao/uploaded_file_dao.py +++ b/app/dao/uploaded_file_dao.py @@ -21,6 +21,10 @@ class UploadedFileDAO: def get_by_id(self, file_id: int) -> UploadedFile | None: return self.db.get(UploadedFile, file_id) + def get_by_sha256(self, sha256: str) -> UploadedFile | None: + stmt = select(UploadedFile).where(UploadedFile.sha256 == sha256).limit(1) + return self.db.scalars(stmt).first() + def list(self, limit: int = 100, offset: int = 0) -> list[UploadedFile]: stmt = ( select(UploadedFile) diff --git a/app/schemas/file.py b/app/schemas/file.py index 136b97d..ef2c3a7 100644 --- a/app/schemas/file.py +++ b/app/schemas/file.py @@ -33,3 +33,14 @@ class FileUploadResponse(BaseModel): class FileListResponse(BaseModel): total: int items: list[UploadedFileOut] + + +class SftpRegisterRequest(BaseModel): + """登记一个通过 SFTP 上传到 ``incoming/`` 下的文件。""" + + filename: str = Field( + ..., + description="文件在 SFTP chroot 下的相对路径,必须落在 incoming/ 之下,例如 incoming/abc.bin", + ) + original_filename: str = Field(..., description="客户端原始文件名") + uploaded_by: str = Field(default="sftp", description="登记者标识,写入 uploaded_by 字段") diff --git a/app/services/sftp_server.py b/app/services/sftp_server.py index 899d173..dbb74fc 100644 --- a/app/services/sftp_server.py +++ b/app/services/sftp_server.py @@ -119,6 +119,8 @@ async def _run() -> None: upload_root = settings.resolved_upload_dir() upload_root.mkdir(parents=True, exist_ok=True) + # SFTP 客户端登记前的暂存目录;register-sftp 只接受此目录下的路径。 + (upload_root / "incoming").mkdir(parents=True, exist_ok=True) host_key_path = (PROJECT_ROOT / settings.sftp.host_key_path).resolve() _ensure_host_key(host_key_path) diff --git a/app/services/upload_service.py b/app/services/upload_service.py index acc1a6e..e00ee9e 100644 --- a/app/services/upload_service.py +++ b/app/services/upload_service.py @@ -8,13 +8,16 @@ import uuid from datetime import datetime, timezone from pathlib import Path -from fastapi import UploadFile +from fastapi import HTTPException, UploadFile from ..config import get_settings from ..dao.uploaded_file_dao import UploadedFileDAO from ..models.uploaded_file import UploadedFile from ..schemas.file import FileUploadResponse, UploadedFileOut +# SFTP 客户端登记前必须先把文件写入此目录(chroot 内) +SFTP_INCOMING_DIR = "incoming" + class UploadService: def __init__(self, dao: UploadedFileDAO) -> None: @@ -68,14 +71,7 @@ class UploadService: part_path.unlink(missing_ok=True) raise - return FileUploadResponse( - id=saved.id, - filename=saved.original_filename, - size_bytes=saved.size_bytes, - sha256=saved.sha256, - storage_path=saved.storage_path, - uploaded_at=saved.uploaded_at, - ) + return self._to_response(saved) # ---------------- 查询 ---------------- @@ -88,10 +84,62 @@ class UploadService: row = self.dao.get_by_id(file_id) return UploadedFileOut.model_validate(row) if row else None + def find_by_sha256(self, sha256: str) -> UploadedFileOut | None: + row = self.dao.get_by_sha256(sha256) + return UploadedFileOut.model_validate(row) if row else None + def resolve_disk_path(self, file_id: int) -> Path | None: row = self.dao.get_by_id(file_id) return (self.upload_root / row.storage_path).resolve() if row else None + # ---------------- SFTP 登记 ---------------- + + def register_sftp( + self, filename: str, original_filename: str, uploaded_by: str = "sftp", + ) -> FileUploadResponse: + """把一个已经通过 SFTP 落到 incoming/ 下的文件登记入库。 + + 失败 / 拒绝场景: + - filename 路径穿越或不在 incoming/ 下 → 400 + - 文件不存在 / 不是普通文件 → 404 + - sha256 已存在 → 返回已有行(去重),同时删除新上传的副本 + - 否则把文件从 incoming/ 原子改名到 YYYY/MM/.,落 DB 行(source='sftp') + """ + src_abs = self._validate_incoming_path(filename) + size, digest = self._hash_disk_file(src_abs) + + existing = self.dao.get_by_sha256(digest) + if existing is not None: + # 内容已存在 -> 丢弃新副本,避免 uploads/ 越积越多。 + src_abs.unlink(missing_ok=True) + return self._to_response(existing) + + rel_dir = self._relative_dir() + (self.upload_root / rel_dir).mkdir(parents=True, exist_ok=True) + ext = self._safe_ext(original_filename or filename) + rel_path = rel_dir / f"{uuid.uuid4().hex}{ext}" + abs_path = self.upload_root / rel_path + + entity = UploadedFile( + storage_path=str(rel_path), + original_filename=os.path.basename(original_filename or src_abs.name), + content_type="", + size_bytes=size, + sha256=digest, + source="sftp", + uploaded_by=uploaded_by, + ) + saved = self.dao.create(entity) + + try: + os.replace(src_abs, abs_path) + except Exception: + # 改名失败 -> 回滚 DB 行,避免出现孤儿元数据 + self.dao.delete(saved.id) + raise + + return self._to_response(saved) + # ---------------- 内部 ---------------- @staticmethod @@ -103,6 +151,44 @@ class UploadService: def _safe_ext(filename: str) -> str: return os.path.splitext(os.path.basename(filename))[1] + @staticmethod + def _to_response(row: UploadedFile) -> FileUploadResponse: + return FileUploadResponse( + id=row.id, + filename=row.original_filename, + size_bytes=row.size_bytes, + sha256=row.sha256, + storage_path=row.storage_path, + uploaded_at=row.uploaded_at, + ) + + def _validate_incoming_path(self, filename: str) -> Path: + """把客户端传来的相对路径解析成 upload_root 下的绝对路径。 + + 要求落在 ``upload_root/incoming/`` 之下且为普通文件,否则抛 HTTPException。 + """ + if not filename or filename.startswith(("/", "\\")): + raise HTTPException(400, "filename 必须是 incoming/ 下的相对路径") + incoming_root = (self.upload_root / SFTP_INCOMING_DIR).resolve() + try: + abs_path = (self.upload_root / filename).resolve() + abs_path.relative_to(incoming_root) + except ValueError: + raise HTTPException(400, f"filename 必须落在 {SFTP_INCOMING_DIR}/ 之下") + if not abs_path.is_file(): + raise HTTPException(404, f"文件不存在或不是普通文件:{filename}") + return abs_path + + def _hash_disk_file(self, path: Path) -> tuple[int, str]: + """以流式方式读盘上文件,返回 (size, sha256)。""" + h = hashlib.sha256() + size = 0 + with path.open("rb") as fh: + while chunk := fh.read(self.chunk_bytes): + size += len(chunk) + h.update(chunk) + return size, h.hexdigest() + def _write_part(self, file: UploadFile, part_path: Path) -> tuple[int, str]: """流式写到 part_path 并 fsync;返回 (size, sha256)。失败时清理残品。""" hasher = hashlib.sha256() if self.hash_on_upload else None