Files
7z-encrypt/transfer.py
T

439 lines
18 KiB
Python

"""
传输客户端模块 (transfer)
模块:客户端 / 传输客户端
输入:init json + 卷文件 + 服务器配置
输出:服务端回执 (status/file_id)
协议 (共享契约):
POST /api/transfer/init {init json} -> {transfer_id}
GET /api/transfer/{id}/chunks -> {received: [n...]}
PUT /api/transfer/{id}/chunk/{n} 卷密文 -> {ok}
POST /api/transfer/{id}/complete -> {status, file_id}
断点续传: 先查已收卷集合只补缺卷; 单卷失败重试 max_retries 次, 仍失败抛 TransferError。
"""
import argparse
import base64
import json
import shutil
import subprocess
import sys
import time
from pathlib import Path
from typing import IO, Any
from tqdm import tqdm
from urllib import error as urlerror
from urllib import parse as urlparse
from urllib import request as urlrequest
DEFAULT_MAX_RETRIES = 3
DEFAULT_RETRY_DELAY = 1.0 # 秒
class TransferError(RuntimeError):
"""传输失败 (服务端错误/网络错误/重试耗尽)"""
class TransferClient:
"""HTTP 分片上传客户端 (标准库 urllib, 零依赖)"""
def __init__(
self,
base_url: str,
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 封装 ----------
def _request(
self, method: str, path: str, body: bytes | None = None, timeout: int | float | None = None
) -> tuple[int, dict[str, Any]]:
"""发请求, 返回 (HTTP状态码, JSON 载荷)。4xx/5xx 不抛, 由调用方判断"""
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()
code = resp.status
except urlerror.HTTPError as e:
code = e.code
data = e.read()
except urlerror.URLError as e:
raise TransferError(f"网络错误: {e.reason}") from e
except TimeoutError:
# urllib 读响应阶段的超时是裸抛, 不包 URLError
raise TransferError(f"请求超时 ({timeout or self.timeout}s)") from None
try:
payload: dict[str, Any] = json.loads(data) if data else {}
except json.JSONDecodeError:
payload = {}
return code, payload
# ---------- 协议端点 ----------
def init_transfer(self, init_json: dict[str, Any]) -> str:
"""POST /init: 建任务, 返回 transfer_id"""
code, payload = self._request(
"POST", "/api/transfer/init", json.dumps(init_json).encode()
)
if code != 200 or "transfer_id" not in payload:
raise TransferError(f"init 失败 (HTTP {code}): {payload}")
return payload["transfer_id"]
def get_received_chunks(self, transfer_id: str) -> set[int]:
"""GET /chunks: 查询服务端已收卷 (断点续传依据)"""
code, payload = self._request("GET", f"/api/transfer/{transfer_id}/chunks")
if code != 200:
raise TransferError(f"查询已收卷失败 (HTTP {code})")
return set(payload.get("received", []))
def upload_chunk(self, transfer_id: str, index: int, chunk_path: Path) -> None:
"""PUT /chunk/{n}: 上传单卷, 失败重试 max_retries 次"""
data = chunk_path.read_bytes()
code: int = 0
for attempt in range(1, self.max_retries + 1):
code, payload = self._request(
"PUT", f"/api/transfer/{transfer_id}/chunk/{index}", body=data
)
if code == 200 and payload.get("ok"):
return
if attempt < self.max_retries:
time.sleep(self.retry_delay)
raise TransferError(
f"卷 {index} 上传失败 (HTTP {code}, 重试 {self.max_retries} 次耗尽)"
)
def complete(self, transfer_id: str) -> dict[str, Any]:
"""POST /complete: 触发服务端汇聚合并入库, 返回回执
服务端零合并 (卷齐即入库), 正常秒回; 超时 600s 兜底。
"""
sys.stdout.write("[等待] 服务端确认入库...\n")
sys.stdout.flush()
code, payload = self._request(
"POST", f"/api/transfer/{transfer_id}/complete", timeout=600
)
if code != 200:
raise TransferError(f"complete 失败 (HTTP {code}): {payload}")
return payload
def list_files(self) -> list[dict[str, Any]]:
"""GET /api/files: 服务端文件列表 (ls)"""
code, payload = self._request("GET", "/api/files")
if code != 200:
raise TransferError(f"文件列表失败 (HTTP {code}): {payload}")
return payload.get("files", [])
def delete_file(self, file_id: str) -> None:
"""DELETE /api/files/{id}: 删除服务端文件"""
code, payload = self._request("DELETE", f"/api/files/{file_id}")
if code != 200:
raise TransferError(f"删除失败 (HTTP {code}): {payload}")
if not payload.get("ok"):
raise TransferError(f"删除失败: {payload}")
def get_quota(self) -> dict[str, Any]:
"""GET /api/quota: 空间配额 (used/quota/remain/percent)"""
code, payload = self._request("GET", "/api/quota")
if code != 200:
raise TransferError(f"配额查询失败 (HTTP {code}): {payload}")
return payload
def download(self, file_id: str, dest: Path | None = None) -> tuple[Path, dict[str, Any]]:
"""逐卷下载密文 (零合并: 服务端只存卷, 客户端本地拼接), 返回 (密文路径, 解密参数)。
先拉卷 1 拿响应头 (X-Enc-Params/X-Chunk-Count/文件名), 再循环 2..N 追加写。
dest 缺省用服务端文件名; 是目录则拼文件名。
"""
try:
with self._open_chunk(file_id, 1) as resp:
enc_params: dict[str, Any] = {}
ep = resp.headers.get("X-Enc-Params") or ""
if ep:
enc_params = json.loads(base64.b64decode(ep).decode())
chunk_count = max(int(resp.headers.get("X-Chunk-Count") or 1), 1)
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():
dest = dest / name
assert dest is not None
outer = tqdm(total=chunk_count, desc="[下载]", unit="卷",
position=0, leave=False)
with open(dest, "wb") as f:
self._download_chunk_to(resp, f, outer, 1, chunk_count)
for idx in range(2, chunk_count + 1):
with self._open_chunk(file_id, idx) as resp:
with open(dest, "ab") as f:
self._download_chunk_to(resp, f, outer, idx, chunk_count)
outer.close()
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
except urlerror.URLError as e:
raise TransferError(f"网络错误: {e.reason}") from e
except TimeoutError:
raise TransferError(f"下载超时 ({self.timeout}s)") from None
def _open_chunk(self, file_id: str, idx: int):
"""打开单卷响应流 (带认证)"""
req = urlrequest.Request(
self.base_url + f"/api/files/{file_id}/chunk/{idx}", method="GET"
)
if self.token:
req.add_header("Authorization", f"Bearer {self.token}")
return urlrequest.urlopen(req, timeout=self.timeout)
def _show_chunk_progress(self, cur: int, total_chunks: int) -> None:
"""按卷粒度显示下载进度 (与上传日志对称)"""
pct = cur * 100 // total_chunks
sys.stdout.write(f"\r[下载] 卷 {cur}/{total_chunks} ({pct}%)")
sys.stdout.flush()
def _download_chunk_to(
self, resp, f: IO[bytes], outer: Any, idx: int, total_chunks: int
) -> int:
"""卷内字节流写入 + 进度 (内层字节条), 返回写入字节数"""
size = int(resp.headers.get("Content-Length") or 0)
inner = tqdm(total=size, desc=f" 卷 {idx}/{total_chunks}", unit="B",
unit_scale=True, position=1, leave=False)
done = 0
while True:
chunk = resp.read(1 << 20)
if not chunk:
break
f.write(chunk)
done += len(chunk)
inner.update(len(chunk))
inner.close()
outer.update(1)
return done
# ---------- 全流程 ----------
def transfer(
self, init_json: dict[str, Any], chunk_files: dict[int, Path]
) -> dict[str, Any]:
"""init -> 查已收 -> 逐卷上传(跳过已收) -> complete
Args:
init_json: metadata.build_init_json 的产物
chunk_files: {卷序号: 卷文件路径}
Returns:
dict: 服务端回执 {status, file_id}
"""
transfer_id = self.init_transfer(init_json)
print(f"[传输] transfer_id: {transfer_id}")
received = self.get_received_chunks(transfer_id)
total = init_json["chunk_count"]
pending = [i for i in range(1, total + 1) if i not in received]
if received:
print(f"[传输] 断点续传: 服务端已收 {len(received)}/{total} 卷, 补传 {len(pending)} 卷")
pbar = tqdm(total=len(pending), desc="[传输] 上传", unit="卷", leave=False)
for done, index in enumerate(pending, start=1):
try:
self.upload_chunk(transfer_id, index, chunk_files[index])
except TransferError:
raise # 重试耗尽, 由调用方决定 (记录断点状态待补传)
pbar.update(1)
pbar.set_description(f"[传输] 上传 {done}/{len(pending)} 卷")
pbar.close()
receipt = self.complete(transfer_id)
print(f"[传输] 完成: {receipt}")
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")
ap.add_argument("--init", help="init json 路径 (上传模式)")
ap.add_argument("--manifest", help="卷清单路径 (上传模式)")
ap.add_argument("--list", action="store_true", help="列出服务端文件")
ap.add_argument("--quota", action="store_true", help="显示空间配额")
ap.add_argument("--download", metavar="FILE_ID", help="下载文件")
ap.add_argument("--delete", metavar="FILE_ID", 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()
token = args.token if args.token is not None else _load_token()
client = TransferClient(args.server, token=token)
if args.list:
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
print(f" {f['file_id']} {f['file_name']} ({size_mb:.1f} MB) {f['created_at']}")
return 0
if args.download:
try:
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())
if enc_params.get("compressed"):
# 解密还原的是 7z 流 -> 解压成原始文件
tmp7z = tmp_enc.with_suffix(".7z")
pbar = tqdm(total=enc_path.stat().st_size, desc="[解密]", unit="B",
unit_scale=True, leave=False)
with open(enc_path, "rb") as src, open(tmp7z, "wb") as dst:
engine.decrypt_stream(
src, dst,
enc_params["key_id"],
enc_params.get("context", "7z-encrypt:v1").encode(),
progress=lambda n: pbar.update(n),
)
pbar.close()
enc_path.unlink()
if dest_arg and not dest_arg.is_dir():
out_dir = dest_arg.parent
else:
out_dir = (dest_arg if dest_arg else Path("."))
out_dir.mkdir(parents=True, exist_ok=True)
print(f"[解压] 7z x -o{out_dir} ...")
r = subprocess.run(
["7z", "x", "-y", "-bd", f"-o{out_dir}", str(tmp7z)],
capture_output=True, text=True,
)
if r.returncode != 0:
print(f"[错误] 7z 解压失败: {r.stderr[-300:]}")
return 1
tmp7z.unlink()
if dest_arg and not dest_arg.is_dir():
# 指定了输出文件名: 包内原名 -> 改名为目标
shutil.move(out_dir / name, dest_arg)
final = dest_arg
print(f"[下载] 完成 -> {final} ({final.stat().st_size / 1048576:.1f} MB, 已解密+解压)")
else:
pbar = tqdm(total=enc_path.stat().st_size, desc="[解密]", unit="B",
unit_scale=True, leave=False)
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(),
progress=lambda n: pbar.update(n),
)
pbar.close()
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
return 0
if args.quota:
try:
q = client.get_quota()
except TransferError as e:
print(f"[错误] {e}")
return 1
print(f"[配额] 已用 {q['used_bytes'] / 1048576:.1f} MB / "
f"{q['quota_bytes'] / 1048576:.1f} MB ({q['percent']}%) "
f"剩余 {q['remain_bytes'] / 1048576:.1f} MB")
return 0
if args.delete:
try:
client.delete_file(args.delete)
except TransferError as e:
print(f"[错误] {e}")
return 1
print(f"[删除] 已删除: {args.delete}")
return 0
if not (args.init and args.manifest):
ap.error("需指定 --list / --quota / --download / --delete / (--init + --manifest) 之一")
init_json = json.loads(Path(args.init).read_text(encoding="utf-8"))
manifest = json.loads(Path(args.manifest).read_text(encoding="utf-8"))
chunk_files = {
ch["index"]: Path(ch["filename"]).resolve()
for ch in manifest
}
missing = [i for i, p in chunk_files.items() if not p.exists()]
if missing:
print(f"[错误] 卷文件缺失: {missing}")
return 1
try:
client.transfer(init_json, chunk_files)
except TransferError as e:
print(f"[错误] {e}")
return 1
return 0
if __name__ == "__main__":
sys.exit(main())