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:
@@ -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. │
|
||||
╰──────────────────────────────────────────────────────────────────────────────╯
|
||||
```
|
||||
|
||||
@@ -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. │
|
||||
╰──────────────────────────────────────────────────────────────────────────────╯
|
||||
```
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
141
src/prisma/fleet.py
Normal 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")
|
||||
@@ -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
|
||||
|
||||
@@ -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
199
tests/prisma/test_fleet.py
Normal 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"]
|
||||
285
tests/prisma/test_run_fleet_stages.py
Normal file
285
tests/prisma/test_run_fleet_stages.py
Normal 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
|
||||
Reference in New Issue
Block a user