1use super::task::{EvalRunResult, RunOutcome};
2use serde::{Deserialize, Serialize};
3
4#[derive(Debug, Clone, Serialize, Deserialize)]
5pub struct EvalMetric {
6 pub(crate) pass_at_k: f64,
7 pub(crate) pass_all_k: f64,
8 pub(crate) total_runs: u32,
9 pub(crate) passed_runs: u32,
10 pub(crate) task_id: String,
11}
12
13pub fn compute_metric(task_id: &str, results: &[EvalRunResult]) -> EvalMetric {
14 let total = results.len() as u32;
15 let passed = results.iter().filter(|r| r.outcome == RunOutcome::Pass).count() as u32;
16 let pass_at_k = if total > 0 { passed as f64 / total as f64 } else { 0.0 };
17 let pass_all_k = if total > 0 && passed == total { 1.0 } else { 0.0 };
18 EvalMetric {
19 pass_at_k,
20 pass_all_k,
21 total_runs: total,
22 passed_runs: passed,
23 task_id: task_id.into(),
24 }
25}
26
27pub fn aggregate_metrics(metrics: &[EvalMetric]) -> EvalMetric {
28 if metrics.is_empty() {
29 return EvalMetric {
30 pass_at_k: 0.0,
31 pass_all_k: 0.0,
32 total_runs: 0,
33 passed_runs: 0,
34 task_id: "aggregate".into(),
35 };
36 }
37 let total_runs: u32 = metrics.iter().map(|m| m.total_runs).sum();
38 let passed_runs: u32 = metrics.iter().map(|m| m.passed_runs).sum();
39 let pass_at_k = if total_runs > 0 {
40 passed_runs as f64 / total_runs as f64
41 } else {
42 0.0
43 };
44 let pass_all_k = if metrics.iter().all(|m| m.pass_all_k > 0.0) {
45 1.0
46 } else {
47 0.0
48 };
49 EvalMetric {
50 pass_at_k,
51 pass_all_k,
52 total_runs,
53 passed_runs,
54 task_id: "aggregate".into(),
55 }
56}
57
58pub fn pass_at_k(results: &[EvalRunResult]) -> f64 {
59 let total = results.len() as f64;
60 if total == 0.0 {
61 return 0.0;
62 }
63 let passed = results.iter().filter(|r| r.outcome == RunOutcome::Pass).count() as f64;
64 passed / total
65}
66
67pub fn pass_all_k(results: &[EvalRunResult]) -> f64 {
68 if results.is_empty() {
69 return 0.0;
70 }
71 if results.iter().all(|r| r.outcome == RunOutcome::Pass) {
72 1.0
73 } else {
74 0.0
75 }
76}
77
78#[cfg(test)]
79mod tests {
80 use super::*;
81 use crate::task::{EvalRunResult, RunOutcome};
82
83 fn r(outcome: RunOutcome) -> EvalRunResult {
84 EvalRunResult {
85 task_id: "t".into(),
86 outcome,
87 error_message: None,
88 duration_secs: 0.0,
89 attempt: 1,
90 cost_usd: None,
91 transcript_path: None,
92 }
93 }
94
95 #[test]
96 fn compute_metric_pass_at_k() {
97 let results = vec![r(RunOutcome::Pass), r(RunOutcome::Fail), r(RunOutcome::Error)];
98 let m = compute_metric("t", &results);
99 assert_eq!(m.total_runs, 3);
100 assert_eq!(m.passed_runs, 1);
101 assert!((m.pass_at_k - 1.0 / 3.0).abs() < 1e-9);
102 assert_eq!(m.pass_all_k, 0.0);
103 }
104
105 #[test]
106 fn compute_metric_all_pass() {
107 let results = vec![r(RunOutcome::Pass), r(RunOutcome::Pass)];
108 let m = compute_metric("t", &results);
109 assert_eq!(m.pass_all_k, 1.0);
110 assert!((m.pass_at_k - 1.0).abs() < 1e-9);
111 }
112
113 #[test]
114 fn compute_metric_empty() {
115 let m = compute_metric("t", &[]);
116 assert_eq!(m.total_runs, 0);
117 assert_eq!(m.pass_at_k, 0.0);
118 }
119
120 #[test]
121 fn aggregate_combines_runs() {
122 let a = EvalMetric {
123 pass_at_k: 1.0,
124 pass_all_k: 1.0,
125 total_runs: 1,
126 passed_runs: 1,
127 task_id: "a".into(),
128 };
129 let b = EvalMetric {
130 pass_at_k: 0.0,
131 pass_all_k: 0.0,
132 total_runs: 1,
133 passed_runs: 0,
134 task_id: "b".into(),
135 };
136 let agg = aggregate_metrics(&[a, b]);
137 assert_eq!(agg.total_runs, 2);
138 assert_eq!(agg.passed_runs, 1);
139 assert!((agg.pass_at_k - 0.5).abs() < 1e-9);
140 assert_eq!(agg.pass_all_k, 0.0);
141 }
142
143 #[test]
144 fn aggregate_empty() {
145 let agg = aggregate_metrics(&[]);
146 assert_eq!(agg.total_runs, 0);
147 assert_eq!(agg.pass_at_k, 0.0);
148 }
149}