Files
7z-encrypt/tests/test_transfer.py
T

241 lines
9.4 KiB
Python
Raw Normal View History

"""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
import urllib.parse
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"
self.files: list[dict] = [] # ls/下载测试: [{"file_id","file_name","size","data"}]
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
if self.path == "/api/files":
# 列表只含元信息, 不含 data (bytes 不可 JSON 序列化)
self._json(200, {
"files": [
{k: v for k, v in f.items() if k != "data"}
for f in srv.files
]
})
return
m = re.fullmatch(r"/api/files/([^/]+)", self.path)
if m:
fid = m.group(1)
rec = next((f for f in srv.files if f["file_id"] == fid), None)
if rec is None:
self._json(404, {"error": "文件不存在"})
return
body = rec["data"]
self.send_response(200)
self.send_header("Content-Type", "application/octet-stream")
self.send_header("Content-Length", str(len(body)))
self.send_header(
"Content-Disposition",
f"attachment; filename*=utf-8''{urllib.parse.quote(rec['file_name'])}",
)
self.end_headers()
self.wfile.write(body)
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.tmp = Path(self._td.name)
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")
class TestListDownload(TransferTestBase):
"""ls 列表 + 下载 (真实 HTTP 服务器)"""
def setUp(self):
super().setUp()
self.srv.files = [
{"file_id": "F-1", "file_name": "报告.pdf", "size": 6, "data": b"hello"},
]
def test_list_files(self):
files = self.client.list_files()
self.assertEqual(len(files), 1)
self.assertEqual(files[0]["file_id"], "F-1")
self.assertEqual(files[0]["file_name"], "报告.pdf")
def test_download_to_path(self):
dest = self.tmp / "out.bin"
got = self.client.download("F-1", dest)
self.assertEqual(got, dest)
self.assertEqual(dest.read_bytes(), b"hello")
def test_download_uses_server_filename(self):
got = self.client.download("F-1")
self.assertEqual(got.name, "报告.pdf")
self.assertEqual(got.read_bytes(), b"hello")
got.unlink()
def test_download_unknown_404(self):
with self.assertRaises(TransferError):
self.client.download("NOPE")
if __name__ == "__main__":
unittest.main()