1use serde::{Deserialize, Serialize};
7
8#[derive(Debug, Clone, Serialize, Deserialize)]
12pub struct EvalMemory {
13 pub key: String,
15 pub text: String,
17 #[serde(default)]
22 pub valid_to: Option<String>,
23 #[serde(default)]
29 pub superseded_by_key: Option<String>,
30}
31
32#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
38#[serde(rename_all = "snake_case")]
39pub enum CaseKind {
40 #[default]
42 Recall,
43 KnowledgeUpdate,
46 Contradiction,
48 Temporal,
50 MultiSession,
52}
53
54#[derive(Debug, Clone, Serialize, Deserialize)]
56pub struct EvalCase {
57 pub query: String,
58 pub relevant: Vec<String>,
61 #[serde(default)]
63 pub family: String,
64 #[serde(default)]
67 pub kind: CaseKind,
68 #[serde(default)]
72 pub stale: Vec<String>,
73}
74
75#[derive(Debug, Clone, Serialize, Deserialize)]
77pub struct EvalFixture {
78 pub memories: Vec<EvalMemory>,
79 pub cases: Vec<EvalCase>,
80}
81
82pub fn recall_at_k(ranked: &[String], relevant: &[String], k: usize) -> f64 {
90 if relevant.is_empty() {
91 return 0.0;
92 }
93 if k == 0 || ranked.is_empty() {
94 return 0.0;
95 }
96 let window = &ranked[..k.min(ranked.len())];
97 let relevant: std::collections::HashSet<_> = relevant.iter().collect();
98 let found = relevant
99 .iter()
100 .filter(|r| window.iter().any(|w| w == **r))
101 .count();
102 found as f64 / relevant.len() as f64
103}
104
105#[derive(Debug, Clone, Default, Serialize, Deserialize)]
108pub struct EvaluationMetrics {
109 pub positive_count: usize,
110 pub negative_count: usize,
111 pub recall_at_2: Option<f64>,
112 pub recall_at_4: Option<f64>,
113 pub hit_at_2: Option<f64>,
114 pub hit_at_4: Option<f64>,
115 pub mrr: Option<f64>,
116 pub negative_accuracy: Option<f64>,
117 pub false_injection_rate: Option<f64>,
118 pub quality: Option<f64>,
121 pub mean_final_bound: f64,
122}
123
124pub fn summarize_deliveries(
125 cases: &[&EvalCase],
126 ranked: &[Vec<String>],
127 final_bounds: &[u32],
128) -> Result<EvaluationMetrics, String> {
129 if cases.len() != ranked.len() || cases.len() != final_bounds.len() {
130 return Err("every case requires one delivery and final cost measurement".into());
131 }
132 let positives: Vec<_> = cases
133 .iter()
134 .zip(ranked)
135 .filter(|(c, _)| !c.relevant.is_empty())
136 .collect();
137 let negatives: Vec<_> = cases
138 .iter()
139 .zip(ranked)
140 .filter(|(c, _)| c.relevant.is_empty())
141 .collect();
142 let avg = |values: Vec<f64>| (!values.is_empty()).then(|| mean(&values));
143 let recall = |k| {
144 avg(positives
145 .iter()
146 .map(|(c, r)| recall_at_k(r, &c.relevant, k))
147 .collect())
148 };
149 let hit = |k| {
150 avg(positives
151 .iter()
152 .map(|(c, r)| f64::from(recall_at_k(r, &c.relevant, k) > 0.0))
153 .collect())
154 };
155 let mrr = avg(positives.iter().map(|(c, r)| mrr(r, &c.relevant)).collect());
156 let negative_accuracy = avg(negatives
157 .iter()
158 .map(|(_, r)| f64::from(r.is_empty()))
159 .collect());
160 let quality = avg(mrr.into_iter().chain(negative_accuracy).collect());
161 Ok(EvaluationMetrics {
162 positive_count: positives.len(),
163 negative_count: negatives.len(),
164 recall_at_2: recall(2),
165 recall_at_4: recall(4),
166 hit_at_2: hit(2),
167 hit_at_4: hit(4),
168 mrr,
169 negative_accuracy,
170 false_injection_rate: negative_accuracy.map(|a| 1.0 - a),
171 quality,
172 mean_final_bound: mean(
173 &final_bounds
174 .iter()
175 .map(|b| f64::from(*b))
176 .collect::<Vec<_>>(),
177 ),
178 })
179}
180
181pub fn mrr(ranked: &[String], relevant: &[String]) -> f64 {
186 if relevant.is_empty() || ranked.is_empty() {
187 return 0.0;
188 }
189 for (idx, key) in ranked.iter().enumerate() {
190 if relevant.iter().any(|r| r == key) {
191 return 1.0 / (idx as f64 + 1.0);
192 }
193 }
194 0.0
195}
196
197pub fn mean(values: &[f64]) -> f64 {
201 if values.is_empty() {
202 return 0.0;
203 }
204 values.iter().sum::<f64>() / values.len() as f64
205}
206
207pub fn stale_hit_rate(ranked: &[String], stale: &[String], k: usize) -> f64 {
213 if stale.is_empty() || k == 0 || ranked.is_empty() {
214 return 0.0;
215 }
216 let window = &ranked[..k.min(ranked.len())];
217 if stale.iter().any(|s| window.iter().any(|w| w == s)) {
218 1.0
219 } else {
220 0.0
221 }
222}
223
224pub fn resolution_correct(ranked: &[String], relevant: &[String], stale: &[String]) -> bool {
234 if relevant.is_empty() {
235 return false;
236 }
237 let best_relevant = ranked
239 .iter()
240 .enumerate()
241 .find(|(_, k)| relevant.iter().any(|r| r == *k))
242 .map(|(i, _)| i);
243
244 let best_relevant = match best_relevant {
245 Some(pos) => pos,
246 None => return false, };
248
249 let best_stale = ranked
251 .iter()
252 .enumerate()
253 .find(|(_, k)| stale.iter().any(|s| s == *k))
254 .map(|(i, _)| i);
255
256 match best_stale {
257 None => true, Some(stale_pos) => best_relevant < stale_pos,
259 }
260}
261
262#[cfg(test)]
265mod tests {
266 use super::*;
267
268 fn s(v: &[&str]) -> Vec<String> {
269 v.iter().map(|x| x.to_string()).collect()
270 }
271
272 #[test]
275 fn recall_at_k_negative_has_no_vacuous_quality_credit() {
276 assert_eq!(recall_at_k(&s(&["a", "b"]), &[], 4), 0.0);
277 assert_eq!(recall_at_k(&[], &[], 4), 0.0);
278 }
279
280 #[test]
281 fn delivered_metrics_separate_fraction_hit_and_known_negative_accuracy() {
282 let positive: EvalCase =
283 serde_json::from_value(serde_json::json!({"query":"two facts","relevant":["a","b"]}))
284 .unwrap();
285 let negative: EvalCase =
286 serde_json::from_value(serde_json::json!({"query":"unanswerable","relevant":[]}))
287 .unwrap();
288 let metrics =
289 summarize_deliveries(&[&positive, &negative], &[s(&["a"]), vec![]], &[512, 256])
290 .unwrap();
291 assert_eq!(metrics.recall_at_2, Some(0.5));
292 assert_eq!(metrics.hit_at_2, Some(1.0));
293 assert_eq!(metrics.mrr, Some(1.0));
294 assert_eq!(metrics.negative_accuracy, Some(1.0));
295 assert_eq!(metrics.quality, Some(1.0));
296 assert_eq!(metrics.mean_final_bound, 384.0);
297 let injecting = summarize_deliveries(
298 &[&positive, &negative],
299 &[s(&["a"]), s(&["junk"])],
300 &[512, 512],
301 )
302 .unwrap();
303 assert_eq!(injecting.quality, Some(0.5));
304 let only_negative = summarize_deliveries(&[&negative], &[vec![]], &[256]).unwrap();
305 assert_eq!(only_negative.mrr, None);
306 assert_eq!(only_negative.recall_at_2, None);
307 assert!(summarize_deliveries(&[&positive], &[], &[]).is_err());
308 }
309
310 #[test]
311 fn recall_at_k_zero_k_is_zero() {
312 assert_eq!(recall_at_k(&s(&["a", "b"]), &s(&["a"]), 0), 0.0);
313 }
314
315 #[test]
316 fn recall_at_k_k_larger_than_ranked_uses_full_list() {
317 let ranked = s(&["a", "b"]);
319 let relevant = s(&["a", "b", "c"]);
320 let r = recall_at_k(&ranked, &relevant, 100);
322 assert!((r - 2.0 / 3.0).abs() < 1e-9);
323 }
324
325 #[test]
326 fn recall_at_k_exact_hits() {
327 let ranked = s(&["a", "b", "c", "d"]);
328 let relevant = s(&["b", "d"]);
329 assert!((recall_at_k(&ranked, &relevant, 2) - 0.5).abs() < 1e-9);
331 assert_eq!(recall_at_k(&ranked, &relevant, 4), 1.0);
333 }
334
335 #[test]
336 fn recall_at_k_duplicates_in_ranked_count_once() {
337 let ranked = s(&["a", "a", "b"]);
339 let relevant = s(&["a", "b"]);
340 assert_eq!(recall_at_k(&ranked, &relevant, 3), 1.0);
342 assert!((recall_at_k(&ranked, &relevant, 1) - 0.5).abs() < 1e-9);
344 }
345
346 #[test]
347 fn recall_at_k_no_hits_is_zero() {
348 let ranked = s(&["x", "y", "z"]);
349 let relevant = s(&["a", "b"]);
350 assert_eq!(recall_at_k(&ranked, &relevant, 5), 0.0);
351 }
352
353 #[test]
356 fn mrr_first_position_is_one() {
357 let ranked = s(&["a", "b", "c"]);
358 let relevant = s(&["a"]);
359 assert_eq!(mrr(&ranked, &relevant), 1.0);
360 }
361
362 #[test]
363 fn mrr_second_position_is_half() {
364 let ranked = s(&["x", "a", "b"]);
365 let relevant = s(&["a"]);
366 assert!((mrr(&ranked, &relevant) - 0.5).abs() < 1e-9);
367 }
368
369 #[test]
370 fn mrr_third_position_is_one_third() {
371 let ranked = s(&["x", "y", "a"]);
372 let relevant = s(&["a"]);
373 assert!((mrr(&ranked, &relevant) - 1.0 / 3.0).abs() < 1e-9);
374 }
375
376 #[test]
377 fn mrr_absent_is_zero() {
378 let ranked = s(&["x", "y", "z"]);
379 let relevant = s(&["a"]);
380 assert_eq!(mrr(&ranked, &relevant), 0.0);
381 }
382
383 #[test]
384 fn mrr_empty_relevant_is_zero() {
385 let ranked = s(&["a", "b"]);
386 assert_eq!(mrr(&ranked, &[]), 0.0);
387 }
388
389 #[test]
390 fn mrr_empty_ranked_is_zero() {
391 assert_eq!(mrr(&[], &s(&["a"]),), 0.0);
392 }
393
394 #[test]
395 fn mrr_uses_first_hit_when_multiple_relevant() {
396 let ranked = s(&["x", "b", "a"]);
398 let relevant = s(&["a", "b"]);
399 assert!((mrr(&ranked, &relevant) - 0.5).abs() < 1e-9);
400 }
401
402 #[test]
405 fn mean_empty_is_zero() {
406 assert_eq!(mean(&[]), 0.0);
407 }
408
409 #[test]
410 fn mean_single() {
411 assert!((mean(&[0.75]) - 0.75).abs() < 1e-9);
412 }
413
414 #[test]
415 fn mean_normal() {
416 let v = [0.0, 0.5, 1.0];
417 assert!((mean(&v) - 0.5).abs() < 1e-9);
418 }
419
420 #[test]
421 fn mean_all_ones() {
422 assert!((mean(&[1.0, 1.0, 1.0]) - 1.0).abs() < 1e-9);
423 }
424
425 #[test]
428 fn stale_hit_rate_no_stale_is_zero() {
429 assert_eq!(stale_hit_rate(&s(&["a", "b", "c"]), &[], 4), 0.0);
431 assert_eq!(stale_hit_rate(&[], &[], 4), 0.0);
432 }
433
434 #[test]
435 fn stale_hit_rate_stale_in_top_k_is_one() {
436 let ranked = s(&["a", "b", "c", "d"]);
438 let stale = s(&["b"]);
439 assert_eq!(stale_hit_rate(&ranked, &stale, 4), 1.0);
440 }
441
442 #[test]
443 fn stale_hit_rate_stale_beyond_k_is_zero() {
444 let ranked = s(&["a", "b", "c", "d"]);
446 let stale = s(&["d"]);
447 assert_eq!(stale_hit_rate(&ranked, &stale, 2), 0.0);
448 }
449
450 #[test]
451 fn stale_hit_rate_stale_absent_is_zero() {
452 let ranked = s(&["a", "b", "c"]);
453 let stale = s(&["z"]);
454 assert_eq!(stale_hit_rate(&ranked, &stale, 4), 0.0);
455 }
456
457 #[test]
460 fn resolution_correct_relevant_above_stale_is_true() {
461 let ranked = s(&["new", "x", "old"]);
463 assert!(resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
464 }
465
466 #[test]
467 fn resolution_correct_stale_above_relevant_is_false() {
468 let ranked = s(&["old", "x", "new"]);
470 assert!(!resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
471 }
472
473 #[test]
474 fn resolution_correct_stale_absent_is_true() {
475 let ranked = s(&["new", "x", "y"]);
477 assert!(resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
478 }
479
480 #[test]
481 fn resolution_correct_relevant_absent_is_false() {
482 let ranked = s(&["old", "x", "y"]);
484 assert!(!resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
485 }
486
487 #[test]
488 fn resolution_correct_empty_relevant_is_false() {
489 let ranked = s(&["new", "old"]);
490 assert!(!resolution_correct(&ranked, &[], &s(&["old"])));
491 }
492}