From 54737be9dbf6a25e8ea5063251a5741693db7caf Mon Sep 17 00:00:00 2001 From: Eric Curtin Date: Sun, 30 Aug 2026 23:18:31 +0100 Subject: [PATCH 1/2] feat: resolve oci:// model references via llmman serve Lets --model-id point at a model published as a CNCF ModelPack OCI artifact: infinity_emb v2 --model-id oci://ghcr.io/org/model:tag Model distribution is increasingly moving to OCI registries, which lets a deployment reuse the registry, credentials, mirroring and air-gap tooling it already has for container images. Acquisition is delegated to a running `llmman serve`, which already implements the ModelPack media types, registry auth, resumable blob download and a content-addressed store. The daemon does the pull (POST /api/pull, streamed as NDJSON so a multi-gigabyte fetch is not silent, and an error arriving in-band at HTTP 200 is caught) but deliberately exposes no local path, so `llmman resolve --no-pull` reports where the bytes landed. The client is stdlib-only, so no new dependency. EngineArgs.__post_init__ is the single dispatch point, resolved first so the rest of that method and the loading strategy only ever see a local path -- every engine then loads it exactly as it would a local directory. served_model_name keeps the reference the user typed rather than the store path, unless one was given explicitly. An explicit oci:// scheme is required rather than sniffing a bare registry/name:tag: that shape is indistinguishable from a HuggingFace repo id, so guessing would silently hijack existing deployments. Signed-off-by: Eric Curtin --- libs/infinity_emb/infinity_emb/args.py | 14 ++ libs/infinity_emb/infinity_emb/llmman.py | 222 ++++++++++++++++++ libs/infinity_emb/infinity_emb/oci.py | 56 +++++ .../tests/unit_test/test_llmman.py | 117 +++++++++ libs/infinity_emb/tests/unit_test/test_oci.py | 105 +++++++++ 5 files changed, 514 insertions(+) create mode 100644 libs/infinity_emb/infinity_emb/llmman.py create mode 100644 libs/infinity_emb/infinity_emb/oci.py create mode 100644 libs/infinity_emb/tests/unit_test/test_llmman.py create mode 100644 libs/infinity_emb/tests/unit_test/test_oci.py diff --git a/libs/infinity_emb/infinity_emb/args.py b/libs/infinity_emb/infinity_emb/args.py index fde57081d..27481419b 100644 --- a/libs/infinity_emb/infinity_emb/args.py +++ b/libs/infinity_emb/infinity_emb/args.py @@ -8,6 +8,7 @@ from copy import deepcopy +from infinity_emb import oci from infinity_emb._optional_imports import CHECK_PYDANTIC from infinity_emb.env import MANAGER from infinity_emb.primitives import ( @@ -74,6 +75,19 @@ class EngineArgs: _loading_strategy: Optional[LoadingStrategy] = None def __post_init__(self): + # A CNCF ModelPack artifact is pulled through an llmman daemon and + # extracted to a local directory, which every engine then loads exactly + # as it would a local path. Done first so the rest of this method, and + # the loading strategy, only ever see a local path. + if oci.is_oci_ref(self.model_name_or_path): + if not self.served_model_name: + # Keep the reference the user typed as the served name; the + # resolved path is an implementation detail of the store. + object.__setattr__(self, "served_model_name", self.model_name_or_path) + object.__setattr__( + self, "model_name_or_path", oci.resolve(self.model_name_or_path) + ) + # convert the following strings to enums # so they don't need to be exported to the external interface if not isinstance(self.engine, InferenceEngine): diff --git a/libs/infinity_emb/infinity_emb/llmman.py b/libs/infinity_emb/infinity_emb/llmman.py new file mode 100644 index 000000000..e18359197 --- /dev/null +++ b/libs/infinity_emb/infinity_emb/llmman.py @@ -0,0 +1,222 @@ +"""Client for a running ``llmman serve`` daemon. + +Used to acquire models published as CNCF ModelPack +(https://github.com/modelpack/model-spec) OCI artifacts. The daemon owns the +registry work -- ModelPack media types, registry auth, resumable blob download +and a content-addressed store -- so it is not reimplemented here. + +Contract (from llmman's src/cmd/serve.rs and src/daemon.rs): + - LLMMAN_HOST is ``[scheme://]host[:port][/path]``, default 127.0.0.1:17434. + A wildcard bind host (0.0.0.0, ::) is rewritten to loopback, since a client + cannot connect to "every interface". + - ``GET /api/version`` -> ``{"version":..., "exe":..., "pid":...}``. + - ``POST /api/pull`` ``{"model": ref}`` -> NDJSON stream of ``{"status":...}`` + objects, terminated by ``{"status":"success"}`` or ``{"error":"..."}``. + An error can arrive in-band at HTTP 200. + - ``llmman resolve --no-pull `` -> one line of JSON carrying ``path``. +""" + +import ipaddress +import json +import logging +import os +import shutil +import subprocess +import urllib.error +import urllib.request + +logger = logging.getLogger(__name__) + +HOST_ENV = "LLMMAN_HOST" +BIN_ENV = "INFINITY_LLMMAN_BIN" + +DEFAULT_HOST = "127.0.0.1" +DEFAULT_PORT = 17434 + +PROBE_TIMEOUT_SECONDS = 5 + + +def _connectable_host(host: str) -> str: + """Rewrite a wildcard bind host to its loopback equivalent.""" + try: + ip = ipaddress.ip_address(host.strip("[]")) + except ValueError: + return host + if not ip.is_unspecified: + return host + return "127.0.0.1" if ip.version == 4 else "::1" + + +def endpoint() -> str: + """The http origin of the llmman daemon, honouring LLMMAN_HOST.""" + raw = os.getenv(HOST_ENV, "").strip().strip("\"'") + if not raw: + return f"http://{DEFAULT_HOST}:{DEFAULT_PORT}" + + if "://" in raw: + raw = raw.split("://", 1)[1] + raw = raw.split("/", 1)[0] + + host, port = raw, DEFAULT_PORT + if raw.startswith("["): # bracketed IPv6, optionally with :port + close = raw.find("]") + if close != -1: + host = raw[: close + 1] + rest = raw[close + 1 :] + if rest.startswith(":") and rest[1:].isdigit(): + port = int(rest[1:]) + elif raw.count(":") == 1: + maybe_host, maybe_port = raw.rsplit(":", 1) + if maybe_port.isdigit(): + host, port = maybe_host, int(maybe_port) + + host = host or DEFAULT_HOST + resolved = _connectable_host(host) + if ":" in resolved and not resolved.startswith("["): + resolved = f"[{resolved}]" + return f"http://{resolved}:{port}" + + +def llmman_bin() -> str: + """The llmman executable name, overridable per project.""" + return os.getenv(BIN_ENV, "").strip() or "llmman" + + +def check_daemon(base: str) -> None: + """Confirm an llmman daemon is listening and is actually llmman.""" + url = base + "/api/version" + try: + with urllib.request.urlopen(url, timeout=PROBE_TIMEOUT_SECONDS) as resp: + if resp.status != 200: + raise RuntimeError( + f"llmman daemon at {base} answered /api/version with HTTP {resp.status}" + ) + payload = json.loads(resp.read().decode("utf-8")) + except urllib.error.URLError as exc: + raise RuntimeError( + f"no llmman daemon reachable at {base} ({exc.reason}). Start one with " + f"`llmman serve`, or point {HOST_ENV} at an existing daemon." + ) from exc + except json.JSONDecodeError as exc: + raise RuntimeError( + f"the server at {base} is not an llmman daemon (unparseable /api/version)" + ) from exc + + if not isinstance(payload, dict) or not payload.get("version"): + raise RuntimeError( + f"the server at {base} is not an llmman daemon (no version in /api/version)" + ) + + +def pull(base: str, reference: str, progress=None) -> None: + """Stream POST /api/pull until the daemon reports success. + + ``progress`` receives ``(status, completed, total)``. An error can arrive + in-band at HTTP 200, and a stream that ends without ``success`` is also a + failure -- neither is treated as a completed pull. + """ + body = json.dumps({"model": reference}).encode("utf-8") + req = urllib.request.Request( + base + "/api/pull", + data=body, + headers={"Content-Type": "application/json"}, + method="POST", + ) + + succeeded = False + try: + with urllib.request.urlopen(req) as resp: + if resp.status != 200: + raise RuntimeError(f"llmman pull of {reference!r} failed: HTTP {resp.status}") + for raw_line in resp: + line = raw_line.decode("utf-8").strip() + if not line: + continue + try: + obj = json.loads(line) + except json.JSONDecodeError: + # Tolerate a non-JSON diagnostic rather than aborting a + # pull that may still be progressing. + continue + if not isinstance(obj, dict): + continue + if obj.get("error"): + raise RuntimeError(f"llmman pull of {reference!r} failed: {obj['error']}") + status = obj.get("status") + if status == "success": + succeeded = True + continue + if progress is not None and status: + progress(status, obj.get("completed", 0), obj.get("total", 0)) + except urllib.error.HTTPError as exc: + raise RuntimeError(f"llmman pull of {reference!r} failed: HTTP {exc.code}") from exc + except urllib.error.URLError as exc: + raise RuntimeError(f"llmman pull of {reference!r} failed: {exc.reason}") from exc + + if not succeeded: + raise RuntimeError(f"llmman pull of {reference!r} ended without reporting success") + + +def parse_resolve_output(stdout: str, reference: str) -> str: + """Parse ``llmman resolve`` stdout into the resolved local path.""" + lines = [line.strip() for line in stdout.splitlines() if line.strip()] + if not lines: + raise RuntimeError(f"llmman resolve {reference!r}: no output on stdout") + + try: + payload = json.loads(lines[-1]) + except json.JSONDecodeError as exc: + raise RuntimeError( + f"llmman resolve {reference!r}: could not parse output as JSON: {lines[-1]}" + ) from exc + + if not isinstance(payload, dict): + # A protocol violation rather than a caller type error, so RuntimeError + # keeps every llmman failure one exception type for callers. + msg = f"llmman resolve {reference!r}: expected a JSON object, got {lines[-1]}" + raise RuntimeError(msg) # noqa: TRY004 + + path = payload.get("path") + if not isinstance(path, str) or not path.strip(): + raise RuntimeError(f"llmman resolve {reference!r}: returned an empty path") + if not os.path.exists(path): + raise RuntimeError(f"llmman resolve {reference!r}: reported path {path!r} does not exist") + return path + + +def resolve(reference: str) -> str: + """Ask the CLI where the daemon's pull left the model on disk. + + ``--no-pull`` guarantees this only reports on bytes ``/api/pull`` already + fetched, so the daemon stays the only thing that touches the network. + """ + binary = llmman_bin() + if shutil.which(binary) is None and not os.path.isfile(binary): + raise RuntimeError( + f"{binary!r} not found. Install llmman " + "(https://github.com/llmmanorg/llmman) and put it on PATH, or set " + f"{BIN_ENV} to its location." + ) + + completed = subprocess.run( + [binary, "resolve", "--no-pull", reference], + capture_output=True, + stdin=subprocess.DEVNULL, + text=True, + check=False, + ) + if completed.returncode != 0: + raise RuntimeError( + f"`{binary} resolve --no-pull {reference}` failed with exit code " + f"{completed.returncode}: {completed.stderr.strip()}" + ) + return parse_resolve_output(completed.stdout, reference) + + +def pull_and_resolve(reference: str, progress=None) -> str: + """Full acquisition: probe the daemon, pull through it, report the path.""" + base = endpoint() + check_daemon(base) + logger.info("Pulling %s via llmman daemon at %s", reference, base) + pull(base, reference, progress) + return resolve(reference) diff --git a/libs/infinity_emb/infinity_emb/oci.py b/libs/infinity_emb/infinity_emb/oci.py new file mode 100644 index 000000000..dd5193302 --- /dev/null +++ b/libs/infinity_emb/infinity_emb/oci.py @@ -0,0 +1,56 @@ +"""Resolve ``oci://`` model references to a local path. + +A model published as a CNCF ModelPack (https://github.com/modelpack/model-spec) +artifact lives in an ordinary container registry, so it reuses the registry, +credentials, mirroring and air-gap tooling a deployment already has for +container images. + +Acquisition is delegated to a running ``llmman serve`` +(https://github.com/llmmanorg/llmman), which already implements the ModelPack +media types, registry auth, resumable blob download and a content-addressed +store. The daemon does the pull (POST /api/pull, streamed so a multi-gigabyte +fetch is not silent) but deliberately exposes no local path, so +``llmman resolve --no-pull`` reports where the bytes landed. + +An explicit ``oci://`` scheme is required rather than sniffing a bare +``registry/name:tag``: that shape is indistinguishable from a HuggingFace repo +id (``org/model``), so guessing would silently hijack existing deployments. +""" + +import logging + +from infinity_emb import llmman + +logger = logging.getLogger("infinity_emb") + +SCHEME = "oci://" + + +def is_oci_ref(model_name_or_path) -> bool: + """Whether the reference carries the ``oci://`` scheme.""" + if not model_name_or_path: + return False + return str(model_name_or_path).lower().startswith(SCHEME) + + +def strip_scheme(model_name_or_path) -> str: + """Drop the ``oci://`` prefix, leaving the bare registry reference.""" + text = str(model_name_or_path) + if is_oci_ref(text): + return text[len(SCHEME) :] + return text + + +def resolve(model_name_or_path) -> str: + """Pull an ``oci://`` reference through llmman and return the local path.""" + reference = strip_scheme(model_name_or_path).strip() + if not reference: + raise ValueError(f"empty OCI model reference: {model_name_or_path!r}") + + def _progress(status, completed, total): + if total: + logger.info("llmman: %s (%s/%s bytes)", status, completed, total) + else: + logger.info("llmman: %s", status) + + return llmman.pull_and_resolve(reference, progress=_progress) diff --git a/libs/infinity_emb/tests/unit_test/test_llmman.py b/libs/infinity_emb/tests/unit_test/test_llmman.py new file mode 100644 index 000000000..189a325a9 --- /dev/null +++ b/libs/infinity_emb/tests/unit_test/test_llmman.py @@ -0,0 +1,117 @@ +"""The `llmman serve` client: the daemon protocol behind oci:// model paths. + +Exercised against a real HTTP server on a loopback port rather than mocks, so +the NDJSON streaming contract is genuinely tested. +""" + +import http.server +import json +import socketserver +import threading + +import pytest + +from infinity_emb import llmman + + +def _ndjson(*objs): + return "".join(json.dumps(o) + "\n" for o in objs) + + +class _FakeDaemon: + """A minimal stand-in for `llmman serve`, on a real loopback port.""" + + def __init__(self): + self.version = {"version": "0.1.0", "pid": 1} + self.pull_body = _ndjson({"status": "success"}) + self.pull_status = 200 + self.last_request = None + daemon = self + + class Handler(http.server.BaseHTTPRequestHandler): + def log_message(self, *args): + pass + + def _send(self, status, body, ctype): + raw = body.encode() + self.send_response(status) + self.send_header("Content-Type", ctype) + self.send_header("Content-Length", str(len(raw))) + self.end_headers() + self.wfile.write(raw) + + def do_GET(self): + self._send(200, json.dumps(daemon.version), "application/json") + + def do_POST(self): + length = int(self.headers.get("Content-Length", 0)) + daemon.last_request = json.loads(self.rfile.read(length)) + self._send(daemon.pull_status, daemon.pull_body, "application/x-ndjson") + + self._server = socketserver.TCPServer(("127.0.0.1", 0), Handler) + self.url = f"http://127.0.0.1:{self._server.server_address[1]}" + threading.Thread(target=self._server.serve_forever, daemon=True).start() + + def close(self): + self._server.shutdown() + self._server.server_close() + + +@pytest.fixture +def daemon(): + d = _FakeDaemon() + yield d + d.close() + + +def test_accepts_a_llmman_daemon(daemon): + llmman.check_daemon(daemon.url) + + +def test_rejects_a_non_llmman_server(daemon): + daemon.version = {"hello": "world"} + with pytest.raises(RuntimeError, match="not an llmman daemon"): + llmman.check_daemon(daemon.url) + + +def test_reports_nothing_listening_actionably(): + with pytest.raises(RuntimeError, match="llmman serve"): + llmman.check_daemon("http://127.0.0.1:1") + + +def test_pull_succeeds_and_forwards_progress(daemon): + daemon.pull_body = _ndjson( + {"status": "pulling manifest"}, + {"status": "pulling blobs", "completed": 50, "total": 100}, + {"status": "success"}, + ) + seen = [] + llmman.pull(daemon.url, "ghcr.io/org/model:tag", lambda *a: seen.append(a)) + + assert daemon.last_request == {"model": "ghcr.io/org/model:tag"} + assert seen == [("pulling manifest", 0, 0), ("pulling blobs", 50, 100)] + + +def test_reports_an_in_band_error_at_http_200(daemon): + # The daemon streams errors in-band, so a 200 does not mean success. + daemon.pull_body = _ndjson({"status": "pulling"}, {"error": "unauthorized"}) + with pytest.raises(RuntimeError, match="unauthorized"): + llmman.pull(daemon.url, "ref") + + +def test_rejects_a_stream_that_ends_without_success(daemon): + daemon.pull_body = _ndjson({"status": "pulling blobs"}) + with pytest.raises(RuntimeError, match="without reporting success"): + llmman.pull(daemon.url, "ref") + + +def test_reports_a_non_ok_status(daemon): + daemon.pull_status = 400 + daemon.pull_body = '{"error":"bad request"}' + with pytest.raises(RuntimeError): + llmman.pull(daemon.url, "ref") + + +def test_tolerates_a_non_json_diagnostic_line(daemon): + daemon.pull_body = "not json\n" + _ndjson({"status": "success"}) + llmman.pull(daemon.url, "ref") diff --git a/libs/infinity_emb/tests/unit_test/test_oci.py b/libs/infinity_emb/tests/unit_test/test_oci.py new file mode 100644 index 000000000..77483b1fc --- /dev/null +++ b/libs/infinity_emb/tests/unit_test/test_oci.py @@ -0,0 +1,105 @@ +"""`oci://` model references resolve to a local path. + +The scheme is explicit on purpose: a bare `registry/name:tag` is the same shape +as a HuggingFace repo id, so sniffing would hijack existing deployments. +""" + +import os +from unittest import mock + +import pytest + +from infinity_emb import llmman +from infinity_emb.oci import is_oci_ref, resolve, strip_scheme + + +def test_recognizes_the_oci_scheme(): + assert is_oci_ref("oci://ghcr.io/org/model:tag") + assert is_oci_ref("OCI://ghcr.io/org/model:tag") + + +@pytest.mark.parametrize( + "value", + [ + "michaelfeil/bge-small-en-v1.5", + "ghcr.io/org/model:tag", + "/local/path/to/model", + "s3://bucket/key", + "", + None, + ], +) +def test_leaves_every_other_shape_alone(value): + # A bare HF repo id must never be claimed. + assert not is_oci_ref(value) + + +def test_strips_the_scheme_only_when_present(): + assert strip_scheme("oci://ghcr.io/org/model:tag") == "ghcr.io/org/model:tag" + assert strip_scheme("OCI://ghcr.io/org/model:tag") == "ghcr.io/org/model:tag" + assert strip_scheme("michaelfeil/bge") == "michaelfeil/bge" + + +@pytest.mark.parametrize("ref", ["oci://", "oci:// "]) +def test_rejects_an_empty_reference(ref): + with pytest.raises(ValueError): + resolve(ref) + + +def test_hands_the_bare_reference_to_the_daemon(): + with mock.patch( + "infinity_emb.oci.llmman.pull_and_resolve", return_value="/resolved" + ) as acquire: + assert resolve("oci://ghcr.io/org/model:tag") == "/resolved" + assert acquire.call_args[0][0] == "ghcr.io/org/model:tag" + assert acquire.call_args[1]["progress"] is not None + + +@pytest.mark.parametrize( + "host,want", + [ + ("", "http://127.0.0.1:17434"), + ("1.2.3.4:9999", "http://1.2.3.4:9999"), + ("1.2.3.4", "http://1.2.3.4:17434"), + # A wildcard bind is meaningful to the server but not to a client. + ("0.0.0.0:9999", "http://127.0.0.1:9999"), + ("[::]:9999", "http://[::1]:9999"), + ], +) +def test_endpoint_parsing(host, want): + with mock.patch.dict(os.environ, {llmman.HOST_ENV: host}): + assert llmman.endpoint() == want + + +class TestEngineArgsIntegration: + """EngineArgs resolves the reference before anything else reads it.""" + + def test_rewrites_model_name_or_path_and_keeps_the_served_name(self): + from infinity_emb.args import EngineArgs + + with mock.patch("infinity_emb.oci.resolve", return_value="/resolved"): + args = EngineArgs(model_name_or_path="oci://ghcr.io/org/model:tag") + + assert args.model_name_or_path == "/resolved" + # The served name stays the reference the user typed, not the store path. + assert args.served_model_name == "oci://ghcr.io/org/model:tag" + + def test_an_explicit_served_name_wins(self): + from infinity_emb.args import EngineArgs + + with mock.patch("infinity_emb.oci.resolve", return_value="/resolved"): + args = EngineArgs( + model_name_or_path="oci://ghcr.io/org/model:tag", + served_model_name="my-model", + ) + + assert args.served_model_name == "my-model" + + def test_a_hf_repo_id_is_untouched(self): + from infinity_emb.args import EngineArgs + + with mock.patch("infinity_emb.oci.resolve") as resolver: + args = EngineArgs(model_name_or_path="michaelfeil/bge-small-en-v1.5") + + resolver.assert_not_called() + assert args.model_name_or_path == "michaelfeil/bge-small-en-v1.5" From 0e7f8d71bddc7a74854533809984a083f06f50d6 Mon Sep 17 00:00:00 2001 From: Eric Curtin Date: Tue, 29 Sep 2026 12:59:33 +0100 Subject: [PATCH 2/2] fix: bound llmman pull/resolve, keep https, check binary --- libs/infinity_emb/infinity_emb/llmman.py | 173 ++++++++++-------- .../tests/unit_test/test_llmman.py | 61 ++++++ libs/infinity_emb/tests/unit_test/test_oci.py | 4 + 3 files changed, 166 insertions(+), 72 deletions(-) diff --git a/libs/infinity_emb/infinity_emb/llmman.py b/libs/infinity_emb/infinity_emb/llmman.py index e18359197..9a8df48a8 100644 --- a/libs/infinity_emb/infinity_emb/llmman.py +++ b/libs/infinity_emb/infinity_emb/llmman.py @@ -16,6 +16,7 @@ - ``llmman resolve --no-pull `` -> one line of JSON carrying ``path``. """ +import contextlib import ipaddress import json import logging @@ -24,6 +25,7 @@ import subprocess import urllib.error import urllib.request +from typing import Any, Iterator, Optional logger = logging.getLogger(__name__) @@ -34,6 +36,10 @@ DEFAULT_PORT = 17434 PROBE_TIMEOUT_SECONDS = 5 +# Pulls run for hours, so bound silence on the stream, not total time. +PULL_STALL_TIMEOUT_SECONDS = 300 +# Extracting a large model into the cache can be slow. +RESOLVE_TIMEOUT_SECONDS = 1800 def _connectable_host(host: str) -> str: @@ -48,13 +54,16 @@ def _connectable_host(host: str) -> str: def endpoint() -> str: - """The http origin of the llmman daemon, honouring LLMMAN_HOST.""" + """The http(s) origin of the llmman daemon, honouring LLMMAN_HOST.""" raw = os.getenv(HOST_ENV, "").strip().strip("\"'") if not raw: return f"http://{DEFAULT_HOST}:{DEFAULT_PORT}" + scheme = "http" if "://" in raw: - raw = raw.split("://", 1)[1] + given, raw = raw.split("://", 1) + if given.lower() == "https": + scheme = "https" raw = raw.split("/", 1)[0] host, port = raw, DEFAULT_PORT @@ -74,7 +83,7 @@ def endpoint() -> str: resolved = _connectable_host(host) if ":" in resolved and not resolved.startswith("["): resolved = f"[{resolved}]" - return f"http://{resolved}:{port}" + return f"{scheme}://{resolved}:{port}" def llmman_bin() -> str: @@ -82,25 +91,46 @@ def llmman_bin() -> str: return os.getenv(BIN_ENV, "").strip() or "llmman" +@contextlib.contextmanager +def _open( + what: str, url: str, timeout: float, payload: Optional[dict] = None, hint: str = "" +) -> Iterator[Any]: + """Open ``url`` (JSON POST if ``payload``); any transport failure is a RuntimeError.""" + data = None if payload is None else json.dumps(payload).encode("utf-8") + req = urllib.request.Request( + url, + data=data, + headers={"Content-Type": "application/json"} if data else {}, + method="GET" if data is None else "POST", + ) + try: + with urllib.request.urlopen(req, timeout=timeout) as resp: + yield resp + except urllib.error.HTTPError as exc: + raise RuntimeError(f"{what}: HTTP {exc.code}{hint}") from exc + except OSError as exc: # refused, timed out, reset + raise RuntimeError(f"{what}: {getattr(exc, 'reason', exc)}{hint}") from exc + + def check_daemon(base: str) -> None: """Confirm an llmman daemon is listening and is actually llmman.""" - url = base + "/api/version" - try: - with urllib.request.urlopen(url, timeout=PROBE_TIMEOUT_SECONDS) as resp: - if resp.status != 200: - raise RuntimeError( - f"llmman daemon at {base} answered /api/version with HTTP {resp.status}" - ) + hint = f". Start one with `llmman serve`, or point {HOST_ENV} at an existing daemon." + with _open( + f"no llmman daemon reachable at {base}", + base + "/api/version", + PROBE_TIMEOUT_SECONDS, + hint=hint, + ) as resp: + if resp.status != 200: + raise RuntimeError( + f"llmman daemon at {base} answered /api/version with HTTP {resp.status}" + ) + try: payload = json.loads(resp.read().decode("utf-8")) - except urllib.error.URLError as exc: - raise RuntimeError( - f"no llmman daemon reachable at {base} ({exc.reason}). Start one with " - f"`llmman serve`, or point {HOST_ENV} at an existing daemon." - ) from exc - except json.JSONDecodeError as exc: - raise RuntimeError( - f"the server at {base} is not an llmman daemon (unparseable /api/version)" - ) from exc + except ValueError as exc: + raise RuntimeError( + f"the server at {base} is not an llmman daemon (unparseable /api/version)" + ) from exc if not isinstance(payload, dict) or not payload.get("version"): raise RuntimeError( @@ -113,45 +143,34 @@ def pull(base: str, reference: str, progress=None) -> None: ``progress`` receives ``(status, completed, total)``. An error can arrive in-band at HTTP 200, and a stream that ends without ``success`` is also a - failure -- neither is treated as a completed pull. + failure -- neither is treated as a completed pull. A stream silent for + ``PULL_STALL_TIMEOUT_SECONDS`` is abandoned. """ - body = json.dumps({"model": reference}).encode("utf-8") - req = urllib.request.Request( - base + "/api/pull", - data=body, - headers={"Content-Type": "application/json"}, - method="POST", - ) - + what = f"llmman pull of {reference!r} failed" succeeded = False - try: - with urllib.request.urlopen(req) as resp: - if resp.status != 200: - raise RuntimeError(f"llmman pull of {reference!r} failed: HTTP {resp.status}") - for raw_line in resp: - line = raw_line.decode("utf-8").strip() - if not line: - continue - try: - obj = json.loads(line) - except json.JSONDecodeError: - # Tolerate a non-JSON diagnostic rather than aborting a - # pull that may still be progressing. - continue - if not isinstance(obj, dict): - continue - if obj.get("error"): - raise RuntimeError(f"llmman pull of {reference!r} failed: {obj['error']}") - status = obj.get("status") - if status == "success": - succeeded = True - continue - if progress is not None and status: - progress(status, obj.get("completed", 0), obj.get("total", 0)) - except urllib.error.HTTPError as exc: - raise RuntimeError(f"llmman pull of {reference!r} failed: HTTP {exc.code}") from exc - except urllib.error.URLError as exc: - raise RuntimeError(f"llmman pull of {reference!r} failed: {exc.reason}") from exc + with _open(what, base + "/api/pull", PULL_STALL_TIMEOUT_SECONDS, {"model": reference}) as resp: + if resp.status != 200: + raise RuntimeError(f"{what}: HTTP {resp.status}") + for raw_line in resp: + line = raw_line.decode("utf-8").strip() + if not line: + continue + try: + obj = json.loads(line) + except json.JSONDecodeError: + # Tolerate a non-JSON diagnostic rather than aborting a + # pull that may still be progressing. + continue + if not isinstance(obj, dict): + continue + if obj.get("error"): + raise RuntimeError(f"{what}: {obj['error']}") + status = obj.get("status") + if status == "success": + succeeded = True + continue + if progress is not None and status: + progress(status, obj.get("completed", 0), obj.get("total", 0)) if not succeeded: raise RuntimeError(f"llmman pull of {reference!r} ended without reporting success") @@ -184,12 +203,8 @@ def parse_resolve_output(stdout: str, reference: str) -> str: return path -def resolve(reference: str) -> str: - """Ask the CLI where the daemon's pull left the model on disk. - - ``--no-pull`` guarantees this only reports on bytes ``/api/pull`` already - fetched, so the daemon stays the only thing that touches the network. - """ +def require_bin() -> str: + """The llmman executable, or a RuntimeError saying how to get it.""" binary = llmman_bin() if shutil.which(binary) is None and not os.path.isfile(binary): raise RuntimeError( @@ -197,18 +212,31 @@ def resolve(reference: str) -> str: "(https://github.com/llmmanorg/llmman) and put it on PATH, or set " f"{BIN_ENV} to its location." ) + return binary - completed = subprocess.run( - [binary, "resolve", "--no-pull", reference], - capture_output=True, - stdin=subprocess.DEVNULL, - text=True, - check=False, - ) + +def resolve(reference: str, binary: Optional[str] = None) -> str: + """Ask the CLI where the daemon's pull left the model on disk. + + ``--no-pull`` guarantees this only reports on bytes ``/api/pull`` already + fetched, so the daemon stays the only thing that touches the network. + """ + binary = binary or require_bin() + cmd = f"`{binary} resolve --no-pull {reference}`" + try: + completed = subprocess.run( + [binary, "resolve", "--no-pull", reference], + capture_output=True, + stdin=subprocess.DEVNULL, + text=True, + check=False, + timeout=RESOLVE_TIMEOUT_SECONDS, + ) + except subprocess.TimeoutExpired as exc: + raise RuntimeError(f"{cmd} timed out after {RESOLVE_TIMEOUT_SECONDS}s") from exc if completed.returncode != 0: raise RuntimeError( - f"`{binary} resolve --no-pull {reference}` failed with exit code " - f"{completed.returncode}: {completed.stderr.strip()}" + f"{cmd} failed with exit code {completed.returncode}: {completed.stderr.strip()}" ) return parse_resolve_output(completed.stdout, reference) @@ -217,6 +245,7 @@ def pull_and_resolve(reference: str, progress=None) -> str: """Full acquisition: probe the daemon, pull through it, report the path.""" base = endpoint() check_daemon(base) + binary = require_bin() # fail before a multi-gigabyte pull, not after logger.info("Pulling %s via llmman daemon at %s", reference, base) pull(base, reference, progress) - return resolve(reference) + return resolve(reference, binary) diff --git a/libs/infinity_emb/tests/unit_test/test_llmman.py b/libs/infinity_emb/tests/unit_test/test_llmman.py index 189a325a9..4fd99fc2a 100644 --- a/libs/infinity_emb/tests/unit_test/test_llmman.py +++ b/libs/infinity_emb/tests/unit_test/test_llmman.py @@ -7,7 +7,10 @@ import http.server import json import socketserver +import subprocess +import sys import threading +from unittest import mock import pytest @@ -26,6 +29,8 @@ def __init__(self): self.pull_body = _ndjson({"status": "success"}) self.pull_status = 200 self.last_request = None + self.stall = None # "headers": never answer, "stream": answer then go quiet + self.release = threading.Event() daemon = self class Handler(http.server.BaseHTTPRequestHandler): @@ -41,11 +46,21 @@ def _send(self, status, body, ctype): self.wfile.write(raw) def do_GET(self): + if daemon.stall == "headers": + return daemon.release.wait(10) self._send(200, json.dumps(daemon.version), "application/json") def do_POST(self): length = int(self.headers.get("Content-Length", 0)) daemon.last_request = json.loads(self.rfile.read(length)) + if daemon.stall == "headers": + return daemon.release.wait(10) + if daemon.stall == "stream": + self.send_response(200) + self.send_header("Content-Length", "1000") + self.end_headers() + self.wfile.write(_ndjson({"status": "pulling manifest"}).encode()) + return daemon.release.wait(10) self._send(daemon.pull_status, daemon.pull_body, "application/x-ndjson") self._server = socketserver.TCPServer(("127.0.0.1", 0), Handler) @@ -53,6 +68,7 @@ def do_POST(self): threading.Thread(target=self._server.serve_forever, daemon=True).start() def close(self): + self.release.set() self._server.shutdown() self._server.server_close() @@ -115,3 +131,48 @@ def test_reports_a_non_ok_status(daemon): def test_tolerates_a_non_json_diagnostic_line(daemon): daemon.pull_body = "not json\n" + _ndjson({"status": "success"}) llmman.pull(daemon.url, "ref") + + +@pytest.fixture +def fast_timeouts(monkeypatch): + monkeypatch.setattr(llmman, "PROBE_TIMEOUT_SECONDS", 0.3) + monkeypatch.setattr(llmman, "PULL_STALL_TIMEOUT_SECONDS", 0.3) + + +def test_gives_up_on_a_stalled_probe(daemon, fast_timeouts): + daemon.stall = "headers" + with pytest.raises(RuntimeError, match="llmman serve"): + llmman.check_daemon(daemon.url) + + +@pytest.mark.parametrize("stall", ["headers", "stream"]) +def test_pull_gives_up_on_a_stalled_daemon(daemon, fast_timeouts, stall): + daemon.stall = stall + with pytest.raises(RuntimeError, match="timed out"): + llmman.pull(daemon.url, "ref") + + +def test_resolve_bounds_the_subprocess(): + hang = subprocess.TimeoutExpired("llmman", 1) + with mock.patch.object(llmman.subprocess, "run", side_effect=hang) as run: + with pytest.raises(RuntimeError, match="timed out"): + llmman.resolve("ref", binary="llmman") + assert run.call_args.kwargs["timeout"] == llmman.RESOLVE_TIMEOUT_SECONDS + + +def test_a_missing_binary_fails_before_the_pull(daemon, monkeypatch, tmp_path): + monkeypatch.setenv(llmman.HOST_ENV, daemon.url) + monkeypatch.setenv(llmman.BIN_ENV, str(tmp_path / "no-such-llmman")) + with pytest.raises(RuntimeError, match="not found"): + llmman.pull_and_resolve("ref") + assert daemon.last_request is None + + +def test_pull_and_resolve_pulls_then_resolves(daemon, monkeypatch, tmp_path): + monkeypatch.setenv(llmman.HOST_ENV, daemon.url) + monkeypatch.setenv(llmman.BIN_ENV, sys.executable) + done = subprocess.CompletedProcess([], 0, stdout=json.dumps({"path": str(tmp_path)})) + with mock.patch.object(llmman.subprocess, "run", return_value=done) as run: + assert llmman.pull_and_resolve("ref") == str(tmp_path) + assert daemon.last_request == {"model": "ref"} + assert run.call_args.args[0] == [sys.executable, "resolve", "--no-pull", "ref"] diff --git a/libs/infinity_emb/tests/unit_test/test_oci.py b/libs/infinity_emb/tests/unit_test/test_oci.py index 77483b1fc..018a59ad6 100644 --- a/libs/infinity_emb/tests/unit_test/test_oci.py +++ b/libs/infinity_emb/tests/unit_test/test_oci.py @@ -64,6 +64,10 @@ def test_hands_the_bare_reference_to_the_daemon(): # A wildcard bind is meaningful to the server but not to a client. ("0.0.0.0:9999", "http://127.0.0.1:9999"), ("[::]:9999", "http://[::1]:9999"), + # A configured scheme is kept, so a TLS daemon is not downgraded. + ("https://example.com:9999", "https://example.com:9999"), + ("HTTPS://1.2.3.4", "https://1.2.3.4:17434"), + ("http://1.2.3.4:9999/", "http://1.2.3.4:9999"), ], ) def test_endpoint_parsing(host, want):