Skip to main content

vtcode_eval/
metric.rs

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}