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