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.
This commit is contained in:
ethernet
2026-09-04 02:33:49 -04:00
parent 5890f3bbe3
commit d967e9ed2b
2 changed files with 244 additions and 179 deletions

View File

@@ -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}

View File

@@ -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"