feat(evidence): evaluate retrieval with a small fixture
This commit is contained in:
@@ -0,0 +1,185 @@
|
||||
import pytest
|
||||
|
||||
|
||||
def _fixture(path):
|
||||
path.write_text(
|
||||
"""
|
||||
schema_version: 1
|
||||
queries:
|
||||
- id: lexical-code
|
||||
query: Qual è il codice ICD-10?
|
||||
profile: lexical
|
||||
purpose: schema_linking
|
||||
expected: [evidence:icd]
|
||||
- id: semantic-age
|
||||
query: Come distinguo i pazienti pediatrici?
|
||||
profile: semantic
|
||||
purpose: sql_generation
|
||||
expected: [evidence:fascia-pediatrica]
|
||||
- id: mixed-formula
|
||||
query: Formula per patient.birth_date?
|
||||
profile: mixed
|
||||
purpose: sql_generation
|
||||
expected: [evidence:fascia-pediatrica]
|
||||
""".strip(),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
class _Embedder:
|
||||
def embed_query(self, query):
|
||||
return [float(len(query))]
|
||||
|
||||
|
||||
class _Hit:
|
||||
def __init__(self, evidence_id, kind, score):
|
||||
self.id = evidence_id + ":fragment"
|
||||
self.similarity = score
|
||||
self.metadata = {"evidence_id": evidence_id, "evidence_kind": kind}
|
||||
|
||||
|
||||
class _Searcher:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
def search(self, collections, embedding, **kwargs):
|
||||
self.calls.append((collections, embedding, kwargs))
|
||||
mode = kwargs["retrieval_mode"]
|
||||
if mode == "dense":
|
||||
return [_Hit("evidence:fascia-pediatrica", "formula", 0.9)]
|
||||
if mode == "bm25":
|
||||
return [_Hit("evidence:icd", "enum", 0.8)]
|
||||
return [
|
||||
_Hit("evidence:fascia-pediatrica", "formula", 0.95),
|
||||
_Hit("evidence:icd", "enum", 0.8),
|
||||
]
|
||||
|
||||
|
||||
def test_evaluation_reports_branch_and_fused_ranks_without_turning_diagnostics_into_gates(tmp_path):
|
||||
from tht.evidence.evaluation import evaluate_retrieval, load_evaluation_fixture
|
||||
|
||||
fixture_path = tmp_path / "evaluation.yaml"
|
||||
_fixture(fixture_path)
|
||||
searcher = _Searcher()
|
||||
|
||||
report = evaluate_retrieval(
|
||||
load_evaluation_fixture(fixture_path),
|
||||
workspace_revision="a" * 40,
|
||||
document_generations={"doc:one": "gen:" + "1" * 32},
|
||||
workspace_id="psd-clinical",
|
||||
language="italian",
|
||||
searcher=searcher,
|
||||
embedder=_Embedder(),
|
||||
)
|
||||
|
||||
assert report.passed is True
|
||||
assert report.workspace_revision == "a" * 40
|
||||
assert report.vector_generation == "gen:" + "1" * 32
|
||||
assert report.rrf == {"algorithm": "rrf", "k": 60, "prefetch_limit_multiplier": 2}
|
||||
assert report.queries[0].hit_at_5 is True
|
||||
assert report.queries[0].hit_at_10 is True
|
||||
assert report.queries[0].missing_expected == ()
|
||||
assert report.queries[0].expected[0].dense_rank is None
|
||||
assert report.queries[0].expected[0].bm25_rank == 1
|
||||
assert report.queries[0].expected[0].fused_rank == 2
|
||||
assert report.counts_by_expected_kind == {"enum": 1, "formula": 2}
|
||||
assert len(searcher.calls) == 9
|
||||
assert {call[2]["retrieval_mode"] for call in searcher.calls} == {"dense", "bm25", "fused"}
|
||||
assert all(call[2]["metadata_filter"] == {
|
||||
"workspace_id": "psd-clinical",
|
||||
"vector_generation": "gen:" + "1" * 32,
|
||||
"document_ids": ["doc:one"],
|
||||
"purpose": call[2]["metadata_filter"]["purpose"],
|
||||
"required_kinds": [],
|
||||
"required_concepts": [],
|
||||
"required_tables": [],
|
||||
"required_columns": [],
|
||||
} for call in searcher.calls)
|
||||
|
||||
|
||||
def test_evaluation_fails_only_when_a_query_has_no_expected_fused_hit_in_top_ten(tmp_path):
|
||||
from tht.evidence.evaluation import evaluate_retrieval, load_evaluation_fixture
|
||||
|
||||
fixture_path = tmp_path / "evaluation.yaml"
|
||||
_fixture(fixture_path)
|
||||
|
||||
class MissingExpectedSearcher(_Searcher):
|
||||
def search(self, collections, embedding, **kwargs):
|
||||
self.calls.append((collections, embedding, kwargs))
|
||||
return [_Hit("evidence:other", "domain", 1.0)]
|
||||
|
||||
report = evaluate_retrieval(
|
||||
load_evaluation_fixture(fixture_path),
|
||||
workspace_revision="a" * 40,
|
||||
document_generations={"doc:one": "gen:" + "1" * 32},
|
||||
workspace_id="psd-clinical",
|
||||
language="italian",
|
||||
searcher=MissingExpectedSearcher(),
|
||||
embedder=_Embedder(),
|
||||
expected_kinds={"evidence:icd": "enum", "evidence:fascia-pediatrica": "formula"},
|
||||
)
|
||||
|
||||
assert report.passed is False
|
||||
assert all(query.hit_at_5 is False and query.hit_at_10 is False for query in report.queries)
|
||||
assert all(query.empty_result is False for query in report.queries)
|
||||
assert report.queries[0].missing_expected == ("evidence:icd",)
|
||||
assert report.counts_by_expected_kind == {"enum": 1, "formula": 2}
|
||||
|
||||
|
||||
def test_evaluation_fixture_requires_all_retrieval_profiles(tmp_path):
|
||||
from tht.evidence.evaluation import EvaluationFixtureError, load_evaluation_fixture
|
||||
|
||||
path = tmp_path / "evaluation.yaml"
|
||||
path.write_text(
|
||||
"""
|
||||
schema_version: 1
|
||||
queries:
|
||||
- id: lexical-code
|
||||
query: Qual è il codice ICD-10?
|
||||
profile: lexical
|
||||
purpose: schema_linking
|
||||
expected: [evidence:icd]
|
||||
- id: semantic-age
|
||||
query: Come distinguo i pazienti pediatrici?
|
||||
profile: semantic
|
||||
purpose: sql_generation
|
||||
expected: [evidence:fascia-pediatrica]
|
||||
""".strip(),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
with pytest.raises(EvaluationFixtureError, match="mixed"):
|
||||
load_evaluation_fixture(path)
|
||||
|
||||
|
||||
def test_evaluation_fixture_rejects_duplicate_ids_empty_expectations_and_private_purposes(tmp_path):
|
||||
from tht.evidence.evaluation import EvaluationFixtureError, load_evaluation_fixture
|
||||
|
||||
path = tmp_path / "evaluation.yaml"
|
||||
path.write_text(
|
||||
"""
|
||||
schema_version: 1
|
||||
queries:
|
||||
- id: duplicate
|
||||
query: a
|
||||
profile: lexical
|
||||
purpose: private
|
||||
expected: []
|
||||
- id: duplicate
|
||||
query: b
|
||||
profile: semantic
|
||||
purpose: sql_generation
|
||||
expected: [evidence:b]
|
||||
- id: mixed
|
||||
query: c
|
||||
profile: mixed
|
||||
purpose: rewriting
|
||||
expected: [evidence:c]
|
||||
""".strip(),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
with pytest.raises(EvaluationFixtureError) as failure:
|
||||
load_evaluation_fixture(path)
|
||||
|
||||
assert {"duplicate", "expected", "purpose"} <= set(str(failure.value).split())
|
||||
Reference in New Issue
Block a user