evaluation.py 3.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105
  1. from __future__ import annotations
  2. from dataclasses import dataclass
  3. @dataclass(frozen=True)
  4. class EvaluationCase:
  5. case_id: str
  6. case_type: str
  7. expected_source_uids: frozenset[str]
  8. allowed_business_domains: frozenset[str]
  9. must_refuse: bool = False
  10. @dataclass(frozen=True)
  11. class EvaluationObservation:
  12. retrieved_source_uids: tuple[str, ...]
  13. retrieved_business_domains: tuple[str | None, ...]
  14. cited_source_uids: tuple[str, ...] = ()
  15. cited_freshness: tuple[str, ...] = ()
  16. answer_status: str = "no_answer"
  17. latency_ms: int = 0
  18. @dataclass(frozen=True)
  19. class EvaluationSummary:
  20. case_count: int
  21. source_recall: float
  22. citation_precision: float
  23. refusal_accuracy: float
  24. permission_leak_count: int
  25. stale_citation_count: int
  26. p95_latency_ms: int
  27. passed: bool
  28. def _percentile_95(values: list[int]) -> int:
  29. if not values:
  30. return 0
  31. ordered = sorted(values)
  32. index = max(0, (95 * len(ordered) + 99) // 100 - 1)
  33. return ordered[index]
  34. def evaluate(
  35. results: list[tuple[EvaluationCase, EvaluationObservation]],
  36. *,
  37. minimum_source_recall: float = 0.90,
  38. minimum_citation_precision: float = 0.95,
  39. minimum_refusal_accuracy: float = 0.90,
  40. maximum_p95_latency_ms: int = 1500,
  41. ) -> EvaluationSummary:
  42. recalls: list[float] = []
  43. citation_hits = 0
  44. citation_total = 0
  45. refusal_cases = 0
  46. correct_refusals = 0
  47. permission_leaks = 0
  48. stale_citations = 0
  49. latencies = []
  50. for case, observation in results:
  51. expected = case.expected_source_uids
  52. retrieved = set(observation.retrieved_source_uids)
  53. recalls.append(len(expected & retrieved) / len(expected) if expected else 1.0)
  54. citation_total += len(observation.cited_source_uids)
  55. citation_hits += sum(
  56. 1 for uid in observation.cited_source_uids if uid in expected
  57. )
  58. if case.must_refuse:
  59. refusal_cases += 1
  60. if observation.answer_status == "no_answer":
  61. correct_refusals += 1
  62. permission_leaks += sum(
  63. 1
  64. for domain in observation.retrieved_business_domains
  65. if domain is not None and domain not in case.allowed_business_domains
  66. )
  67. stale_citations += sum(
  68. 1 for freshness in observation.cited_freshness if freshness == "stale"
  69. )
  70. latencies.append(observation.latency_ms)
  71. source_recall = sum(recalls) / len(recalls) if recalls else 0.0
  72. citation_precision = citation_hits / citation_total if citation_total else 1.0
  73. refusal_accuracy = correct_refusals / refusal_cases if refusal_cases else 1.0
  74. p95_latency = _percentile_95(latencies)
  75. passed = all(
  76. (
  77. source_recall >= minimum_source_recall,
  78. citation_precision >= minimum_citation_precision,
  79. refusal_accuracy >= minimum_refusal_accuracy,
  80. permission_leaks == 0,
  81. stale_citations == 0,
  82. p95_latency <= maximum_p95_latency_ms,
  83. )
  84. )
  85. return EvaluationSummary(
  86. case_count=len(results),
  87. source_recall=source_recall,
  88. citation_precision=citation_precision,
  89. refusal_accuracy=refusal_accuracy,
  90. permission_leak_count=permission_leaks,
  91. stale_citation_count=stale_citations,
  92. p95_latency_ms=p95_latency,
  93. passed=passed,
  94. )