Files
stack/tests/llm/test_evidence.py

760 lines
27 KiB
Python

"""llm.evidence — valuation prompt block, SSE payload, code-cited rule sources."""
from __future__ import annotations
from concurrent.futures import ThreadPoolExecutor
from dataclasses import replace
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import duckdb
import pytest
from llm import evidence
from llm.config import LlmConfig
from llm.evidence import (
ValuationEvidence,
code_cited_sources,
merge_sources,
valuation_evidence,
warm,
)
from pfs.valuation import ValuationRow
CFG = LlmConfig(
ollama_hosts=("http://h1:11434",),
embed_model="e",
instruct_model="c",
embed_dim=768,
build_ann_index=False,
pg_host="x",
pg_port=5432,
pg_db="llm",
pg_user="llm",
duckdb_replica="/nonexistent/aco.ro.duckdb",
valuation_years=4,
code_cited_per_code=2,
)
@pytest.fixture(autouse=True)
def _clear_replica_cache():
"""The replica handle is process-global; no test may inherit another's."""
evidence._REPLICA = None
yield
evidence._REPLICA = None
def _row(**kw) -> ValuationRow:
base = dict(
code="G0556",
description="Adv prim care mgmt lvl 1",
vintage="CY2026 final",
year=2026,
proposed=False,
status="A",
work=0.25,
pe_nf=0.22,
pe_f=0.06,
mp=0.02,
total_nf=0.49,
total_f=0.33,
cf=33.4009,
pay_nf=16.37,
pay_f=11.02,
label="[PFS CY2026 Addendum B]",
citation="90 FR 49266",
url="https://www.federalregister.gov/citation/90-FR-49266",
)
base.update(kw)
return ValuationRow(**base)
class TestEvidenceText:
def test_prompt_block_exact(self):
ev = ValuationEvidence(
codes=("G0556",), families=("APCM",), rows=(_row(),), unpriced=("Z9999",)
)
assert ev.prompt_block() == (
"Valuation (authoritative for RVUs, conversion factors and payments; "
"national unadjusted, no GPCI):\n"
'[PFS CY2026 Addendum B] G0556 "Adv prim care mgmt lvl 1" — status A; work 0.25; '
"PE 0.22 non-fac / 0.06 fac; MP 0.02; total 0.49 non-fac / 0.33 fac; "
"CF $33.4009 → payment $16.37 non-fac / $11.02 fac\n"
"Not priced in the indexed fee schedule: Z9999"
)
def test_prompt_block_labels_the_cf_and_an_unpaid_status(self):
ev = ValuationEvidence(
("G0559",),
(),
(
_row(
code="G0559",
status="B",
status_note="bundled — no separate PFS payment",
cf_note="non-APM standard CF; QP CF $33.5675",
pay_nf=None,
pay_f=None,
),
),
(),
)
line = ev.prompt_block().splitlines()[1]
assert "— status B; bundled — no separate PFS payment; work" in line
assert "CF $33.4009 [non-APM standard CF; QP CF $33.5675] → payment" in line
assert "payment n/a non-fac / n/a fac" in line
def test_prompt_block_null_components(self):
ev = ValuationEvidence(
("G9999",),
(),
(
_row(
code="G9999",
status="I",
work=None,
pe_nf=None,
pe_f=None,
mp=None,
total_nf=None,
total_f=None,
pay_nf=None,
pay_f=None,
),
),
(),
)
line = ev.prompt_block().splitlines()[1]
assert (
"status I; work n/a; PE n/a non-fac / n/a fac" in line
and "payment n/a" in line
)
def test_payload_shape(self):
ev = ValuationEvidence(
("G0556",),
("APCM",),
(
_row(),
_row(
year=2027,
proposed=True,
vintage="CY2027 proposed",
label="[CY2027 NPRM Addendum B]",
citation="91 FR 43842",
url="https://www.federalregister.gov/citation/91-FR-43842",
),
),
(),
)
p = ev.payload()
assert (
p["type"] == "valuation"
and p["codes"] == ["G0556"]
and p["families"] == ["APCM"]
)
assert p["rows"][0]["pay_nf"] == 16.37 and p["rows"][1]["proposed"] is True
assert p["rows"][0]["status_note"] == "" and p["rows"][0]["cf_note"] == ""
assert p["provenance"] == [
{
"label": "[PFS CY2026 Addendum B]",
"vintage": "CY2026 final",
"citation": "90 FR 49266",
"url": "https://www.federalregister.gov/citation/90-FR-49266",
},
{
"label": "[CY2027 NPRM Addendum B]",
"vintage": "CY2027 proposed",
"citation": "91 FR 43842",
"url": "https://www.federalregister.gov/citation/91-FR-43842",
},
]
assert p["unpriced"] == []
class TestValuationEvidence:
def test_none_when_no_codes(self):
assert (
valuation_evidence("what did commenters say about telehealth?", CFG) is None
)
@patch("llm.evidence.valuation", return_value=([_row()], []))
@patch("llm.evidence.duckdb.connect")
def test_detects_and_queries_replica_read_only(self, mock_connect, mock_val):
ev = valuation_evidence("How is APCM valued?", CFG)
mock_connect.assert_called_once_with(
"/nonexistent/aco.ro.duckdb", read_only=True
)
assert mock_val.call_args.args[1] == ["G0556", "G0557", "G0558"]
assert mock_val.call_args.kwargs == {"years": 4}
assert (
ev is not None
and ev.codes == ("G0556", "G0557", "G0558")
and ev.families == ("APCM",)
)
mock_connect.return_value.close.assert_not_called() # cached, not closed
@patch("llm.evidence.duckdb.connect", side_effect=OSError("no replica"))
def test_unopenable_replica_yields_none(self, _c, caplog):
assert valuation_evidence("G0556?", CFG) is None
assert "valuation evidence skipped" in caplog.text
@patch("llm.evidence.valuation", return_value=([_row()], []))
@patch("llm.evidence.duckdb.connect")
def test_replica_connection_is_reused(self, mock_connect, _val):
valuation_evidence("G0556?", CFG)
valuation_evidence("G0557?", CFG)
mock_connect.assert_called_once()
mock_connect.return_value.close.assert_not_called()
@patch("llm.evidence.valuation", return_value=([_row()], []))
@patch("llm.evidence.duckdb.connect")
def test_republished_replica_reopens_and_closes_the_old_handle(
self, mock_connect, _val, monkeypatch
):
first, second = MagicMock(), MagicMock()
mock_connect.side_effect = [first, second]
mtimes = iter([1, 2])
monkeypatch.setattr(evidence, "_mtime", lambda _p: next(mtimes))
valuation_evidence("G0556?", CFG)
valuation_evidence("G0557?", CFG)
assert mock_connect.call_count == 2
first.close.assert_called_once()
second.close.assert_not_called()
@patch("llm.evidence.valuation", return_value=([], []))
@patch("llm.evidence.duckdb.connect")
def test_many_codes_shorten_the_window(self, _c, mock_val):
valuation_evidence("How is CCM valued?", CFG) # six-code family
assert mock_val.call_args.kwargs == {"years": 2}
class TestWarm:
"""#699 ruling B3: warm the replica connection ahead of the first
real request — at API startup and again (a no-op once cached) at the
top of every ``valuation_evidence`` call."""
def test_warm_opens_once_and_is_a_noop_on_replay(self, tmp_path):
path = tmp_path / "aco.ro.duckdb"
path.touch() # warm() only needs the file to exist, not be a real db
cfg = replace(CFG, duckdb_replica=str(path))
with patch("llm.evidence.duckdb.connect") as mock_connect:
mock_connect.return_value = MagicMock()
warm(cfg)
warm(cfg)
mock_connect.assert_called_once_with(str(path), read_only=True)
def test_warm_is_a_noop_when_the_replica_file_is_missing(self):
with patch("llm.evidence.duckdb.connect") as mock_connect:
warm(CFG) # CFG.duckdb_replica ("/nonexistent/...") doesn't exist
mock_connect.assert_not_called()
def test_warm_swallows_connect_failures(self, tmp_path, caplog):
path = tmp_path / "aco.ro.duckdb"
path.touch()
cfg = replace(CFG, duckdb_replica=str(path))
with patch("llm.evidence.duckdb.connect", side_effect=OSError("boom")):
warm(cfg) # must not raise
assert "replica warm-up skipped" in caplog.text
@patch("llm.evidence.detect_codes")
@patch("llm.evidence.warm")
def test_valuation_evidence_warms_before_detecting_codes(
self, mock_warm, mock_detect
):
order: list[str] = []
mock_warm.side_effect = lambda cfg: order.append("warm")
def _detect(question):
order.append("detect")
return SimpleNamespace(codes=())
mock_detect.side_effect = _detect
assert valuation_evidence("anything", CFG) is None
assert order == ["warm", "detect"]
mock_warm.assert_called_once_with(CFG)
class TestConnectRefreshesFamilies:
"""#699: derived families (``pfs.code_family``) live only on the
replica — ``_connect`` is the chat's one hook to pick them up, on
every open and every reopen (a republished replica)."""
@patch("llm.evidence.refresh_from")
@patch("llm.evidence.duckdb.connect")
def test_connect_refreshes_families_on_open_and_reopen(
self, mock_connect, mock_refresh, monkeypatch
):
mtimes = iter([1, 1, 2])
monkeypatch.setattr(evidence, "_mtime", lambda _p: next(mtimes))
mock_connect.side_effect = [MagicMock(), MagicMock()]
con1 = evidence._connect("/x")
assert mock_refresh.call_count == 1
mock_refresh.assert_called_with(con1)
con2 = evidence._connect("/x") # same (path, mtime) — cached, no refresh
assert con2 is con1
assert mock_refresh.call_count == 1
con3 = evidence._connect("/x") # mtime changed — reopen, refresh again
assert mock_refresh.call_count == 2
mock_refresh.assert_called_with(con3)
@patch("llm.evidence.refresh_from", side_effect=RuntimeError("boom"))
@patch("llm.evidence.duckdb.connect")
def test_refresh_failure_does_not_break_the_handle(
self, _mock_connect, _mock_refresh, caplog
):
con = evidence._connect("/x")
assert con is not None
assert "family refresh skipped" in caplog.text
RVU_COLS = (
"hcpcs VARCHAR, mod VARCHAR, description VARCHAR, status_code VARCHAR, "
"work_rvu DOUBLE, non_fac_pe_rvu DOUBLE, fac_pe_rvu DOUBLE, mp_rvu DOUBLE, "
"non_fac_total DOUBLE, fac_total DOUBLE, conv_factor DOUBLE, year INTEGER"
)
PROPOSED_COLS = (
"hcpcs VARCHAR, mod VARCHAR, description VARCHAR, status_code VARCHAR, "
"work_rvu DOUBLE, non_fac_pe_rvu DOUBLE, fac_pe_rvu DOUBLE, mp_rvu DOUBLE, "
"cms_rule_id VARCHAR"
)
class TestConcurrentChats:
"""/chat streams run in Starlette's threadpool, so several turns share
the cached replica handle. Without a per-call cursor their result sets
cross-talk and the broad except swallows it as "no valuation"."""
@pytest.fixture
def replica(self, tmp_path):
db = tmp_path / "aco.ro.duckdb"
con = duckdb.connect(str(db))
con.execute("CREATE SCHEMA pfs")
con.execute(f"CREATE TABLE pfs.rvu ({RVU_COLS})")
con.execute(f"CREATE TABLE pfs.rvu_proposed ({PROPOSED_COLS})")
# two unfamilied codes with different row counts, so a swapped
# result set is visible as a wrong count
con.executemany(
"INSERT INTO pfs.rvu VALUES (?,?,?,?,?,?,?,?,?,?,?,?)",
[
(
"99213",
None,
"Office visit",
"A",
1.0,
1.0,
0.5,
0.05,
2.05,
1.55,
30.0,
2026,
),
*[
(
"99214",
None,
"Office visit",
"A",
1.5,
1.0,
0.5,
0.05,
2.55,
2.05,
30.0,
y,
)
for y in (2024, 2025, 2026)
],
],
)
con.close()
return db
def test_parallel_turns_share_one_handle_without_cross_talk(self, replica):
cfg = replace(CFG, duckdb_replica=str(replica))
real_connect, opens = duckdb.connect, []
def counting_connect(*a, **kw):
opens.append(a)
return real_connect(*a, **kw)
expected = {"99213": 1, "99214": 3}
questions = ["What does 99213 pay?", "What does 99214 pay?"]
def one(i):
q = questions[i % 2]
ev = valuation_evidence(q, cfg)
assert ev is not None, "valuation dropped under concurrency"
(code,) = ev.codes
return code, len(ev.rows)
with patch("llm.evidence.duckdb.connect", counting_connect):
with ThreadPoolExecutor(max_workers=8) as pool:
results = list(pool.map(one, range(160)))
assert all(n == expected[code] for code, n in results), sorted(set(results))
assert len(opens) == 1 # one handle, shared
def test_cursor_is_closed_but_the_handle_is_kept(self, replica):
cfg = replace(CFG, duckdb_replica=str(replica))
assert valuation_evidence("99213?", cfg) is not None
assert evidence._REPLICA is not None
# the cached parent still answers — it was never closed
assert valuation_evidence("99214?", cfg) is not None
def _engine(rows):
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
conn.execute.return_value.fetchall.return_value = rows
return engine, conn
def _rule_md(p_id, codes, date="2024-12-09"):
return {
"kind": "rule",
"item_key": "R1",
"p_id": str(p_id),
"page": "97710",
"ordinal": "1",
"html_url": "https://www.federalregister.gov/d/2024-25382",
"fr_volume": "89",
"date": date,
"title": "CY2025 PFS final rule",
"codes": codes,
}
def _comment_md(seq, codes, item_key="C1", date="2024-12-09"):
return {
"kind": "comment",
"item_key": item_key,
"seq": str(seq),
"comment_id": "CMS-2026-2377-3438",
"date": date,
"codes": codes,
"families": "",
}
def _corpus_md(seq, codes="", families="", item_key="CPT-ED", date="2024-12-09"):
return {
"kind": "corpus",
"item_key": item_key,
"seq": str(seq),
"url": "https://example.org/cpt.pdf",
"title": "CPT Manual",
"date": date,
"codes": codes,
"families": families,
}
def _multi_engine(by_code=None, by_family=None):
"""A fake engine whose rows depend on which collection *and* which of
the two SQL statements (codes vs. families) is being executed —
distinguished by which bind parameter the call carries."""
engine = MagicMock()
conn = engine.begin.return_value.__enter__.return_value
by_code = by_code or {}
by_family = by_family or {}
def _execute(_sql, params):
collection = params["collection"]
table = by_code if "codes" in params else by_family
result = MagicMock()
result.fetchall.return_value = table.get(collection, [])
return result
conn.execute.side_effect = _execute
return engine, conn
class TestCodeCitedSources:
def test_sql_and_per_code_cap_and_dedupe(self):
engine, conn = _engine(
[
("para one G0556", _rule_md(1, "G0556")),
("para two G0556", _rule_md(2, "G0556")),
("para three G0556", _rule_md(3, "G0556")), # over the cap for G0556
(
"para four G0557",
_rule_md(4, "G0557 G0556"),
), # G0557 still under cap
("para one again", _rule_md(1, "G0556")), # duplicate (item_key, p_id)
]
)
out = code_cited_sources(engine, ["G0556", "G0557"], per_code=2)
sql = str(conn.execute.call_args.args[0])
assert "string_to_array" in sql and "langchain_pg_collection" in sql
assert "PARTITION BY" in sql
# the collection is a scalar subquery, not a join to filter after
assert "JOIN langchain_pg_collection" not in sql
assert "e.collection_id = (SELECT uuid FROM langchain_pg_collection" in sql
assert conn.execute.call_args.args[1]["collection"] == "rules"
assert conn.execute.call_args.args[1]["codes"] == ["G0556", "G0557"]
assert conn.execute.call_args.args[1]["window"] == 4
assert [s["snippet"] for s in out] == [
"para one G0556",
"para two G0556",
"para four G0557",
]
assert all(s["kind"] == "rule" and s["score"] == 0.0 for s in out)
def test_snippets_are_capped(self):
engine, _ = _engine([("x" * 900, _rule_md(1, "G0556"))])
(out,) = code_cited_sources(engine, ["G0556"], per_code=2)
assert len(out["snippet"]) == 250
def test_engine_error_yields_empty(self, caplog):
engine = MagicMock()
engine.begin.side_effect = RuntimeError("pg down")
assert code_cited_sources(engine, ["G0556"], per_code=2) == []
assert "code-cited sources skipped" in caplog.text
def test_no_codes_returns_empty_without_querying(self):
engine = MagicMock()
assert code_cited_sources(engine, [], per_code=2) == []
engine.begin.assert_not_called()
class TestCodeCitedSourcesMultiCollection:
def test_concatenated_rules_then_comments_then_corpus(self):
engine, _conn = _multi_engine(
by_code={
"rules": [("rule chunk", _rule_md(1, "G0556"))],
"comments": [("comment chunk", _comment_md(1, "G0556"))],
"corpus": [("corpus chunk", _corpus_md(1, "G0556"))],
}
)
out = code_cited_sources(
engine, ["G0556"], per_code=2, collections=("rules", "comments", "corpus")
)
assert [s["kind"] for s in out] == ["rule", "comment", "corpus"]
def test_per_collection_cap_applies_independently(self):
engine, _conn = _multi_engine(
by_code={
"rules": [
("r1", _rule_md(1, "G0556")),
("r2", _rule_md(2, "G0556")),
("r3", _rule_md(3, "G0556")), # over the per-collection cap
],
"comments": [
("c1", _comment_md(1, "G0556")),
("c2", _comment_md(2, "G0556")),
],
}
)
out = code_cited_sources(
engine, ["G0556"], per_code=2, collections=("rules", "comments")
)
# Capped to 2/collection, then round-robin interleaved across
# collections (rules, comments), not simply concatenated.
assert [s["snippet"] for s in out] == ["r1", "c1", "r2", "c2"]
def test_dedupe_uses_seq_for_non_rule_kinds(self):
engine, _conn = _multi_engine(
by_code={
"comments": [
("c1", _comment_md(1, "G0556", item_key="C1")),
("c1 dup", _comment_md(1, "G0556", item_key="C1")), # same seq
("c2", _comment_md(2, "G0556", item_key="C1")), # different seq
]
}
)
out = code_cited_sources(
engine, ["G0556"], per_code=5, collections=("comments",)
)
assert [s["snippet"] for s in out] == ["c1", "c2"]
class TestFamilyCitedSources:
def test_skipped_entirely_when_no_families(self):
engine, conn = _multi_engine(by_code={"rules": [("r1", _rule_md(1, "G0556"))]})
code_cited_sources(engine, ["G0556"], per_code=2, collections=("rules",))
assert conn.execute.call_count == 1 # codes query only, no family round trip
def test_sql_shape_and_params(self):
engine, conn = _multi_engine(by_family={"rules": []})
code_cited_sources(
engine, [], per_code=2, collections=("rules",), families=["CCM"]
)
sql = str(conn.execute.call_args.args[0])
assert "&&" in sql
# one chunk per (family, item_key), then the survivors ranked per family
assert sql.count("PARTITION BY") == 2
assert "PARTITION BY w.family, e.cmetadata->>'item_key'" in sql
assert "PARTITION BY family" in sql
assert conn.execute.call_args.args[1] == {
"collection": "rules",
"families": ["CCM"],
"window": 4,
}
def test_called_once_per_collection_with_same_params(self):
engine, conn = _multi_engine(
by_family={"rules": [], "comments": [], "corpus": []}
)
code_cited_sources(
engine,
[],
per_code=3,
collections=("rules", "comments", "corpus"),
families=["CCM"],
)
assert conn.execute.call_count == 3
for call, collection in zip(
conn.execute.call_args_list, ("rules", "comments", "corpus")
):
assert call.args[1] == {
"collection": collection,
"families": ["CCM"],
"window": 6,
}
def test_family_rows_appended_after_code_rows(self):
engine, _conn = _multi_engine(
by_code={"rules": [("code hit", _rule_md(1, "G0556"))]},
by_family={
"corpus": [
(
"CCM guideline",
_corpus_md(1, codes="", families="CCM", item_key="CPT-ED"),
)
]
},
)
out = code_cited_sources(
engine,
["G0556"],
per_code=2,
collections=("rules", "corpus"),
families=["CCM"],
)
assert [s["snippet"] for s in out] == ["code hit", "CCM guideline"]
assert out[1]["kind"] == "corpus"
def test_family_dedupes_against_a_code_cited_chunk(self):
md = _rule_md(1, "G0556")
md["families"] = "APCM"
engine, _conn = _multi_engine(
by_code={"rules": [("shared", md)]}, by_family={"rules": [("shared", md)]}
)
out = code_cited_sources(
engine, ["G0556"], per_code=2, collections=("rules",), families=["APCM"]
)
assert [s["snippet"] for s in out] == ["shared"]
class TestMaxTotal:
"""``max_total`` interleaves code/family rows round-robin across
collections, then truncates (#budget)."""
def test_interleaves_three_collections_round_robin(self):
engine, _conn = _multi_engine(
by_code={
"rules": [("r1", _rule_md(1, "G0556")), ("r2", _rule_md(2, "G0556"))],
"comments": [
("c1", _comment_md(1, "G0556")),
("c2", _comment_md(2, "G0556")),
],
"corpus": [
("x1", _corpus_md(1, codes="G0556", item_key="X1")),
("x2", _corpus_md(2, codes="G0556", item_key="X2")),
],
}
)
out = code_cited_sources(
engine,
["G0556"],
per_code=5,
collections=("rules", "comments", "corpus"),
max_total=0,
)
assert [s["snippet"] for s in out] == ["r1", "c1", "x1", "r2", "c2", "x2"]
def test_family_rows_interleaved_then_appended_after_code_rows(self):
engine, _conn = _multi_engine(
by_code={
"rules": [("r1", _rule_md(1, "G0556"))],
"comments": [("c1", _comment_md(1, "G0556"))],
},
by_family={
"rules": [
(
"f-rules",
_corpus_md(9, codes="", families="CCM", item_key="FR"),
)
],
"comments": [
(
"f-comments",
_corpus_md(10, codes="", families="CCM", item_key="FC"),
)
],
},
)
out = code_cited_sources(
engine,
["G0556"],
per_code=5,
collections=("rules", "comments"),
families=["CCM"],
max_total=0,
)
assert [s["snippet"] for s in out] == ["r1", "c1", "f-rules", "f-comments"]
def test_truncates_to_max_total(self):
engine, _conn = _multi_engine(
by_code={
"rules": [("r1", _rule_md(1, "G0556")), ("r2", _rule_md(2, "G0556"))],
"comments": [
("c1", _comment_md(1, "G0556")),
("c2", _comment_md(2, "G0556")),
],
}
)
out = code_cited_sources(
engine,
["G0556"],
per_code=5,
collections=("rules", "comments"),
max_total=3,
)
assert [s["snippet"] for s in out] == ["r1", "c1", "r2"]
def test_max_total_zero_or_negative_is_unlimited(self):
rows = [(f"r{i}", _rule_md(i, "G0556")) for i in range(1, 15)]
engine, _conn = _multi_engine(by_code={"rules": rows})
for unlimited in (0, -1):
out = code_cited_sources(
engine,
["G0556"],
per_code=20,
collections=("rules",),
max_total=unlimited,
)
assert len(out) == 14
def test_default_max_total_is_twelve(self):
rows = [(f"r{i}", _rule_md(i, "G0556")) for i in range(1, 15)]
engine, _conn = _multi_engine(by_code={"rules": rows})
out = code_cited_sources(engine, ["G0556"], per_code=20, collections=("rules",))
assert len(out) == 12
class TestMergeSources:
def test_keeps_order_and_dedupes_by_label(self):
a = [{"label": "X", "score": 0.9}, {"label": "Y", "score": 0.8}]
b = [{"label": "Y", "score": 0.0}, {"label": "Z", "score": 0.0}]
assert [s["label"] for s in merge_sources(a, b)] == ["X", "Y", "Z"]
assert merge_sources(a, b)[1]["score"] == 0.8