feat(llm): multi-host Ollama pool — least-loaded embed fan-out (refs #568)

This commit is contained in:
kert
2026-07-17 08:57:23 -04:00
parent e8e15c61da
commit f1f5b35e58
3 changed files with 212 additions and 0 deletions

View File

@@ -18,3 +18,5 @@ Spec: docs/superpowers/specs/2026-07-16-llm-module-design.md
from llm.config import LlmConfig as LlmConfig from llm.config import LlmConfig as LlmConfig
from llm.config import load as load from llm.config import load as load
from llm.config import pg_url as pg_url from llm.config import pg_url as pg_url
from llm.pool import HostPool as HostPool
from llm.pool import PoolEmbeddings as PoolEmbeddings

115
src/llm/pool.py Normal file
View File

@@ -0,0 +1,115 @@
"""Least-loaded fan-out across Ollama hosts.
Remote GPUs are plain Ollama endpoints (rig 4090 / laptop 5080 run
``ollama serve``); this pool spreads embed calls across them while all
bib/pgvector I/O stays on the server. Uses the raw ``/api/embed`` HTTP API
rather than langchain-ollama so one pool can juggle N base URLs;
``PoolEmbeddings`` adapts it back to the langchain interface for PGVector.
"""
from __future__ import annotations
import threading
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from typing import Iterator, Sequence
import httpx
from langchain_core.embeddings import Embeddings
_TIMEOUT = httpx.Timeout(120.0, connect=5.0)
class HostPool:
"""Tracks in-flight requests per host; hands out the idlest one."""
def __init__(self, hosts: Sequence[str]) -> None:
self._lock = threading.Lock()
self._in_flight: dict[str, int] = {h.rstrip("/"): 0 for h in hosts}
@classmethod
def from_config(cls, cfg) -> "HostPool":
return cls(cfg.ollama_hosts)
@property
def hosts(self) -> list[str]:
return list(self._in_flight)
def check(self, model: str) -> list[str]:
"""Keep only hosts that are up and serve ``model``.
Ollama tags models ``name:latest``; match on the bare prefix.
"""
alive: list[str] = []
with httpx.Client(timeout=_TIMEOUT) as client:
for host in self.hosts:
try:
resp = client.get(f"{host}/api/tags")
names = {
m["name"].split(":")[0] for m in resp.json().get("models", [])
}
if model.split(":")[0] in names:
alive.append(host)
except httpx.HTTPError:
continue
with self._lock:
self._in_flight = {h: 0 for h in alive}
if not alive:
raise RuntimeError(
f"no Ollama host in pool serves {model!r} — "
f"pull it or fix LLM_OLLAMA_HOSTS"
)
return alive
@contextmanager
def acquire(self) -> Iterator[str]:
with self._lock:
host = min(self._in_flight, key=self._in_flight.__getitem__)
self._in_flight[host] += 1
try:
yield host
finally:
with self._lock:
self._in_flight[host] -= 1
def embed_texts(
pool: HostPool,
model: str,
texts: Sequence[str],
*,
batch_size: int = 64,
) -> list[list[float]]:
"""Embed ``texts`` in order, fanning batches across the pool."""
batches = [
(i, list(texts[i : i + batch_size])) for i in range(0, len(texts), batch_size)
]
results: dict[int, list[list[float]]] = {}
def run(start: int, batch: list[str]) -> None:
with pool.acquire() as host, httpx.Client(timeout=_TIMEOUT) as client:
resp = client.post(
f"{host}/api/embed", json={"model": model, "input": batch}
)
resp.raise_for_status()
results[start] = resp.json()["embeddings"]
workers = max(1, min(len(pool.hosts) * 2, len(batches)))
with ThreadPoolExecutor(max_workers=workers) as ex:
for future in [ex.submit(run, s, b) for s, b in batches]:
future.result()
return [vec for start in sorted(results) for vec in results[start]]
class PoolEmbeddings(Embeddings):
"""langchain adapter so PGVector can query through the pool."""
def __init__(self, pool: HostPool, model: str) -> None:
self._pool = pool
self._model = model
def embed_documents(self, texts: list[str]) -> list[list[float]]:
return embed_texts(self._pool, self._model, texts)
def embed_query(self, text: str) -> list[float]:
return self.embed_documents([text])[0]

95
tests/llm/test_pool.py Normal file
View File

@@ -0,0 +1,95 @@
"""llm.pool — least-loaded fan-out across Ollama hosts."""
from unittest.mock import MagicMock, patch
import pytest
from llm.pool import HostPool, PoolEmbeddings, embed_texts
H1, H2 = "http://h1:11434", "http://h2:11434"
def _resp(payload, status=200):
r = MagicMock()
r.status_code = status
r.json.return_value = payload
return r
class TestCheck:
@patch("llm.pool.httpx.Client")
def test_drops_host_missing_model(self, MockClient):
client = MockClient.return_value.__enter__.return_value
client.get.side_effect = [
_resp({"models": [{"name": "nomic-embed-text:latest"}]}),
_resp({"models": [{"name": "other:latest"}]}),
]
pool = HostPool([H1, H2])
assert pool.check("nomic-embed-text") == [H1]
@patch("llm.pool.httpx.Client")
def test_drops_unreachable_host(self, MockClient):
import httpx
client = MockClient.return_value.__enter__.return_value
client.get.side_effect = [
_resp({"models": [{"name": "m:latest"}]}),
httpx.ConnectError("down"),
]
pool = HostPool([H1, H2])
assert pool.check("m") == [H1]
@patch("llm.pool.httpx.Client")
def test_no_hosts_left_raises(self, MockClient):
client = MockClient.return_value.__enter__.return_value
client.get.return_value = _resp({"models": []})
with pytest.raises(RuntimeError, match="no Ollama host"):
HostPool([H1]).check("m")
class TestAcquire:
def test_prefers_least_in_flight(self):
pool = HostPool([H1, H2])
with pool.acquire() as first, pool.acquire() as second:
assert {first, second} == {H1, H2}
def test_released_host_is_reused(self):
pool = HostPool([H1, H2])
with pool.acquire() as first:
pass
with pool.acquire() as a, pool.acquire() as b:
assert {a, b} == {H1, H2}
assert first in {H1, H2}
class TestEmbedTexts:
@patch("llm.pool.httpx.Client")
def test_order_preserved_across_batches(self, MockClient):
client = MockClient.return_value.__enter__.return_value
def fake_post(url, json=None, **kw):
return _resp({"embeddings": [[float(len(t))] for t in json["input"]]})
client.post.side_effect = fake_post
pool = HostPool([H1, H2])
out = embed_texts(pool, "m", ["a", "bb", "ccc"], batch_size=1)
assert out == [[1.0], [2.0], [3.0]]
@patch("llm.pool.httpx.Client")
def test_http_error_raises(self, MockClient):
client = MockClient.return_value.__enter__.return_value
resp = _resp({}, status=500)
resp.raise_for_status.side_effect = Exception("boom")
client.post.return_value = resp
with pytest.raises(Exception, match="boom"):
embed_texts(HostPool([H1]), "m", ["a"])
class TestPoolEmbeddings:
@patch("llm.pool.embed_texts")
def test_adapter_delegates(self, mock_embed):
mock_embed.return_value = [[0.1]]
pool = HostPool([H1])
emb = PoolEmbeddings(pool, "m")
assert emb.embed_documents(["x"]) == [[0.1]]
assert emb.embed_query("x") == [0.1]