Files
stack/tests/llm/test_pool.py
kert cbc55c95d4 fix(llm): drop a pool host that dies mid-run and retry the batch on the survivors (refs #796)
check() filters the pool at start, but a host can vanish afterwards
(rig.local stopped resolving during the 2026-09-22 CMS-2026-2377
re-index; same on 2026-09-04) and embed_texts let the first transport
error abort the whole run. Now an httpx.TransportError drops that host
(HostPool.drop) and the batch is retried elsewhere; only when no host
is left does a RuntimeError name the last failure. acquire() and
acquire_generation() on an empty pool raise the same clear error
instead of min()/max()'s ValueError. HTTP error responses still raise
through raise_for_status — that is a model/config problem, not a dead
host.
2026-09-22 15:52:07 -04:00

305 lines
11 KiB
Python

"""llm.pool — least-loaded fan-out across Ollama hosts."""
from unittest.mock import MagicMock, patch
import httpx
import pytest
from llm.config import LlmConfig
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")
@patch("llm.pool.httpx.Client")
def test_drops_host_with_malformed_json(self, MockClient):
"""A host that's up but returns a non-JSON/malformed body is
dropped, not treated as a fatal error for the whole run."""
client = MockClient.return_value.__enter__.return_value
bad = MagicMock()
bad.json.side_effect = ValueError("not JSON")
client.get.side_effect = [bad, _resp({"models": [{"name": "m:latest"}]})]
pool = HostPool([H1, H2])
assert pool.check("m") == [H2]
@patch("llm.pool.httpx.Client")
def test_drops_host_missing_name_key(self, MockClient):
"""A model dict missing 'name' raises KeyError inside check();
that host should be dropped, not abort the whole check()."""
client = MockClient.return_value.__enter__.return_value
client.get.side_effect = [
_resp({"models": [{"no_name": "whatever"}]}),
_resp({"models": [{"name": "m:latest"}]}),
]
pool = HostPool([H1, H2])
assert pool.check("m") == [H2]
class TestFromConfig:
def test_from_config_uses_ollama_hosts(self):
cfg = LlmConfig(
ollama_hosts=(H1, H2),
embed_model="m",
instruct_model="g",
embed_dim=768,
build_ann_index=False,
pg_host="x",
pg_port=5432,
pg_db="llm",
pg_user="llm",
)
pool = HostPool.from_config(cfg)
assert pool.hosts == list(cfg.ollama_hosts)
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_dying_host_is_dropped_and_batch_retried_on_survivor(self, MockClient):
"""#796: a transport error from one host must not abort the run —
the host leaves the pool and the batch completes elsewhere."""
client = MockClient.return_value.__enter__.return_value
calls: list[str] = []
def fake_post(url, json=None, **kw):
calls.append(url)
if url.startswith(H1):
raise httpx.ConnectError("[Errno -2] Name or service not known")
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]]
assert pool.hosts == [H2]
assert any(u.startswith(H1) for u in calls) # H1 was tried, then dropped
assert pool._in_flight == {H2: 0} # dropped host never left a count behind
@patch("llm.pool.httpx.Client")
def test_every_host_dying_raises_runtime_error(self, MockClient):
client = MockClient.return_value.__enter__.return_value
client.post.side_effect = httpx.ReadTimeout("timed out")
pool = HostPool([H1, H2])
with pytest.raises(RuntimeError, match="every Ollama host failed"):
embed_texts(pool, "m", ["a"])
assert pool.hosts == []
def test_acquire_on_empty_pool_is_a_clear_error(self):
pool = HostPool([])
with pytest.raises(RuntimeError, match="no live Ollama host"):
with pool.acquire():
pass
with pytest.raises(RuntimeError, match="no live Ollama host"):
with pool.acquire_generation():
pass
@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 TestCheckAcquireInterleave:
@patch("llm.pool.httpx.Client")
def test_check_drops_held_host_without_corruption(self, MockClient):
client = MockClient.return_value.__enter__.return_value
pool = HostPool([H1, H2])
with pool.acquire() as held:
survivor = H2 if held == H1 else H1
def fake_get(url, **kw):
if url.startswith(held):
return _resp({"models": []})
return _resp({"models": [{"name": "m:latest"}]})
client.get.side_effect = fake_get
assert pool.check("m") == [survivor]
# exiting acquire() above must not raise despite the dropped host
with pool.acquire() as again:
assert again == survivor
assert pool._in_flight[survivor] == 0
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]
class TestGeneration:
def _pool(self):
return HostPool([H1, H2], vram_gb={H1: 12, H2: 24})
@patch("llm.pool.httpx.Client")
def test_acquire_generation_prefers_largest_live_host(self, MockClient):
client = MockClient.return_value.__enter__.return_value
client.get.side_effect = [
_resp({"models": [{"name": "chat:latest"}]}),
_resp({"models": [{"name": "chat:latest"}, {"name": "big:latest"}]}),
]
pool = self._pool()
pool.check("chat")
with pool.acquire_generation() as host:
assert host == H2
assert pool._in_flight[H2] == 1
assert pool._in_flight[H2] == 0
@patch("llm.pool.httpx.Client")
def test_falls_back_when_largest_is_down(self, MockClient):
import httpx
client = MockClient.return_value.__enter__.return_value
client.get.side_effect = [
_resp({"models": [{"name": "chat:latest"}]}),
httpx.ConnectError("rig down"),
]
pool = self._pool()
pool.check("chat")
with pool.acquire_generation() as host:
assert host == H1
def test_tie_breaks_on_least_in_flight(self):
pool = HostPool([H1, H2], vram_gb={H1: 24, H2: 24})
pool._in_flight[H1] = 3
with pool.acquire_generation() as host:
assert host == H2
def test_vram_defaults_to_zero(self):
assert HostPool([H1]).vram(H1) == 0.0
@patch("llm.pool.httpx.Client")
def test_serves_reflects_last_check(self, MockClient):
client = MockClient.return_value.__enter__.return_value
client.get.side_effect = [
_resp({"models": [{"name": "chat:latest"}, {"name": "big:latest"}]}),
]
pool = HostPool([H1])
pool.check("chat")
assert pool.serves(H1, "big") is True
assert pool.serves(H1, "big:latest") is True
assert pool.serves(H1, "nope") is False
assert pool.reachable(H1) is True
assert pool.reachable(H2) is False
assert pool.status() == [
{"host": H1, "vram_gb": 0.0, "models": ["big:latest", "chat:latest"]}
]
class TestPickModel:
def _cfg(self, **kw):
base = dict(
ollama_hosts=(H1, H2),
embed_model="e",
instruct_model="chat",
instruct_model_large="big",
large_min_vram_gb=20,
embed_dim=768,
build_ann_index=False,
pg_host="x",
pg_port=5432,
pg_db="llm",
pg_user="llm",
)
base.update(kw)
return LlmConfig(**base)
def _pool_serving(self, models_by_host):
pool = HostPool(list(models_by_host), vram_gb={H1: 12, H2: 24})
pool._models = {h: set(ms) for h, ms in models_by_host.items()}
return pool
def test_large_on_big_host_that_serves_it(self):
from llm.pool import pick_model
pool = self._pool_serving({H1: {"chat"}, H2: {"chat", "big"}})
assert pick_model(self._cfg(), pool, H2) == "big"
def test_baseline_on_small_host(self):
from llm.pool import pick_model
pool = self._pool_serving({H1: {"chat", "big"}, H2: {"chat", "big"}})
assert pick_model(self._cfg(), pool, H1) == "chat"
def test_baseline_when_big_host_lacks_large_model(self):
from llm.pool import pick_model
pool = self._pool_serving({H1: {"chat"}, H2: {"chat"}})
assert pick_model(self._cfg(), pool, H2) == "chat"
def test_baseline_when_no_large_configured(self):
from llm.pool import pick_model
pool = self._pool_serving({H2: {"chat", "big"}})
assert pick_model(self._cfg(instruct_model_large=""), pool, H2) == "chat"