feat: 增加 /api/files/exists 与 /api/files/register-sftp 接口

- DAO 增加 get_by_sha256
- UploadService 抽出 _to_response 复用;新增 find_by_sha256、register_sftp
- 控制器新增两条路由(注意排在 /{file_id} 之前以避免被吞)
- 引入 SftpRegisterRequest schema
- SFTP 服务启动时自动创建 uploads/incoming/
- 行为:路径穿越 400、文件不存在 404、sha256 已存在则去重返回已有行
This commit is contained in:
zikai
2026-06-23 16:22:03 +00:00
parent 80b96d236f
commit 6c7d3f1272
5 changed files with 154 additions and 10 deletions

View File

@@ -10,7 +10,12 @@ from sqlalchemy.orm import Session
from ..database import get_db from ..database import get_db
from ..dao.uploaded_file_dao import UploadedFileDAO 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 from ..services.upload_service import UploadService
router = APIRouter(prefix="/api/files", tags=["files"]) router = APIRouter(prefix="/api/files", tags=["files"])
@@ -49,6 +54,42 @@ def list_files(
return FileListResponse(total=total, items=items) 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/<name>``,再用本接口登记入库。"
"服务端会计算 sha256已存在则去重、把文件原子改名到 ``YYYY/MM/<uuid>.<ext>``、写 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="查询单个文件元数据") @router.get("/{file_id}", response_model=UploadedFileOut, summary="查询单个文件元数据")
def get_file(file_id: int, service: UploadService = Depends(_service)) -> UploadedFileOut: def get_file(file_id: int, service: UploadService = Depends(_service)) -> UploadedFileOut:
out = service.get_out(file_id) out = service.get_out(file_id)

View File

@@ -21,6 +21,10 @@ class UploadedFileDAO:
def get_by_id(self, file_id: int) -> UploadedFile | None: def get_by_id(self, file_id: int) -> UploadedFile | None:
return self.db.get(UploadedFile, file_id) 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]: def list(self, limit: int = 100, offset: int = 0) -> list[UploadedFile]:
stmt = ( stmt = (
select(UploadedFile) select(UploadedFile)

View File

@@ -33,3 +33,14 @@ class FileUploadResponse(BaseModel):
class FileListResponse(BaseModel): class FileListResponse(BaseModel):
total: int total: int
items: list[UploadedFileOut] 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 字段")

View File

@@ -119,6 +119,8 @@ async def _run() -> None:
upload_root = settings.resolved_upload_dir() upload_root = settings.resolved_upload_dir()
upload_root.mkdir(parents=True, exist_ok=True) 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() host_key_path = (PROJECT_ROOT / settings.sftp.host_key_path).resolve()
_ensure_host_key(host_key_path) _ensure_host_key(host_key_path)

View File

@@ -8,13 +8,16 @@ import uuid
from datetime import datetime, timezone from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from fastapi import UploadFile from fastapi import HTTPException, UploadFile
from ..config import get_settings from ..config import get_settings
from ..dao.uploaded_file_dao import UploadedFileDAO from ..dao.uploaded_file_dao import UploadedFileDAO
from ..models.uploaded_file import UploadedFile from ..models.uploaded_file import UploadedFile
from ..schemas.file import FileUploadResponse, UploadedFileOut from ..schemas.file import FileUploadResponse, UploadedFileOut
# SFTP 客户端登记前必须先把文件写入此目录chroot 内)
SFTP_INCOMING_DIR = "incoming"
class UploadService: class UploadService:
def __init__(self, dao: UploadedFileDAO) -> None: def __init__(self, dao: UploadedFileDAO) -> None:
@@ -68,14 +71,7 @@ class UploadService:
part_path.unlink(missing_ok=True) part_path.unlink(missing_ok=True)
raise raise
return FileUploadResponse( return self._to_response(saved)
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,
)
# ---------------- 查询 ---------------- # ---------------- 查询 ----------------
@@ -88,10 +84,62 @@ class UploadService:
row = self.dao.get_by_id(file_id) row = self.dao.get_by_id(file_id)
return UploadedFileOut.model_validate(row) if row else None 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: def resolve_disk_path(self, file_id: int) -> Path | None:
row = self.dao.get_by_id(file_id) row = self.dao.get_by_id(file_id)
return (self.upload_root / row.storage_path).resolve() if row else None 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/<uuid>.<ext>,落 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 @staticmethod
@@ -103,6 +151,44 @@ class UploadService:
def _safe_ext(filename: str) -> str: def _safe_ext(filename: str) -> str:
return os.path.splitext(os.path.basename(filename))[1] 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]: def _write_part(self, file: UploadFile, part_path: Path) -> tuple[int, str]:
"""流式写到 part_path 并 fsync返回 (size, sha256)。失败时清理残品。""" """流式写到 part_path 并 fsync返回 (size, sha256)。失败时清理残品。"""
hasher = hashlib.sha256() if self.hash_on_upload else None hasher = hashlib.sha256() if self.hash_on_upload else None