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