diff --git a/hermes_cli/local_runtime/binaries.py b/hermes_cli/local_runtime/binaries.py index d18985646c..c304f06ffa 100644 --- a/hermes_cli/local_runtime/binaries.py +++ b/hermes_cli/local_runtime/binaries.py @@ -9,6 +9,7 @@ import os import platform import shutil import subprocess +import tempfile import urllib.request import zipfile from dataclasses import dataclass, field @@ -185,19 +186,29 @@ def _download(url: str, dest: Path, """Stream url -> dest. ``progress(done_bytes, total_bytes)`` ticks per chunk (total 0 when the server sends no Content-Length) — a several-hundred-MB archive must never look hung.""" logger.info("downloading %s", url) - tmp = dest.with_suffix(dest.suffix + ".part") - with urllib.request.urlopen(url, timeout=120) as r, open(tmp, "wb") as f: - total = int(r.headers.get("Content-Length") or 0) - done = 0 - while True: - chunk = r.read(1 << 20) - if not chunk: - break - f.write(chunk) - done += len(chunk) - if progress is not None: - progress(done, total) - tmp.replace(dest) + staging = tempfile.NamedTemporaryFile( + mode="wb", dir=dest.parent, prefix=f"{dest.name}.", suffix=".part", delete=False) + tmp = Path(staging.name) + try: + with staging as f, urllib.request.urlopen(url, timeout=120) as r: + length = r.headers.get("Content-Length") + total = int(length) if length is not None else 0 + done = 0 + while True: + chunk = r.read(1 << 20) + if not chunk: + break + f.write(chunk) + done += len(chunk) + if progress is not None: + progress(done, total) + # Chunked reads can return EOF without raising IncompleteRead. + if length is not None and done != total: + raise BinaryResolutionError( + f"incomplete download for {dest.name}: expected {total} bytes, got {done}") + tmp.replace(dest) + finally: + tmp.unlink(missing_ok=True) def _extract(archive: Path, dest: Path, diff --git a/tests/hermes_cli/test_local_runtime_downloads.py b/tests/hermes_cli/test_local_runtime_downloads.py new file mode 100644 index 0000000000..c0ac8558e4 --- /dev/null +++ b/tests/hermes_cli/test_local_runtime_downloads.py @@ -0,0 +1,96 @@ +"""Runtime downloads must not turn a transient transfer failure into a poisoned cache.""" + +import threading +from concurrent.futures import ThreadPoolExecutor +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from urllib.error import HTTPError + +import pytest + +from hermes_cli.local_runtime import binaries + + +@pytest.fixture +def asset_server(): + responses = [] + requests = [] + + class Handler(BaseHTTPRequestHandler): + def do_GET(self): + requests.append(self.path) + body, length, status = responses.pop(0) + self.send_response(status) + if length is not None: + self.send_header("Content-Length", str(length)) + self.end_headers() + self.wfile.write(body) + self.close_connection = True + + def log_message(self, *args): + pass + + with ThreadingHTTPServer(("127.0.0.1", 0), Handler) as server: + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield f"http://127.0.0.1:{server.server_port}/{{asset}}", responses, requests + finally: + server.shutdown() + thread.join(timeout=5) + + +@pytest.mark.parametrize("failure", ["short", "callback"]) +@pytest.mark.parametrize("known_length", [True, False]) +def test_download_only_publishes_complete_transfers(tmp_path, asset_server, failure, known_length): + url, responses, requests = asset_server + dest = tmp_path / "runtime.zip" + payload = b"runtime archive bytes" + responses.append((payload[:-1], len(payload), 200)) + + def progress(done, total): + if failure == "callback": + raise RuntimeError("cancelled") + + error = RuntimeError if failure == "callback" else binaries.BinaryResolutionError + with pytest.raises(error): + binaries._download(url.format(asset=dest.name), dest, progress=progress) + assert not dest.exists() + assert not list(tmp_path.glob("*.part")) + + # A retry works, including servers that omit Content-Length. + responses.append((payload, len(payload) if known_length else None, 200)) + ticks = [] + binaries._download(url.format(asset=dest.name), dest, + progress=lambda done, total: ticks.append((done, total))) + assert dest.read_bytes() == payload + assert ticks[-1] == (len(payload), len(payload) if known_length else 0) + assert len(requests) == 2 + + +def test_failed_request_preserves_another_active_download(tmp_path, asset_server): + url, responses, requests = asset_server + dest = tmp_path / "runtime.zip" + payload = b"x" * (2 << 20) + responses.extend([(payload, len(payload), 200), (b"", 0, 503)]) + started = threading.Event() + resume = threading.Event() + + def pause_download(done, total): + started.set() + assert resume.wait(10), "active download was not resumed" + + with ThreadPoolExecutor(max_workers=1) as pool: + active = pool.submit(binaries._download, url.format(asset=dest.name), dest, + progress=pause_download) + try: + assert started.wait(10), "active download did not write its first chunk" + with pytest.raises(HTTPError) as failed: + binaries._download(url.format(asset=dest.name), dest) + assert failed.value.code == 503 + finally: + resume.set() + active.result(timeout=10) + + assert dest.read_bytes() == payload + assert not list(tmp_path.glob("*.part")) + assert len(requests) == 2