756 lines
31 KiB
Python
756 lines
31 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 os
|
|
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, json_body: dict[str, Any] | None = None,
|
|
) -> tuple[int, dict[str, Any]]:
|
|
"""发请求, 返回 (HTTP状态码, JSON 载荷)。4xx/5xx 不抛, 由调用方判断"""
|
|
data = body
|
|
req = urlrequest.Request(self.base_url + path, data=data, method=method)
|
|
if json_body is not None:
|
|
data = json.dumps(json_body, ensure_ascii=False).encode()
|
|
req = urlrequest.Request(self.base_url + path, data=data, method=method)
|
|
req.add_header("Content-Type", "application/json")
|
|
elif data 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 whoami(self) -> str:
|
|
"""GET /api/auth/me: 验证 token 并返回当前用户名"""
|
|
code, payload = self._request("GET", "/api/auth/me")
|
|
if code != 200:
|
|
raise TransferError(f"登录验证失败 (HTTP {code}): {payload}")
|
|
return str(payload.get("username", "?"))
|
|
|
|
def register(self, username: str, password: str) -> dict[str, Any]:
|
|
"""POST /api/auth/register: 注册账号 (成功即登录)"""
|
|
code, payload = self._request(
|
|
"POST", "/api/auth/register",
|
|
json.dumps({"username": username, "password": password}).encode(),
|
|
)
|
|
if code != 200 or "token" not in payload:
|
|
raise TransferError(f"注册失败 (HTTP {code}): {payload}")
|
|
return payload
|
|
|
|
def login(self, username: str, password: str) -> dict[str, Any]:
|
|
"""POST /api/auth/login: 登录, 返回 {username, token}"""
|
|
code, payload = self._request(
|
|
"POST", "/api/auth/login",
|
|
json.dumps({"username": username, "password": password}).encode(),
|
|
)
|
|
if code != 200 or "token" not in payload:
|
|
raise TransferError(f"登录失败 (HTTP {code}): {payload}")
|
|
return payload
|
|
|
|
def logout(self) -> None:
|
|
"""POST /api/auth/logout: 注销当前 token"""
|
|
code, payload = self._request("POST", "/api/auth/logout")
|
|
if code != 200:
|
|
raise TransferError(f"注销失败 (HTTP {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)
|
|
name = _parse_disposition(resp.headers.get("Content-Disposition") or "")
|
|
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)
|
|
outer.close()
|
|
# 卷 2..N: aria2 并发拉取 (打满带宽); 无 aria2 回退逐卷
|
|
if chunk_count > 1:
|
|
try:
|
|
self._download_aria2(file_id, dest, chunk_count)
|
|
except TransferError as e:
|
|
print(f"[aria2] 不可用, 回退逐卷下载: {e}")
|
|
outer = tqdm(total=chunk_count - 1, desc="[下载]", unit="卷",
|
|
position=0, leave=False)
|
|
with open(dest, "ab") as f:
|
|
for idx in range(2, chunk_count + 1):
|
|
with self._open_chunk(file_id, idx) as resp:
|
|
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 _aria2_pull(self, file_id: str, chunk_range: range, workdir: Path) -> list[Path]:
|
|
"""aria2 并发拉取卷文件 (多任务打满带宽), 返回按序的卷文件列表
|
|
|
|
每卷是独立 URL 的独立文件, aria2 多连接并发下载; 服务端无需 Range。
|
|
"""
|
|
import shutil
|
|
aria2 = shutil.which("aria2c")
|
|
if aria2 is None:
|
|
raise TransferError("未安装 aria2 (apt install aria2)")
|
|
urls_file = workdir / "urls.txt"
|
|
with open(urls_file, "w", encoding="utf-8") as f:
|
|
for i in chunk_range:
|
|
f.write(f"{self.base_url}/api/files/{file_id}/chunk/{i}\n")
|
|
cmd = [
|
|
aria2, "-i", str(urls_file), "-d", str(workdir),
|
|
"--max-concurrent-downloads=8", "--split=4", "--max-connection-per-server=4",
|
|
"--min-split-size=1M", "--file-allocation=none", "--allow-overwrite=true",
|
|
"--auto-file-renaming=false", "--console-log-level=warn",
|
|
"--summary-interval=1", "--quiet=false", "--no-conf=true",
|
|
]
|
|
if self.token:
|
|
cmd += ["--header", f"Authorization: Bearer {self.token}"]
|
|
r = subprocess.run(cmd, timeout=self.timeout * max(len(chunk_range), 1) + 60)
|
|
if r.returncode != 0:
|
|
raise TransferError(f"aria2 下载失败 (exit {r.returncode})")
|
|
files = sorted(workdir.glob("chunk_*"), key=lambda p: int(p.name.split("_")[1]))
|
|
if len(files) != len(chunk_range):
|
|
raise TransferError(f"aria2 卷文件不完整: {len(files)}/{len(chunk_range)}")
|
|
return files
|
|
|
|
def _download_aria2(self, file_id: str, dest: Path, chunk_count: int) -> None:
|
|
"""aria2 并发拉卷 2..N 追加到 dest (卷 1 已在 requests 流中下载)"""
|
|
import tempfile
|
|
with tempfile.TemporaryDirectory(prefix="sz-aria2-") as td:
|
|
files = self._aria2_pull(file_id, range(2, chunk_count + 1), Path(td))
|
|
with open(dest, "ab") as f:
|
|
for p in files:
|
|
with open(p, "rb") as src:
|
|
shutil.copyfileobj(src, f, 1 << 20)
|
|
|
|
# ---------- aria2 反向拉上传 ----------
|
|
|
|
def _serve_chunks(self, transfer_id: str, chunk_files: dict[int, Path],
|
|
pull_token: str, port: int = 0):
|
|
"""起临时 HTTP 服务提供卷文件 (token 校验), 返回 httpd (serve_forever 已起)
|
|
|
|
端口段 19000-19099 固定 (PC 防火墙放行此段, 手机 ZeroTermux 无防火墙)。
|
|
"""
|
|
import secrets
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
if pull_token is None:
|
|
pull_token = secrets.token_hex(8)
|
|
|
|
class Handler(BaseHTTPRequestHandler):
|
|
def do_GET(self) -> None: # noqa: N802 (HTTP 方法名)
|
|
import urllib.parse as up
|
|
parsed = up.urlparse(self.path)
|
|
qs = up.parse_qs(parsed.query)
|
|
if qs.get("token", [""])[0] != pull_token:
|
|
self.send_response(403)
|
|
self.end_headers()
|
|
return
|
|
parts = parsed.path.split("/")
|
|
if len(parts) != 4 or parts[1] != "pull" or parts[2] != transfer_id:
|
|
self.send_response(404)
|
|
self.end_headers()
|
|
return
|
|
try:
|
|
f = chunk_files.get(int(parts[3]))
|
|
except ValueError:
|
|
f = None
|
|
if f is None or not f.exists():
|
|
self.send_response(404)
|
|
self.end_headers()
|
|
return
|
|
data = f.read_bytes()
|
|
self.send_response(200)
|
|
self.send_header("Content-Length", str(len(data)))
|
|
self.end_headers()
|
|
self.wfile.write(data)
|
|
|
|
def log_message(self, format: str, *args: Any) -> None:
|
|
pass
|
|
|
|
httpd: ThreadingHTTPServer | None = None
|
|
if port == 0:
|
|
for p in range(19000, 19100):
|
|
try:
|
|
httpd = ThreadingHTTPServer(("0.0.0.0", p), Handler)
|
|
break
|
|
except OSError:
|
|
continue
|
|
if httpd is None:
|
|
httpd = ThreadingHTTPServer(("0.0.0.0", port), Handler)
|
|
import threading
|
|
threading.Thread(target=httpd.serve_forever, daemon=True).start()
|
|
return httpd, pull_token
|
|
|
|
@staticmethod
|
|
def _local_ip() -> str:
|
|
"""局域网 IP: UDP connect 拿本地出口地址"""
|
|
import socket
|
|
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
|
try:
|
|
s.connect(("10.255.255.255", 1))
|
|
return s.getsockname()[0]
|
|
except OSError:
|
|
return "127.0.0.1"
|
|
finally:
|
|
s.close()
|
|
|
|
def _pull_upload(self, transfer_id: str, chunk_files: dict[int, Path]) -> bool:
|
|
"""aria2 反向拉上传: 临时 HTTP 服务 -> 服务端 pull 并发拉卷
|
|
|
|
成功返回 True; 服务端无 aria2/拉取失败返回 False (调用方回退逐卷 PUT)。
|
|
"""
|
|
import secrets
|
|
httpd, pull_token = self._serve_chunks(
|
|
transfer_id, chunk_files, secrets.token_hex(8))
|
|
try:
|
|
base_url = f"http://{self._local_ip()}:{httpd.server_address[1]}"
|
|
code, payload = self._request(
|
|
"POST", f"/api/transfer/{transfer_id}/pull",
|
|
json_body={"base_url": base_url, "token": pull_token},
|
|
timeout=60,
|
|
)
|
|
if code != 200:
|
|
print(f"[aria2] 服务端未启用反向拉取 (HTTP {code}): {payload.get('error', '')}")
|
|
return False
|
|
pulled, total = payload.get("pulled", 0), payload.get("total", 0)
|
|
print(f"[aria2] 反向拉取完成: {pulled}/{total} 卷")
|
|
return pulled == total
|
|
except TransferError as e:
|
|
print(f"[aria2] 反向拉取失败, 回退逐卷: {e}")
|
|
return False
|
|
finally:
|
|
httpd.shutdown()
|
|
httpd.server_close()
|
|
|
|
# ---------- 全流程 ----------
|
|
|
|
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)} 卷")
|
|
|
|
# aria2 反向拉上传: 优先 (打满上行带宽), 失败回退逐卷 PUT
|
|
if pending and self._pull_upload(transfer_id, chunk_files):
|
|
pending = []
|
|
if 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 _safe_name(name: str) -> str:
|
|
"""下载文件名消毒: 只保留 basename, 去路径穿越/控制字符 (防恶意服务端)"""
|
|
name = Path(name).name # 去掉目录部分 (.. / ../../ 无效化)
|
|
return "".join(c for c in name if c.isprintable() and c not in "/\\\x00").strip() or "download.bin"
|
|
|
|
|
|
def _decrypt_name(token: str) -> str | None:
|
|
"""解密 'enc:' 前缀文件名, 失败返回 None (明文名原样返回)"""
|
|
if not token.startswith("enc:"):
|
|
return None
|
|
try:
|
|
from crypto import CryptoEngine
|
|
return CryptoEngine(_load_keyring_path()).decrypt_name(token)
|
|
except Exception:
|
|
return None
|
|
|
|
|
|
def _parse_disposition(cd: str) -> str:
|
|
"""解析 Content-Disposition 文件名 (utf-8 优先, 解密 enc: 前缀)"""
|
|
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('"')
|
|
plain = _decrypt_name(name)
|
|
return _safe_name(plain if plain is not None else name)
|
|
|
|
|
|
def _config_path() -> Path:
|
|
"""配置文件路径: SZ_CONFIG env -> XDG (~/.config/7z-encrypt/config.json)
|
|
-> 兼容旧式相对 config/config.json (源码开发模式)
|
|
|
|
安装版 (deb) cwd 不可写, 配置必须走用户级 XDG 目录。
|
|
"""
|
|
env = os.environ.get("SZ_CONFIG")
|
|
if env:
|
|
return Path(env).expanduser()
|
|
xdg = os.environ.get("XDG_CONFIG_HOME") or str(Path.home() / ".config")
|
|
xdg_cfg = Path(xdg) / "7z-encrypt" / "config.json"
|
|
if xdg_cfg.exists():
|
|
return xdg_cfg
|
|
legacy = Path("config/config.json")
|
|
return xdg_cfg if not legacy.exists() else legacy
|
|
|
|
|
|
def _config_server() -> str:
|
|
"""从配置文件读 server.url (安装版 CLI 默认服务器)"""
|
|
try:
|
|
cfg = _config_path()
|
|
return str(json.loads(cfg.read_text(encoding="utf-8")).get("server", {}).get("url", ""))
|
|
except (OSError, json.JSONDecodeError):
|
|
return ""
|
|
|
|
|
|
def _load_token() -> str:
|
|
"""从配置文件读 server.token (CLI 默认来源, XDG 优先)"""
|
|
try:
|
|
cfg = _config_path()
|
|
return str(json.loads(cfg.read_text(encoding="utf-8")).get("server", {}).get("token", ""))
|
|
except (OSError, json.JSONDecodeError):
|
|
return ""
|
|
|
|
|
|
def _save_token(token: str) -> None:
|
|
"""登录/注册成功后 token 持久化到配置文件 server.token"""
|
|
cfg = _config_path()
|
|
cfg.parent.mkdir(parents=True, exist_ok=True)
|
|
data = json.loads(cfg.read_text(encoding="utf-8")) if cfg.exists() else {}
|
|
data.setdefault("server", {})["token"] = token
|
|
cfg.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")
|
|
|
|
|
|
def _load_keyring_path() -> Path:
|
|
"""密钥库路径: 配置显式指定用指定, 否则系统默认 (SZ_KEYRING env/XDG)"""
|
|
try:
|
|
cfg = _config_path()
|
|
raw = str(json.loads(cfg.read_text(encoding="utf-8")).get("crypto", {}).get("keyring_path", ""))
|
|
if raw:
|
|
return Path(raw)
|
|
except (OSError, json.JSONDecodeError):
|
|
pass
|
|
from crypto import DEFAULT_KEYRING
|
|
return DEFAULT_KEYRING
|
|
|
|
|
|
def _resolve_id(client: TransferClient, ref: str) -> str:
|
|
"""文件引用解析: 纯数字 = 列表序号 (1-based), 否则视为 file_id"""
|
|
if ref.isdigit():
|
|
files = client.list_files()
|
|
idx = int(ref) - 1
|
|
if not 0 <= idx < len(files):
|
|
raise TransferError(f"序号越界: {ref} (共 {len(files)} 个文件)")
|
|
return files[idx]["file_id"]
|
|
return ref
|
|
|
|
|
|
def main() -> int:
|
|
ap = argparse.ArgumentParser(description="传输客户端: 上传 / 文件列表 / 下载")
|
|
ap.add_argument("--server", default=None, help="服务端地址 (默认读配置 server.url)")
|
|
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)")
|
|
ap.add_argument("--register", nargs=2, metavar=("用户名", "密码"), help="注册账号 (成功即登录)")
|
|
ap.add_argument("--login", nargs=2, metavar=("用户名", "密码"), help="登录账号")
|
|
ap.add_argument("--logout", action="store_true", help="注销当前账号 (token 失效)")
|
|
ap.add_argument("--whoami", action="store_true", help="显示当前登录用户")
|
|
args = ap.parse_args()
|
|
|
|
server = args.server or _config_server()
|
|
if not server:
|
|
ap.error("未指定 --server 且配置无 server.url (先 sz-config --set-server)")
|
|
token = args.token if args.token is not None else _load_token()
|
|
client = TransferClient(server, token=token)
|
|
|
|
if args.whoami:
|
|
try:
|
|
uname = client.whoami()
|
|
except TransferError as e:
|
|
print(f"[错误] {e}")
|
|
return 1
|
|
print(f"[用户] {uname}")
|
|
return 0
|
|
|
|
if args.register:
|
|
try:
|
|
payload = client.register(args.register[0], args.register[1])
|
|
_save_token(payload["token"])
|
|
except TransferError as e:
|
|
print(f"[错误] {e}")
|
|
return 1
|
|
print(f"[注册] 成功, 已登录: {payload['username']} (token 已保存)")
|
|
return 0
|
|
|
|
if args.login:
|
|
try:
|
|
payload = client.login(args.login[0], args.login[1])
|
|
_save_token(payload["token"])
|
|
except TransferError as e:
|
|
print(f"[错误] {e}")
|
|
return 1
|
|
print(f"[登录] 成功: {payload['username']} (token 已保存)")
|
|
return 0
|
|
|
|
if args.logout:
|
|
try:
|
|
client.logout()
|
|
_save_token("")
|
|
except TransferError as e:
|
|
print(f"[错误] {e}")
|
|
return 1
|
|
print("[注销] 已退出 (token 已清除)")
|
|
return 0
|
|
|
|
if args.list:
|
|
try:
|
|
files = client.list_files()
|
|
except TransferError as e:
|
|
print(f"[错误] {e}")
|
|
return 1
|
|
print(f"[文件] 共 {len(files)} 个:")
|
|
for i, f in enumerate(files, 1):
|
|
size_mb = f["size"] / 1048576
|
|
name = _decrypt_name(f["file_name"]) or f["file_name"]
|
|
if name.startswith("enc:"):
|
|
name = name[:22] + "..." # 加密名截断 (keyring 无对应密钥)
|
|
print(f" {i:>2}. {name} ({size_mb:.1f} MB)")
|
|
return 0
|
|
|
|
if args.download:
|
|
try:
|
|
args.download = _resolve_id(client, args.download)
|
|
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:
|
|
args.delete = _resolve_id(client, args.delete)
|
|
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())
|