Files
ThothII/harness/tht/evidence/evaluation.py

304 lines
11 KiB
Python

"""Read-only retrieval evaluation for published and candidate Evidence generations."""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import yaml
from tht.evidence.canonical import EVIDENCE_PURPOSES
from tht.evidence.search import EvidenceSearchContext, render_evidence_query
_PROFILES = frozenset({"lexical", "semantic", "mixed"})
class EvaluationFixtureError(ValueError):
"""The versioned retrieval fixture is not safe to use as a publication gate."""
class EvaluationError(RuntimeError):
"""The configured Evidence generation could not be evaluated safely."""
@dataclass(frozen=True)
class EvaluationQuery:
query_id: str
query: str
profile: str
purpose: str
expected: tuple[str, ...]
@dataclass(frozen=True)
class EvaluationFixture:
queries: tuple[EvaluationQuery, ...]
@dataclass(frozen=True)
class ExpectedEvidenceReport:
evidence_id: str
kind: str | None
dense_rank: int | None
bm25_rank: int | None
fused_rank: int | None
@dataclass(frozen=True)
class EvaluationQueryReport:
query_id: str
profile: str
purpose: str
hit_at_5: bool
hit_at_10: bool
missing_expected: tuple[str, ...]
empty_result: bool
expected: tuple[ExpectedEvidenceReport, ...]
@dataclass(frozen=True)
class EvaluationReport:
workspace_revision: str
vector_generation: str | None
rrf: dict[str, object]
passed: bool
queries: tuple[EvaluationQueryReport, ...]
counts_by_expected_kind: dict[str, int]
def model_dump(self) -> dict[str, object]:
return {
"workspaceRevision": self.workspace_revision,
"vectorGeneration": self.vector_generation,
"rrf": self.rrf,
"passed": self.passed,
"countsByExpectedKind": self.counts_by_expected_kind,
"queries": [
{
"id": query.query_id,
"profile": query.profile,
"purpose": query.purpose,
"hitAt5": query.hit_at_5,
"hitAt10": query.hit_at_10,
"missingExpected": list(query.missing_expected),
"emptyResult": query.empty_result,
"expected": [
{
"evidenceId": expected.evidence_id,
"kind": expected.kind,
"denseRank": expected.dense_rank,
"bm25Rank": expected.bm25_rank,
"fusedRank": expected.fused_rank,
}
for expected in query.expected
],
}
for query in self.queries
],
}
def load_evaluation_fixture(path: Path) -> EvaluationFixture:
"""Load the deliberately small, complete v1 evaluation fixture."""
try:
raw = yaml.safe_load(path.read_text(encoding="utf-8"))
except (OSError, UnicodeError, yaml.YAMLError) as error:
raise EvaluationFixtureError("evaluation fixture is unreadable") from error
if not isinstance(raw, dict) or set(raw) != {"schema_version", "queries"}:
raise EvaluationFixtureError("evaluation fixture schema is invalid")
if raw.get("schema_version") != 1 or not isinstance(raw.get("queries"), list):
raise EvaluationFixtureError("evaluation fixture schema is invalid")
queries: list[EvaluationQuery] = []
errors: set[str] = set()
ids: set[str] = set()
profiles: set[str] = set()
for entry in raw["queries"]:
entry_errors: set[str] = set()
if not isinstance(entry, dict) or set(entry) != {"id", "query", "profile", "purpose", "expected"}:
errors.add("schema")
continue
query_id = entry["id"]
query = entry["query"]
profile = entry["profile"]
purpose = entry["purpose"]
expected = entry["expected"]
if not isinstance(query_id, str) or not query_id.strip() or query_id in ids:
entry_errors.add("duplicate")
else:
ids.add(query_id)
if not isinstance(query, str) or not query.strip():
entry_errors.add("query")
if profile not in _PROFILES:
entry_errors.add("profile")
else:
profiles.add(profile)
if purpose not in EVIDENCE_PURPOSES:
entry_errors.add("purpose")
if (
not isinstance(expected, list)
or not expected
or any(not isinstance(value, str) or not value.strip() for value in expected)
):
entry_errors.add("expected")
errors.update(entry_errors)
if not entry_errors:
queries.append(EvaluationQuery(query_id, query, profile, purpose, tuple(expected)))
missing_profiles = _PROFILES - profiles
if missing_profiles:
errors.update(missing_profiles)
if errors:
raise EvaluationFixtureError(" ".join(sorted(errors)))
return EvaluationFixture(tuple(queries))
def _ranked_evidence(hits) -> tuple[dict[str, int], dict[str, str]]:
ranks: dict[str, int] = {}
kinds: dict[str, str] = {}
for hit in hits:
metadata = getattr(hit, "metadata", None)
if not isinstance(metadata, dict):
raise EvaluationError("evaluation search returned malformed Evidence payload")
evidence_id = metadata.get("evidence_id")
kind = metadata.get("evidence_kind")
if not isinstance(evidence_id, str) or not evidence_id or not isinstance(kind, str) or not kind:
raise EvaluationError("evaluation search returned malformed Evidence payload")
if evidence_id not in ranks:
ranks[evidence_id] = len(ranks) + 1
kinds[evidence_id] = kind
return ranks, kinds
def _search_generation(
searcher,
embedding: list[float],
*,
rendered_query: str,
purpose: str,
workspace_id: str,
generation: str,
document_ids: list[str],
language: str,
retrieval_mode: str,
):
return searcher.search(
["evidence"],
embedding,
limit=10,
kinds=["evidence"],
query_text=rendered_query,
query_language=language,
retrieval_mode=retrieval_mode,
metadata_filter={
"workspace_id": workspace_id,
"vector_generation": generation,
"document_ids": document_ids,
"purpose": purpose,
"required_kinds": [],
"required_concepts": [],
"required_tables": [],
"required_columns": [],
},
)
def evaluate_retrieval(
fixture: EvaluationFixture,
*,
workspace_revision: str,
document_generations: dict[str, str],
workspace_id: str,
language: str,
searcher,
embedder,
vector_generation: str | None = None,
expected_kinds: dict[str, str] | None = None,
) -> EvaluationReport:
"""Evaluate a generation with the runtime hybrid request plus branch diagnostics."""
if not document_generations:
raise EvaluationError("evaluation requires indexed Evidence documents")
by_generation: dict[str, list[str]] = {}
for document_id, generation in document_generations.items():
if not isinstance(document_id, str) or not isinstance(generation, str) or not generation:
raise EvaluationError("evaluation document generations are invalid")
by_generation.setdefault(generation, []).append(document_id)
for document_ids in by_generation.values():
document_ids.sort()
reports: list[EvaluationQueryReport] = []
expected_kinds = expected_kinds or {}
for query in fixture.queries:
rendered = render_evidence_query(query.query, EvidenceSearchContext())
embedding = embedder.embed_query(rendered)
branch_hits = {"dense": [], "bm25": [], "fused": []}
for generation, document_ids in sorted(by_generation.items()):
for mode, hits in branch_hits.items():
hits.extend(_search_generation(
searcher,
embedding,
rendered_query=rendered,
purpose=query.purpose,
workspace_id=workspace_id,
generation=generation,
document_ids=document_ids,
language=language,
retrieval_mode=mode,
))
ranks_by_branch: dict[str, dict[str, int]] = {}
kinds_by_branch: dict[str, dict[str, str]] = {}
for mode, hits in branch_hits.items():
ordered = sorted(hits, key=lambda hit: (-float(hit.similarity), str(hit.id)))
ranks_by_branch[mode], kinds_by_branch[mode] = _ranked_evidence(ordered)
expected = []
for evidence_id in query.expected:
kind = expected_kinds.get(evidence_id) or next((
kinds_by_branch[mode][evidence_id]
for mode in ("fused", "dense", "bm25")
if evidence_id in kinds_by_branch[mode]
), None)
expected.append(ExpectedEvidenceReport(
evidence_id=evidence_id,
kind=kind,
dense_rank=ranks_by_branch["dense"].get(evidence_id),
bm25_rank=ranks_by_branch["bm25"].get(evidence_id),
fused_rank=ranks_by_branch["fused"].get(evidence_id),
))
fused_ranks = ranks_by_branch["fused"]
missing = tuple(item.evidence_id for item in expected if item.fused_rank is None)
reports.append(EvaluationQueryReport(
query_id=query.query_id,
profile=query.profile,
purpose=query.purpose,
hit_at_5=any(item.fused_rank is not None and item.fused_rank <= 5 for item in expected),
hit_at_10=any(item.fused_rank is not None and item.fused_rank <= 10 for item in expected),
missing_expected=missing,
empty_result=not fused_ranks,
expected=tuple(expected),
))
counts: dict[str, int] = {}
for query in reports:
for expected in query.expected:
if expected.kind is not None:
counts[expected.kind] = counts.get(expected.kind, 0) + 1
evaluated = vector_generation or (next(iter(by_generation)) if len(by_generation) == 1 else None)
return EvaluationReport(
workspace_revision=workspace_revision,
vector_generation=evaluated,
rrf={"algorithm": "rrf", "k": 60, "prefetch_limit_multiplier": 2},
passed=all(query.hit_at_10 for query in reports),
queries=tuple(reports),
counts_by_expected_kind=dict(sorted(counts.items())),
)
__all__ = [
"EvaluationError",
"EvaluationFixture",
"EvaluationFixtureError",
"EvaluationQuery",
"EvaluationQueryReport",
"EvaluationReport",
"ExpectedEvidenceReport",
"evaluate_retrieval",
"load_evaluation_fixture",
]