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