| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105 |
- from __future__ import annotations
- from dataclasses import dataclass
- @dataclass(frozen=True)
- class EvaluationCase:
- case_id: str
- case_type: str
- expected_source_uids: frozenset[str]
- allowed_business_domains: frozenset[str]
- must_refuse: bool = False
- @dataclass(frozen=True)
- class EvaluationObservation:
- retrieved_source_uids: tuple[str, ...]
- retrieved_business_domains: tuple[str | None, ...]
- cited_source_uids: tuple[str, ...] = ()
- cited_freshness: tuple[str, ...] = ()
- answer_status: str = "no_answer"
- latency_ms: int = 0
- @dataclass(frozen=True)
- class EvaluationSummary:
- case_count: int
- source_recall: float
- citation_precision: float
- refusal_accuracy: float
- permission_leak_count: int
- stale_citation_count: int
- p95_latency_ms: int
- passed: bool
- def _percentile_95(values: list[int]) -> int:
- if not values:
- return 0
- ordered = sorted(values)
- index = max(0, (95 * len(ordered) + 99) // 100 - 1)
- return ordered[index]
- def evaluate(
- results: list[tuple[EvaluationCase, EvaluationObservation]],
- *,
- minimum_source_recall: float = 0.90,
- minimum_citation_precision: float = 0.95,
- minimum_refusal_accuracy: float = 0.90,
- maximum_p95_latency_ms: int = 1500,
- ) -> EvaluationSummary:
- recalls: list[float] = []
- citation_hits = 0
- citation_total = 0
- refusal_cases = 0
- correct_refusals = 0
- permission_leaks = 0
- stale_citations = 0
- latencies = []
- for case, observation in results:
- expected = case.expected_source_uids
- retrieved = set(observation.retrieved_source_uids)
- recalls.append(len(expected & retrieved) / len(expected) if expected else 1.0)
- citation_total += len(observation.cited_source_uids)
- citation_hits += sum(
- 1 for uid in observation.cited_source_uids if uid in expected
- )
- if case.must_refuse:
- refusal_cases += 1
- if observation.answer_status == "no_answer":
- correct_refusals += 1
- permission_leaks += sum(
- 1
- for domain in observation.retrieved_business_domains
- if domain is not None and domain not in case.allowed_business_domains
- )
- stale_citations += sum(
- 1 for freshness in observation.cited_freshness if freshness == "stale"
- )
- latencies.append(observation.latency_ms)
- source_recall = sum(recalls) / len(recalls) if recalls else 0.0
- citation_precision = citation_hits / citation_total if citation_total else 1.0
- refusal_accuracy = correct_refusals / refusal_cases if refusal_cases else 1.0
- p95_latency = _percentile_95(latencies)
- passed = all(
- (
- source_recall >= minimum_source_recall,
- citation_precision >= minimum_citation_precision,
- refusal_accuracy >= minimum_refusal_accuracy,
- permission_leaks == 0,
- stale_citations == 0,
- p95_latency <= maximum_p95_latency_ms,
- )
- )
- return EvaluationSummary(
- case_count=len(results),
- source_recall=source_recall,
- citation_precision=citation_precision,
- refusal_accuracy=refusal_accuracy,
- permission_leak_count=permission_leaks,
- stale_citation_count=stale_citations,
- p95_latency_ms=p95_latency,
- passed=passed,
- )
|