transfer 传输客户端: HTTP 分片上传/断点续传/重试 + 真实协议测试 (29 passed)

This commit is contained in:
lou
2026-08-09 22:55:58 +08:00
parent 2e85e9a033
commit dc87ed483f
3 changed files with 363 additions and 0 deletions
+2
View File
@@ -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
+178
View File
@@ -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()
+183
View File
@@ -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())