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:
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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 字段")
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user