#!/usr/bin/env python3 """Download and safely extract the published HandEdit LoRA archive.""" from __future__ import annotations import argparse import hashlib import json import shutil import zipfile from pathlib import Path from huggingface_hub import hf_hub_download DEFAULT_REPO_ID = "HandEdit/HandEdit-LoRA" ARCHIVE_NAME = "checkpoints.zip" def sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: for chunk in iter(lambda: handle.read(8 * 1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def safe_extract(archive: Path, destination: Path) -> None: root = destination.resolve() with zipfile.ZipFile(archive) as bundle: for member in bundle.infolist(): target = (root / member.filename).resolve() if target != root and root not in target.parents: raise RuntimeError(f"Unsafe archive member: {member.filename}") bundle.extractall(root) def main() -> None: parser = argparse.ArgumentParser( description="Download the sanitized HandEdit LoRA checkpoints." ) parser.add_argument("--repo-id", default=DEFAULT_REPO_ID) parser.add_argument( "--output-dir", type=Path, default=Path("."), help="Repository root where checkpoints/ will be extracted.", ) parser.add_argument( "--keep-archive", action="store_true", help="Keep checkpoints.zip after successful extraction.", ) args = parser.parse_args() output_dir = args.output_dir.expanduser().resolve() output_dir.mkdir(parents=True, exist_ok=True) archive = Path( hf_hub_download( repo_id=args.repo_id, filename=ARCHIVE_NAME, repo_type="model", local_dir=output_dir, ) ) manifest_path = Path( hf_hub_download( repo_id=args.repo_id, filename="weights_manifest.json", repo_type="model", local_dir=output_dir, ) ) manifest = json.loads(manifest_path.read_text(encoding="utf-8")) expected = manifest["archive"]["sha256"] actual = sha256(archive) if actual != expected: raise RuntimeError( f"Archive checksum mismatch: expected {expected}, received {actual}" ) safe_extract(archive, output_dir) print(f"[OK] Extracted sanitized weights to: {output_dir / 'checkpoints'}") if not args.keep_archive: archive.unlink() print(f"[OK] Removed downloaded archive: {archive}") if __name__ == "__main__": main()