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.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}