Files

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()