From cade33cf7a7787019c3d9035577d53a83271e025 Mon Sep 17 00:00:00 2001 From: ethernet Date: Sat, 12 Sep 2026 10:32:06 -0400 Subject: [PATCH] perf(ci): archive all unique pinned inputs concurrently --- scripts/ci/archive_inputs.py | 19 ++++++-- tests/scripts/test_archive_inputs.py | 71 ++++++++++++++++++++++++++++ 2 files changed, 85 insertions(+), 5 deletions(-) diff --git a/scripts/ci/archive_inputs.py b/scripts/ci/archive_inputs.py index 2578f34848..ac4cf1c68b 100644 --- a/scripts/ci/archive_inputs.py +++ b/scripts/ci/archive_inputs.py @@ -2,6 +2,7 @@ from __future__ import annotations import argparse +from concurrent.futures import ThreadPoolExecutor, as_completed from dataclasses import dataclass import json from pathlib import Path, PurePosixPath @@ -148,15 +149,19 @@ def _download_upstream(pin: InputPin, local: Path) -> str: def stage_inputs(pins: list[InputPin], *, archive: Archive, store: Store | None = None, payload: Path | None = None) -> int: - """Archive each digest once; optionally seed the actual stagers' inputs.""" + """Archive all unique digests concurrently; optionally seed the stagers' inputs.""" from scripts.termux.stage_runtime_libs import download_path groups: dict[str, list[InputPin]] = {} for pin in pins: groups.setdefault(pin.sha256, []).append(pin) - with tempfile.TemporaryDirectory(prefix="hermes-inputs-") as temporary: - local = Path(temporary) / "input" - for digest, references in groups.items(): + if not groups: + return 0 + + def stage_digest(digest: str, references: list[InputPin]) -> None: + # Each worker owns its downloads and cleanup, including on failure. + with tempfile.TemporaryDirectory(prefix="hermes-inputs-") as temporary: + local = Path(temporary) / "input" origin = archive.fetch(references[0], local) for pin in references: if pin.kind != "tool" and payload is not None: @@ -176,8 +181,12 @@ def stage_inputs(pins: list[InputPin], *, archive: Archive, store: Store | None if entry.exists(): shutil.rmtree(entry) store.publish(staged, entry.name) - local.unlink() print(f" {references[0].name}: {origin} -> {object_key(digest)}", flush=True) + + with ThreadPoolExecutor(max_workers=len(groups)) as pool: + futures = [pool.submit(stage_digest, digest, references) for digest, references in groups.items()] + for future in as_completed(futures): + future.result() return len(groups) diff --git a/tests/scripts/test_archive_inputs.py b/tests/scripts/test_archive_inputs.py index 4ff73eb636..bbe8ed5e74 100644 --- a/tests/scripts/test_archive_inputs.py +++ b/tests/scripts/test_archive_inputs.py @@ -155,6 +155,77 @@ def test_racing_misses_verify_the_immutable_winner(tmp_path, upstream, r2_server assert all(p.read_bytes() == body for p in paths) +def test_all_digests_start_together_and_seed_every_reference(tmp_path, upstream, r2_server, monkeypatch): + from collections import Counter + from pm.store import Store + from scripts.termux.stage_runtime_libs import download_path + + server, root = upstream + bodies = {f"input-{i}": f"distinct pinned bytes {i}".encode() for i in range(9)} + pins = [] + for name, body in bodies.items(): + (root / f"{name}.deb").write_bytes(body) + digest = hashlib.sha256(body).hexdigest() + for kind, label in (("tool", name), ("library", name), ("library", f"{name}-alias")): + pins.append(inputs.InputPin(label, f"{server.url}/{name}.deb", digest, kind)) + + # Every unique input must reach the real HTTP path before any can finish. + barrier = threading.Barrier(len(bodies)) + original = r2.signed_request + def together(method, url, **kwargs): + if method == "HEAD": + barrier.wait(timeout=10) + return original(method, url, **kwargs) + monkeypatch.setattr(r2, "signed_request", together) + store = Store(tmp_path / "tools") + payload = tmp_path / "payload" + assert inputs.stage_inputs(pins, archive=inputs.Archive(*r2.credentials()), + store=store, payload=payload) == len(bodies) + puts = Counter(path for method, path, _ in r2_server.requests if method == "PUT") + assert len(puts) == len(bodies) and set(puts.values()) == {1} + for name, body in bodies.items(): + digest = hashlib.sha256(body).hexdigest() + assert r2_server.store[object_key(digest)][0] == body + assert (store.entry(f"fetch-{digest}") / f"{name}.deb").read_bytes() == body + for label in (name, f"{name}-alias"): + assert download_path(payload, label).read_bytes() == body + + +def test_parallel_readback_failure_reaches_cli_and_preserves_destination(tmp_path, upstream, r2_server, monkeypatch, capsys): + from pm import paths + from scripts.termux.stage_runtime_libs import download_path + + server, root = upstream + packages, libs = {}, {} + for name in ("good", "bad"): + body = name.encode() + (root / f"{name}.deb").write_bytes(body) + row = {"url": f"{server.url}/{name}.deb", "sha256": hashlib.sha256(body).hexdigest()} + libs[name] = row + packages[name] = {"version": "1", "artifacts": {"any": row}} + r2_server.store[object_key(row["sha256"])] = (body if name == "good" else b"xxx", '"etag"') + repo = tmp_path / "repo" + write_pins(repo, packages, {"libs": libs}) + monkeypatch.setattr(paths, "repo_root", lambda: repo) + payload = tmp_path / "payload" + preserved = download_path(payload, "bad") + preserved.parent.mkdir(parents=True) + preserved.write_bytes(b"existing destination") + barrier = threading.Barrier(len(libs)) + original = r2.signed_request + def together(method, url, **kwargs): + if method == "HEAD": + barrier.wait(timeout=10) + return original(method, url, **kwargs) + monkeypatch.setattr(r2, "signed_request", together) + with pytest.raises(ValueError, match="checksum mismatch"): + inputs.main(["--payload", str(payload), "--store", str(tmp_path / "tools")]) + assert preserved.read_bytes() == b"existing destination" + assert not (tmp_path / "tools" / f"fetch-{libs['bad']['sha256']}").exists() + assert download_path(payload, "good").read_bytes() == b"good" + assert "Verified " not in capsys.readouterr().out + + def test_historical_recovery_keeps_the_original_digest(tmp_path, upstream, r2_server, monkeypatch): server, root = upstream body = (root / "lib.deb").read_bytes()