From d967e9ed2b8cd58fc55d4f25167dedbec5b30713 Mon Sep 17 00:00:00 2001 From: ethernet Date: Fri, 4 Sep 2026 02:33:49 -0400 Subject: [PATCH] fix(local-models): restore pm.downloader router integration clobbered by upstream merge Upstream's 43e67d872f (feat: local models) reintroduced the old bespoke ranged-parallel machinery (download_file/_probe_range_support/ _DOWNLOAD_CONNECTIONS) into the local-models router and deleted the pm.downloader-based _download_job plus the pause/resume routes. Re-swap all 3 call sites (model download, runtime-install leg, browsed download) onto pm.downloader.Download via _download_job: one resumable 8-way parallel job per plan, progress via the shared tick callback, partials in the managed partials root. Restore /download/pause + /download/resume routes with the _RUNNING handle registry (dl handle kept off the JSON job dict; resume re-runs the job body). Restore the pre-clobber route tests (FakeRangeOpener stands in for pm.downloader._OPENER with honest Range support) incl. the pause / resume / finished-job-releases-handles coverage. --- hermes_cli/web_routers/local_models.py | 235 +++++++------------ tests/hermes_cli/test_local_models_routes.py | 188 +++++++++++++-- 2 files changed, 244 insertions(+), 179 deletions(-) diff --git a/hermes_cli/web_routers/local_models.py b/hermes_cli/web_routers/local_models.py index 055e80a2db..5bc3d8d1c8 100644 --- a/hermes_cli/web_routers/local_models.py +++ b/hermes_cli/web_routers/local_models.py @@ -28,6 +28,8 @@ from typing import Any, Dict, Optional from fastapi import APIRouter, HTTPException from pydantic import BaseModel +from pm.downloader import Download, DownloadPaused, Source + from hermes_cli.local_runtime.endpoint import _state_endpoint logger = logging.getLogger(__name__) @@ -37,6 +39,9 @@ router = APIRouter() _GIB = 1 << 30 _JOBS: Dict[str, Dict[str, Any]] = {} _JOBS_LOCK = threading.Lock() +# Live download handles + the resume callable per job_id. Kept OFF the job +# dict so the JSON-serializable poll payload never carries a Download object. +_RUNNING: Dict[str, Dict[str, Any]] = {} def _human_gb(n: int | float) -> str: @@ -62,39 +67,6 @@ def _job(kind: str, target: str, model_id: str | None = None) -> Dict[str, Any]: return job -# ── fast download: ranged parallel streams ─────────────────── - -# One TCP stream to a CDN rarely fills a fast line; 8 ranged connections -# writing into a preallocated file saturate consumer gigabit. -_DOWNLOAD_CONNECTIONS = 8 -_CHUNK = 4 << 20 - - -def _probe_range_support(url: str) -> int: - """Total size when the server honors Range requests, else 0. - - Auth-shaped failures raise with a plain-language message — a 401/403 - from the CDN means the repo is gated or the catalog entry names a - wrong repo, and the user deserves better than a bare status code. - """ - req = urllib.request.Request(url, headers={"Range": "bytes=0-0"}) - try: - with urllib.request.urlopen(req, timeout=60) as r: - if r.status == 206: - content_range = r.headers.get("Content-Range", "") - if "/" in content_range: - return int(content_range.rsplit("/", 1)[1]) - except urllib.error.HTTPError as exc: - if exc.code in (401, 403): - raise RuntimeError( - "The model host refused the download (gated or moved). " - "This is a catalog problem, not yours — please report it.") from exc - raise - except Exception: # noqa: BLE001 - pass - return 0 - - def _model_id_for(gguf: Path) -> str: """Variant model id for a staged file (strips split-part suffixes).""" import re @@ -121,107 +93,33 @@ def _variant_files_on_disk(model_id: str) -> "list[Path]": return files -def download_file(url: str, dest: Path, job: Dict[str, Any], - *, - base_done: int = 0, keep_totals: bool = False) -> None: - """Download url -> dest with byte progress on ``job``. +def _download_job(job: Dict[str, Any], plan) -> None: + """Run a plan of (url, dest, size) downloads as ONE resumable Download. - Ranged-parallel when the server supports it, single-stream fallback - otherwise. There is no integrity check against the CATALOG by - design: catalog sizes may lag an upstream re-upload, and a - newer file than we know about must download fine. Completeness is - checked only against what the SERVER declared for this transfer - (range-probe total / Content-Length) — self-consistent and always - current — so a dropped connection still errors instead of staging a - truncated file. Never leaves a .part behind. - - Multi-file variants: ``base_done`` offsets the progress so this file's - bytes accumulate onto the files before it, and ``keep_totals=True`` - stops the per-file size from overwriting the variant's total. + Completion and range progress feed the job dict (done_bytes / + total_bytes / ranges — the bar reads the aggregate, the detail view + reads the bitmap; both come from the same callback). Sources carry NO + hash by design: catalog sizes may lag an upstream re-upload, so + completeness is judged by the downloader against the server's declared + total, never the catalog. On pause the downloader raises + DownloadPaused; the job is marked 'paused' with its partials intact + for a later resume. """ - import shutil - import threading as _threading + dl = Download([Source(url, dest) for url, dest, _ in plan]) + _RUNNING.setdefault(job["job_id"], {})["dl"] = dl - tmp = dest.with_suffix(".part") - dest.parent.mkdir(parents=True, exist_ok=True) - file_done = [0] - progress_lock = _threading.Lock() - - def bump(n: int) -> None: - with progress_lock: - file_done[0] += n - job["done_bytes"] = base_done + file_done[0] + def tick(done: int, total: int, ranges: dict) -> None: + job["done_bytes"] = done + job["total_bytes"] = total + job["ranges"] = ranges try: - # The probe and the preallocation both take real seconds on a - # 20+ GB file — narrate them, or the pane shows a dead '— of X GB' - # until the first ranged byte lands. - job["detail"] = "Connecting" - total = _probe_range_support(url) - if total: - if not keep_totals: - job["total_bytes"] = total - # Preallocate so each worker writes at its own offset. - job["detail"] = f"Reserving {_human_gb(total)} of disk space" - with open(tmp, "wb") as f: - f.truncate(total) - job["detail"] = "" - errors: list[Exception] = [] - bounds = [(i * total // _DOWNLOAD_CONNECTIONS, - (i + 1) * total // _DOWNLOAD_CONNECTIONS - 1) - for i in range(_DOWNLOAD_CONNECTIONS)] - - def fetch_range(start: int, end: int) -> None: - try: - req = urllib.request.Request( - url, headers={"Range": f"bytes={start}-{end}"}) - with urllib.request.urlopen(req, timeout=120) as r, \ - open(tmp, "r+b") as f: - f.seek(start) - while True: - chunk = r.read(_CHUNK) - if not chunk: - break - f.write(chunk) - bump(len(chunk)) - except Exception as exc: # noqa: BLE001 - errors.append(exc) - - threads = [_threading.Thread(target=fetch_range, args=b, daemon=True, - name=f"lm-dl-{i}") - for i, b in enumerate(bounds)] - for t in threads: - t.start() - for t in threads: - t.join() - if errors: - raise errors[0] - if file_done[0] != total: - raise RuntimeError( - f"download incomplete ({file_done[0]} of {total} bytes)") - else: - # No range support: single stream, large chunks. Completeness - # is judged by the server's own Content-Length when it sent - # one — never by the catalog, which may lag a re-upload. - with urllib.request.urlopen(url, timeout=120) as r, open(tmp, "wb") as f: - length = int(r.headers.get("Content-Length") or 0) - if length and not keep_totals: - job["total_bytes"] = length - while True: - chunk = r.read(_CHUNK) - if not chunk: - break - f.write(chunk) - bump(len(chunk)) - if length and file_done[0] != length: - raise RuntimeError( - f"Download ended at {file_done[0]:,} bytes but the server " - f"said {length:,} — connection dropped? Removed; try again") - - shutil.move(str(tmp), str(dest)) - except Exception: - tmp.unlink(missing_ok=True) + dl.run(progress=tick) + except DownloadPaused: + job["status"] = "paused" raise + finally: + _RUNNING.get(job["job_id"], {}).pop("dl", None) def _models_dir() -> Path: @@ -820,17 +718,7 @@ async def local_models_download(body: ModelDownloadBody): try: job["phase"] = "downloading" job["detail"] = f"{entry.display_name} — {_human_gb(total)}" - done_before = 0 - for url, dest, size in plan: - if dest.exists(): - done_before += size - job["done_bytes"] = done_before - continue - download_file(url, dest, job, - base_done=done_before, keep_totals=True) - job["phase"] = "downloading" - done_before += size - job["done_bytes"] = done_before + _download_job(job, plan) job["phase"] = "done" job["status"] = "done" job["detail"] = f"{entry.display_name} ready" @@ -843,15 +731,57 @@ async def local_models_download(body: ModelDownloadBody): refresh_local_runtime() except Exception: # noqa: BLE001 logger.debug("post-download runtime refresh skipped", exc_info=True) + except DownloadPaused: + pass # status already "paused"; partials kept for resume except Exception as exc: # noqa: BLE001 logger.warning("model download failed: %s", exc) job["status"] = "error" job["error"] = str(exc) + finally: + if job["status"] != "paused": + _RUNNING.pop(job["job_id"], None) + _RUNNING[job["job_id"]] = {"resume": _run, "dl": None} threading.Thread(target=_run, daemon=True, name="lr-model-download").start() return {"job_id": job["job_id"], "model_id": variant.model_id} +class JobIdBody(BaseModel): + job_id: str + + +@router.post("/api/local-models/download/pause") +async def local_models_download_pause(body: JobIdBody): + """Pause a running model download. The downloader's live handle is + checked off the job dict (which stays JSON-serializable for the poll + route); when none is active (already done/paused) this is a no-op.""" + running = _RUNNING.get(body.job_id) + if running is None: + raise HTTPException(status_code=404, detail="unknown download job") + dl = running.get("dl") + if dl is None: + return {"ok": True, "paused": False} + dl.pause() + return {"ok": True, "paused": True} + + +@router.post("/api/local-models/download/resume") +async def local_models_download_resume(body: JobIdBody): + """Resume a paused model download. Completed files are skipped and + partials resume from their byte-range bitmap (the downloader owns that + — the route just re-runs the job body).""" + job = _JOBS.get(body.job_id) + running = _RUNNING.get(body.job_id) + if job is None or running is None or running.get("resume") is None: + raise HTTPException(status_code=404, detail="unknown or finished download job") + if job["status"] != "paused": + return {"ok": True, "resumed": False} + job["status"] = "running" + threading.Thread(target=running["resume"], daemon=True, + name="lm-dl-resume").start() + return {"ok": True, "resumed": True} + + @router.delete("/api/local-models/models/{model_id}") async def local_models_delete(model_id: str): """Remove a staged model: every split part plus its private assets. @@ -1017,17 +947,7 @@ async def local_models_quickstart(body: QuickstartBody): job["done_bytes"] = 0 job["total_bytes"] = total job["detail"] = f"{entry.display_name} — {_human_gb(total)}" - done_before = 0 - for url, dest, size in download_plan: - if dest.exists(): - done_before += size - job["done_bytes"] = done_before - continue - download_file(url, dest, job, - base_done=done_before, keep_totals=True) - job["phase"] = "downloading" - done_before += size - job["done_bytes"] = done_before + _download_job(job, download_plan) # Activate: same sequence as /activate's job body. from hermes_cli.config import load_config, save_config @@ -1361,16 +1281,13 @@ async def local_models_download_browsed(body: BrowsedDownloadBody): def _run(): try: job["phase"] = "downloading" + plan = [] for p in paths: url = (f"https://huggingface.co/{body.repo}" f"/resolve/main/{urllib.parse.quote(p)}") dest = _models_dir() / p.rsplit("/", 1)[-1] - if dest.exists(): - continue - download_file(url, dest, job, - base_done=int(job.get("done_bytes") or 0), - keep_totals=bool(job.get("total_bytes"))) - job["phase"] = "downloading" + plan.append((url, dest, 0)) + _download_job(job, plan) job["phase"] = "done" job["status"] = "done" job["detail"] = f"{model_id} ready" @@ -1380,10 +1297,16 @@ async def local_models_download_browsed(body: BrowsedDownloadBody): refresh_local_runtime() except Exception: # noqa: BLE001 logger.debug("post-download runtime refresh skipped", exc_info=True) + except DownloadPaused: + pass # status already "paused"; partials kept for resume except Exception as exc: # noqa: BLE001 job["status"] = "error" job["error"] = str(exc) + finally: + if job["status"] != "paused": + _RUNNING.pop(job["job_id"], None) + _RUNNING[job["job_id"]] = {"resume": _run, "dl": None} threading.Thread(target=_run, daemon=True, name="lm-download-browsed").start() return {"job_id": job["job_id"], "model_id": model_id} diff --git a/tests/hermes_cli/test_local_models_routes.py b/tests/hermes_cli/test_local_models_routes.py index 6ee98aae1e..79ebc2b400 100644 --- a/tests/hermes_cli/test_local_models_routes.py +++ b/tests/hermes_cli/test_local_models_routes.py @@ -119,6 +119,50 @@ def test_catalog_never_hides_unaffordable_models(client, monkeypatch): # ── downloads ──────────────────────────────────────────────── +class _FakeRangeOpener: + """Stands in for pm.downloader._OPENER: serves `body` with honest Range + support (the downloader probes bytes=0-0 and then fetches 8-way ranges).""" + + def __init__(self, body: bytes, content_length: int | None = None): + self._body = body + self._length = content_length if content_length is not None else len(body) + + def open(self, req, timeout=None): + parent = self + + class _Resp(io.BytesIO): + status = 200 + headers = {"Content-Length": str(parent._length)} + + def __enter__(self): + return self + + def __exit__(self, *a): + return False + + rng = req.headers.get("Range") if req.headers else None + if rng and rng.startswith("bytes="): + lo_hi = rng[len("bytes="):] + if lo_hi == "0-0": + # probe: 206 with full Content-Range so range support is detected + probe = _Resp(b"") + probe.status = 206 + probe.headers = { + "Content-Range": f"bytes 0-0/{parent._length}", + "Content-Length": "1", + } + return probe + lo, hi = (int(x) for x in lo_hi.split("-")) + part = parent._body[lo:hi + 1] + resp = _Resp(part) + resp.headers = { + "Content-Range": f"bytes {lo}-{hi}/{parent._length}", + "Content-Length": str(len(part)), + } + return resp + return _Resp(parent._body) + + def test_download_unknown_model_404s(client): r = client.post("/api/local-models/download", json={"model_id": "nope"}) assert r.status_code == 404 @@ -131,18 +175,10 @@ def test_download_short_of_server_length_errors_and_cleans_up(client, monkeypatc fewer bytes than the server promised means a dropped connection, so the job errors and nothing is staged.""" - class FakeResponse(io.BytesIO): - # Body is 17 bytes; the server promises 32 — a truncated stream. - headers = {"Content-Length": "32"} - - def __enter__(self): - return self - - def __exit__(self, *a): - return False - - monkeypatch.setattr("urllib.request.urlopen", - lambda *a, **k: FakeResponse(b"not the real body")) + # Body is 17 bytes; the server promises 32 — a truncated stream. + monkeypatch.setattr( + "pm.downloader._OPENER", + _FakeRangeOpener(b"not the real body", content_length=32)) # Pin a generous budget: variant selection prices against the machine # running the test, and a GPU-less CI runner honestly refuses every @@ -250,17 +286,8 @@ def test_download_tolerates_stale_catalog_size(client, monkeypatch): body = b"x" * 48 # server-consistent: Content-Length == body length - class FakeResponse(io.BytesIO): - headers = {"Content-Length": str(len(body))} - - def __enter__(self): - return self - - def __exit__(self, *a): - return False - - monkeypatch.setattr("urllib.request.urlopen", - lambda *a, **k: FakeResponse(body)) + monkeypatch.setattr( + "pm.downloader._OPENER", _FakeRangeOpener(body)) from hermes_cli.local_runtime.estimator import HardwareBudget @@ -291,3 +318,118 @@ def test_download_tolerates_stale_catalog_size(client, monkeypatch): break time.sleep(0.05) assert status is not None and status["status"] == "done", status.get("error") + + +def test_download_pause_and_resume_unknown_404(client): + assert client.post("/api/local-models/download/pause", + json={"job_id": "deadbeef"}).status_code == 404 + assert client.post("/api/local-models/download/resume", + json={"job_id": "deadbeef"}).status_code == 404 + + +def test_download_finished_job_releases_handles(client, monkeypatch): + """A finished download drops its live download + resume handles, so the + job dict stays JSON-serializable and nothing lingers to pause/resume.""" + body = b"x" * 48 + + monkeypatch.setattr("pm.downloader._OPENER", _FakeRangeOpener(body)) + from hermes_cli.local_runtime.estimator import HardwareBudget + + budget = HardwareBudget(usable_vram_bytes=64 << 30, + total_device_bytes=64 << 30, + ram_available_bytes=64 << 30) + monkeypatch.setattr("hermes_cli.local_runtime.hardware.probe_budget", + lambda **kw: budget) + monkeypatch.setattr( + "hermes_cli.local_runtime.bootstrap.refresh_local_runtime", + lambda: False) + from hermes_cli.local_runtime.catalog import CATALOG + + r = client.post("/api/local-models/download", json={"model_id": CATALOG[0].id}) + job_id = r.json()["job_id"] + deadline = time.time() + 10 + status = None + while time.time() < deadline: + status = client.get(f"/api/local-models/jobs/{job_id}").json() + if status["status"] in ("done", "error"): + break + time.sleep(0.05) + assert status is not None and status["status"] == "done" + + from hermes_cli.web_routers import local_models as lm + + assert job_id not in lm._RUNNING + + +def test_download_pause_reaches_paused_status(client, monkeypatch): + """A running download can be paused; the job lands in a 'paused' state + (partials preserved for resume) instead of erroring.""" + import threading + + gate = threading.Event() + + class _BlockingOpener: + """Stands in for pm.downloader._OPENER: the probe answers honestly + (1 MiB, range-supported) and every body read blocks on the gate, so + the worker is mid-download when the test pauses it.""" + + def open(self, req, timeout=None): + parent = self + + class _Resp: + status = 206 + headers = {"Content-Range": "bytes 0-0/1048576", + "Content-Length": "1"} + + def __enter__(self): + return self + + def __exit__(self, *a): + return False + + def read(self, size=-1): + gate.wait(timeout=15) # hold the worker until let go + return b"" + + rng = req.headers.get("Range") if req.headers else None + if rng and rng.startswith("bytes=") and rng != "bytes=0-0": + # body range: answer the probe shape with full length so the + # worker blocks inside read() + _Resp.headers = {"Content-Range": "bytes 0-1048575/1048576", + "Content-Length": "1048576"} + return _Resp() + + monkeypatch.setattr("pm.downloader._OPENER", _BlockingOpener()) + from hermes_cli.local_runtime.catalog import CATALOG + from hermes_cli.local_runtime.estimator import HardwareBudget + + budget = HardwareBudget(usable_vram_bytes=64 << 30, + total_device_bytes=64 << 30, + ram_available_bytes=64 << 30) + monkeypatch.setattr("hermes_cli.local_runtime.hardware.probe_budget", + lambda **kw: budget) + + r = client.post("/api/local-models/download", json={"model_id": CATALOG[0].id}) + job_id = r.json()["job_id"] + + # Pause as soon as the live download handle is in place (idempotent). + deadline = time.time() + 10 + paused = False + while time.time() < deadline: + pr = client.post("/api/local-models/download/pause", + json={"job_id": job_id}) + if pr.status_code == 200 and pr.json()["paused"] is True: + paused = True + break + time.sleep(0.05) + assert paused + gate.set() + + deadline = time.time() + 10 + status = None + while time.time() < deadline: + status = client.get(f"/api/local-models/jobs/{job_id}").json() + if status["status"] in ("paused", "done", "error"): + break + time.sleep(0.05) + assert status is not None and status["status"] == "paused"