Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions libs/infinity_emb/infinity_emb/args.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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):
Expand Down
251 changes: 251 additions & 0 deletions libs/infinity_emb/infinity_emb/llmman.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,251 @@
"""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 <ref>`` -> one line of JSON carrying ``path``.
"""

import contextlib
import ipaddress
import json
import logging
import os
import shutil
import subprocess
import urllib.error
import urllib.request
from typing import Any, Iterator, Optional

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
# 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:
"""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(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:
given, raw = raw.split("://", 1)
if given.lower() == "https":
scheme = "https"
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"{scheme}://{resolved}:{port}"


def llmman_bin() -> str:
"""The llmman executable name, overridable per project."""
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."""
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 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(
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. A stream silent for
``PULL_STALL_TIMEOUT_SECONDS`` is abandoned.
"""
what = f"llmman pull of {reference!r} failed"
succeeded = False
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")


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 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(
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."
)
return binary


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"{cmd} failed with exit code {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)
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)
Comment thread
ericcurtin marked this conversation as resolved.
return resolve(reference, binary)
56 changes: 56 additions & 0 deletions libs/infinity_emb/infinity_emb/oci.py
Original file line number Diff line number Diff line change
@@ -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)
Loading