零知识: transfer 带 token + download 返回 enc_params + --decrypt 本地解密还原; config 加 server.token; 测试 62 项

This commit is contained in:
lou
2026-08-10 01:06:33 +08:00
parent a2beb85957
commit 3e92c811b9
3 changed files with 161 additions and 14 deletions
+23 -2
View File
@@ -18,7 +18,7 @@ DEFAULT_CONFIG = Path(__file__).resolve().parent / "config" / "config.json"
# 默认值 (用户 config.json 逐层覆盖)
DEFAULTS: dict[str, Any] = {
"server": {"url": "", "timeout": 30, "max_retries": 3, "retry_delay": 1.0},
"server": {"url": "", "token": "", "timeout": 30, "max_retries": 3, "retry_delay": 1.0},
"crypto": {"key_id": "default_key", "keyring_path": "config/keyring.json"},
"splitter": {"chunk_size_mb": 10},
"state": {"db_path": "config/tasks.db", "cleanup_days": 7},
@@ -95,6 +95,11 @@ class AppConfig:
def retry_delay(self) -> float:
return self.data["server"]["retry_delay"]
@property
def token(self) -> str:
"""服务端认证 token (Bearer)"""
return self.data["server"]["token"]
@property
def key_id(self) -> str:
return self.data["crypto"]["key_id"]
@@ -144,6 +149,7 @@ def main() -> int:
ap.add_argument("--init", action="store_true", help="生成默认 config.json 模板")
ap.add_argument("--path", default=str(DEFAULT_CONFIG), help="配置文件路径")
ap.add_argument("--set-server", metavar="URL", help="更新 server.url 并保存")
ap.add_argument("--set-token", metavar="TOKEN", help="更新 server.token 并保存")
args = ap.parse_args()
if args.init:
@@ -157,6 +163,7 @@ def main() -> int:
{
"server": {
"url": "http://127.0.0.1:8000",
"token": "",
"timeout": 30,
"max_retries": 3,
"retry_delay": 1.0,
@@ -188,13 +195,27 @@ def main() -> int:
print(f"[config] server.url -> {args.set_server} ({path})")
return 0
if args.set_token:
path = Path(args.path)
if not path.exists():
print(f"[错误] 配置文件不存在: {path}, 先 --init 生成")
return 1
data = json.loads(path.read_text(encoding="utf-8"))
data.setdefault("server", {})["token"] = args.set_token
path.write_text(
json.dumps(data, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
print(f"[config] server.token 已更新 ({path})")
return 0
try:
cfg = AppConfig(args.path)
except (ConfigError, json.JSONDecodeError) as e:
print(f"[错误] 配置无效: {e}")
return 1
print(f"[config] 已加载: {cfg.config_path}")
print(f" server: {cfg.server_url} (timeout={cfg.timeout}, retries={cfg.max_retries})")
print(f" server: {cfg.server_url} (token={'已设置' if cfg.token else '未设置'}, timeout={cfg.timeout}, retries={cfg.max_retries})")
print(f" crypto: key_id={cfg.key_id}, keyring={cfg.keyring_path}")
print(f" splitter: 卷大小 {cfg.data['splitter']['chunk_size_mb']} MB")
print(f" state: db={cfg.db_path}, 清理保留 {cfg.cleanup_days}")
+58 -3
View File
@@ -4,6 +4,8 @@
运行: cd /home/lou/文档/7z-encrypt && .venv/bin/python -m unittest discover -s tests -v
"""
import base64
import io
import json
import os
import re
@@ -38,6 +40,8 @@ class ProtocolServer:
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.require_token = "" # 非空时要求 Authorization: Bearer <此值>
def make_handler(self):
srv = self
@@ -77,6 +81,9 @@ class ProtocolServer:
self._json(404, {"error": "not found"})
def do_GET(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/transfer/([^/]+)/chunks", self.path)
if m:
self._json(200, {"received": sorted(srv.received)})
@@ -107,6 +114,8 @@ class ProtocolServer:
"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.end_headers()
self.wfile.write(body)
return
@@ -219,6 +228,11 @@ class TestListDownload(TransferTestBase):
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()
@@ -228,12 +242,13 @@ class TestListDownload(TransferTestBase):
def test_download_to_path(self):
dest = self.tmp / "out.bin"
got = self.client.download("F-1", dest)
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")
got, _ = self.client.download("F-1")
self.assertEqual(got.name, "报告.pdf")
self.assertEqual(got.read_bytes(), b"hello")
got.unlink()
@@ -242,10 +257,50 @@ class TestListDownload(TransferTestBase):
# dest 是目录 -> 自动拼服务端文件名
dest_dir = self.tmp / "dl"
dest_dir.mkdir()
got = self.client.download("F-1", dest_dir)
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_complete_uses_long_timeout(self):
# 客户端 timeout=1 但 complete 固定 600s: 服务器 0.3s 慢响应不超时
srv = ProtocolServer()
+80 -9
View File
@@ -14,6 +14,7 @@
"""
import argparse
import base64
import json
import sys
import time
@@ -40,11 +41,13 @@ class TransferClient:
timeout: int | float = 30,
max_retries: int = DEFAULT_MAX_RETRIES,
retry_delay: float = DEFAULT_RETRY_DELAY,
token: str = "",
) -> None:
self.base_url = base_url.rstrip("/")
self.timeout = timeout
self.max_retries = max_retries
self.retry_delay = retry_delay
self.token = token
# ---------- HTTP 封装 ----------
@@ -55,6 +58,8 @@ class TransferClient:
req = urlrequest.Request(self.base_url + path, data=body, method=method)
if body is not None:
req.add_header("Content-Type", "application/octet-stream")
if self.token:
req.add_header("Authorization", f"Bearer {self.token}")
try:
with urlrequest.urlopen(req, timeout=timeout or self.timeout) as resp:
data = resp.read()
@@ -127,18 +132,29 @@ class TransferClient:
raise TransferError(f"文件列表失败 (HTTP {code}): {payload}")
return payload.get("files", [])
def download(self, file_id: str, dest: Path | None = None) -> Path:
"""GET /api/files/{id}: 流式下载, 实时进度。dest 缺省用服务端文件名; 是目录则拼文件名。"""
def download(self, file_id: str, dest: Path | None = None) -> tuple[Path, dict[str, Any]]:
"""GET /api/files/{id}: 流式下载密文, 实时进度, 返回 (密文路径, 解密参数)。
dest 缺省用服务端文件名; 是目录则拼文件名。
解密参数来自响应头 X-Enc-Params (base64 JSON: alg/key_id/context), 客户端本地解密用。
"""
req = urlrequest.Request(self.base_url + f"/api/files/{file_id}", method="GET")
if self.token:
req.add_header("Authorization", f"Bearer {self.token}")
try:
with urlrequest.urlopen(req, timeout=self.timeout) as resp:
# 从 Content-Disposition 取文件名 (filename*=utf-8''... 或 filename=...)
# 解密参数头
enc_params: dict[str, Any] = {}
ep = resp.headers.get("X-Enc-Params") or ""
if ep:
enc_params = json.loads(base64.b64decode(ep).decode())
cd = resp.headers.get("Content-Disposition") or ""
name = "download.bin"
if "filename*=utf-8''" in cd:
name = urlparse.unquote(cd.split("filename*=utf-8''")[1].split(";")[0])
elif "filename=" in cd:
name = cd.split("filename=")[1].split(";")[0].strip('"')
enc_params.setdefault("file_name", name) # 解密时还原用
if dest is None:
dest = Path(name)
elif dest.is_dir():
@@ -162,8 +178,10 @@ class TransferClient:
if total:
sys.stdout.write("\n")
sys.stdout.flush()
return dest
return dest, enc_params
except urlerror.HTTPError as e:
if e.code == 401:
raise TransferError("未授权: token 无效 (检查 config server.token)") from e
if e.code == 404:
raise TransferError("文件不存在 (HTTP 404)") from e
raise TransferError(f"下载失败 (HTTP {e.code})") from e
@@ -214,6 +232,23 @@ class TransferClient:
return receipt
def _load_token() -> str:
"""从 config/config.json 读 server.token (CLI 默认来源)"""
try:
cfg = Path("config/config.json")
return str(json.loads(cfg.read_text(encoding="utf-8")).get("server", {}).get("token", ""))
except (OSError, json.JSONDecodeError):
return ""
def _load_keyring_path() -> Path:
try:
cfg = Path("config/config.json")
return Path(str(json.loads(cfg.read_text(encoding="utf-8")).get("crypto", {}).get("keyring_path", "config/keyring.json")))
except (OSError, json.JSONDecodeError):
return Path("config/keyring.json")
def main() -> int:
ap = argparse.ArgumentParser(description="传输客户端: 上传 / 文件列表 / 下载")
ap.add_argument("--server", required=True, help="服务端地址, 如 http://127.0.0.1:8000")
@@ -221,13 +256,20 @@ def main() -> int:
ap.add_argument("--manifest", help="卷清单路径 (上传模式)")
ap.add_argument("--list", action="store_true", help="列出服务端文件")
ap.add_argument("--download", metavar="FILE_ID", help="下载文件")
ap.add_argument("--out", help="下载保存路径 (默认当前目录, 用服务端文件名)")
ap.add_argument("--out", help="保存路径 (目录或完整文件名)")
ap.add_argument("--decrypt", action="store_true", help="下载后用本地 keyring 解密还原明文")
ap.add_argument("--token", help="服务端 token (默认读 config/config.json)")
args = ap.parse_args()
client = TransferClient(args.server)
token = args.token if args.token is not None else _load_token()
client = TransferClient(args.server, token=token)
if args.list:
files = client.list_files()
try:
files = client.list_files()
except TransferError as e:
print(f"[错误] {e}")
return 1
print(f"[文件] 共 {len(files)} 个:")
for f in files:
size_mb = f["size"] / 1048576
@@ -236,11 +278,40 @@ def main() -> int:
if args.download:
try:
dest = client.download(args.download, Path(args.out) if args.out else None)
dest_arg = Path(args.out) if args.out else None
if args.decrypt:
# 密文先落临时 .enc, 本地解密后删
if dest_arg and not dest_arg.is_dir():
tmp_enc = Path(str(dest_arg) + ".enc")
else:
tmp_enc = (dest_arg if dest_arg else Path(".")) / ".sz-download.enc"
enc_path, enc_params = client.download(args.download, tmp_enc)
if enc_params.get("key_id") is None:
print("[错误] 响应缺解密参数 (X-Enc-Params)")
return 1
name = str(enc_params.get("file_name", "download.bin"))
if dest_arg and dest_arg.is_dir():
final = dest_arg / name
elif dest_arg:
final = dest_arg
else:
final = Path(name)
from crypto import CryptoEngine
engine = CryptoEngine(_load_keyring_path())
with open(enc_path, "rb") as src, open(final, "wb") as dst:
engine.decrypt_stream(
src, dst,
enc_params["key_id"],
enc_params.get("context", "7z-encrypt:v1").encode(),
)
enc_path.unlink()
print(f"[下载] 完成 -> {final} ({final.stat().st_size / 1048576:.1f} MB, 已解密)")
else:
dest, _ = client.download(args.download, dest_arg)
print(f"[下载] 完成 -> {dest} ({dest.stat().st_size / 1048576:.1f} MB, 密文)")
except TransferError as e:
print(f"[错误] {e}")
return 1
print(f"[下载] 完成 -> {dest} ({dest.stat().st_size / 1048576:.1f} MB)")
return 0
if not (args.init and args.manifest):