compare.py 8.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255
  1. """Compare: load and aggregate drill results across backends and runs."""
  2. from __future__ import annotations
  3. import json
  4. from dataclasses import dataclass
  5. from pathlib import Path
  6. from typing import Any
  7. from drill.stats import wilson_ci
  8. from drill.verifier import Verdict
  9. @dataclass
  10. class BackendResult:
  11. backend: str
  12. total_runs: int
  13. passed_runs: int
  14. errored_runs: int
  15. avg_turns: float
  16. criterion_counts: dict[str, tuple[int, int]] # criterion -> (passed, total)
  17. sweep_id: str | None
  18. timestamp: str | None
  19. partial: bool
  20. @property
  21. def pass_rate(self) -> float:
  22. if self.total_runs == 0:
  23. return 0.0
  24. return self.passed_runs / self.total_runs
  25. def load_scenario_results(
  26. scenario_dir: Path,
  27. *,
  28. sweep_id: str | None = None,
  29. ) -> dict[str, BackendResult]:
  30. results: dict[str, BackendResult] = {}
  31. for backend_dir in sorted(scenario_dir.iterdir()):
  32. if not backend_dir.is_dir():
  33. continue
  34. timestamp_dirs = sorted(backend_dir.iterdir())
  35. if not timestamp_dirs:
  36. continue
  37. target_dir: Path | None = None
  38. if sweep_id:
  39. for d in timestamp_dirs:
  40. rg_path = d / "run-group.json"
  41. if rg_path.exists():
  42. rg = json.loads(rg_path.read_text())
  43. if rg.get("sweep_id") == sweep_id:
  44. target_dir = d
  45. break
  46. else:
  47. target_dir = timestamp_dirs[-1]
  48. if target_dir is None:
  49. continue
  50. result = _load_backend_result(backend_dir.name, target_dir)
  51. if result is not None:
  52. results[backend_dir.name] = result
  53. return results
  54. def _load_backend_result(backend_name: str, timestamp_dir: Path) -> BackendResult | None:
  55. rg_path = timestamp_dir / "run-group.json"
  56. if rg_path.exists():
  57. return _load_new_format(backend_name, timestamp_dir, rg_path)
  58. elif (timestamp_dir / "verdict.json").exists():
  59. return _load_old_format(backend_name, timestamp_dir)
  60. return None
  61. def _load_new_format(backend_name: str, timestamp_dir: Path, rg_path: Path) -> BackendResult:
  62. rg: dict[str, Any] = json.loads(rg_path.read_text())
  63. run_dirs = sorted(
  64. d for d in timestamp_dir.iterdir() if d.is_dir() and d.name.startswith("run-")
  65. )
  66. verdicts: list[Verdict] = []
  67. metas: list[dict[str, Any]] = []
  68. for run_dir in run_dirs:
  69. verdict_path = run_dir / "verdict.json"
  70. meta_path = run_dir / "meta.json"
  71. if verdict_path.exists():
  72. verdicts.append(Verdict.model_validate_json(verdict_path.read_text()))
  73. if meta_path.exists():
  74. metas.append(json.loads(meta_path.read_text()))
  75. passed_runs = sum(1 for v in verdicts if v.passed)
  76. errored_runs = sum(1 for r in rg.get("runs", []) if r.get("status") == "error")
  77. avg_turns = sum(m.get("actor_turns", 0) for m in metas) / len(metas) if metas else 0.0
  78. criterion_counts: dict[str, tuple[int, int]] = {}
  79. for v in verdicts:
  80. for c in v.criteria:
  81. prev_passed, prev_total = criterion_counts.get(c.criterion, (0, 0))
  82. criterion_counts[c.criterion] = (
  83. prev_passed + (1 if c.verdict == "pass" else 0),
  84. prev_total + 1,
  85. )
  86. return BackendResult(
  87. backend=backend_name,
  88. total_runs=len(verdicts),
  89. passed_runs=passed_runs,
  90. errored_runs=errored_runs,
  91. avg_turns=round(avg_turns, 1),
  92. criterion_counts=criterion_counts,
  93. sweep_id=rg.get("sweep_id"),
  94. timestamp=rg.get("timestamp"),
  95. partial=rg.get("partial", False),
  96. )
  97. def _load_old_format(backend_name: str, timestamp_dir: Path) -> BackendResult:
  98. verdict = Verdict.model_validate_json((timestamp_dir / "verdict.json").read_text())
  99. meta: dict[str, Any] = {}
  100. meta_path = timestamp_dir / "meta.json"
  101. if meta_path.exists():
  102. meta = json.loads(meta_path.read_text())
  103. criterion_counts: dict[str, tuple[int, int]] = {}
  104. for c in verdict.criteria:
  105. criterion_counts[c.criterion] = (1 if c.verdict == "pass" else 0, 1)
  106. return BackendResult(
  107. backend=backend_name,
  108. total_runs=1,
  109. passed_runs=1 if verdict.passed else 0,
  110. errored_runs=0,
  111. avg_turns=float(meta.get("actor_turns", 0)),
  112. criterion_counts=criterion_counts,
  113. sweep_id=None,
  114. timestamp=None,
  115. partial=False,
  116. )
  117. def format_compare_output(
  118. scenario: str,
  119. results: dict[str, BackendResult],
  120. ) -> str:
  121. if not results:
  122. return f"No results found for: {scenario}"
  123. lines: list[str] = []
  124. is_multi_run = any(r.total_runs > 1 for r in results.values())
  125. if is_multi_run:
  126. first = next(iter(results.values()))
  127. lines.append(f"Scenario: {scenario}")
  128. if first.sweep_id:
  129. sweep_label = f"Sweep: {first.sweep_id}"
  130. if first.timestamp:
  131. date_str = first.timestamp.split("T")[0]
  132. sweep_label += f" | {date_str}"
  133. lines.append(sweep_label)
  134. lines.append("")
  135. header = f"{'':40s}"
  136. sub_header = f"{'':40s}"
  137. for name, r in results.items():
  138. header += f" {name:>12s}"
  139. sub_header += f" {'(n=' + str(r.total_runs) + ')':>12s}"
  140. lines.append(header)
  141. lines.append(sub_header)
  142. lines.append("-" * len(header))
  143. rate_line = f"{'Overall pass rate':40s}"
  144. ci_line = f"{' 95% CI':40s}"
  145. for r in results.values():
  146. pct = f"{r.pass_rate * 100:.1f}%"
  147. rate_line += f" {pct:>12s}"
  148. lo, hi = wilson_ci(r.passed_runs, r.total_runs)
  149. ci_str = f"[{lo * 100:.0f}, {hi * 100:.0f}]"
  150. ci_line += f" {ci_str:>12s}"
  151. lines.append(rate_line)
  152. lines.append(ci_line)
  153. lines.append("")
  154. all_criteria: list[str] = []
  155. seen: set[str] = set()
  156. for r in results.values():
  157. for crit in r.criterion_counts:
  158. if crit not in seen:
  159. all_criteria.append(crit)
  160. seen.add(crit)
  161. for crit in all_criteria:
  162. crit_line = f"{crit[:40]:40s}"
  163. for r in results.values():
  164. passed, total = r.criterion_counts.get(crit, (0, 0))
  165. crit_line += f" {str(passed) + '/' + str(total):>12s}"
  166. lines.append(crit_line)
  167. lines.append("")
  168. avg_line = f"{'Avg turns':40s}"
  169. err_line = f"{'Errors':40s}"
  170. for r in results.values():
  171. avg_line += f" {str(r.avg_turns):>12s}"
  172. err_line += f" {str(r.errored_runs):>12s}"
  173. lines.append(avg_line)
  174. lines.append(err_line)
  175. if any(r.total_runs < 10 for r in results.values()):
  176. lines.append("")
  177. lines.append("Note: CI is wide due to small sample size; consider --n 10+")
  178. if any(r.partial for r in results.values()):
  179. lines.append("")
  180. lines.append("Warning: Sweep was interrupted — results are incomplete.")
  181. else:
  182. lines.append(f"Scenario: {scenario}")
  183. lines.append("")
  184. lines.append(f"{'Backend':20s} {'Result':8s} {'Score':7s} {'Turns':5s}")
  185. lines.append("-" * 42)
  186. for name, r in results.items():
  187. result_str = "PASS" if r.passed_runs == r.total_runs else "FAIL"
  188. total_criteria = sum(t for _, t in r.criterion_counts.values())
  189. passed_criteria = sum(p for p, _ in r.criterion_counts.values())
  190. score = f"{passed_criteria}/{total_criteria}"
  191. turns_str = (
  192. str(int(r.avg_turns)) if r.avg_turns == int(r.avg_turns) else str(r.avg_turns)
  193. )
  194. lines.append(f"{name:20s} {result_str:8s} {score:7s} {turns_str:5s}")
  195. all_criteria = []
  196. seen = set()
  197. for r in results.values():
  198. for crit in r.criterion_counts:
  199. if crit not in seen:
  200. all_criteria.append(crit)
  201. seen.add(crit)
  202. lines.append("")
  203. header = f"{'':40s}"
  204. for name in results:
  205. header += f" {name:>12s}"
  206. lines.append(header)
  207. lines.append("-" * len(header))
  208. for crit in all_criteria:
  209. crit_line = f"{crit[:40]:40s}"
  210. for r in results.values():
  211. p, t = r.criterion_counts.get(crit, (0, 0))
  212. icon = "PASS" if p == t and t > 0 else "FAIL"
  213. crit_line += f" {icon:>12s}"
  214. lines.append(crit_line)
  215. return "\n".join(lines)