runifold_testkit/evaluation/
runner.rs1use super::{
2 Arc, BTreeMap, BTreeSet, EvaluationCase, EvaluationCaseResult, EvaluationDataset,
3 EvaluationError, EvaluationFailure, EvaluationFailureStage, EvaluationReport, EvaluationScore,
4 EvaluationScoreSummary, EvaluationScorer, EvaluationTarget, NonZeroUsize, StreamExt,
5 ensure_not_empty, ensure_ratio, fmt, stream,
6};
7
8pub struct EvaluationRunner {
10 target: Arc<dyn EvaluationTarget>,
11 scorers: Vec<Arc<dyn EvaluationScorer>>,
12 concurrency: NonZeroUsize,
13}
14
15impl EvaluationRunner {
16 pub fn new(target: impl EvaluationTarget + 'static) -> Self {
18 Self {
19 target: Arc::new(target),
20 scorers: Vec::new(),
21 concurrency: NonZeroUsize::MIN,
22 }
23 }
24
25 #[must_use]
27 pub fn with_scorer(mut self, scorer: impl EvaluationScorer + 'static) -> Self {
28 self.scorers.push(Arc::new(scorer));
29 self
30 }
31
32 #[must_use]
34 pub const fn with_concurrency(mut self, concurrency: NonZeroUsize) -> Self {
35 self.concurrency = concurrency;
36 self
37 }
38
39 pub async fn run(
48 &self,
49 dataset: &EvaluationDataset,
50 candidate_version: impl Into<String>,
51 ) -> Result<EvaluationReport, EvaluationError> {
52 let candidate_version = candidate_version.into();
53 ensure_not_empty("candidate version", &candidate_version)?;
54 validate_scorers(&self.scorers)?;
55 let target = Arc::clone(&self.target);
56 let scorers = self.scorers.clone();
57 let mut indexed = stream::iter(dataset.cases.iter().cloned().enumerate())
58 .map(|(index, case)| {
59 let target = Arc::clone(&target);
60 let scorers = scorers.clone();
61 async move { (index, evaluate_case(target, scorers, case).await) }
62 })
63 .buffer_unordered(self.concurrency.get())
64 .collect::<Vec<_>>()
65 .await;
66 indexed.sort_by_key(|(index, _)| *index);
67 let cases = indexed
68 .into_iter()
69 .map(|(_, result)| result)
70 .collect::<Vec<_>>();
71 Ok(build_report(dataset, candidate_version, cases))
72 }
73}
74
75impl fmt::Debug for EvaluationRunner {
76 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
77 formatter
78 .debug_struct("EvaluationRunner")
79 .field("scorers", &self.scorers.len())
80 .field("concurrency", &self.concurrency)
81 .finish_non_exhaustive()
82 }
83}
84
85async fn evaluate_case(
86 target: Arc<dyn EvaluationTarget>,
87 scorers: Vec<Arc<dyn EvaluationScorer>>,
88 case: EvaluationCase,
89) -> EvaluationCaseResult {
90 let output = match target.execute(case.clone()).await {
91 Ok(output) => output,
92 Err(error) => {
93 return EvaluationCaseResult {
94 case_id: case.id,
95 run_id: None,
96 metrics: None,
97 scores: Vec::new(),
98 failures: vec![EvaluationFailure {
99 stage: EvaluationFailureStage::Target,
100 scorer: None,
101 message: error.to_string(),
102 }],
103 };
104 }
105 };
106 let run_id = output.run_id;
107 let metrics = output.metrics.clone();
108 let scorer_concurrency = scorers.len().max(1);
109 let scored = stream::iter(scorers)
110 .map(|scorer| {
111 let case = case.clone();
112 let output = output.clone();
113 async move {
114 let name = scorer.name().to_owned();
115 let threshold = scorer.threshold();
116 let result = scorer.score(case, output).await;
117 (name, threshold, result)
118 }
119 })
120 .buffer_unordered(scorer_concurrency)
121 .collect::<Vec<_>>()
122 .await;
123 let mut case_scores = Vec::new();
124 let mut failures = Vec::new();
125 for (name, threshold, result) in scored {
126 match result {
127 Ok(score) => case_scores.push(EvaluationScore {
128 name,
129 value: score.value,
130 threshold,
131 passed: score.value >= threshold,
132 rationale: score.rationale,
133 }),
134 Err(error) => failures.push(EvaluationFailure {
135 stage: EvaluationFailureStage::Scorer,
136 scorer: Some(name),
137 message: error.to_string(),
138 }),
139 }
140 }
141 case_scores.sort_by(|left, right| left.name.cmp(&right.name));
142 failures.sort_by(|left, right| left.scorer.cmp(&right.scorer));
143 EvaluationCaseResult {
144 case_id: case.id,
145 run_id,
146 metrics,
147 scores: case_scores,
148 failures,
149 }
150}
151
152fn build_report(
153 dataset: &EvaluationDataset,
154 candidate_version: String,
155 cases: Vec<EvaluationCaseResult>,
156) -> EvaluationReport {
157 let total_cases = cases.len();
158 let total_cases_ratio = cases.iter().fold(0.0, |total, _| total + 1.0);
159 let successful = cases.iter().fold(0.0, |total, result| {
160 if result
161 .failures
162 .iter()
163 .any(|failure| failure.stage == EvaluationFailureStage::Target)
164 {
165 total
166 } else {
167 total + 1.0
168 }
169 });
170 let mut aggregate = BTreeMap::<String, (usize, f64, f64, f64)>::new();
171 for score in cases.iter().flat_map(|result| &result.scores) {
172 let entry = aggregate.entry(score.name.clone()).or_default();
173 entry.0 += 1;
174 entry.1 += score.value;
175 entry.2 += 1.0;
176 entry.3 += if score.passed { 1.0 } else { 0.0 };
177 }
178 let summaries = aggregate
179 .into_iter()
180 .map(
181 |(name, (scored_cases, total, scored_cases_ratio, passed))| EvaluationScoreSummary {
182 name,
183 scored_cases,
184 total_cases,
185 mean: total / scored_cases_ratio,
186 pass_rate: passed / total_cases_ratio,
187 },
188 )
189 .collect();
190 EvaluationReport {
191 dataset_name: dataset.name.clone(),
192 dataset_version: dataset.version.clone(),
193 candidate_version,
194 execution_success_rate: successful / total_cases_ratio,
195 cases,
196 summaries,
197 }
198}
199
200fn validate_scorers(scorers: &[Arc<dyn EvaluationScorer>]) -> Result<(), EvaluationError> {
201 if scorers.is_empty() {
202 return Err(EvaluationError::NoScorers);
203 }
204 let mut names = BTreeSet::new();
205 for scorer in scorers {
206 ensure_not_empty("scorer name", scorer.name())?;
207 ensure_ratio("score threshold", scorer.threshold())?;
208 if !names.insert(scorer.name()) {
209 return Err(EvaluationError::DuplicateScorer {
210 scorer: scorer.name().to_owned(),
211 });
212 }
213 }
214 Ok(())
215}