feat(llm): multi-host Ollama pool — least-loaded embed fan-out (refs #568)
This commit is contained in:
@@ -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
115
src/llm/pool.py
Normal 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
95
tests/llm/test_pool.py
Normal 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]
|
||||||
Reference in New Issue
Block a user