1use std::collections::HashMap;
30
31#[derive(Debug, Clone)]
35pub struct EvalCase {
36 pub id: String,
38 pub prompt: String,
40 pub reference_answer: Option<String>,
42 pub tags: Vec<String>,
44 pub difficulty: u8,
46}
47
48#[derive(Debug, Clone)]
52pub enum EvalMetric {
53 ExactMatch,
55 ContainsAnswer,
57 WordOverlap(f64),
59 LengthRatio { min: f64, max: f64 },
61 Custom(String),
63}
64
65#[derive(Debug, Clone)]
69pub struct EvalResult {
70 pub case_id: String,
72 pub response: String,
74 pub metrics: HashMap<String, f64>,
76 pub passed: bool,
78 pub latency_ms: u64,
80 pub cost_usd: f64,
82}
83
84#[derive(Debug, Clone)]
88pub struct EvalReport {
89 pub strategy_name: String,
91 pub results: Vec<EvalResult>,
93 pub pass_rate: f64,
95 pub avg_latency_ms: f64,
97 pub avg_cost_usd: f64,
99 pub total_cost_usd: f64,
101}
102
103pub struct EvalHarness {
107 cases: Vec<EvalCase>,
108 primary_metric: EvalMetric,
110}
111
112impl Default for EvalHarness {
113 fn default() -> Self {
114 Self::new()
115 }
116}
117
118impl EvalHarness {
119 pub fn new() -> Self {
121 Self {
122 cases: Vec::new(),
123 primary_metric: EvalMetric::ExactMatch,
124 }
125 }
126
127 pub fn with_metric(metric: EvalMetric) -> Self {
129 Self {
130 cases: Vec::new(),
131 primary_metric: metric,
132 }
133 }
134
135 pub fn add_case(&mut self, case: EvalCase) {
137 self.cases.push(case);
138 }
139
140 pub fn run_eval(
150 &self,
151 strategy_name: &str,
152 responses: Vec<(String, u64, f64)>,
153 ) -> EvalReport {
154 let mut results: Vec<EvalResult> = Vec::new();
155
156 for (case, (response, latency_ms, cost_usd)) in
157 self.cases.iter().zip(responses)
158 {
159 let primary_score = self.score(&response, case, &self.primary_metric);
160 let passed = primary_score >= 0.5;
161
162 let mut metrics = HashMap::new();
163 metrics.insert(metric_name(&self.primary_metric), primary_score);
164
165 results.push(EvalResult {
166 case_id: case.id.clone(),
167 response,
168 metrics,
169 passed,
170 latency_ms,
171 cost_usd,
172 });
173 }
174
175 let n = results.len() as f64;
176 let pass_rate = if n == 0.0 {
177 0.0
178 } else {
179 results.iter().filter(|r| r.passed).count() as f64 / n
180 };
181 let avg_latency_ms = if n == 0.0 {
182 0.0
183 } else {
184 results.iter().map(|r| r.latency_ms as f64).sum::<f64>() / n
185 };
186 let total_cost_usd = results.iter().map(|r| r.cost_usd).sum::<f64>();
187 let avg_cost_usd = if n == 0.0 { 0.0 } else { total_cost_usd / n };
188
189 EvalReport {
190 strategy_name: strategy_name.to_string(),
191 results,
192 pass_rate,
193 avg_latency_ms,
194 avg_cost_usd,
195 total_cost_usd,
196 }
197 }
198
199 pub fn score(&self, response: &str, case: &EvalCase, metric: &EvalMetric) -> f64 {
201 match metric {
202 EvalMetric::ExactMatch => {
203 let reference = case
204 .reference_answer
205 .as_deref()
206 .unwrap_or("")
207 .trim()
208 .to_lowercase();
209 let resp = response.trim().to_lowercase();
210 if resp == reference { 1.0 } else { 0.0 }
211 }
212 EvalMetric::ContainsAnswer => {
213 let reference = case
214 .reference_answer
215 .as_deref()
216 .unwrap_or("")
217 .trim()
218 .to_lowercase();
219 if reference.is_empty() {
220 return 0.0;
221 }
222 if response.to_lowercase().contains(&reference) {
223 1.0
224 } else {
225 0.0
226 }
227 }
228 EvalMetric::WordOverlap(_threshold) => {
229 let ref_words = word_set(case.reference_answer.as_deref().unwrap_or(""));
230 let resp_words = word_set(response);
231 if ref_words.is_empty() && resp_words.is_empty() {
232 return 1.0;
233 }
234 let intersection = ref_words.iter().filter(|w| resp_words.contains(*w)).count();
235 let union = ref_words.len() + resp_words.len() - intersection;
236 if union == 0 { 0.0 } else { intersection as f64 / union as f64 }
237 }
238 EvalMetric::LengthRatio { min, max } => {
239 let ref_len = case.reference_answer.as_deref().unwrap_or("").len();
240 if ref_len == 0 {
241 return 0.0;
242 }
243 let ratio = response.len() as f64 / ref_len as f64;
244 if ratio >= *min && ratio <= *max { 1.0 } else { 0.0 }
245 }
246 EvalMetric::Custom(_name) => 0.5,
247 }
248 }
249
250 pub fn compare_strategies(
254 reports: &[EvalReport],
255 ) -> Vec<(&EvalReport, f64)> {
256 if reports.is_empty() {
257 return Vec::new();
258 }
259 let max_cost = reports
260 .iter()
261 .map(|r| r.avg_cost_usd)
262 .fold(0.0_f64, f64::max);
263 let max_lat = reports
264 .iter()
265 .map(|r| r.avg_latency_ms)
266 .fold(0.0_f64, f64::max);
267
268 let mut scored: Vec<(&EvalReport, f64)> = reports
269 .iter()
270 .map(|r| {
271 let cost_norm = if max_cost == 0.0 {
272 0.0
273 } else {
274 r.avg_cost_usd / max_cost
275 };
276 let lat_norm = if max_lat == 0.0 {
277 0.0
278 } else {
279 r.avg_latency_ms / max_lat
280 };
281 let composite = r.pass_rate * 0.6 - cost_norm * 0.2 - lat_norm * 0.2;
282 (r, composite)
283 })
284 .collect();
285
286 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
287 scored
288 }
289
290 pub fn filter_by_tag(&self, tag: &str) -> Vec<&EvalCase> {
292 self.cases
293 .iter()
294 .filter(|c| c.tags.iter().any(|t| t == tag))
295 .collect()
296 }
297
298 pub fn difficulty_breakdown(&self, report: &EvalReport) -> HashMap<u8, f64> {
302 let difficulty_map: HashMap<&str, u8> = self
304 .cases
305 .iter()
306 .map(|c| (c.id.as_str(), c.difficulty))
307 .collect();
308
309 let mut totals: HashMap<u8, (u64, u64)> = HashMap::new(); for result in &report.results {
311 if let Some(&diff) = difficulty_map.get(result.case_id.as_str()) {
312 let entry = totals.entry(diff).or_insert((0, 0));
313 entry.1 += 1;
314 if result.passed {
315 entry.0 += 1;
316 }
317 }
318 }
319
320 totals
321 .into_iter()
322 .map(|(diff, (passed, total))| {
323 (diff, if total == 0 { 0.0 } else { passed as f64 / total as f64 })
324 })
325 .collect()
326 }
327}
328
329fn metric_name(metric: &EvalMetric) -> String {
332 match metric {
333 EvalMetric::ExactMatch => "exact_match".into(),
334 EvalMetric::ContainsAnswer => "contains_answer".into(),
335 EvalMetric::WordOverlap(_) => "word_overlap".into(),
336 EvalMetric::LengthRatio { .. } => "length_ratio".into(),
337 EvalMetric::Custom(name) => name.clone(),
338 }
339}
340
341fn word_set(text: &str) -> std::collections::HashSet<String> {
342 text.to_lowercase()
343 .split_whitespace()
344 .map(|w| w.trim_matches(|c: char| !c.is_alphanumeric()).to_string())
345 .filter(|w| !w.is_empty())
346 .collect()
347}
348
349#[cfg(test)]
352mod tests {
353 use super::*;
354
355 fn make_case(id: &str, reference: &str, tags: &[&str], difficulty: u8) -> EvalCase {
356 EvalCase {
357 id: id.into(),
358 prompt: format!("prompt for {}", id),
359 reference_answer: Some(reference.into()),
360 tags: tags.iter().map(|t| t.to_string()).collect(),
361 difficulty,
362 }
363 }
364
365 #[test]
366 fn test_exact_match_pass() {
367 let harness = EvalHarness::new();
368 let case = make_case("c1", "Paris", &[], 1);
369 assert_eq!(harness.score("Paris", &case, &EvalMetric::ExactMatch), 1.0);
370 }
371
372 #[test]
373 fn test_exact_match_fail() {
374 let harness = EvalHarness::new();
375 let case = make_case("c1", "Paris", &[], 1);
376 assert_eq!(harness.score("London", &case, &EvalMetric::ExactMatch), 0.0);
377 }
378
379 #[test]
380 fn test_exact_match_case_insensitive() {
381 let harness = EvalHarness::new();
382 let case = make_case("c1", "Paris", &[], 1);
383 assert_eq!(harness.score("paris", &case, &EvalMetric::ExactMatch), 1.0);
384 }
385
386 #[test]
387 fn test_contains_answer_pass() {
388 let harness = EvalHarness::new();
389 let case = make_case("c1", "42", &[], 1);
390 assert_eq!(
391 harness.score("The answer is 42.", &case, &EvalMetric::ContainsAnswer),
392 1.0
393 );
394 }
395
396 #[test]
397 fn test_contains_answer_fail() {
398 let harness = EvalHarness::new();
399 let case = make_case("c1", "42", &[], 1);
400 assert_eq!(
401 harness.score("The answer is 43.", &case, &EvalMetric::ContainsAnswer),
402 0.0
403 );
404 }
405
406 #[test]
407 fn test_word_overlap() {
408 let harness = EvalHarness::new();
409 let case = make_case("c1", "the quick brown fox", &[], 1);
410 let score = harness.score(
411 "the quick brown dog",
412 &case,
413 &EvalMetric::WordOverlap(0.5),
414 );
415 assert!((score - 0.6).abs() < 0.01);
417 }
418
419 #[test]
420 fn test_length_ratio_pass() {
421 let harness = EvalHarness::new();
422 let case = make_case("c1", "hello", &[], 1); assert_eq!(
425 harness.score("world", &case, &EvalMetric::LengthRatio { min: 0.8, max: 1.2 }),
426 1.0
427 );
428 }
429
430 #[test]
431 fn test_length_ratio_fail() {
432 let harness = EvalHarness::new();
433 let case = make_case("c1", "hi", &[], 1); let long_resp: String = "x".repeat(100);
436 assert_eq!(
437 harness.score(&long_resp, &case, &EvalMetric::LengthRatio { min: 0.8, max: 1.2 }),
438 0.0
439 );
440 }
441
442 #[test]
443 fn test_run_eval_pass_rate() {
444 let mut harness = EvalHarness::new();
445 harness.add_case(make_case("c1", "Paris", &[], 1));
446 harness.add_case(make_case("c2", "Berlin", &[], 2));
447 harness.add_case(make_case("c3", "Rome", &[], 3));
448
449 let responses = vec![
450 ("Paris".into(), 100, 0.001), ("Madrid".into(), 200, 0.001), ("Rome".into(), 150, 0.001), ];
454 let report = harness.run_eval("strategy-a", responses);
455 assert!((report.pass_rate - 2.0 / 3.0).abs() < 0.01);
456 assert_eq!(report.strategy_name, "strategy-a");
457 assert_eq!(report.results.len(), 3);
458 }
459
460 #[test]
461 fn test_avg_latency_and_cost() {
462 let mut harness = EvalHarness::new();
463 harness.add_case(make_case("c1", "x", &[], 1));
464 harness.add_case(make_case("c2", "y", &[], 1));
465
466 let report = harness.run_eval(
467 "s",
468 vec![("x".into(), 100, 0.01), ("y".into(), 200, 0.02)],
469 );
470 assert!((report.avg_latency_ms - 150.0).abs() < 0.01);
471 assert!((report.total_cost_usd - 0.03).abs() < 0.0001);
472 assert!((report.avg_cost_usd - 0.015).abs() < 0.0001);
473 }
474
475 #[test]
476 fn test_filter_by_tag() {
477 let mut harness = EvalHarness::new();
478 harness.add_case(make_case("c1", "x", &["math"], 1));
479 harness.add_case(make_case("c2", "y", &["science"], 2));
480 harness.add_case(make_case("c3", "z", &["math", "hard"], 3));
481
482 let math_cases = harness.filter_by_tag("math");
483 assert_eq!(math_cases.len(), 2);
484 }
485
486 #[test]
487 fn test_difficulty_breakdown() {
488 let mut harness = EvalHarness::new();
489 harness.add_case(make_case("c1", "Paris", &[], 1));
490 harness.add_case(make_case("c2", "Berlin", &[], 1));
491 harness.add_case(make_case("c3", "Rome", &[], 2));
492
493 let responses = vec![
494 ("Paris".into(), 100, 0.0), ("Madrid".into(), 100, 0.0), ("Rome".into(), 100, 0.0), ];
498 let report = harness.run_eval("s", responses);
499 let breakdown = harness.difficulty_breakdown(&report);
500 assert!((breakdown[&1] - 0.5).abs() < 0.01);
501 assert!((breakdown[&2] - 1.0).abs() < 0.01);
502 }
503
504 #[test]
505 fn test_compare_strategies() {
506 let mut harness = EvalHarness::new();
507 harness.add_case(make_case("c1", "Paris", &[], 1));
508
509 let r1 = harness.run_eval("good", vec![("Paris".into(), 50, 0.001)]);
510 let r2 = harness.run_eval("bad", vec![("London".into(), 200, 0.01)]);
511
512 let reports = vec![r1, r2];
513 let ranked = EvalHarness::compare_strategies(&reports);
514 assert_eq!(ranked.len(), 2);
515 assert_eq!(ranked[0].0.strategy_name, "good");
516 }
517
518 #[test]
519 fn test_custom_metric_returns_half() {
520 let harness = EvalHarness::new();
521 let case = make_case("c1", "any", &[], 1);
522 let score = harness.score("anything", &case, &EvalMetric::Custom("my_metric".into()));
523 assert!((score - 0.5).abs() < 0.01);
524 }
525}