""" 任务管理模块 (task_manager) 模块: 服务端 / 任务管理 输入: init 请求 / 卷到达事件 输出: transfer 状态机 状态机: uploading -> storing -> done / failed (零合并: 无 assembling/decrypting 阶段, 卷齐即入库) transfer_id 生命周期贯穿接收->入库; 孤儿任务由清理逻辑回收 (见 server/data/tmp)。 """ import uuid from typing import Any from db import ServerDB STATUS_UPLOADING = "uploading" STATUS_STORING = "storing" STATUS_DONE = "done" STATUS_FAILED = "failed" class TaskManager: """transfer 生命周期管理 (基于 ServerDB)""" def __init__(self, db: ServerDB) -> None: self.db = db def create(self, init_json: dict[str, Any], user_id: str, dir_id: int | None = None) -> str: """POST /init: 建任务, 返回 transfer_id (输入清洗: 文件名/卷数上限) 文件名只做可打印字符过滤 + 长度截断。不删 / \\ 等路径分隔符: 加密文件名是 URL-safe base64 (可能含 - _ =), 路径安全由下载端 _safe_name 兜底 """ # 文件名字段清洗: 长度上限 + 去路径分隔符/控制字符 raw = str(init_json.get("file_name", "")) clean = "".join(c for c in raw if c.isprintable())[:255] init_json["file_name"] = clean or "unnamed" # 卷数上限: 防超大清单放大 (1MB init / 120B 每卷 spec ≈ 8000) chunks = init_json.get("chunks") or [] if len(chunks) > 10000: raise ValueError(f"卷数超上限: {len(chunks)} > 10000") init_json["chunk_count"] = len(chunks) transfer_id = uuid.uuid4().hex[:12] self.db.create_transfer(transfer_id, init_json, user_id, dir_id) return transfer_id def get(self, transfer_id: str) -> dict[str, Any] | None: return self.db.get_transfer(transfer_id) def chunk_spec(self, transfer_id: str, idx: int) -> dict[str, Any] | None: return self.db.get_chunk_spec(transfer_id, idx) def is_chunk_received(self, transfer_id: str, idx: int) -> bool: return self.db.is_chunk_received(transfer_id, idx) def mark_chunk_received(self, transfer_id: str, idx: int) -> None: """事件: 卷到达 (校验通过后)""" self.db.mark_chunk_received(transfer_id, idx) def received(self, transfer_id: str) -> set[int]: """已收卷集合 (GET /chunks 依据)""" return self.db.get_received_chunks(transfer_id) def set_status(self, transfer_id: str, status: str) -> None: self.db.set_status(transfer_id, status)