From dc87ed483f65b462ec568f395a4d59af5500eda3 Mon Sep 17 00:00:00 2001 From: lou Date: Sun, 9 Aug 2026 22:55:58 +0800 Subject: [PATCH] =?UTF-8?q?transfer=20=E4=BC=A0=E8=BE=93=E5=AE=A2=E6=88=B7?= =?UTF-8?q?=E7=AB=AF:=20HTTP=20=E5=88=86=E7=89=87=E4=B8=8A=E4=BC=A0/?= =?UTF-8?q?=E6=96=AD=E7=82=B9=E7=BB=AD=E4=BC=A0/=E9=87=8D=E8=AF=95=20+=20?= =?UTF-8?q?=E7=9C=9F=E5=AE=9E=E5=8D=8F=E8=AE=AE=E6=B5=8B=E8=AF=95=20(29=20?= =?UTF-8?q?passed)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- metadata.py | 2 + tests/test_transfer.py | 178 +++++++++++++++++++++++++++++++++++++++ transfer.py | 183 +++++++++++++++++++++++++++++++++++++++++ 3 files changed, 363 insertions(+) create mode 100644 tests/test_transfer.py create mode 100644 transfer.py diff --git a/metadata.py b/metadata.py index b60d0ed..cde6e32 100644 --- a/metadata.py +++ b/metadata.py @@ -159,6 +159,8 @@ def main() -> int: # 2. 分卷 -> 卷清单 with open(enc_path, "rb") as f: manifest = split_stream(f, chunk_size, str(enc_path), total_size=enc_path.stat().st_size) + with open(manifest_path, "w", encoding="utf-8") as f: + json.dump(manifest, f, indent=2, ensure_ascii=False) print(f"[2/4] 分卷完成: {len(manifest)} 卷, 清单 -> {manifest_path.name}") # 3. 密文整体 SHA-256 diff --git a/tests/test_transfer.py b/tests/test_transfer.py new file mode 100644 index 0000000..7a56825 --- /dev/null +++ b/tests/test_transfer.py @@ -0,0 +1,178 @@ +"""transfer 模块行为测试 (真实 HTTP 协议交互, 非 mock) + +内置线程 HTTP 服务器实现协议端点, 验证: 全流程/断点续传/重试/失败。 + +运行: cd /home/lou/文档/7z-encrypt && .venv/bin/python -m unittest discover -s tests -v +""" +import json +import os +import re +import sys +import tempfile +import threading +import unittest +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from transfer import TransferClient, TransferError + +CHUNK_COUNT = 3 + + +class ProtocolServer: + """协议测试服务器: 内存状态, 支持预置已收卷/失败卷""" + + def __init__(self) -> None: + self.chunks: dict[int, bytes] = {} + self.init_json: dict = {} + self.received: set[int] = set() + self.fail_once: set[int] = set() # 首次返回 500 的卷 + self.fail_always: set[int] = set() # 始终 500 的卷 + self.transfer_id = "T-1" + self.completed = False + self.file_id = "F-42" + + def make_handler(self): + srv = self + + class Handler(BaseHTTPRequestHandler): + def log_message(self, format: str, *args) -> None: # 静默 + pass + + def _json(self, code: int, obj: dict) -> None: + body = json.dumps(obj).encode() + self.send_response(code) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def _read_body(self) -> bytes: + length = int(self.headers.get("Content-Length", 0)) + return self.rfile.read(length) if length else b"" + + def do_POST(self): + if self.path == "/api/transfer/init": + srv.init_json = json.loads(self._read_body()) + self._json(200, {"transfer_id": srv.transfer_id}) + return + m = re.fullmatch(r"/api/transfer/([^/]+)/complete", self.path) + if m: + if len(srv.received) == srv.init_json["chunk_count"]: + srv.completed = True + self._json(200, {"status": "done", "file_id": srv.file_id}) + else: + self._json(409, {"status": "incomplete", + "received": sorted(srv.received)}) + return + self._json(404, {"error": "not found"}) + + def do_GET(self): + m = re.fullmatch(r"/api/transfer/([^/]+)/chunks", self.path) + if m: + self._json(200, {"received": sorted(srv.received)}) + return + self._json(404, {"error": "not found"}) + + def do_PUT(self): + m = re.fullmatch(r"/api/transfer/([^/]+)/chunk/(\d+)", self.path) + if not m: + self._json(404, {"error": "not found"}) + return + idx = int(m.group(2)) + if idx in srv.fail_always: + self._json(500, {"error": "disk full"}) + return + if idx in srv.fail_once: + srv.fail_once.discard(idx) + self._json(500, {"error": "transient"}) + return + srv.chunks[idx] = self._read_body() + srv.received.add(idx) + self._json(200, {"ok": True}) + + return Handler + + +class TransferTestBase(unittest.TestCase): + def setUp(self): + self.srv = ProtocolServer() + self.httpd = ThreadingHTTPServer(("127.0.0.1", 0), self.srv.make_handler()) + self.thread = threading.Thread(target=self.httpd.serve_forever, daemon=True) + self.thread.start() + self.addCleanup(self.httpd.shutdown) + self.addCleanup(self.httpd.server_close) + self.client = TransferClient( + f"http://127.0.0.1:{self.httpd.server_address[1]}", timeout=10, retry_delay=0.01 + ) + self._td = tempfile.TemporaryDirectory(prefix='hermes-test-') + self.addCleanup(self._td.cleanup) + self.chunk_files: dict[int, Path] = {} + for i in range(1, CHUNK_COUNT + 1): + p = Path(self._td.name) / f"chunk{i}.bin" + p.write_bytes(b"chunk-%d-data" % i) + self.chunk_files[i] = p + self.init_json = { + "file_name": "a.mp4", "file_size": 100, "chunk_count": CHUNK_COUNT, + "total_sha256": "a" * 64, + "enc": {"alg": "cobblestone-aes256gcm", "key_id": "k", "context": "ctx"}, + "chunks": [{"index": i, "size": 13, "sha256": "a" * 64} for i in range(1, CHUNK_COUNT + 1)], + } + + +class TestFullFlow(TransferTestBase): + def test_all_chunks_uploaded(self): + receipt = self.client.transfer(self.init_json, self.chunk_files) + self.assertEqual(receipt["status"], "done") + self.assertEqual(receipt["file_id"], self.srv.file_id) + self.assertEqual(self.srv.received, {1, 2, 3}) + # 服务端收到的字节与本地卷一致 + for i in range(1, CHUNK_COUNT + 1): + self.assertEqual(self.srv.chunks[i], self.chunk_files[i].read_bytes()) + self.assertTrue(self.srv.completed) + + +class TestResume(TransferTestBase): + def test_skips_received_chunks(self): + self.srv.received.add(1) # 预置: 服务端已收卷 1 + receipt = self.client.transfer(self.init_json, self.chunk_files) + self.assertEqual(receipt["status"], "done") + self.assertEqual(self.srv.received, {1, 2, 3}) # 只补传 2,3 + self.assertNotIn(1, self.srv.chunks) # 卷1 未重新上传 + + +class TestRetry(TransferTestBase): + def test_transient_failure_retried(self): + self.srv.fail_once.add(2) # 卷2 第一次 500, 之后成功 + receipt = self.client.transfer(self.init_json, self.chunk_files) + self.assertEqual(receipt["status"], "done") + self.assertEqual(self.srv.received, {1, 2, 3}) + + def test_persistent_failure_raises(self): + self.srv.fail_always.add(2) # 卷2 永远 500 + with self.assertRaises(TransferError): + self.client.transfer(self.init_json, self.chunk_files) + # 重试耗尽后其他卷不阻塞, 但 complete 未触发 + self.assertFalse(self.srv.completed) + + +class TestProtocolErrors(TransferTestBase): + def test_bad_server(self): + bad = TransferClient("http://127.0.0.1:1", timeout=2) # 无服务端口 + with self.assertRaises(TransferError): + bad.init_transfer(self.init_json) + + def test_incomplete_complete(self): + self.client.init_transfer(self.init_json) # 先建任务 + self.srv.received.add(1) # 只收 1 卷就 complete + code, payload = self.client._request( + "POST", f"/api/transfer/{self.srv.transfer_id}/complete" + ) + self.assertEqual(code, 409) + self.assertEqual(payload["status"], "incomplete") + + +if __name__ == '__main__': + unittest.main() diff --git a/transfer.py b/transfer.py new file mode 100644 index 0000000..0492191 --- /dev/null +++ b/transfer.py @@ -0,0 +1,183 @@ +""" +传输客户端模块 (transfer) +模块:客户端 / 传输客户端 +输入:init json + 卷文件 + 服务器配置 +输出:服务端回执 (status/file_id) + +协议 (共享契约): + POST /api/transfer/init {init json} -> {transfer_id} + GET /api/transfer/{id}/chunks -> {received: [n...]} + PUT /api/transfer/{id}/chunk/{n} 卷密文 -> {ok} + POST /api/transfer/{id}/complete -> {status, file_id} + +断点续传: 先查已收卷集合只补缺卷; 单卷失败重试 max_retries 次, 仍失败抛 TransferError。 +""" + +import argparse +import json +import sys +import time +from pathlib import Path +from typing import Any +from urllib import error as urlerror +from urllib import request as urlrequest + +DEFAULT_MAX_RETRIES = 3 +DEFAULT_RETRY_DELAY = 1.0 # 秒 + + +class TransferError(RuntimeError): + """传输失败 (服务端错误/网络错误/重试耗尽)""" + + +class TransferClient: + """HTTP 分片上传客户端 (标准库 urllib, 零依赖)""" + + def __init__( + self, + base_url: str, + timeout: int = 30, + max_retries: int = DEFAULT_MAX_RETRIES, + retry_delay: float = DEFAULT_RETRY_DELAY, + ) -> None: + self.base_url = base_url.rstrip("/") + self.timeout = timeout + self.max_retries = max_retries + self.retry_delay = retry_delay + + # ---------- HTTP 封装 ---------- + + def _request( + self, method: str, path: str, body: bytes | None = None + ) -> tuple[int, dict[str, Any]]: + """发请求, 返回 (HTTP状态码, JSON 载荷)。4xx/5xx 不抛, 由调用方判断""" + req = urlrequest.Request(self.base_url + path, data=body, method=method) + if body is not None: + req.add_header("Content-Type", "application/octet-stream") + try: + with urlrequest.urlopen(req, timeout=self.timeout) as resp: + data = resp.read() + code = resp.status + except urlerror.HTTPError as e: + code = e.code + data = e.read() + except urlerror.URLError as e: + raise TransferError(f"网络错误: {e.reason}") from e + try: + payload: dict[str, Any] = json.loads(data) if data else {} + except json.JSONDecodeError: + payload = {} + return code, payload + + # ---------- 协议端点 ---------- + + def init_transfer(self, init_json: dict[str, Any]) -> str: + """POST /init: 建任务, 返回 transfer_id""" + code, payload = self._request( + "POST", "/api/transfer/init", json.dumps(init_json).encode() + ) + if code != 200 or "transfer_id" not in payload: + raise TransferError(f"init 失败 (HTTP {code}): {payload}") + return payload["transfer_id"] + + def get_received_chunks(self, transfer_id: str) -> set[int]: + """GET /chunks: 查询服务端已收卷 (断点续传依据)""" + code, payload = self._request("GET", f"/api/transfer/{transfer_id}/chunks") + if code != 200: + raise TransferError(f"查询已收卷失败 (HTTP {code})") + return set(payload.get("received", [])) + + def upload_chunk(self, transfer_id: str, index: int, chunk_path: Path) -> None: + """PUT /chunk/{n}: 上传单卷, 失败重试 max_retries 次""" + data = chunk_path.read_bytes() + code: int = 0 + for attempt in range(1, self.max_retries + 1): + code, payload = self._request( + "PUT", f"/api/transfer/{transfer_id}/chunk/{index}", body=data + ) + if code == 200 and payload.get("ok"): + return + if attempt < self.max_retries: + time.sleep(self.retry_delay) + raise TransferError( + f"卷 {index} 上传失败 (HTTP {code}, 重试 {self.max_retries} 次耗尽)" + ) + + def complete(self, transfer_id: str) -> dict[str, Any]: + """POST /complete: 触发服务端汇聚解密入库, 返回回执""" + code, payload = self._request("POST", f"/api/transfer/{transfer_id}/complete") + if code != 200: + raise TransferError(f"complete 失败 (HTTP {code}): {payload}") + return payload + + # ---------- 全流程 ---------- + + def transfer( + self, init_json: dict[str, Any], chunk_files: dict[int, Path] + ) -> dict[str, Any]: + """init -> 查已收 -> 逐卷上传(跳过已收) -> complete + + Args: + init_json: metadata.build_init_json 的产物 + chunk_files: {卷序号: 卷文件路径} + + Returns: + dict: 服务端回执 {status, file_id} + """ + transfer_id = self.init_transfer(init_json) + print(f"[传输] transfer_id: {transfer_id}") + + received = self.get_received_chunks(transfer_id) + total = init_json["chunk_count"] + pending = [i for i in range(1, total + 1) if i not in received] + if received: + print(f"[传输] 断点续传: 服务端已收 {len(received)}/{total} 卷, 补传 {len(pending)} 卷") + + for done, index in enumerate(pending, start=1): + try: + self.upload_chunk(transfer_id, index, chunk_files[index]) + except TransferError: + raise # 重试耗尽, 由调用方决定 (记录断点状态待补传) + pct = done * 100 // len(pending) + sys.stdout.write( + f"\r[传输] 上传进度: {done}/{len(pending)} 卷 ({pct}%)" + ) + sys.stdout.flush() + if pending: + sys.stdout.write("\n") + sys.stdout.flush() + + receipt = self.complete(transfer_id) + print(f"[传输] 完成: {receipt}") + return receipt + + +def main() -> int: + ap = argparse.ArgumentParser(description="传输客户端: 上传 init json + 卷文件") + ap.add_argument("--server", required=True, help="服务端地址, 如 http://127.0.0.1:8000") + ap.add_argument("--init", required=True, help="init json 路径 (metadata 产物)") + ap.add_argument("--manifest", required=True, help="卷清单路径 (splitter 产物)") + args = ap.parse_args() + + init_json = json.loads(Path(args.init).read_text(encoding="utf-8")) + manifest = json.loads(Path(args.manifest).read_text(encoding="utf-8")) + chunk_files = { + ch["index"]: Path(ch["filename"]).resolve() + for ch in manifest + } + missing = [i for i, p in chunk_files.items() if not p.exists()] + if missing: + print(f"[错误] 卷文件缺失: {missing}") + return 1 + + client = TransferClient(args.server) + try: + client.transfer(init_json, chunk_files) + except TransferError as e: + print(f"[错误] {e}") + return 1 + return 0 + + +if __name__ == "__main__": + sys.exit(main())