transfer 传输客户端: HTTP 分片上传/断点续传/重试 + 真实协议测试 (29 passed)
This commit is contained in:
@@ -159,6 +159,8 @@ def main() -> int:
|
|||||||
# 2. 分卷 -> 卷清单
|
# 2. 分卷 -> 卷清单
|
||||||
with open(enc_path, "rb") as f:
|
with open(enc_path, "rb") as f:
|
||||||
manifest = split_stream(f, chunk_size, str(enc_path), total_size=enc_path.stat().st_size)
|
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}")
|
print(f"[2/4] 分卷完成: {len(manifest)} 卷, 清单 -> {manifest_path.name}")
|
||||||
|
|
||||||
# 3. 密文整体 SHA-256
|
# 3. 密文整体 SHA-256
|
||||||
|
|||||||
@@ -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
@@ -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())
|
||||||
Reference in New Issue
Block a user