"""分片上传服务:会话管理 + 分片落盘 + 拼接 + sha256 去重 + 原子入库。 存储布局:: uploads/.work//0.part 分片暂存 uploads/.work//1.part ... uploads/2026/07/. complete 后的正式文件 complete 流程: 1. 校验分片齐全; 2. 顺序拼接为 .part,流式算 sha256; 3. 复用 UploadService.dedup_or_commit:按 sha256 去重命中则删会话返回旧行, 否则写 DB 行 + os.replace 原子改名到正式路径; 4. 清理 .work//。 """ from __future__ import annotations import io import logging import os import shutil import uuid from collections.abc import Iterator from pathlib import Path from fastapi import HTTPException from ..config import get_settings from ..dao.upload_session_dao import UploadSessionDAO from ..dao.uploaded_file_dao import UploadedFileDAO from ..models.upload_session import UploadSession from ..models.uploaded_file import UploadedFile from ..schemas.chunk import ( CreateSessionRequest, CreateSessionResponse, SessionStatusResponse, ) from ..schemas.file import FileUploadResponse from .upload_service import UploadService, hash_stream logger = logging.getLogger("zikai.chunk") # 会话暂存目录名(位于 upload_root 下) SESSION_WORK_DIR = ".work" class ChunkUploadService: def __init__( self, session_dao: UploadSessionDAO, file_dao: UploadedFileDAO, ) -> None: s = get_settings() self.session_dao = session_dao self.file_dao = file_dao self.upload_root = s.resolved_upload_dir() self.chunk_bytes = s.storage.chunk_bytes # 会话暂存目录:配置给出的是相对 upload_root 的子目录名 session_sub = os.path.basename(s.storage.chunk_session_dir) or SESSION_WORK_DIR self.work_root = (self.upload_root / session_sub).resolve() self.work_root.mkdir(parents=True, exist_ok=True) self.session_ttl = s.storage.chunk_session_ttl_seconds # 复用 UploadService 的存储路径生成 / 落库提交 / sha256 去重逻辑 self._upload = UploadService(file_dao) # ---------------- 会话生命周期 ---------------- def create_session(self, body: CreateSessionRequest) -> CreateSessionResponse: upload_id = uuid.uuid4().hex session = UploadSession( upload_id=upload_id, filename=body.filename, size_bytes=body.size_bytes, chunk_size=body.chunk_size, total_chunks=body.total_chunks, uploaded_chunks=[], status="pending", ) self.session_dao.create(session) self._session_dir(upload_id).mkdir(parents=True, exist_ok=True) return CreateSessionResponse( upload_id=upload_id, filename=body.filename, size_bytes=body.size_bytes, chunk_size=body.chunk_size, total_chunks=body.total_chunks, ) def get_status(self, upload_id: str) -> SessionStatusResponse: session = self._require_session(upload_id) return SessionStatusResponse( upload_id=session.upload_id, filename=session.filename, size_bytes=session.size_bytes, chunk_size=session.chunk_size, total_chunks=session.total_chunks, uploaded_chunks=list(session.uploaded_chunks or []), completed=(session.status == "completed"), file_id=session.file_id, ) # ---------------- 分片写入 ---------------- def write_chunk( self, upload_id: str, index: int, data: bytes, ) -> list[int]: session = self._require_session(upload_id) self._validate_index(session, index) session_dir = self._session_dir(upload_id) session_dir.mkdir(parents=True, exist_ok=True) chunk_path = session_dir / f"{index}.part" # 幂等覆盖:同一分片重传时直接覆盖旧文件 try: with chunk_path.open("wb") as out: out.write(data) out.flush() os.fsync(out.fileno()) except Exception: chunk_path.unlink(missing_ok=True) raise uploaded = list(session.uploaded_chunks or []) if index not in uploaded: uploaded.append(index) self.session_dao.mark_uploaded(session, uploaded) return sorted(uploaded) # ---------------- 拼接入库 ---------------- def complete(self, upload_id: str) -> FileUploadResponse: session = self._require_session(upload_id) # 幂等:已 complete 直接复用结果 if session.status == "completed" and session.file_id is not None: row = self.file_dao.get_by_id(session.file_id) if row is not None: return UploadService.to_response(row, deduplicated=True) uploaded = set(session.uploaded_chunks or []) missing = [i for i in range(session.total_chunks) if i not in uploaded] if missing: raise HTTPException( 409, f"分片不齐全:缺失 {len(missing)} 个,例如 {sorted(missing)[:10]}", ) session_dir = self._session_dir(upload_id) part_path = session_dir / "_assembled.part" size, digest = self._assemble_and_hash(session, session_dir, part_path) # 复用 UploadService 的「sha256 去重 + 落库 + 原子改名」公共尾部 entity = UploadedFile( storage_path="", # 由 dedup_or_commit 内部生成 original_filename=os.path.basename(session.filename), content_type="", size_bytes=size, sha256=digest, source="chunk", uploaded_by="web", ) resp = self._upload.dedup_or_commit(entity, part_path) self.session_dao.mark_completed(session, resp.id) # 会话目录里的分片已拼接走,清理残留 self._cleanup_session_dir(upload_id) return resp # ---------------- 过期会话清理 ---------------- def reap_stale_sessions(self) -> int: """清理被放弃的会话:删 .work// 目录 + DB 记录。 判定标准:status=pending 且 updated_at 距今超过 session_ttl 秒。 返回清理的会话数。不依赖 start.sh,由后台任务周期调用。 """ stale = self.session_dao.list_stale(self.session_ttl) for session in stale: self._cleanup_session_dir(session.upload_id) self.session_dao.delete(session) logger.info("清理过期分片会话 upload_id=%s file=%s", session.upload_id, session.filename) return len(stale) def _cleanup_session_dir(self, upload_id: str) -> None: """删除 .work// 目录;失败只记日志,不抛异常。""" session_dir = self._session_dir(upload_id) try: if session_dir.exists(): shutil.rmtree(session_dir, ignore_errors=True) except Exception as exc: # pragma: no cover logger.warning("清理会话目录失败 upload_id=%s: %s", upload_id, exc) # ---------------- 内部 ---------------- def _require_session(self, upload_id: str) -> UploadSession: if not upload_id: raise HTTPException(400, "upload_id 不能为空") session = self.session_dao.get_by_upload_id(upload_id) if session is None: raise HTTPException(404, f"会话不存在或已过期:{upload_id}") return session @staticmethod def _validate_index(session: UploadSession, index: int) -> None: if index < 0 or index >= session.total_chunks: raise HTTPException( 400, f"分片下标越界:{index} 不在 [0, {session.total_chunks})", ) def _session_dir(self, upload_id: str) -> Path: return self.work_root / upload_id def _assemble_and_hash( self, session: UploadSession, session_dir: Path, out_path: Path, ) -> tuple[int, str]: """按 index 顺序拼接全部分片为 out_path,流式计算 (size, sha256)。""" try: with out_path.open("wb") as out: size, digest = hash_stream( self._iter_assembled_bytes(session, session_dir, out) ) out.flush() os.fsync(out.fileno()) except Exception: out_path.unlink(missing_ok=True) raise return size, digest def _iter_assembled_bytes( self, session: UploadSession, session_dir: Path, out: io.BufferedWriter, ) -> Iterator[bytes]: """按 index 顺序读各分片,写入 out 同时 yield 每个字节块(供 hash_stream)。""" for index in range(session.total_chunks): chunk_path = session_dir / f"{index}.part" if not chunk_path.is_file(): raise HTTPException(409, f"拼接时发现分片缺失:{index}.part") with chunk_path.open("rb") as src: while buf := src.read(self.chunk_bytes): out.write(buf) yield buf