diff --git a/crypto.py b/crypto.py index bef2d0e..5392e01 100644 --- a/crypto.py +++ b/crypto.py @@ -16,7 +16,7 @@ import json import os from pathlib import Path -from typing import IO, Any +from typing import IO, Any, Callable from cryptography.cobblestone import Cobblestone256Decryptor, Cobblestone256Encryptor from cryptography.exceptions import InvalidTag @@ -82,7 +82,12 @@ class CryptoEngine: # ---------- 加解密 ---------- def encrypt_stream( - self, src: IO[bytes], dst: IO[bytes], key_id: str, context: bytes = CONTEXT + self, + src: IO[bytes], + dst: IO[bytes], + key_id: str, + context: bytes = CONTEXT, + progress: Callable[[int], Any] | None = None, ) -> dict[str, Any]: """步骤 A & F: 流式加密。src 读明文 -> dst 写密文, 返回加密参数 @@ -91,32 +96,43 @@ class CryptoEngine: dst: 可写二进制流 key_id: 密钥标识 context: 应用绑定上下文 (默认 7z-encrypt:v1) - - Returns: - dict: 加密参数 {alg, key_id, context}, 供 init json 传给服务端 + progress: 进度回调, 每读块后调 progress(已读字节) """ key = self.get_key(key_id) enc = Cobblestone256Encryptor(key, context) + done = 0 while True: chunk = src.read(1 << 16) # 64KiB 读块 if not chunk: break dst.write(enc.update(chunk)) + done += len(chunk) + if progress is not None: + progress(done) dst.write(enc.finalize()) return {"alg": "cobblestone-aes256gcm", "key_id": key_id, "context": context.decode()} def decrypt_stream( - self, src: IO[bytes], dst: IO[bytes], key_id: str, context: bytes = CONTEXT + self, + src: IO[bytes], + dst: IO[bytes], + key_id: str, + context: bytes = CONTEXT, + progress: Callable[[int], Any] | None = None, ) -> None: - """流式解密。src 读密文 -> dst 写明文。完整性验证失败抛 IntegrityError""" + """解密流 (进度回调: 每读块后调 progress(已读字节))""" key = self.get_key(key_id) - dec = Cobblestone256Decryptor(key, context) try: + dec = Cobblestone256Decryptor(key, context) + done = 0 while True: chunk = src.read(1 << 16) if not chunk: break dst.write(dec.update(chunk)) # 块级认证: 篡改在这里就抛 InvalidTag + done += len(chunk) + if progress is not None: + progress(done) dst.write(dec.finalize()) # finalize 返回剩余明文 except InvalidTag as e: raise IntegrityError("解密验证失败: 密钥不匹配或密文被篡改") from e diff --git a/metadata.py b/metadata.py index 2fec532..f7568f5 100644 --- a/metadata.py +++ b/metadata.py @@ -19,6 +19,8 @@ import sys from pathlib import Path from typing import Any +from tqdm import tqdm + from crypto import CryptoEngine, KeyNotFoundError from splitter import split_stream @@ -156,12 +158,21 @@ def main() -> int: engine.generate_key(args.key_id) print(f"[密钥库] 已生成新密钥 -> {engine.keyring_path}") with open(src_path, "rb") as src, open(enc_path, "wb") as dst: - enc_params = engine.encrypt_stream(src, dst, args.key_id) + pbar = tqdm(total=src_path.stat().st_size, desc="[加密]", unit="B", + unit_scale=True, leave=False) + enc_params = engine.encrypt_stream(src, dst, args.key_id, + progress=lambda n: pbar.update(n)) + pbar.close() print(f"[1/4] 加密完成 -> {enc_path.name} ({enc_path.stat().st_size / 1048576:.1f} MB)") # 2. 分卷 -> 卷清单 with open(enc_path, "rb") as f: - manifest = split_stream(f, chunk_size, str(enc_path), total_size=enc_path.stat().st_size) + pbar = tqdm(total=enc_path.stat().st_size, desc="[分卷]", unit="B", + unit_scale=True, leave=False) + manifest = split_stream(f, chunk_size, str(enc_path), + total_size=enc_path.stat().st_size, + progress=lambda n: pbar.update(n)) + pbar.close() with open(manifest_path, "w", encoding="utf-8") as f: json.dump(manifest, f, indent=2, ensure_ascii=False) print(f"[2/4] 分卷完成: {len(manifest)} 卷, 清单 -> {manifest_path.name}") diff --git a/splitter.py b/splitter.py index a1a8447..52d6892 100644 --- a/splitter.py +++ b/splitter.py @@ -15,7 +15,7 @@ import argparse import hashlib import json import sys -from typing import IO, Any +from typing import IO, Any, Callable def _read_block(stream: IO[bytes], size: int) -> bytes: @@ -32,7 +32,11 @@ def _read_block(stream: IO[bytes], size: int) -> bytes: def split_stream( - stream: IO[bytes], chunk_size: int, output_prefix: str, total_size: int | None = None + stream: IO[bytes], + chunk_size: int, + output_prefix: str, + total_size: int | None = None, + progress: Callable[[int], Any] | None = None, ) -> list[dict[str, Any]]: """从输入流切分卷文件, 返回清单 [{index, size, sha256, filename}] @@ -40,7 +44,8 @@ def split_stream( stream: 二进制输入流 (open(...,'rb') 或 sys.stdin.buffer) chunk_size: 每卷最大字节数 output_prefix: 输出前缀, 卷文件命名为 .partXXXX - total_size: 输入总字节数 (文件模式传入以显示百分比进度, stdin 模式为 None) + total_size: 输入总大小 (进度条用, 可 None) + progress: 进度回调, 每写一卷后调 progress(已写字节) Returns: list: 清单 @@ -72,8 +77,10 @@ def split_stream( part_num += 1 done += size - # 实时进度 (不做黑盒等待) - if total_size: + # 实时进度: 有回调交给调用方 (tqdm), 否则内部打印 + if progress is not None: + progress(done) + elif total_size: pct = done * 100 // total_size if pct != last_pct: last_pct = pct diff --git a/transfer.py b/transfer.py index 9d6a540..e268b2c 100644 --- a/transfer.py +++ b/transfer.py @@ -21,7 +21,9 @@ import subprocess import sys import time from pathlib import Path -from typing import Any +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 @@ -175,25 +177,16 @@ class TransferClient: 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: - while True: - chunk = resp.read(1 << 20) - if not chunk: - break - f.write(chunk) - self._show_chunk_progress(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: - while True: - chunk = resp.read(1 << 20) - if not chunk: - break - f.write(chunk) - self._show_chunk_progress(idx, chunk_count) - sys.stdout.write("\n") - sys.stdout.flush() - return dest, enc_params + 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 @@ -220,6 +213,25 @@ class TransferClient: 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( @@ -243,19 +255,15 @@ class TransferClient: 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 # 重试耗尽, 由调用方决定 (记录断点状态待补传) - pct = done * 100 // len(pending) - sys.stdout.write( - f"\r[传输] 上传进度: {done}/{len(pending)} 卷 ({pct}%)" - ) - sys.stdout.flush() - if pending: - sys.stdout.write("\n") - sys.stdout.flush() + pbar.update(1) + pbar.set_description(f"[传输] 上传 {done}/{len(pending)} 卷") + pbar.close() receipt = self.complete(transfer_id) print(f"[传输] 完成: {receipt}") @@ -333,12 +341,16 @@ def main() -> int: 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 @@ -360,12 +372,16 @@ def main() -> int: 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: