feat(pfs,llm): derived families reach the chat — indexed detect_codes/family_of, refresh on replica open (refs #699)

This commit is contained in:
kert
2026-09-10 09:16:59 -04:00
parent a60cb6ea7e
commit 5f31d7b645
5 changed files with 295 additions and 20 deletions

View File

@@ -206,6 +206,7 @@ ignore = ["E501", "E741"]
testpaths = ["tests"] testpaths = ["tests"]
markers = [ markers = [
"stub: marks tests that report stub vs implemented status (deselect with '-m \"not stub\"')", "stub: marks tests that report stub vs implemented status (deselect with '-m \"not stub\"')",
"live: marks tests that need a live DuckDB replica on disk (skipped without one)",
] ]
filterwarnings = [ filterwarnings = [
"error::ResourceWarning", "error::ResourceWarning",

View File

@@ -25,7 +25,7 @@ from sqlalchemy import text
from llm.config import LlmConfig from llm.config import LlmConfig
from llm.links import as_source from llm.links import as_source
from pfs.families import detect_codes from pfs.families import detect_codes, refresh_from
from pfs.valuation import ValuationRow, valuation from pfs.valuation import ValuationRow, valuation
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@@ -129,6 +129,14 @@ def _connect(path: str) -> Any:
log.warning("stale replica handle not closed: %s", e) log.warning("stale replica handle not closed: %s", e)
con = duckdb.connect(path, read_only=True) con = duckdb.connect(path, read_only=True)
_REPLICA = (path, mtime, con) _REPLICA = (path, mtime, con)
# Derived families (pfs.code_family) live only on the replica —
# every (re)open is the chat's one chance to pick them up. Never
# let a bad/missing table take the chat down over it.
try:
n = refresh_from(con)
log.info("families refreshed from replica %s: %d", path, n)
except Exception as e: # noqa: BLE001 — the chat must still get its handle
log.warning("family refresh skipped (%s): %s", path, e)
return con return con

View File

@@ -10,6 +10,7 @@ from __future__ import annotations
import logging import logging
import re import re
import threading
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Mapping, Sequence from typing import Any, Mapping, Sequence
@@ -110,13 +111,128 @@ HAND_FAMILIES: dict[str, Family] = {
#: follow-up issue — the chat sees hand families only. #: follow-up issue — the chat sees hand families only.
FAMILIES: dict[str, Family] = dict(HAND_FAMILIES) FAMILIES: dict[str, Family] = dict(HAND_FAMILIES)
#: Guards the swap of ``FAMILIES``/``_PHRASE_RE``/``_PHRASE_INDEX``/
#: ``_CODE_INDEX`` in ``refresh_from`` — the chat runs each turn in a
#: threadpool worker, and a reader must never see a half-updated index.
_REGISTRY_LOCK = threading.Lock()
#: A derived family is matched by name only when its name reads as a real
#: phrase, not a bare CPT heading word (``GENE``, ``ADM``, ``Introduction``).
_MIN_PHRASE_LEN = 12
_MIN_PHRASE_WORDS = 2
#: Matches nothing — the safe ``_PHRASE_RE`` value when the registry (or a
#: monkeypatched ``FAMILIES``) has no qualifying phrase at all, so
#: ``detect_codes`` never has to special-case an empty alternation.
_NEVER_RE = re.compile(r"(?!x)x")
def _phrase_qualifies(name: str) -> bool:
return len(name) >= _MIN_PHRASE_LEN and len(name.split()) >= _MIN_PHRASE_WORDS
def _trie_alternation(phrases: Sequence[str]) -> str:
"""A regex alternation over *phrases*, sharing common prefixes as a
character trie instead of one flat ``a|b|c|...``.
The live registry contributes ~7,600 qualifying name phrases; a flat
``\\b(?:p1|p2|...)\\b`` alternation makes ``re`` retry every phrase at
every non-matching text position (~5ms on a 300-character question —
over the 5ms budget). A trie collapses that to one branch per
distinct next character (mostly <= 26), the standard fix for large
alternations in Python's backtracking ``re`` engine. Because a
shared-prefix node's continuation is optional (``(?:...)?``) and
``?`` is greedy, the trie already prefers the longest phrase that
matches at a position — no separate "longest first" sort needed the
way a flat alternation requires."""
root: dict[str, Any] = {}
for p in phrases:
node = root
for ch in p:
node = node.setdefault(ch, {})
node["\0"] = True # end-of-phrase marker; never a real dict key
def _compile(node: dict[str, Any]) -> str:
end = "\0" in node
chars = sorted(k for k in node if k != "\0")
if not chars:
return "" # end() only — caller makes it optional
alts = [re.escape(ch) + _compile(node[ch]) for ch in chars]
body = alts[0] if len(alts) == 1 else "(?:" + "|".join(alts) + ")"
return f"{body}?" if end else body
return _compile(root)
def _build_registry_index(
families: Mapping[str, Family],
) -> tuple[re.Pattern[str], dict[str, tuple[str, ...]], dict[str, tuple[str, ...]]]:
"""The compiled phrase alternation, its phrase -> family-keys lookup,
and the code -> family-keys index for *families* — the three pieces
``rebuild_index``/``refresh_from`` swap in together.
A hand family is matched by any of its curated ``synonyms``; a
derived (non-hand) family is matched by its own ``name`` too, but
only when the name is a real phrase (``_phrase_qualifies``) — a
one-word CPT heading title must never become a chat trigger.
"""
phrase_families: dict[str, set[str]] = {}
for key, fam in families.items():
phrases = list(fam.synonyms)
if key not in HAND_FAMILIES and _phrase_qualifies(fam.name):
phrases.append(fam.name.lower())
for p in phrases:
phrase_families.setdefault(p, set()).add(key)
if phrase_families:
alt = _trie_alternation(sorted(phrase_families))
phrase_re = re.compile(rf"\b(?:{alt})\b", re.IGNORECASE)
else:
phrase_re = _NEVER_RE
phrase_index = {p: tuple(sorted(keys)) for p, keys in phrase_families.items()}
code_families: dict[str, set[str]] = {}
for key, fam in families.items():
for code in fam.codes:
code_families.setdefault(code, set()).add(key)
def _key_order(k: str) -> tuple[int, str]:
return (0, k) if k in HAND_FAMILIES else (1, k)
code_index = {
code: tuple(sorted(keys, key=_key_order))
for code, keys in code_families.items()
}
return phrase_re, phrase_index, code_index
#: Compiled once at import (seeded from ``HAND_FAMILIES`` below) and
#: rebuilt by ``rebuild_index``/``refresh_from`` whenever the registry
#: changes — never recomputed per ``detect_codes``/``family_of`` call.
_PHRASE_RE: re.Pattern[str] = _NEVER_RE
_PHRASE_INDEX: dict[str, tuple[str, ...]] = {}
_CODE_INDEX: dict[str, tuple[str, ...]] = {}
def rebuild_index() -> None:
"""Recompile ``_PHRASE_RE`` and ``_CODE_INDEX`` from the live
``FAMILIES`` registry. Call after mutating ``FAMILIES`` directly (the
``stack pfs`` CLI does, in tests); ``refresh_from`` builds its own
replacement registry and index off to the side and swaps everything
under ``_REGISTRY_LOCK`` instead of calling this."""
global _PHRASE_RE, _PHRASE_INDEX, _CODE_INDEX
_PHRASE_RE, _PHRASE_INDEX, _CODE_INDEX = _build_registry_index(FAMILIES)
rebuild_index() # seed the index from HAND_FAMILIES
def family_of(code: str) -> Family | None: def family_of(code: str) -> Family | None:
code = code.upper() """O(1) via ``_CODE_INDEX``. When *code* belongs to more than one
for fam in FAMILIES.values(): family, the first key in sorted order wins — hand families before
if code in fam.codes: derived ones, then lexicographic."""
return fam keys = _CODE_INDEX.get(code.upper())
return None return FAMILIES.get(keys[0]) if keys else None
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -128,20 +244,21 @@ class Detection:
def detect_codes(text: str) -> Detection: def detect_codes(text: str) -> Detection:
"""Codes a question is about: explicit codes plus every code of any """Codes a question is about: explicit codes plus every code of any
family named (by synonym) or touched (by one member code).""" family named (by synonym, or by a qualifying derived name) or touched
(by one member code). Runs off the precompiled ``_PHRASE_RE``/
``_CODE_INDEX`` — no per-family scan."""
explicit = find_codes(text) explicit = find_codes(text)
lowered = text.lower() lowered = text.lower()
families: set[str] = set() families: set[str] = set()
for fam in FAMILIES.values(): for m in _PHRASE_RE.finditer(lowered):
if any(re.search(rf"\b{re.escape(s)}\b", lowered) for s in fam.synonyms): families.update(_PHRASE_INDEX.get(m.group(0), ()))
families.add(fam.key)
for code in explicit: for code in explicit:
fam = family_of(code) families.update(_CODE_INDEX.get(code, ()))
if fam is not None:
families.add(fam.key)
codes = set(explicit) codes = set(explicit)
for key in families: for key in families:
codes.update(FAMILIES[key].codes) fam = FAMILIES.get(key)
if fam is not None:
codes.update(fam.codes)
return Detection( return Detection(
codes=tuple(sorted(codes)), codes=tuple(sorted(codes)),
families=tuple(sorted(families)), families=tuple(sorted(families)),
@@ -799,12 +916,23 @@ def refresh_from(con: Any) -> int:
derived (superset) codes — ``load_families`` already carries the derived (superset) codes — ``load_families`` already carries the
hand's own name/synonyms forward for that key (Ruling 17: a bare hand's own name/synonyms forward for that key (Ruling 17: a bare
``FAMILIES.clear()`` used to wipe every hand family the moment any ``FAMILIES.clear()`` used to wipe every hand family the moment any
derived row loaded).""" derived row loaded).
Idempotent and safe to call from a chat thread: the new registry and
its phrase/code index are built off to the side first, then swapped
into ``FAMILIES``/``_PHRASE_RE``/``_CODE_INDEX`` together under
``_REGISTRY_LOCK`` — a concurrent ``detect_codes``/``family_of`` call
always sees either the old registry+index or the new one, never a
mix."""
derived = load_families(con) derived = load_families(con)
if not derived: if not derived:
return 0 return 0
merged = dict(HAND_FAMILIES) merged = dict(HAND_FAMILIES)
merged.update(derived) merged.update(derived)
phrase_re, phrase_index, code_index = _build_registry_index(merged)
global _PHRASE_RE, _PHRASE_INDEX, _CODE_INDEX
with _REGISTRY_LOCK:
FAMILIES.clear() FAMILIES.clear()
FAMILIES.update(merged) FAMILIES.update(merged)
_PHRASE_RE, _PHRASE_INDEX, _CODE_INDEX = phrase_re, phrase_index, code_index
return len(FAMILIES) return len(FAMILIES)

View File

@@ -228,6 +228,42 @@ class TestValuationEvidence:
assert mock_val.call_args.kwargs == {"years": 2} assert mock_val.call_args.kwargs == {"years": 2}
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 = ( RVU_COLS = (
"hcpcs VARCHAR, mod VARCHAR, description VARCHAR, status_code VARCHAR, " "hcpcs VARCHAR, mod VARCHAR, description VARCHAR, status_code VARCHAR, "
"work_rvu DOUBLE, non_fac_pe_rvu DOUBLE, fac_pe_rvu DOUBLE, mp_rvu DOUBLE, " "work_rvu DOUBLE, non_fac_pe_rvu DOUBLE, fac_pe_rvu DOUBLE, mp_rvu DOUBLE, "

View File

@@ -22,6 +22,7 @@ from pfs.families import (
FAMILIES, FAMILIES,
HAND_FAMILIES, HAND_FAMILIES,
Detection, Detection,
Family,
_cpt_edges, _cpt_edges,
cpt_groups, cpt_groups,
derive_families, derive_families,
@@ -29,6 +30,7 @@ from pfs.families import (
family_of, family_of,
find_codes, find_codes,
load_families, load_families,
rebuild_index,
refresh_from, refresh_from,
stem_tokens, stem_tokens,
) )
@@ -36,9 +38,11 @@ from pfs.families import (
@pytest.fixture @pytest.fixture
def restore_families(): def restore_families():
"""Snapshot ``pfs.families.FAMILIES`` and restore it in teardown, even """Snapshot ``pfs.families.FAMILIES`` and restore it (and the
if the test body raises — a test that calls ``refresh_from`` must not compiled phrase/code index built off it) in teardown, even if the
leave the live registry clobbered for the rest of the session.""" test body raises — a test that calls ``refresh_from``/mutates
``FAMILIES`` directly must not leave the live registry or its index
clobbered for the rest of the session."""
from pfs import families as mod from pfs import families as mod
before = dict(mod.FAMILIES) before = dict(mod.FAMILIES)
@@ -47,6 +51,7 @@ def restore_families():
finally: finally:
mod.FAMILIES.clear() mod.FAMILIES.clear()
mod.FAMILIES.update(before) mod.FAMILIES.update(before)
mod.rebuild_index()
class TestFindCodes: class TestFindCodes:
@@ -120,6 +125,103 @@ class TestDetectCodes:
) )
class TestDetectDerivedFamilies:
"""#699: the chat must see derived (``pfs.code_family``) families,
not only the five hand families — by member code always, by name
only when the name reads as a real phrase."""
def test_detect_derived_family_by_member_code(self, restore_families):
restore_families.FAMILIES["SUTURE-REMOVAL"] = Family(
"SUTURE-REMOVAL", "Removal of Sutures", ("15850", "15851"), ()
)
rebuild_index()
d = detect_codes("removal of sutures 15850")
assert "SUTURE-REMOVAL" in d.families
assert d.codes == ("15850", "15851")
def test_derived_family_name_phrase_requires_two_words(self, restore_families):
restore_families.FAMILIES["GENE"] = Family("GENE", "GENE", ("81400",), ())
restore_families.FAMILIES["WOUND-CARE"] = Family(
"WOUND-CARE", "Complex Wound Debridement Services", ("11042",), ()
)
rebuild_index()
# The single-word heading name never fires, even though the text
# contains it verbatim.
assert "GENE" not in detect_codes("the GENE panel result").families
# A qualifying (>= 12 chars, >= 2 words) derived name does.
d = detect_codes("How are Complex Wound Debridement Services valued?")
assert "WOUND-CARE" in d.families
assert "11042" in d.codes
def test_family_of_is_indexed(self, restore_families):
synthetic = {
f"F{i:04d}": Family(
f"F{i:04d}", f"Synthetic Family {i}", (f"{10000 + i}",), ()
)
for i in range(5000)
}
restore_families.FAMILIES.clear()
restore_families.FAMILIES.update(synthetic)
rebuild_index()
start = time.perf_counter()
for i in range(10000):
code = f"{10000 + (i % 5000)}"
fam = family_of(code)
assert fam is not None and fam.key == f"F{i % 5000:04d}"
elapsed = time.perf_counter() - start
assert elapsed < 0.2, elapsed
def test_hand_families_regression(self):
# Unchanged from TestDetectCodes — the phrase/index refactor must
# not alter a single hand-family detection.
assert detect_codes("What is APCM and how is it valued?") == Detection(
codes=("G0556", "G0557", "G0558"), families=("APCM",), explicit=()
)
d = detect_codes("Tell me about Advanced Primary Care Management.")
assert d.families == ("APCM",)
assert detect_codes("the apcmx code").families == ()
d = detect_codes("How much does G0557 pay?")
assert d.explicit == ("G0557",)
assert d.families == ("APCM",)
assert d.codes == ("G0556", "G0557", "G0558")
assert detect_codes("value of 99213") == Detection(
codes=("99213",), families=(), explicit=("99213",)
)
d = detect_codes("compare CCM and TCM")
assert d.families == ("CCM", "TCM")
assert d.codes == tuple(sorted(FAMILIES["CCM"].codes + FAMILIES["TCM"].codes))
assert detect_codes("what did commenters say about telehealth?") == Detection(
(), (), ()
)
class TestDetectCodesLive:
@pytest.mark.live
def test_detect_codes_under_5ms_on_live_registry(self, restore_families):
from conf import ROOT
replica = ROOT / "data" / "replica" / "aco.ro.duckdb"
if not replica.exists():
pytest.skip(f"no live replica at {replica}")
con = duckdb.connect(str(replica), read_only=True)
try:
n = refresh_from(con)
finally:
con.close()
assert n > 5000, f"expected the full derived registry, got {n} families"
question = (
"Given the history of chronic care management services and "
"99490, how does CY2026 valuation compare to CY2025 for "
"transitional care management, and what changed for advance "
"care planning codes 99497 and 99498 under the proposed rule? "
"Please also note principal care management billing rules."
)[:300]
start = time.perf_counter()
detect_codes(question)
elapsed_ms = (time.perf_counter() - start) * 1000
assert elapsed_ms < 5.0, f"{elapsed_ms:.3f} ms on {n} families"
def _el(code, type_, value, detail=""): def _el(code, type_, value, detail=""):
return ElementRow(code, 2025, type_, value, detail, "", "", 0, 0, "fr") return ElementRow(code, 2025, type_, value, detail, "", "", 0, 0, "fr")