From da087591ed6d979564d3acceec1773707b78b65f Mon Sep 17 00:00:00 2001 From: lou Date: Mon, 10 Aug 2026 02:10:07 +0800 Subject: [PATCH] =?UTF-8?q?task=5Fmanager:=20=E6=96=87=E4=BB=B6=E5=90=8D?= =?UTF-8?q?=E6=B8=85=E6=B4=97=E4=B8=8D=E5=88=A0=E8=B7=AF=E5=BE=84=E5=88=86?= =?UTF-8?q?=E9=9A=94=E7=AC=A6=20(=E5=8A=A0=E5=AF=86=E5=90=8D=20URL-safe=20?= =?UTF-8?q?base64,=20=E8=B7=AF=E5=BE=84=E5=AE=89=E5=85=A8=E7=94=B1?= =?UTF-8?q?=E5=AE=A2=E6=88=B7=E7=AB=AF=E5=85=9C=E5=BA=95);=20=E6=B5=8B?= =?UTF-8?q?=E8=AF=95=209=20=E9=A1=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- task_manager.py | 10 +++++++--- tests/test_api.py | 6 +++--- 2 files changed, 10 insertions(+), 6 deletions(-) diff --git a/task_manager.py b/task_manager.py index f64c37e..f515e5c 100644 --- a/task_manager.py +++ b/task_manager.py @@ -27,10 +27,14 @@ class TaskManager: self.db = db def create(self, init_json: dict[str, Any]) -> str: - """POST /init: 建任务, 返回 transfer_id (输入清洗: 文件名/卷数上限)""" - # 文件名字段清洗: 长度上限 + 去路径分隔符/控制字符 + """POST /init: 建任务, 返回 transfer_id (输入清洗: 文件名/卷数上限) + + 文件名只做可打印字符过滤 + 长度截断。不删 / \\ 等路径分隔符: + 加密文件名是 URL-safe base64 (可能含 - _), 路径安全由客户端 _safe_name + 和存储目录 (transfer_id) 保证。 + """ raw = str(init_json.get("file_name", "unnamed")) - clean = "".join(c for c in raw if c.isprintable() and c not in "/\\\x00")[:255] + 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 [] diff --git a/tests/test_api.py b/tests/test_api.py index a9beb39..18bcced 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -132,7 +132,7 @@ class ServerPipelineTest(unittest.TestCase): self.assertEqual(st["status"], "done") def test_init_sanitizes_file_name(self): - # 恶意文件名: 路径分隔符/控制字符清洗, 长度截断 + # 恶意文件名: 控制字符删除 + 长度截断 (路径分隔符保留, 由客户端 _safe_name 兜底) init_json, chunk_files, _ = _prepare( self.client_kr, os.urandom(64 << 10) ) @@ -147,9 +147,9 @@ class ServerPipelineTest(unittest.TestCase): "SELECT file_name FROM files WHERE file_id = ?", (file_id,) ).fetchone() name = row["file_name"] - self.assertNotIn("/", name) - self.assertNotIn("\\", name) + self.assertNotIn("\x00", name) # 控制字符被删 self.assertLessEqual(len(name), 255) + self.assertIn("..", name) # 路径段保留 (加密名含 - _ 不受影响, 明文名由下载端消毒) def test_init_oversized_json_rejected(self): # init json 超大 -> 413