feat(prisma): fleet mode for eligibility and extract via a shared run_pool driver (refs #655)

prisma.fleet.run_pool owns the multi-host thread-pool loop that screen
had inline: at most --per-host in-flight requests per Ollama host, LLM
calls only on the pool, every SQLite touch and per-item commit on the
calling thread, and a host that fails twice in a row is retired for the
rest of the run (its queued items go to the surviving hosts). screen,
eligibility and extract each supply prepare/apply/on_error glue; items
with no retrievable full text resolve without a fleet call. CLI:
prisma eligible/extract gain --hosts/--fleet-model/--num-ctx/--per-host
matching prisma screen. Docs regenerated.
This commit is contained in:
kert
2026-09-11 18:23:37 -04:00
parent e93e35cb87
commit 49a2f6f723
10 changed files with 1011 additions and 80 deletions

View File

@@ -14,16 +14,28 @@ Usage: stack prisma eligible [OPTIONS] NAME
│ * name TEXT Project slug. [required] │
╰──────────────────────────────────────────────────────────────────────────────╯
╭─ Options ────────────────────────────────────────────────────────────────────╮
│ --limit -n INTEGER [default: 0] │
│ --max-chars INTEGER Cap per-item markdown (full text) at │
│ N chars so bounded-context local │
│ models keep the criteria in window. 0 │
│ = no cap. │
│ [default: 0] │
│ --db PATH [default: │
│ data/zotero/data/zotero.sqlite] │
│ --storage PATH [default: data/zotero/data/storage] │
│ --hold --no-hold [default: hold] │
│ --help Show this message and exit. │
│ --limit -n INTEGER [default: 0] │
│ --max-chars INTEGER Cap per-item markdown (full text) │
│ at N chars so bounded-context local │
│ models keep the criteria in window. │
│ 0 = no cap. │
│ [default: 0] │
│ --db PATH [default: │
│ data/zotero/data/zotero.sqlite] │
│ --storage PATH [default: data/zotero/data/storage] │
│ --hold --no-hold [default: hold] │
│ --hosts TEXT Comma-separated Ollama base URLs — │
│ fleet mode (#655): parallel │
│ eligibility via native /api/chat │
│ with per-request num_ctx; bypasses │
│ PRISMA_LLM_* env config. │
│ --fleet-model TEXT Model name on every fleet host. │
│ [default: qwen2.5:14b] │
│ --num-ctx INTEGER Context window per fleet request. │
│ [default: 8192] │
│ --per-host INTEGER Max in-flight requests per fleet │
│ host. │
│ [default: 1] │
│ --help Show this message and exit. │
╰──────────────────────────────────────────────────────────────────────────────╯
```

View File

@@ -14,11 +14,23 @@ Usage: stack prisma extract [OPTIONS] NAME
│ * name TEXT Project slug. [required] │
╰──────────────────────────────────────────────────────────────────────────────╯
╭─ Options ────────────────────────────────────────────────────────────────────╮
│ --limit -n INTEGER [default: 0] │
│ --db PATH [default: │
│ data/zotero/data/zotero.sqlite] │
│ --storage PATH [default: data/zotero/data/storage] │
│ --hold --no-hold [default: hold] │
│ --help Show this message and exit. │
│ --limit -n INTEGER [default: 0] │
│ --db PATH [default: │
│ data/zotero/data/zotero.sqlite] │
│ --storage PATH [default: data/zotero/data/storage] │
│ --hold --no-hold [default: hold] │
│ --hosts TEXT Comma-separated Ollama base URLs — │
│ fleet mode (#655): parallel │
│ extraction via native /api/chat │
│ with per-request num_ctx; bypasses │
│ PRISMA_LLM_* env config. │
│ --fleet-model TEXT Model name on every fleet host. │
│ [default: qwen2.5:14b] │
│ --num-ctx INTEGER Context window per fleet request. │
│ [default: 8192] │
│ --per-host INTEGER Max in-flight requests per fleet │
│ host. │
│ [default: 1] │
│ --help Show this message and exit. │
╰──────────────────────────────────────────────────────────────────────────────╯
```

View File

@@ -177,6 +177,22 @@ def eligible(
db: Path = typer.Option(_DEFAULT_ZOT, "--db"),
storage: Path = typer.Option(_DEFAULT_STORAGE, "--storage"),
hold: bool = typer.Option(True, "--hold/--no-hold"),
hosts: str = typer.Option(
"",
"--hosts",
help="Comma-separated Ollama base URLs — fleet mode (#655): "
"parallel eligibility via native /api/chat with per-request "
"num_ctx; bypasses PRISMA_LLM_* env config.",
),
fleet_model: str = typer.Option(
"qwen2.5:14b", "--fleet-model", help="Model name on every fleet host."
),
num_ctx: int = typer.Option(
8192, "--num-ctx", help="Context window per fleet request."
),
per_host: int = typer.Option(
1, "--per-host", help="Max in-flight requests per fleet host."
),
) -> None:
"""Stage 3 — LLM-driven full-text eligibility pass."""
from prisma import eligibility
@@ -185,12 +201,23 @@ def eligible(
from zot.db import Db
def _go():
provider = make_provider()
with Db(str(db)) as zdb:
project = load(zdb, name)
if hosts:
return eligibility.run_fleet(
zdb,
project,
hosts=[h.strip() for h in hosts.split(",") if h.strip()],
model=fleet_model,
num_ctx=num_ctx,
per_host=per_host,
storage_dir=storage,
limit=limit or None,
max_fulltext_chars=max_chars or None,
)
return eligibility.run(
zdb,
provider,
make_provider(),
project,
storage_dir=storage,
limit=limit or None,
@@ -209,6 +236,22 @@ def extract(
db: Path = typer.Option(_DEFAULT_ZOT, "--db"),
storage: Path = typer.Option(_DEFAULT_STORAGE, "--storage"),
hold: bool = typer.Option(True, "--hold/--no-hold"),
hosts: str = typer.Option(
"",
"--hosts",
help="Comma-separated Ollama base URLs — fleet mode (#655): "
"parallel extraction via native /api/chat with per-request "
"num_ctx; bypasses PRISMA_LLM_* env config.",
),
fleet_model: str = typer.Option(
"qwen2.5:14b", "--fleet-model", help="Model name on every fleet host."
),
num_ctx: int = typer.Option(
8192, "--num-ctx", help="Context window per fleet request."
),
per_host: int = typer.Option(
1, "--per-host", help="Max in-flight requests per fleet host."
),
) -> None:
"""Stage 3+ — structured data extraction on included studies."""
from prisma import extract as _extract
@@ -217,12 +260,22 @@ def extract(
from zot.db import Db
def _go():
provider = make_provider()
with Db(str(db)) as zdb:
project = load(zdb, name)
if hosts:
return _extract.run_fleet(
zdb,
project,
hosts=[h.strip() for h in hosts.split(",") if h.strip()],
model=fleet_model,
num_ctx=num_ctx,
per_host=per_host,
storage_dir=storage,
limit=limit or None,
)
return _extract.run(
zdb,
provider,
make_provider(),
project,
storage_dir=storage,
limit=limit or None,

View File

@@ -174,3 +174,97 @@ def run(
continue
return stats
def run_fleet(
db: Db,
project: Project,
*,
hosts: list[str],
model: str,
storage_dir: Path,
num_ctx: int = 8192,
per_host: int = 1,
limit: int | None = None,
progress: bool = True,
max_fulltext_chars: int | None = None,
) -> dict[str, int]:
"""Run eligibility screening in parallel across a fleet of Ollama
hosts (#655).
Same invariants as :func:`prisma.screen.run_fleet`: the pool
(shared driver in :mod:`prisma.fleet`) only ever runs the LLM
call; every SQLite touch — item loads and decision writes — stays
on the calling thread with a commit per item. Items with no
retrievable full text are a deterministic PRISMA outcome (same as
:func:`run`) and are resolved without ever touching a fleet host.
"""
from prisma.fleet import run_pool
from prisma.llm import LLMResult, OllamaNativeProvider
providers = {
h.rstrip("/"): OllamaNativeProvider(host=h, model=model, num_ctx=num_ctx)
for h in hosts
}
ids = _queue(db, project.name, limit)
stats = {"assessed": 0, "include": 0, "exclude": 0, "uncertain": 0, "errors": 0}
def prepare(zot_id: int):
snap = load_item(db, zot_id, storage_dir)
md = to_markdown(snap, include_fulltext=True)
if "\n# Full text\n" not in md:
# No retrievable full text — resolved directly, no LLM call.
apply_eligibility_decision(
db,
zot_id,
project.name,
{
"decision": "uncertain",
"reasons": ["unavailable-fulltext"],
"rationale": "No full text retrieved by the fetch "
"cascade; cannot assess eligibility.",
},
)
db.commit()
stats["uncertain"] += 1
stats["assessed"] += 1
return None
if max_fulltext_chars and len(md) > max_fulltext_chars:
md = md[:max_fulltext_chars] + _TRUNCATION_NOTE
return _build_call(project, md)
def apply(zot_id: int, result: LLMResult) -> None:
if not result.tool_calls:
log.warning("no tool_call for item %s — skipping", zot_id)
stats["errors"] += 1
return
payload = result.tool_calls[0].get("input") or {}
apply_eligibility_decision(db, zot_id, project.name, payload)
db.commit()
decision = (payload.get("decision") or "uncertain").lower()
stats[decision] = stats.get(decision, 0) + 1
stats["assessed"] += 1
if progress and stats["assessed"] % 10 == 0:
print(
f" assessed {stats['assessed']}/{len(ids)} "
f"(inc={stats['include']} exc={stats['exclude']} "
f"unc={stats['uncertain']})"
)
def on_error(zot_id: int, e: Exception, stage: str) -> None:
if stage == "prepare":
log.warning("load failed on item %s: %s", zot_id, e)
else:
log.warning("eligibility failed on item %s: %s", zot_id, e)
stats["errors"] += 1
run_pool(
ids,
providers,
per_host=per_host,
prepare=prepare,
apply=apply,
on_error=on_error,
)
return stats

View File

@@ -134,3 +134,72 @@ def run(
continue
return stats
def run_fleet(
db: Db,
project: Project,
*,
hosts: list[str],
model: str,
storage_dir: Path,
num_ctx: int = 8192,
per_host: int = 1,
limit: int | None = None,
progress: bool = True,
) -> dict[str, int]:
"""Extract in parallel across a fleet of Ollama hosts (#655).
Same invariants as :func:`prisma.screen.run_fleet`: the pool
(shared driver in :mod:`prisma.fleet`) only ever runs the LLM
call; every SQLite touch — item loads and note writes — stays on
the calling thread with a commit per item.
"""
from prisma.fleet import run_pool
from prisma.llm import LLMResult, OllamaNativeProvider
# Built once, before dispatch: the schema comes from the project's
# extraction-template note, so a broken template should fail the
# run outright rather than erroring on every item in turn.
tool = build_extract_tool(project.extraction_template, project=project.name)
providers = {
h.rstrip("/"): OllamaNativeProvider(host=h, model=model, num_ctx=num_ctx)
for h in hosts
}
ids = _queue(db, project.name, limit)
stats = {"extracted": 0, "errors": 0}
def prepare(zot_id: int):
snap = load_item(db, zot_id, storage_dir)
md = to_markdown(snap, include_fulltext=True)
return _build_call(project, md, tool)
def apply(zot_id: int, result: LLMResult) -> None:
if not result.tool_calls:
log.warning("no tool_call for item %s — skipping", zot_id)
stats["errors"] += 1
return
payload = result.tool_calls[0].get("input") or {}
apply_extraction(db, zot_id, project.name, payload)
db.commit()
stats["extracted"] += 1
if progress and stats["extracted"] % 5 == 0:
print(f" extracted {stats['extracted']}/{len(ids)}")
def on_error(zot_id: int, e: Exception, stage: str) -> None:
if stage == "prepare":
log.warning("load failed on item %s: %s", zot_id, e)
else:
log.warning("extract failed on item %s: %s", zot_id, e)
stats["errors"] += 1
run_pool(
ids,
providers,
per_host=per_host,
prepare=prepare,
apply=apply,
on_error=on_error,
)
return stats

141
src/prisma/fleet.py Normal file
View File

@@ -0,0 +1,141 @@
"""Shared multi-host thread-pool driver for the LLM stages (#655).
``prisma.screen``, ``prisma.eligibility``, and ``prisma.extract`` each
expose a ``run_fleet`` that fans LLM calls out across a small pool of
Ollama hosts. The three loops are identical except for what happens
before dispatch (build the prompt, or resolve some items without an
LLM call at all) and after completion (apply the stage's decision and
update its stats) — this module owns that loop so the three stages
stay in lock-step instead of copy-pasted.
Invariants preserved from the original screen-only implementation
(commit 6102df6):
- LLM calls run on a thread pool with at most ``per_host`` in-flight
requests per host (Ollama serializes per model anyway).
- Every SQLite touch happens on the calling thread — the pool is only
ever used for the outbound HTTP call — so the caller's single
connection is never shared across threads.
- Per-item commits (done by the caller's ``apply``) keep the crash
story identical to the sequential ``run``: a crash loses at most the
in-flight items, which stay queued for the next pass.
New in this module: a host that fails twice in a row is retired for
the rest of the run. Its currently in-flight items still get a chance
to complete (or fail, which is counted like any other error); once
retired, the host is simply never handed new work, so its remaining
queued items are picked up by the surviving hosts. A dead host
therefore loses at most its in-flight items, never the whole queue.
"""
from __future__ import annotations
import logging
from collections.abc import Callable, Iterable, Iterator
from concurrent.futures import FIRST_COMPLETED, Future, ThreadPoolExecutor, wait
from typing import TypeVar
from prisma.llm import LLMCall, LLMResult, Provider
log = logging.getLogger(__name__)
T = TypeVar("T")
# A host that fails this many times in a row is retired for the rest
# of the run.
MAX_CONSECUTIVE_FAILURES = 2
def run_pool(
ids: Iterable[T],
providers: dict[str, Provider],
*,
per_host: int,
prepare: Callable[[T], LLMCall | None],
apply: Callable[[T, LLMResult], None],
on_error: Callable[[T, Exception, str], None],
) -> None:
"""Drive ``ids`` through ``providers`` on a bounded thread pool.
``prepare`` and ``apply`` (and ``on_error``) all run on the
calling thread — only ``provider.complete`` runs in the pool.
- ``prepare(item)`` builds the :class:`~prisma.llm.LLMCall` for an
item. Returning ``None`` means the item was fully handled
without an LLM call (``prepare`` is then responsible for any
writes/stats itself) — the driver moves on without consuming a
host slot.
- ``apply(item, result)`` runs when a dispatched call completes
successfully; it owns the SQLite write, commit, and stats.
- ``on_error(item, exc, stage)`` runs for any exception raised by
``prepare``, the LLM call itself, or ``apply`` — ``stage`` is
one of ``"prepare"``, ``"complete"``, or ``"apply"`` so callers
can log a message matching what actually failed.
Hosts are tracked for consecutive failures of the LLM call itself
(not ``prepare``/``apply`` failures, which are the caller's own
items, not the host's fault): two in a row retires that host for
the rest of this call.
"""
it: Iterator[T] = iter(ids)
slots = {h: per_host for h in providers}
consecutive_failures = dict.fromkeys(providers, 0)
retired: set[str] = set()
pending: dict[Future, tuple[T, str]] = {}
exhausted = False
def live_hosts() -> list[str]:
return [h for h in providers if h not in retired]
with ThreadPoolExecutor(max_workers=max(1, len(providers) * per_host)) as ex:
while True:
# Fill every free slot on a live host (prepare runs here,
# on the calling thread).
while not exhausted:
free = [h for h in live_hosts() if slots[h] > 0]
if not free:
break
try:
item = next(it)
except StopIteration:
exhausted = True
break
try:
call = prepare(item)
except Exception as e: # noqa: BLE001 — one item shouldn't kill the run
on_error(item, e, "prepare")
continue
if call is None:
continue # prepare already fully handled this item
host = free[0]
slots[host] -= 1
pending[ex.submit(providers[host].complete, call)] = (item, host)
if not pending:
# Either the queue is exhausted, or every host is
# retired with items still unattempted — either way
# there's nothing left this driver can do.
break
done, _ = wait(pending, return_when=FIRST_COMPLETED)
for fut in done:
item, host = pending.pop(fut)
slots[host] += 1
try:
result = fut.result()
except Exception as e: # noqa: BLE001
consecutive_failures[host] += 1
if consecutive_failures[host] >= MAX_CONSECUTIVE_FAILURES:
retired.add(host)
log.warning(
"fleet host %s retired after %d consecutive failures",
host,
consecutive_failures[host],
)
on_error(item, e, "complete")
continue
consecutive_failures[host] = 0
try:
apply(item, result)
except Exception as e: # noqa: BLE001
on_error(item, e, "apply")

View File

@@ -189,10 +189,14 @@ def run_fleet(
never shared across threads. Per-item commits keep the crash
story identical to :func:`run`: a dead host or dead run loses at
most its in-flight items, which stay queued for the next pass.
"""
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from prisma.llm import OllamaNativeProvider
The pool driver itself lives in :mod:`prisma.fleet` — shared with
:func:`prisma.eligibility.run_fleet` and
:func:`prisma.extract.run_fleet` — this function only supplies the
per-item prompt and decision-apply glue.
"""
from prisma.fleet import run_pool
from prisma.llm import LLMResult, OllamaNativeProvider
providers = {
h.rstrip("/"): OllamaNativeProvider(host=h, model=model, num_ctx=num_ctx)
@@ -200,63 +204,43 @@ def run_fleet(
}
ids = _queue(db, project.name, limit)
stats = {"screened": 0, "include": 0, "exclude": 0, "uncertain": 0, "errors": 0}
slots = {h: per_host for h in providers}
pending: dict[object, tuple[int, str]] = {}
it = iter(ids)
exhausted = False
with ThreadPoolExecutor(max_workers=max(1, len(providers) * per_host)) as ex:
while True:
# Fill every free slot (item loads happen here, main thread).
while not exhausted:
free = [h for h, s in slots.items() if s > 0]
if not free:
break
try:
zot_id = next(it)
except StopIteration:
exhausted = True
break
try:
snap = load_item(db, zot_id, storage_dir)
call = _build_call(
project, to_markdown(snap, include_fulltext=False)
)
except Exception as e: # noqa: BLE001
log.warning("load failed on item %s: %s", zot_id, e)
stats["errors"] += 1
continue
host = free[0]
slots[host] -= 1
pending[ex.submit(providers[host].complete, call)] = (zot_id, host)
def prepare(zot_id: int):
snap = load_item(db, zot_id, storage_dir)
return _build_call(project, to_markdown(snap, include_fulltext=False))
if not pending:
break
done, _ = wait(pending, return_when=FIRST_COMPLETED)
for fut in done:
zot_id, host = pending.pop(fut)
slots[host] += 1
try:
result = fut.result()
except Exception as e: # noqa: BLE001
log.warning("screen failed on item %s (%s): %s", zot_id, host, e)
stats["errors"] += 1
continue
if not result.tool_calls:
log.warning("no tool_call for item %s — skipping", zot_id)
stats["errors"] += 1
continue
payload = result.tool_calls[0].get("input") or {}
apply_screen_decision(db, zot_id, project.name, payload)
db.commit()
decision = (payload.get("decision") or "uncertain").lower()
stats[decision] = stats.get(decision, 0) + 1
stats["screened"] += 1
if progress and stats["screened"] % 10 == 0:
print(
f" screened {stats['screened']}/{len(ids)} "
f"(inc={stats['include']} exc={stats['exclude']} "
f"unc={stats['uncertain']})"
)
def apply(zot_id: int, result: LLMResult) -> None:
if not result.tool_calls:
log.warning("no tool_call for item %s — skipping", zot_id)
stats["errors"] += 1
return
payload = result.tool_calls[0].get("input") or {}
apply_screen_decision(db, zot_id, project.name, payload)
db.commit()
decision = (payload.get("decision") or "uncertain").lower()
stats[decision] = stats.get(decision, 0) + 1
stats["screened"] += 1
if progress and stats["screened"] % 10 == 0:
print(
f" screened {stats['screened']}/{len(ids)} "
f"(inc={stats['include']} exc={stats['exclude']} "
f"unc={stats['uncertain']})"
)
def on_error(zot_id: int, e: Exception, stage: str) -> None:
if stage == "prepare":
log.warning("load failed on item %s: %s", zot_id, e)
else:
log.warning("screen failed on item %s: %s", zot_id, e)
stats["errors"] += 1
run_pool(
ids,
providers,
per_host=per_host,
prepare=prepare,
apply=apply,
on_error=on_error,
)
return stats

View File

@@ -94,6 +94,49 @@ class TestEligible:
result = runner.invoke(app, ["eligible", "test-proj", "--no-hold"])
assert result.exit_code == 0
mc_run.assert_called_once()
@patch("cli.prisma.subprocess.run")
@patch("zot.db.Db")
@patch("prisma.llm.make_provider")
@patch("prisma.project.load")
@patch("prisma.eligibility.run_fleet", return_value={"assessed": 5})
@patch("prisma.eligibility.run", return_value={"eligible": 5})
def test_fleet_flags_wired_to_driver(
self, mc_run, mc_fleet, mc_load, mc_prov, mc_db_cls, mc_sub
):
mc_db_cls.return_value = _mock_db()
mc_prov.return_value = MagicMock()
project = MagicMock()
mc_load.return_value = project
result = runner.invoke(
app,
[
"eligible",
"test-proj",
"--no-hold",
"--hosts",
"http://h1:11434, http://h2:11434",
"--fleet-model",
"qwen2.5:32b",
"--num-ctx",
"16384",
"--per-host",
"3",
"--max-chars",
"40000",
],
)
assert result.exit_code == 0
mc_run.assert_not_called()
mc_fleet.assert_called_once()
_, kwargs = mc_fleet.call_args
assert kwargs["hosts"] == ["http://h1:11434", "http://h2:11434"]
assert kwargs["model"] == "qwen2.5:32b"
assert kwargs["num_ctx"] == 16384
assert kwargs["per_host"] == 3
assert kwargs["max_fulltext_chars"] == 40000
class TestExtract:
@@ -109,6 +152,45 @@ class TestExtract:
result = runner.invoke(app, ["extract", "test-proj", "--no-hold"])
assert result.exit_code == 0
mc_run.assert_called_once()
@patch("cli.prisma.subprocess.run")
@patch("zot.db.Db")
@patch("prisma.llm.make_provider")
@patch("prisma.project.load")
@patch("prisma.extract.run_fleet", return_value={"extracted": 3})
@patch("prisma.extract.run", return_value={"extracted": 3})
def test_fleet_flags_wired_to_driver(
self, mc_run, mc_fleet, mc_load, mc_prov, mc_db_cls, mc_sub
):
mc_db_cls.return_value = _mock_db()
mc_prov.return_value = MagicMock()
mc_load.return_value = MagicMock()
result = runner.invoke(
app,
[
"extract",
"test-proj",
"--no-hold",
"--hosts",
"http://h1:11434,http://h2:11434",
"--fleet-model",
"qwen2.5:32b",
"--num-ctx",
"16384",
"--per-host",
"2",
],
)
assert result.exit_code == 0
mc_run.assert_not_called()
mc_fleet.assert_called_once()
_, kwargs = mc_fleet.call_args
assert kwargs["hosts"] == ["http://h1:11434", "http://h2:11434"]
assert kwargs["model"] == "qwen2.5:32b"
assert kwargs["num_ctx"] == 16384
assert kwargs["per_host"] == 2
class TestFlow:

199
tests/prisma/test_fleet.py Normal file
View File

@@ -0,0 +1,199 @@
"""Tests for prisma.fleet.run_pool — the shared multi-host driver (#655)."""
from __future__ import annotations
import threading
import time
from dataclasses import dataclass, field
from prisma.fleet import MAX_CONSECUTIVE_FAILURES, run_pool
from prisma.llm import LLMCall, LLMMessage, LLMResult
def _call(item: int) -> LLMCall:
return LLMCall(messages=[LLMMessage(role="user", content=str(item))])
@dataclass
class FakeProvider:
"""Records which host handled each call; optionally fails N times."""
host: str
log: list
lock: threading.Lock
delay: float = 0.0
fail_times: int = 0
calls: int = field(default=0, init=False)
def complete(self, call: LLMCall) -> LLMResult:
with self.lock:
self.calls += 1
n = self.calls
self.log.append(self.host)
if self.delay:
time.sleep(self.delay)
if n <= self.fail_times:
raise RuntimeError(f"{self.host} boom #{n}")
return LLMResult(text="", tool_calls=[{"name": "x", "input": {}}], usage={})
class TestFanOut:
def test_uses_both_hosts(self):
log: list[str] = []
providers = {
"A": FakeProvider(host="A", log=log, lock=threading.Lock()),
"B": FakeProvider(host="B", log=log, lock=threading.Lock()),
}
applied: list[int] = []
errors: list[tuple[int, str]] = []
run_pool(
list(range(6)),
providers,
per_host=1,
prepare=_call,
apply=lambda item, result: applied.append(item),
on_error=lambda item, exc, stage: errors.append((item, stage)),
)
assert errors == []
assert sorted(applied) == list(range(6))
assert set(log) == {"A", "B"}
class TestPerHostCap:
def test_never_exceeds_per_host_in_flight(self):
per_host = 2
lock = threading.Lock()
current = {"A": 0, "B": 0}
max_seen = {"A": 0, "B": 0}
log: list[str] = []
class TrackingProvider:
def __init__(self, host):
self.host = host
def complete(self, call):
with lock:
current[self.host] += 1
max_seen[self.host] = max(max_seen[self.host], current[self.host])
log.append(self.host)
time.sleep(0.02)
with lock:
current[self.host] -= 1
return LLMResult(
text="", tool_calls=[{"name": "x", "input": {}}], usage={}
)
providers = {"A": TrackingProvider("A"), "B": TrackingProvider("B")}
applied: list[int] = []
run_pool(
list(range(8)),
providers,
per_host=per_host,
prepare=_call,
apply=lambda item, result: applied.append(item),
on_error=lambda item, exc, stage: None,
)
assert len(applied) == 8
assert max_seen["A"] <= per_host
assert max_seen["B"] <= per_host
assert set(log) == {"A", "B"}
class TestFailingHostRetired:
def test_dead_host_loses_only_in_flight_items_rest_go_to_survivor(self):
log: list[str] = []
lock = threading.Lock()
# A always fails (and is instant); B always succeeds but is
# deliberately slower so A's failures resolve first and A gets
# retired well before B drains the queue.
providers = {
"A": FakeProvider(host="A", log=log, lock=lock, fail_times=10**6),
"B": FakeProvider(host="B", log=log, lock=lock, delay=0.05),
}
applied: list[int] = []
errors: list[tuple[int, str]] = []
run_pool(
list(range(5)),
providers,
per_host=1,
prepare=_call,
apply=lambda item, result: applied.append(item),
on_error=lambda item, exc, stage: errors.append((item, stage)),
)
a_calls = log.count("A")
b_calls = log.count("B")
# A is retired after exactly MAX_CONSECUTIVE_FAILURES calls —
# never dispatched a third time.
assert a_calls == MAX_CONSECUTIVE_FAILURES
# Every item not sent to A was resolved by B.
assert b_calls == 5 - a_calls
assert len(applied) == b_calls
assert len(errors) == a_calls
assert all(stage == "complete" for _, stage in errors)
class TestPrepareSkipsLLM:
def test_none_from_prepare_skips_dispatch(self):
providers = {"A": FakeProvider(host="A", log=[], lock=threading.Lock())}
handled_directly: list[int] = []
def prepare(item):
if item % 2 == 0:
handled_directly.append(item)
return None
return _call(item)
applied: list[int] = []
run_pool(
list(range(4)),
providers,
per_host=1,
prepare=prepare,
apply=lambda item, result: applied.append(item),
on_error=lambda item, exc, stage: None,
)
assert handled_directly == [0, 2]
assert sorted(applied) == [1, 3]
class TestErrorStages:
def test_prepare_error_tagged(self):
providers = {"A": FakeProvider(host="A", log=[], lock=threading.Lock())}
errors = []
def prepare(item):
raise ValueError("bad item")
run_pool(
[1],
providers,
per_host=1,
prepare=prepare,
apply=lambda item, result: None,
on_error=lambda item, exc, stage: errors.append(stage),
)
assert errors == ["prepare"]
def test_apply_error_tagged(self):
providers = {"A": FakeProvider(host="A", log=[], lock=threading.Lock())}
errors = []
def apply(item, result):
raise ValueError("bad apply")
run_pool(
[1],
providers,
per_host=1,
prepare=_call,
apply=apply,
on_error=lambda item, exc, stage: errors.append(stage),
)
assert errors == ["apply"]

View File

@@ -0,0 +1,285 @@
"""Fleet-mode tests for the stage-2/3/3+ run_fleet entry points (#655).
Mirrors the driver-level tests in test_fleet.py but exercises each
stage's own run_fleet against a real sqlite-backed Db, with
OllamaNativeProvider faked out per host — the same shape as the
sequential run() tests in test_screen_exercise.py / test_stages_exercise.py.
"""
from __future__ import annotations
from pathlib import Path
from unittest.mock import MagicMock, patch
from prisma.llm import LLMResult
from zot.db import Db
from zot.schema import create_db
def _setup(tmp_path):
path = str(tmp_path / "z.sqlite")
con = create_db(path)
con.close()
return path
def _insert_item(db, key, tags):
db.con.execute(
"INSERT INTO items (itemTypeID, libraryID, key, dateAdded, dateModified, clientDateModified) "
"VALUES (2, 1, ?, '', '', '')",
(key,),
)
iid = db.con.execute("SELECT last_insert_rowid()").fetchone()[0]
for t in tags:
row = db.con.execute("SELECT tagID FROM tags WHERE name = ?", (t,)).fetchone()
if row:
tid = row[0]
else:
tid = db.con.execute("INSERT INTO tags (name) VALUES (?)", (t,)).lastrowid
db.con.execute(
"INSERT INTO itemTags (itemID, tagID, type) VALUES (?, ?, 0)", (iid, tid)
)
db.commit()
return iid
class _FakeHostProvider:
"""Stands in for OllamaNativeProvider — records the host it was
built for and either always succeeds or fails a fixed number of
times before succeeding."""
def __init__(self, host, model, num_ctx, *, fail_times=0, payload=None):
self.host = host
self.calls = 0
self.fail_times = fail_times
self.payload = payload or {}
def complete(self, call):
self.calls += 1
if self.calls <= self.fail_times:
raise RuntimeError(f"{self.host} unreachable")
return LLMResult(
text="", tool_calls=[{"name": "x", "input": self.payload}], usage={}
)
def _ctor(per_host):
def _make(*, host, model, num_ctx):
return per_host[host]
return _make
class TestScreenFleet:
@patch("prisma.screen.load_item")
@patch("prisma.screen.to_markdown", return_value="# Item")
def test_fans_out_across_hosts_with_per_item_commits(
self, mc_md, mc_load, tmp_path
):
from prisma.screen import run_fleet
path = _setup(tmp_path)
with Db(path) as db:
for i in range(4):
_insert_item(db, f"K{i}", ["project:test"])
mc_load.return_value = MagicMock()
per_host = {
"http://h1": _FakeHostProvider(
"http://h1", "m", 8192, payload={"decision": "include"}
),
"http://h2": _FakeHostProvider(
"http://h2", "m", 8192, payload={"decision": "include"}
),
}
project = MagicMock()
project.name = "test"
project.criteria = "c"
project.reasons = "r"
with patch("prisma.llm.OllamaNativeProvider", side_effect=_ctor(per_host)):
stats = run_fleet(
db,
project,
hosts=["http://h1", "http://h2"],
model="m",
storage_dir=Path(tmp_path),
per_host=1,
)
assert stats["screened"] == 4
assert stats["include"] == 4
assert stats["errors"] == 0
# Both fleet hosts actually did work.
assert per_host["http://h1"].calls > 0
assert per_host["http://h2"].calls > 0
# Per-item commit: every inserted item now carries the stage tag.
with Db(path) as db2:
tagged = db2.con.execute(
"SELECT COUNT(DISTINCT it.itemID) FROM itemTags it "
"JOIN tags t ON it.tagID = t.tagID "
"WHERE t.name LIKE 'stage:2-%'"
).fetchone()[0]
assert tagged == 4
@patch("prisma.screen.load_item")
@patch("prisma.screen.to_markdown", return_value="# Item")
def test_dead_host_items_still_get_processed(self, mc_md, mc_load, tmp_path):
from prisma.screen import run_fleet
path = _setup(tmp_path)
with Db(path) as db:
for i in range(5):
_insert_item(db, f"K{i}", ["project:test"])
mc_load.return_value = MagicMock()
per_host = {
"http://dead": _FakeHostProvider(
"http://dead", "m", 8192, fail_times=10**6
),
"http://ok": _FakeHostProvider(
"http://ok", "m", 8192, payload={"decision": "uncertain"}
),
}
project = MagicMock()
project.name = "test"
project.criteria = "c"
project.reasons = "r"
with patch("prisma.llm.OllamaNativeProvider", side_effect=_ctor(per_host)):
stats = run_fleet(
db,
project,
hosts=["http://dead", "http://ok"],
model="m",
storage_dir=Path(tmp_path),
per_host=1,
)
# The dead host is retired after 2 consecutive failures — it
# never eats more than 2 items; the rest complete via "ok".
assert per_host["http://dead"].calls == 2
assert stats["errors"] == 2
assert stats["screened"] == 3
class TestEligibilityFleet:
@patch("prisma.eligibility.load_item")
@patch("prisma.eligibility.to_markdown", return_value="# Item\n# Full text\nbody")
def test_fans_out_across_hosts(self, mc_md, mc_load, tmp_path):
from prisma.eligibility import run_fleet
path = _setup(tmp_path)
with Db(path) as db:
for i in range(4):
_insert_item(
db,
f"K{i}",
["project:test", "stage:2-title-abstract", "screen:include"],
)
mc_load.return_value = MagicMock()
per_host = {
"http://h1": _FakeHostProvider(
"http://h1", "m", 8192, payload={"decision": "include"}
),
"http://h2": _FakeHostProvider(
"http://h2", "m", 8192, payload={"decision": "include"}
),
}
project = MagicMock()
project.name = "test"
project.criteria = "c"
project.reasons = "r"
with patch("prisma.llm.OllamaNativeProvider", side_effect=_ctor(per_host)):
stats = run_fleet(
db,
project,
hosts=["http://h1", "http://h2"],
model="m",
storage_dir=Path(tmp_path),
per_host=1,
)
assert stats["assessed"] == 4
assert stats["errors"] == 0
assert per_host["http://h1"].calls > 0
assert per_host["http://h2"].calls > 0
@patch("prisma.eligibility.load_item")
@patch("prisma.eligibility.to_markdown", return_value="# Item (no full text)")
def test_missing_fulltext_resolved_without_llm_call(self, mc_md, mc_load, tmp_path):
"""Items with no retrievable full text bypass the fleet entirely
(same deterministic outcome as the sequential run())."""
from prisma.eligibility import run_fleet
path = _setup(tmp_path)
with Db(path) as db:
_insert_item(
db, "K0", ["project:test", "stage:2-title-abstract", "screen:include"]
)
mc_load.return_value = MagicMock()
per_host = {
"http://h1": _FakeHostProvider("http://h1", "m", 8192),
}
project = MagicMock()
project.name = "test"
project.criteria = "c"
project.reasons = "r"
with patch("prisma.llm.OllamaNativeProvider", side_effect=_ctor(per_host)):
stats = run_fleet(
db,
project,
hosts=["http://h1"],
model="m",
storage_dir=Path(tmp_path),
)
assert stats["assessed"] == 1
assert stats["uncertain"] == 1
assert per_host["http://h1"].calls == 0
class TestExtractFleet:
@patch("prisma.extract.load_item")
@patch("prisma.extract.to_markdown", return_value="# Item")
def test_fans_out_across_hosts(self, mc_md, mc_load, tmp_path):
from prisma.extract import run_fleet
path = _setup(tmp_path)
with Db(path) as db:
for i in range(4):
_insert_item(db, f"K{i}", ["project:test", "stage:4-included"])
mc_load.return_value = MagicMock()
per_host = {
"http://h1": _FakeHostProvider(
"http://h1", "m", 8192, payload={"population": "Medicare"}
),
"http://h2": _FakeHostProvider(
"http://h2", "m", 8192, payload={"population": "Medicare"}
),
}
project = MagicMock()
project.name = "test"
project.extraction_template = (
"- **population** — study population\n- **outcome** — primary outcome"
)
with patch("prisma.llm.OllamaNativeProvider", side_effect=_ctor(per_host)):
stats = run_fleet(
db,
project,
hosts=["http://h1", "http://h2"],
model="m",
storage_dir=Path(tmp_path),
per_host=1,
)
assert stats["extracted"] == 4
assert stats["errors"] == 0
assert per_host["http://h1"].calls > 0
assert per_host["http://h2"].calls > 0