Skip to main content

tokio_prompt_orchestrator/
eval_harness.rs

1//! # Evaluation Harness
2//!
3//! Compare prompt strategies across a shared set of [`EvalCase`] items.
4//! Supports multiple scoring metrics, composite ranking, tag filtering,
5//! and per-difficulty pass-rate breakdowns.
6//!
7//! ## Example
8//!
9//! ```rust
10//! use tokio_prompt_orchestrator::eval_harness::{
11//!     EvalCase, EvalHarness, EvalMetric,
12//! };
13//!
14//! let mut harness = EvalHarness::new();
15//! harness.add_case(EvalCase {
16//!     id: "q1".into(),
17//!     prompt: "What is 2+2?".into(),
18//!     reference_answer: Some("4".into()),
19//!     tags: vec!["math".into()],
20//!     difficulty: 1,
21//! });
22//! let report = harness.run_eval(
23//!     "baseline",
24//!     vec![("4".to_string(), 120, 0.001)],
25//! );
26//! assert!(report.pass_rate > 0.99);
27//! ```
28
29use std::collections::HashMap;
30
31// ── EvalCase ──────────────────────────────────────────────────────────────────
32
33/// A single evaluation case.
34#[derive(Debug, Clone)]
35pub struct EvalCase {
36    /// Unique identifier.
37    pub id: String,
38    /// The prompt sent to the strategy under test.
39    pub prompt: String,
40    /// Optional gold-standard answer for scoring.
41    pub reference_answer: Option<String>,
42    /// Arbitrary tags for filtering (e.g. `"math"`, `"hard"`).
43    pub tags: Vec<String>,
44    /// Difficulty level 1 (easy) – 5 (very hard).
45    pub difficulty: u8,
46}
47
48// ── EvalMetric ────────────────────────────────────────────────────────────────
49
50/// Scoring metric applied to a strategy response.
51#[derive(Debug, Clone)]
52pub enum EvalMetric {
53    /// Score is 1.0 iff response equals the reference answer (trimmed, case-insensitive).
54    ExactMatch,
55    /// Score is 1.0 iff the reference answer appears as a substring of the response.
56    ContainsAnswer,
57    /// Score is the Jaccard word-overlap coefficient; threshold ignored in scoring.
58    WordOverlap(f64),
59    /// Score is 1.0 iff `len(response) / len(reference)` is within `[min, max]`.
60    LengthRatio { min: f64, max: f64 },
61    /// A named custom metric (always scores 0.5 unless overridden externally).
62    Custom(String),
63}
64
65// ── EvalResult ────────────────────────────────────────────────────────────────
66
67/// Result for a single case in a strategy evaluation.
68#[derive(Debug, Clone)]
69pub struct EvalResult {
70    /// ID of the [`EvalCase`] this result corresponds to.
71    pub case_id: String,
72    /// The response produced by the strategy.
73    pub response: String,
74    /// Named metric scores.
75    pub metrics: HashMap<String, f64>,
76    /// `true` if the primary metric score exceeds the pass threshold (0.5).
77    pub passed: bool,
78    /// Wall-clock latency in milliseconds.
79    pub latency_ms: u64,
80    /// API cost in USD.
81    pub cost_usd: f64,
82}
83
84// ── EvalReport ────────────────────────────────────────────────────────────────
85
86/// Aggregated results for one strategy across all cases.
87#[derive(Debug, Clone)]
88pub struct EvalReport {
89    /// Name of the strategy evaluated.
90    pub strategy_name: String,
91    /// Per-case results.
92    pub results: Vec<EvalResult>,
93    /// Fraction of cases that passed (0.0 – 1.0).
94    pub pass_rate: f64,
95    /// Mean latency across all cases.
96    pub avg_latency_ms: f64,
97    /// Mean per-case cost.
98    pub avg_cost_usd: f64,
99    /// Sum of all per-case costs.
100    pub total_cost_usd: f64,
101}
102
103// ── EvalHarness ───────────────────────────────────────────────────────────────
104
105/// Evaluation harness for comparing prompt strategies.
106pub struct EvalHarness {
107    cases: Vec<EvalCase>,
108    /// Primary metric used for pass/fail determination (default: ExactMatch).
109    primary_metric: EvalMetric,
110}
111
112impl Default for EvalHarness {
113    fn default() -> Self {
114        Self::new()
115    }
116}
117
118impl EvalHarness {
119    /// Create a new harness with [`EvalMetric::ExactMatch`] as the primary metric.
120    pub fn new() -> Self {
121        Self {
122            cases: Vec::new(),
123            primary_metric: EvalMetric::ExactMatch,
124        }
125    }
126
127    /// Create a harness with a custom primary metric.
128    pub fn with_metric(metric: EvalMetric) -> Self {
129        Self {
130            cases: Vec::new(),
131            primary_metric: metric,
132        }
133    }
134
135    /// Register a new evaluation case.
136    pub fn add_case(&mut self, case: EvalCase) {
137        self.cases.push(case);
138    }
139
140    /// Run an evaluation for `strategy_name`.
141    ///
142    /// `responses` must be in the same order as cases were added.  Each element
143    /// is `(response_text, latency_ms, cost_usd)`.
144    ///
145    /// # Panics
146    ///
147    /// Does not panic — if `responses` is shorter than `cases`, remaining cases
148    /// are skipped.
149    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    /// Score a single response against a case using the supplied metric.
200    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    /// Rank strategies by composite score: `pass_rate * 0.6 - avg_cost_usd_norm * 0.2 - avg_latency_norm * 0.2`.
251    ///
252    /// Returns references to reports in descending composite score order.
253    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    /// Return all cases tagged with `tag`.
291    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    /// Compute per-difficulty pass rate from a report.
299    ///
300    /// Keys are difficulty levels (1–5); values are pass rates (0.0 – 1.0).
301    pub fn difficulty_breakdown(&self, report: &EvalReport) -> HashMap<u8, f64> {
302        // Build a lookup from case_id to difficulty.
303        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(); // (passed, total)
310        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
329// ── helpers ───────────────────────────────────────────────────────────────────
330
331fn 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// ── Tests ─────────────────────────────────────────────────────────────────────
350
351#[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        // intersection={the,quick,brown}=3, union={the,quick,brown,fox,dog}=5, jaccard=0.6
416        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); // len=5
423        // response len=5, ratio=1.0, within [0.8, 1.2]
424        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); // len=2
434        // response len=100, ratio=50.0, outside [0.8, 1.2]
435        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),  // pass
451            ("Madrid".into(), 200, 0.001), // fail
452            ("Rome".into(), 150, 0.001),   // pass
453        ];
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),  // diff=1, pass
495            ("Madrid".into(), 100, 0.0), // diff=1, fail
496            ("Rome".into(), 100, 0.0),   // diff=2, pass
497        ];
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}