460 lines
19 KiB
Python
460 lines
19 KiB
Python
"""transfer 模块行为测试 (真实 HTTP 协议交互, 非 mock)
|
|
|
|
内置线程 HTTP 服务器实现协议端点, 验证: 全流程/断点续传/重试/失败。
|
|
|
|
运行: cd /home/lou/文档/7z-encrypt && .venv/bin/python -m unittest discover -s tests -v
|
|
"""
|
|
import base64
|
|
import io
|
|
import json
|
|
import os
|
|
import re
|
|
import subprocess
|
|
import sys
|
|
import tempfile
|
|
import threading
|
|
import time
|
|
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"}]
|
|
self.complete_delay = 0.0 # complete 慢响应模拟 (秒)
|
|
self.file_delay = 0.0 # 下载慢响应模拟 (秒)
|
|
self.enc_params_header = "" # 下载响应头 X-Enc-Params (base64)
|
|
self.chunk_count_header = 0 # 下载响应头 X-Chunk-Count (卷数)
|
|
self.require_token = "" # 非空时要求 Authorization: Bearer <此值>
|
|
|
|
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 srv.complete_delay:
|
|
time.sleep(srv.complete_delay)
|
|
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_DELETE(self):
|
|
if srv.require_token and self.headers.get("Authorization") != f"Bearer {srv.require_token}":
|
|
self._json(401, {"error": "未授权"})
|
|
return
|
|
m = re.fullmatch(r"/api/files/([^/]+)", self.path)
|
|
if not m:
|
|
self._json(404, {"error": "not found"})
|
|
return
|
|
fid = m.group(1)
|
|
before = len(srv.files)
|
|
srv.files = [f for f in srv.files if f["file_id"] != fid]
|
|
if len(srv.files) == before:
|
|
self._json(404, {"error": "文件不存在"})
|
|
return
|
|
self._json(200, {"ok": True, "file_id": fid})
|
|
|
|
def do_GET(self):
|
|
if srv.require_token and self.headers.get("Authorization") != f"Bearer {srv.require_token}":
|
|
self._json(401, {"error": "未授权"})
|
|
return
|
|
if self.path == "/api/quota":
|
|
used = sum(f.get("size", len(f.get("data", b""))) for f in srv.files)
|
|
self._json(200, {
|
|
"used_bytes": used,
|
|
"quota_bytes": 1024,
|
|
"remain_bytes": max(1024 - used, 0),
|
|
"percent": round(used * 100 / 1024, 1),
|
|
})
|
|
return
|
|
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/([^/]+)/chunk/(\d+)", self.path)
|
|
if m:
|
|
if srv.file_delay:
|
|
time.sleep(srv.file_delay)
|
|
fid, idx = m.group(1), int(m.group(2))
|
|
rec = next((f for f in srv.files if f["file_id"] == fid), None)
|
|
if rec is None:
|
|
self._json(404, {"error": "文件不存在"})
|
|
return
|
|
chunks = rec.get("chunks") or [rec["data"]]
|
|
if not 1 <= idx <= len(chunks):
|
|
self._json(404, {"error": f"卷号越界: {idx}"})
|
|
return
|
|
body = chunks[idx - 1]
|
|
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'])}",
|
|
)
|
|
if srv.enc_params_header:
|
|
self.send_header("X-Enc-Params", srv.enc_params_header)
|
|
self.send_header("X-Chunk-Count", str(len(chunks)))
|
|
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"},
|
|
]
|
|
# 本地 keyring (下载解密用)
|
|
self._kr = tempfile.mkdtemp(prefix="sz-kr-")
|
|
from crypto import CryptoEngine
|
|
self.engine = CryptoEngine(os.path.join(self._kr, "keyring.json"))
|
|
self.engine.generate_key("k1")
|
|
|
|
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_quota(self):
|
|
self.srv.files = [
|
|
{"file_id": "F-1", "file_name": "a.bin", "size": 300, "data": b"x" * 300},
|
|
{"file_id": "F-2", "file_name": "b.bin", "size": 200, "data": b"y" * 200},
|
|
]
|
|
q = self.client.get_quota()
|
|
self.assertEqual(q["used_bytes"], 500)
|
|
self.assertEqual(q["quota_bytes"], 1024)
|
|
self.assertEqual(q["remain_bytes"], 524)
|
|
|
|
def test_delete_file(self):
|
|
self.srv.files = [
|
|
{"file_id": "F-1", "file_name": "a.bin", "size": 1, "data": b"x"},
|
|
{"file_id": "F-2", "file_name": "b.bin", "size": 1, "data": b"y"},
|
|
]
|
|
self.client.delete_file("F-1")
|
|
files = self.client.list_files()
|
|
self.assertEqual([f["file_id"] for f in files], ["F-2"])
|
|
# 删不存在的 -> TransferError
|
|
with self.assertRaises(TransferError):
|
|
self.client.delete_file("F-1")
|
|
|
|
def test_download_to_path(self):
|
|
dest = self.tmp / "out.bin"
|
|
got, params = self.client.download("F-1", dest)
|
|
self.assertEqual(got, dest)
|
|
self.assertEqual(dest.read_bytes(), b"hello")
|
|
self.assertIsNone(params.get("key_id")) # 服务器没发 X-Enc-Params
|
|
|
|
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_to_dir(self):
|
|
# dest 是目录 -> 自动拼服务端文件名
|
|
dest_dir = self.tmp / "dl"
|
|
dest_dir.mkdir()
|
|
got, _ = self.client.download("F-1", dest_dir)
|
|
self.assertEqual(got, dest_dir / "报告.pdf")
|
|
self.assertEqual(got.read_bytes(), b"hello")
|
|
|
|
def test_sends_token(self):
|
|
# 客户端 token 带在 Authorization 头, 服务器校验通过
|
|
self.srv.require_token = "secret-token"
|
|
client = TransferClient(
|
|
f"http://127.0.0.1:{self.httpd.server_address[1]}", token="secret-token"
|
|
)
|
|
self.assertEqual(
|
|
client.list_files(),
|
|
[{k: v for k, v in f.items() if k != "data"} for f in self.srv.files],
|
|
)
|
|
# 无 token -> 401 包装成 TransferError
|
|
with self.assertRaises(TransferError):
|
|
self.client.list_files()
|
|
|
|
def test_download_decrypt_roundtrip(self):
|
|
# 真实加密 -> 下载密文+参数 -> 本地解密 -> 明文一致 (零知识回环)
|
|
data = os.urandom(256 * 1024)
|
|
enc = io.BytesIO()
|
|
params = self.engine.encrypt_stream(io.BytesIO(data), enc, "k1")
|
|
enc_bytes = enc.getvalue()
|
|
self.srv.files = [
|
|
{"file_id": "F-1", "file_name": "机密.bin", "size": len(enc_bytes), "data": enc_bytes},
|
|
]
|
|
self.srv.enc_params_header = base64.b64encode(
|
|
json.dumps(params, ensure_ascii=False).encode()
|
|
).decode()
|
|
|
|
got, enc_params = self.client.download("F-1")
|
|
self.assertEqual(got.read_bytes(), enc_bytes)
|
|
self.assertEqual(enc_params["key_id"], "k1")
|
|
|
|
# 本地解密还原
|
|
out = self.tmp / "还原.bin"
|
|
with open(got, "rb") as src, open(out, "wb") as dst:
|
|
self.engine.decrypt_stream(
|
|
src, dst, enc_params["key_id"], enc_params["context"].encode()
|
|
)
|
|
self.assertEqual(out.read_bytes(), data)
|
|
got.unlink()
|
|
|
|
def test_download_7z_roundtrip(self):
|
|
# 真实 7z 压缩 -> 加密 -> 下载 -> 解密 -> 解压 -> 原始文件一致
|
|
data = os.urandom(512 * 1024)
|
|
src = self.tmp / "原始.bin"
|
|
src.write_bytes(data)
|
|
src7z = self.tmp / "原始.bin.7z"
|
|
r = subprocess.run(
|
|
["7z", "a", "-mx=9", "-bd", str(src7z), str(src)],
|
|
capture_output=True, text=True,
|
|
)
|
|
self.assertEqual(r.returncode, 0, r.stderr)
|
|
enc = io.BytesIO()
|
|
params = self.engine.encrypt_stream(open(src7z, "rb"), enc, "k1")
|
|
params["compressed"] = True
|
|
enc_bytes = enc.getvalue()
|
|
self.srv.files = [
|
|
{"file_id": "F-1", "file_name": "原始.bin", "size": len(enc_bytes), "data": enc_bytes},
|
|
]
|
|
self.srv.enc_params_header = base64.b64encode(
|
|
json.dumps(params, ensure_ascii=False).encode()
|
|
).decode()
|
|
|
|
got, enc_params = self.client.download("F-1")
|
|
self.assertEqual(got.read_bytes(), enc_bytes)
|
|
self.assertTrue(enc_params.get("compressed"))
|
|
# 解密 -> 7z -> 解压
|
|
tmp7z = self.tmp / "restore.7z"
|
|
with open(got, "rb") as s, open(tmp7z, "wb") as d:
|
|
self.engine.decrypt_stream(
|
|
s, d, enc_params["key_id"], enc_params["context"].encode()
|
|
)
|
|
out_dir = self.tmp / "out"
|
|
out_dir.mkdir()
|
|
r = subprocess.run(
|
|
["7z", "x", "-y", "-bd", f"-o{out_dir}", str(tmp7z)],
|
|
capture_output=True, text=True,
|
|
)
|
|
self.assertEqual(r.returncode, 0, r.stderr)
|
|
self.assertEqual((out_dir / "原始.bin").read_bytes(), data)
|
|
got.unlink()
|
|
|
|
def test_complete_uses_long_timeout(self):
|
|
# 客户端 timeout=1 但 complete 固定 600s: 服务器 0.3s 慢响应不超时
|
|
srv = ProtocolServer()
|
|
srv.complete_delay = 0.3
|
|
httpd = ThreadingHTTPServer(("127.0.0.1", 0), srv.make_handler())
|
|
threading.Thread(target=httpd.serve_forever, daemon=True).start()
|
|
self.addCleanup(httpd.shutdown)
|
|
self.addCleanup(httpd.server_close)
|
|
client = TransferClient(
|
|
f"http://127.0.0.1:{httpd.server_address[1]}", timeout=1, retry_delay=0.01
|
|
)
|
|
# 卷齐 -> complete 成功 (600s 大超时覆盖慢响应)
|
|
srv.init_json = {"chunk_count": 1}
|
|
srv.received = {1}
|
|
self.assertEqual(client.complete("T-1")["status"], "done")
|
|
|
|
def test_download_timeout_wrapped(self):
|
|
# download 用 self.timeout: 超时抛 TransferError, 不裸抛 TimeoutError
|
|
srv = ProtocolServer()
|
|
srv.files = [{"file_id": "F-1", "file_name": "a.bin", "size": 1, "data": b"x"}]
|
|
srv.file_delay = 0.3
|
|
httpd = ThreadingHTTPServer(("127.0.0.1", 0), srv.make_handler())
|
|
threading.Thread(target=httpd.serve_forever, daemon=True).start()
|
|
self.addCleanup(httpd.shutdown)
|
|
self.addCleanup(httpd.server_close)
|
|
client = TransferClient(
|
|
f"http://127.0.0.1:{httpd.server_address[1]}", timeout=0.1, retry_delay=0.01
|
|
)
|
|
with self.assertRaises(TransferError):
|
|
client.download("F-1")
|
|
|
|
def test_download_multichunk(self):
|
|
self.srv.files = [
|
|
{"file_id": "F-1", "file_name": "多卷.bin", "size": 12,
|
|
"chunks": [b"part-one-", b"part-two-", b"part3"]},
|
|
]
|
|
got, _ = self.client.download("F-1")
|
|
self.assertEqual(got.read_bytes(), b"part-one-part-two-part3")
|
|
got.unlink()
|
|
|
|
def test_download_path_traversal_sanitized(self):
|
|
# 恶意服务端返回 ../../ 文件名 -> 消毒成 basename, 不写穿目录
|
|
self.srv.files = [
|
|
{"file_id": "F-1", "file_name": "../../etc/cron.d/evil",
|
|
"size": 4, "data": b"evil"},
|
|
]
|
|
dest = self.tmp / "sub"
|
|
dest.mkdir()
|
|
got, params = self.client.download("F-1", dest)
|
|
self.assertEqual(got, dest / "evil") # basename, 不含 ../
|
|
self.assertEqual(got.read_bytes(), b"evil")
|
|
self.assertFalse((self.tmp / "etc").exists())
|
|
got.unlink()
|
|
|
|
def test_download_unknown_404(self):
|
|
with self.assertRaises(TransferError):
|
|
self.client.download("NOPE")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|