nghorbani's picture
Renamed to Industrial_Foreign_Object_Detection_Walnuts: new id in the card, the pipeline READMEs, fetch.py and manifest.json
01c2b67 verified
Raw History Blame Contribute Delete
12.6 kB
#!/usr/bin/env python3
# /// script
# requires-python = ">=3.11"
# dependencies = []
# ///
"""Download every file of this repository that manifest.json lists and verify each sha256.
Standard library only and no token (the repository is public). Run it inside a checkout, where
it reads the manifest.json next to it, or from anywhere, where it fetches the manifest from the
Hub first:
python fetch.py [--out DIR] [--revision REV] [--repo-id OWNER/NAME] [--repo-type dataset|model]
python fetch.py --selftest
A file whose sha256 already matches the manifest is skipped, so a stopped run resumes with the
files it still misses. A transfer goes to <name>.part and is renamed after the hash check.
Transient errors are retried with a doubling delay. The summary lists what was downloaded,
skipped and rejected; a rejected file (hash mismatch, missing on the server) makes the script
exit 1.
Without --revision the files come from the revision manifest.json records, the commit the manifest
was written for, or from main when it records none.
"""
from __future__ import annotations
import argparse
import contextlib
import hashlib
import json
import shutil
import sys
import tempfile
import time
import urllib.error
import urllib.request
from pathlib import Path
from typing import Callable
REPO_ID = "cubert-gmbh/Industrial_Foreign_Object_Detection_Walnuts"
REPO_TYPE = "model"
HUB = "https://huggingface.co"
CHUNK = 1 << 20
ATTEMPTS = 5
FIRST_DELAY = 2.0
def resolve_base(repo_id: str, repo_type: str, revision: str) -> str:
prefix = "datasets/" if repo_type == "dataset" else ""
return f"{HUB}/{prefix}{repo_id}/resolve/{revision}"
def sha256_of(path: Path) -> str:
digest = hashlib.sha256()
with open(path, "rb") as fh:
for chunk in iter(lambda: fh.read(CHUNK), b""):
digest.update(chunk)
return digest.hexdigest()
def open_url(url: str):
request = urllib.request.Request(url, headers={"User-Agent": "cuvis-ai-fetch/1"})
return urllib.request.urlopen(request, timeout=60)
def retryable(exc: BaseException) -> bool:
if isinstance(exc, urllib.error.HTTPError):
return exc.code == 429 or exc.code >= 500
return isinstance(exc, (urllib.error.URLError, ConnectionError, TimeoutError, OSError))
def with_retries(action: Callable[[], object], sleep: Callable[[float], None] = time.sleep, log=print):
delay = FIRST_DELAY
for attempt in range(1, ATTEMPTS + 1):
try:
return action()
except Exception as exc: # noqa: BLE001 (every failure is reported, the retryable ones retried)
if attempt == ATTEMPTS or not retryable(exc):
raise
log(f" attempt {attempt} failed ({exc}); retrying in {delay:.0f} s")
sleep(delay)
delay *= 2
return None
def download(url: str, dest: Path, sha256: str, opener=open_url, sleep=time.sleep, log=print) -> str:
"""Fetch ``url`` to ``dest`` and return the sha256 of what arrived."""
dest.parent.mkdir(parents=True, exist_ok=True)
part = dest.with_name(dest.name + ".part")
def transfer() -> str:
digest = hashlib.sha256()
with opener(url) as response, open(part, "wb") as out:
for chunk in iter(lambda: response.read(CHUNK), b""):
digest.update(chunk)
out.write(chunk)
return digest.hexdigest()
got = str(with_retries(transfer, sleep=sleep, log=log))
if got == sha256:
shutil.move(str(part), str(dest))
else:
part.unlink(missing_ok=True)
return got
def fetch_all(manifest: dict, base_url: str, out: Path, opener=open_url, sleep=time.sleep, log=print) -> dict:
"""Download every manifest file under ``out``; returns the summary counts and the rejected paths."""
summary = {"downloaded": 0, "skipped": 0, "rejected": []}
files = manifest.get("files") or []
total = sum(int(f.get("size", 0)) for f in files)
log(f"{len(files)} files, {total / 1e6:.1f} MB, from {base_url}")
for entry in files:
rel = str(entry["path"])
# the rule of hfpub_common.repo_path(), repeated here because this file ships on its own
parts = rel.split("/")
if "\\" in rel or ":" in rel or rel.startswith("~") or any(part in ("", ".", "..") for part in parts):
log(f" {rel}: not a relative repository path, rejected")
summary["rejected"].append(rel)
continue
dest = out / rel
expected = str(entry["sha256"])
if dest.is_file() and dest.stat().st_size == int(entry["size"]) and sha256_of(dest) == expected:
summary["skipped"] += 1
continue
url = f"{base_url.rstrip('/')}/{rel}"
log(f" {rel} ({int(entry['size']) / 1e6:.1f} MB)")
try:
got = download(url, dest, expected, opener=opener, sleep=sleep, log=log)
except Exception as exc: # noqa: BLE001
log(f" failed: {exc}")
summary["rejected"].append(rel)
continue
if got != expected:
log(f" sha256 mismatch: expected {expected[:12]}..., got {got[:12]}...")
summary["rejected"].append(rel)
else:
summary["downloaded"] += 1
log(
f"downloaded {summary['downloaded']}, skipped {summary['skipped']} (already verified), "
f"rejected {len(summary['rejected'])}"
)
for rel in summary["rejected"]:
log(f" rejected: {rel}")
return summary
def load_manifest(args: argparse.Namespace, base_url: str) -> dict:
local = Path(__file__).resolve().parent / "manifest.json"
if args.manifest:
return json.loads(Path(args.manifest).read_text(encoding="utf-8"))
if local.is_file() and args.out.resolve() == local.parent:
return json.loads(local.read_text(encoding="utf-8"))
with open_url(f"{base_url}/manifest.json") as response:
return json.loads(response.read().decode("utf-8"))
def pinned_base(args: argparse.Namespace, manifest: dict, base_url: str, log=print) -> str:
"""The download root at the revision manifest.json records, unless the caller chose one."""
pinned = manifest.get("revision")
if args.revision or args.base_url or not pinned:
return base_url
log(f"revision {pinned} (recorded in manifest.json)")
return resolve_base(args.repo_id, args.repo_type, pinned)
def main(argv: list[str] | None = None) -> int:
reconfigure = getattr(sys.stdout, "reconfigure", None)
if reconfigure is not None:
with contextlib.suppress(Exception):
reconfigure(encoding="utf-8", errors="replace")
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--out", type=Path, default=Path("."), help="target folder (default: here)")
parser.add_argument("--repo-id", default=REPO_ID)
parser.add_argument("--repo-type", default=REPO_TYPE, choices=["dataset", "model"])
parser.add_argument(
"--revision",
default=None,
help="commit sha, branch or tag (default: the revision manifest.json records, else main)",
)
parser.add_argument("--base-url", default=None, help="override the download root (testing)")
parser.add_argument("--manifest", default=None, help="a local manifest.json to use instead of fetching it")
parser.add_argument("--selftest", action="store_true")
args = parser.parse_args(argv)
if args.selftest:
return selftest()
base_url = args.base_url or resolve_base(args.repo_id, args.repo_type, args.revision or "main")
manifest = load_manifest(args, base_url)
base_url = pinned_base(args, manifest, base_url)
summary = fetch_all(manifest, base_url, args.out)
return 1 if summary["rejected"] else 0
def selftest() -> int:
failures: list[str] = []
def expect(cond, msg: str) -> None:
if not cond:
failures.append(msg)
quiet = lambda *_a, **_k: None # noqa: E731
with tempfile.TemporaryDirectory(prefix="fetch-selftest-") as tmp:
root = Path(tmp)
remote = root / "remote"
(remote / "b").mkdir(parents=True)
(remote / "a.txt").write_bytes(b"alpha\n")
(remote / "b" / "c.bin").write_bytes(bytes(range(256)) * 3)
files = []
for rel in ("a.txt", "b/c.bin"):
p = remote / rel
files.append({"path": rel, "size": p.stat().st_size, "sha256": sha256_of(p), "role": "data"})
manifest = {"schema_version": 1, "files": files}
base = remote.as_uri()
dest = root / "dest"
s1 = fetch_all(manifest, base, dest, log=quiet)
expect(s1["downloaded"] == 2 and s1["skipped"] == 0 and not s1["rejected"], f"first run {s1}")
expect((dest / "b" / "c.bin").read_bytes() == (remote / "b" / "c.bin").read_bytes(), "bytes arrive intact")
expect(not list(dest.rglob("*.part")), "no .part files left")
s2 = fetch_all(manifest, base, dest, log=quiet)
expect(s2["downloaded"] == 0 and s2["skipped"] == 2, f"resume skips verified files {s2}")
(dest / "a.txt").write_bytes(b"damaged")
s3 = fetch_all(manifest, base, dest, log=quiet)
expect(s3["downloaded"] == 1 and s3["skipped"] == 1, f"a damaged file is fetched again {s3}")
expect((dest / "a.txt").read_bytes() == b"alpha\n", "damaged file replaced")
(remote / "a.txt").write_bytes(b"server changed\n")
(dest / "a.txt").unlink()
s4 = fetch_all(manifest, base, dest, log=quiet)
expect(s4["rejected"] == ["a.txt"], f"hash mismatch rejected {s4}")
expect(not (dest / "a.txt").exists() and not (dest / "a.txt.part").exists(), "mismatch leaves no file")
manifest_missing = {"files": [{"path": "missing.txt", "size": 1, "sha256": "0" * 64, "role": "data"}]}
s5 = fetch_all(manifest_missing, base, dest, sleep=quiet, log=quiet)
expect(s5["rejected"] == ["missing.txt"], f"missing on the server is rejected {s5}")
escapes = ["../escape.txt", "a/name:stream", "~/x.txt"]
manifest_escape = {"files": [{"path": p, "size": 1, "sha256": "0" * 64, "role": "data"} for p in escapes]}
s6 = fetch_all(manifest_escape, base, dest, sleep=quiet, log=quiet)
expect(s6["rejected"] == escapes, f"paths leaving the target are rejected {s6}")
expect(not (root / "escape.txt").exists() and not list(dest.rglob("name*")), "nothing written for them")
ns = argparse.Namespace(revision=None, base_url=None, repo_id="o/n", repo_type="model")
pinned = pinned_base(ns, {"revision": "a" * 40}, "x", log=quiet)
expect(pinned.endswith("/resolve/" + "a" * 40), f"manifest revision pins the root {pinned}")
no_pin = pinned_base(ns, {"revision": None}, "x", log=quiet)
expect(no_pin == "x", "no recorded revision keeps the root")
ns.revision = "main"
explicit = pinned_base(ns, {"revision": "a" * 40}, "x", log=quiet)
expect(explicit == "x", "an explicit --revision wins")
calls = {"n": 0}
sleeps: list[float] = []
def flaky(url: str):
calls["n"] += 1
if calls["n"] < 3:
raise urllib.error.URLError("connection reset")
return open_url(url)
(remote / "a.txt").write_bytes(b"alpha\n")
got = download(f"{base}/a.txt", dest / "a.txt", files[0]["sha256"], opener=flaky, sleep=sleeps.append, log=quiet)
expect(got == files[0]["sha256"] and calls["n"] == 3, "retries after transient errors")
expect(sleeps == [2.0, 4.0], f"doubling delay {sleeps}")
def not_found(url: str):
raise urllib.error.HTTPError(url, 404, "not found", {}, None) # type: ignore[arg-type]
try:
download(f"{base}/a.txt", dest / "x.txt", "0" * 64, opener=not_found, sleep=sleeps.append, log=quiet)
expect(False, "404 raised")
except urllib.error.HTTPError:
expect(len(sleeps) == 2, "404 is not retried")
expect(resolve_base("o/n", "dataset", "main") == "https://huggingface.co/datasets/o/n/resolve/main", "dataset url")
expect(resolve_base("o/n", "model", "abc") == "https://huggingface.co/o/n/resolve/abc", "model url")
if failures:
print("selftest FAILED:")
for f in failures:
print(f" - {f}")
return 1
print("selftest OK: sha256 verification, resume, mismatch, retry, path guard, manifest revision")
return 0
if __name__ == "__main__":
raise SystemExit(main())