270 lines
8.8 KiB
Python
270 lines
8.8 KiB
Python
"""Tests for dev/scripts/llm_golden.py — the golden longitudinal chat
|
|
evaluation (P49 Task 7, refs #692).
|
|
|
|
Loaded by path like the other dev/scripts tests (tests/dev/test_nb_
|
|
issue_filer.py) since dev/scripts/ is not a package. Checkers are
|
|
exercised against canned SSE transcripts (pass and fail cases); the
|
|
YAML set is validated for shape; a live end-to-end test runs entry
|
|
"g2058-replacement" against a real /chat and is skipped unless
|
|
LLM_CHAT_URL is set.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import importlib.util
|
|
import os
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
|
|
_SCRIPT = Path(__file__).resolve().parents[2] / "dev" / "scripts" / "llm_golden.py"
|
|
_spec = importlib.util.spec_from_file_location("_llm_golden", _SCRIPT)
|
|
assert _spec and _spec.loader
|
|
golden = importlib.util.module_from_spec(_spec)
|
|
sys.modules["_llm_golden"] = golden
|
|
_spec.loader.exec_module(golden)
|
|
|
|
Transcript = golden.Transcript
|
|
GOLDEN_SET = Path(__file__).resolve().parent / "golden_lineage.yaml"
|
|
|
|
|
|
def _sources(*rows: dict) -> dict:
|
|
return {"type": "sources", "sources": list(rows)}
|
|
|
|
|
|
def _tokens(*texts: str) -> list[dict]:
|
|
return [{"type": "token", "text": t} for t in texts]
|
|
|
|
|
|
def _lineage(*events: dict) -> dict:
|
|
return {"type": "lineage", "events": list(events)}
|
|
|
|
|
|
# ── check_anchors ────────────────────────────────────────────────
|
|
|
|
|
|
def test_check_anchors_pass_exact_and_range():
|
|
t = Transcript(
|
|
[
|
|
_sources(
|
|
{"item_key": "YBM4IZUS", "p_id": "1578"},
|
|
{"item_key": "DE2VH9PD", "p_id": "1250"},
|
|
)
|
|
]
|
|
)
|
|
result = golden.check_anchors(
|
|
t,
|
|
[
|
|
{"item_key": "YBM4IZUS", "p_id": 1578},
|
|
{"item_key": "DE2VH9PD", "p_id": [1249, 1251]},
|
|
],
|
|
)
|
|
assert result.passed, result.detail
|
|
|
|
|
|
def test_check_anchors_fail_missing():
|
|
t = Transcript([_sources({"item_key": "YBM4IZUS", "p_id": "1578"})])
|
|
result = golden.check_anchors(t, [{"item_key": "DE2VH9PD", "p_id": 1250}])
|
|
assert not result.passed
|
|
assert "DE2VH9PD" in result.detail
|
|
|
|
|
|
def test_check_anchors_fail_pid_outside_range():
|
|
t = Transcript([_sources({"item_key": "DE2VH9PD", "p_id": "1300"})])
|
|
result = golden.check_anchors(t, [{"item_key": "DE2VH9PD", "p_id": [1249, 1251]}])
|
|
assert not result.passed
|
|
|
|
|
|
# ── check_labels ─────────────────────────────────────────────────
|
|
|
|
|
|
def test_check_labels_pass():
|
|
t = Transcript(
|
|
_tokens("As discussed in ", "[CY2021 PFS final ¶1578], G2058 was...")
|
|
)
|
|
result = golden.check_labels(t, [r"CY2021 PFS final ¶1578"])
|
|
assert result.passed, result.detail
|
|
|
|
|
|
def test_check_labels_fail():
|
|
t = Transcript(_tokens("No citations here."))
|
|
result = golden.check_labels(t, [r"CY2021 PFS final"])
|
|
assert not result.passed
|
|
assert "CY2021 PFS final" in result.detail
|
|
|
|
|
|
# ── check_forbidden ──────────────────────────────────────────────
|
|
|
|
|
|
def test_check_forbidden_pass_when_labeled():
|
|
t = Transcript(
|
|
_tokens("The payment is $34.85 [PFS CY2026 Addendum B]. ", "That is all.")
|
|
)
|
|
result = golden.check_forbidden(
|
|
t, [{"pattern": r"\$\d", "unless_label": "Addendum B"}]
|
|
)
|
|
assert result.passed, result.detail
|
|
|
|
|
|
def test_check_forbidden_fail_unlabeled_dollar():
|
|
t = Transcript(_tokens("The payment is roughly $34.85 based on recent rules."))
|
|
result = golden.check_forbidden(
|
|
t, [{"pattern": r"\$\d", "unless_label": "Addendum B"}]
|
|
)
|
|
assert not result.passed
|
|
assert "$" in result.detail or "\\$" in result.detail
|
|
|
|
|
|
# ── check_events ─────────────────────────────────────────────────
|
|
|
|
|
|
def test_check_events_pass_alternatives():
|
|
t = Transcript([_lineage({"code": "99441", "kind": "disappeared", "year": 2025})])
|
|
result = golden.check_events(
|
|
t, [{"code": "99441", "kind": "deleted|disappeared|cpt_deleted", "year": 2025}]
|
|
)
|
|
assert result.passed, result.detail
|
|
|
|
|
|
def test_check_events_fail_wrong_year():
|
|
t = Transcript([_lineage({"code": "G2058", "kind": "replaced_by", "year": 2020})])
|
|
result = golden.check_events(
|
|
t, [{"code": "G2058", "kind": "replaced_by", "year": 2021}]
|
|
)
|
|
assert not result.passed
|
|
|
|
|
|
# ── check_dockets ────────────────────────────────────────────────
|
|
|
|
|
|
def test_check_dockets_pass():
|
|
t = Transcript(
|
|
[
|
|
_sources(
|
|
{"kind": "comment", "docket": "CMS-2023-0121"},
|
|
{"kind": "comment", "docket": "CMS-2025-0304"},
|
|
)
|
|
]
|
|
)
|
|
result = golden.check_dockets(t, ["CMS-2023-0121", "CMS-2025-0304"])
|
|
assert result.passed, result.detail
|
|
|
|
|
|
def test_check_dockets_fail_missing_one():
|
|
t = Transcript([_sources({"kind": "comment", "docket": "CMS-2023-0121"})])
|
|
result = golden.check_dockets(t, ["CMS-2023-0121", "CMS-2025-0304"])
|
|
assert not result.passed
|
|
assert "CMS-2025-0304" in result.detail
|
|
|
|
|
|
# ── check_eras ───────────────────────────────────────────────────
|
|
|
|
|
|
def test_check_eras_pass():
|
|
t = Transcript(
|
|
[
|
|
_sources(
|
|
{"kind": "rule", "date": "2020-11-01"},
|
|
{"kind": "rule", "date": "2015-11-01"},
|
|
{"kind": "comment", "docket": "CMS-2023-0121"},
|
|
{"kind": "corpus", "date": "2025-01-01"},
|
|
)
|
|
]
|
|
)
|
|
result = golden.check_eras(t, 4)
|
|
assert result.passed, result.detail
|
|
|
|
|
|
def test_check_eras_fail_too_few():
|
|
t = Transcript(
|
|
[
|
|
_sources(
|
|
{"kind": "rule", "date": "2020-11-01"},
|
|
{"kind": "rule", "date": "2020-12-01"},
|
|
)
|
|
]
|
|
)
|
|
result = golden.check_eras(t, 2)
|
|
assert not result.passed
|
|
|
|
|
|
def test_check_eras_ignores_undated():
|
|
t = Transcript([_sources({"kind": "corpus", "date": ""})])
|
|
result = golden.check_eras(t, 1)
|
|
assert not result.passed
|
|
|
|
|
|
# ── evaluate: stream error short-circuits ───────────────────────
|
|
|
|
|
|
def test_evaluate_short_circuits_on_error_event():
|
|
t = Transcript([{"type": "error", "message": "boom"}])
|
|
results = golden.evaluate({"expect_labels": ["x"]}, t)
|
|
assert len(results) == 1
|
|
assert results[0].name == "stream"
|
|
assert not results[0].passed
|
|
assert "boom" in results[0].detail
|
|
|
|
|
|
# ── YAML shape ───────────────────────────────────────────────────
|
|
|
|
|
|
def test_golden_set_loads_and_is_well_formed():
|
|
entries = golden.load_set(GOLDEN_SET)
|
|
assert len(entries) == 7
|
|
ids = [e["id"] for e in entries]
|
|
assert len(ids) == len(set(ids)), "duplicate ids"
|
|
for e in entries:
|
|
assert "id" in e and "question" in e
|
|
assert any(k in e for k in golden._EXPECTATION_KEYS), e["id"]
|
|
|
|
|
|
def test_golden_set_expected_ids_present():
|
|
entries = {e["id"] for e in golden.load_set(GOLDEN_SET)}
|
|
assert entries == {
|
|
"ccm-history",
|
|
"g2058-replacement",
|
|
"apcm-vs-ccm-elements",
|
|
"audio-only-em-99441",
|
|
"g2211-commenters-2023-vs-2025",
|
|
"99490-telehealth-steps",
|
|
"g2064-g2065-to-99424-99426-rate",
|
|
}
|
|
|
|
|
|
def test_load_set_rejects_duplicate_ids(tmp_path):
|
|
bad = tmp_path / "bad.yaml"
|
|
bad.write_text(
|
|
"- id: a\n question: q1\n min_eras: 1\n"
|
|
"- id: a\n question: q2\n min_eras: 1\n"
|
|
)
|
|
with pytest.raises(ValueError, match="duplicate id"):
|
|
golden.load_set(bad)
|
|
|
|
|
|
def test_load_set_rejects_entry_without_expectations(tmp_path):
|
|
bad = tmp_path / "bad.yaml"
|
|
bad.write_text("- id: a\n question: q1\n")
|
|
with pytest.raises(ValueError, match="no expectations"):
|
|
golden.load_set(bad)
|
|
|
|
|
|
# ── live (opt-in) ────────────────────────────────────────────────
|
|
|
|
|
|
@pytest.mark.skipif(
|
|
not os.environ.get("LLM_CHAT_URL"),
|
|
reason="LLM_CHAT_URL not set — skipping live chat test",
|
|
)
|
|
def test_live_g2058_replacement():
|
|
url = os.environ["LLM_CHAT_URL"]
|
|
entries = {e["id"]: e for e in golden.load_set(GOLDEN_SET)}
|
|
entry = entries["g2058-replacement"]
|
|
events = golden.stream_chat(
|
|
url, entry["question"], entry.get("mode", "auto"), 180.0
|
|
)
|
|
results = golden.evaluate(entry, Transcript(events))
|
|
failing = [r for r in results if not r.passed]
|
|
assert not failing, [(r.name, r.detail) for r in failing]
|