179 lines
7.0 KiB
Python
179 lines
7.0 KiB
Python
|
|
"""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()
|