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)]
64 pub kind: CaseKind,
65 #[serde(default)]
69 pub stale: Vec<String>,
70}
71
72#[derive(Debug, Clone, Serialize, Deserialize)]
74pub struct EvalFixture {
75 pub memories: Vec<EvalMemory>,
76 pub cases: Vec<EvalCase>,
77}
78
79pub fn recall_at_k(ranked: &[String], relevant: &[String], k: usize) -> f64 {
87 if relevant.is_empty() {
88 return 1.0;
89 }
90 if k == 0 || ranked.is_empty() {
91 return 0.0;
92 }
93 let window = &ranked[..k.min(ranked.len())];
94 let found = relevant
95 .iter()
96 .filter(|r| window.iter().any(|w| w == *r))
97 .count();
98 found as f64 / relevant.len() as f64
99}
100
101pub fn mrr(ranked: &[String], relevant: &[String]) -> f64 {
106 if relevant.is_empty() || ranked.is_empty() {
107 return 0.0;
108 }
109 for (idx, key) in ranked.iter().enumerate() {
110 if relevant.iter().any(|r| r == key) {
111 return 1.0 / (idx as f64 + 1.0);
112 }
113 }
114 0.0
115}
116
117pub fn mean(values: &[f64]) -> f64 {
121 if values.is_empty() {
122 return 0.0;
123 }
124 values.iter().sum::<f64>() / values.len() as f64
125}
126
127pub fn stale_hit_rate(ranked: &[String], stale: &[String], k: usize) -> f64 {
133 if stale.is_empty() || k == 0 || ranked.is_empty() {
134 return 0.0;
135 }
136 let window = &ranked[..k.min(ranked.len())];
137 if stale.iter().any(|s| window.iter().any(|w| w == s)) {
138 1.0
139 } else {
140 0.0
141 }
142}
143
144pub fn resolution_correct(ranked: &[String], relevant: &[String], stale: &[String]) -> bool {
154 if relevant.is_empty() {
155 return false;
156 }
157 let best_relevant = ranked
159 .iter()
160 .enumerate()
161 .find(|(_, k)| relevant.iter().any(|r| r == *k))
162 .map(|(i, _)| i);
163
164 let best_relevant = match best_relevant {
165 Some(pos) => pos,
166 None => return false, };
168
169 let best_stale = ranked
171 .iter()
172 .enumerate()
173 .find(|(_, k)| stale.iter().any(|s| s == *k))
174 .map(|(i, _)| i);
175
176 match best_stale {
177 None => true, Some(stale_pos) => best_relevant < stale_pos,
179 }
180}
181
182#[cfg(test)]
185mod tests {
186 use super::*;
187
188 fn s(v: &[&str]) -> Vec<String> {
189 v.iter().map(|x| x.to_string()).collect()
190 }
191
192 #[test]
195 fn recall_at_k_empty_relevant_is_one() {
196 assert_eq!(recall_at_k(&s(&["a", "b"]), &[], 4), 1.0);
198 assert_eq!(recall_at_k(&[], &[], 4), 1.0);
199 }
200
201 #[test]
202 fn recall_at_k_zero_k_is_zero() {
203 assert_eq!(recall_at_k(&s(&["a", "b"]), &s(&["a"]), 0), 0.0);
204 }
205
206 #[test]
207 fn recall_at_k_k_larger_than_ranked_uses_full_list() {
208 let ranked = s(&["a", "b"]);
210 let relevant = s(&["a", "b", "c"]);
211 let r = recall_at_k(&ranked, &relevant, 100);
213 assert!((r - 2.0 / 3.0).abs() < 1e-9);
214 }
215
216 #[test]
217 fn recall_at_k_exact_hits() {
218 let ranked = s(&["a", "b", "c", "d"]);
219 let relevant = s(&["b", "d"]);
220 assert!((recall_at_k(&ranked, &relevant, 2) - 0.5).abs() < 1e-9);
222 assert_eq!(recall_at_k(&ranked, &relevant, 4), 1.0);
224 }
225
226 #[test]
227 fn recall_at_k_duplicates_in_ranked_count_once() {
228 let ranked = s(&["a", "a", "b"]);
230 let relevant = s(&["a", "b"]);
231 assert_eq!(recall_at_k(&ranked, &relevant, 3), 1.0);
233 assert!((recall_at_k(&ranked, &relevant, 1) - 0.5).abs() < 1e-9);
235 }
236
237 #[test]
238 fn recall_at_k_no_hits_is_zero() {
239 let ranked = s(&["x", "y", "z"]);
240 let relevant = s(&["a", "b"]);
241 assert_eq!(recall_at_k(&ranked, &relevant, 5), 0.0);
242 }
243
244 #[test]
247 fn mrr_first_position_is_one() {
248 let ranked = s(&["a", "b", "c"]);
249 let relevant = s(&["a"]);
250 assert_eq!(mrr(&ranked, &relevant), 1.0);
251 }
252
253 #[test]
254 fn mrr_second_position_is_half() {
255 let ranked = s(&["x", "a", "b"]);
256 let relevant = s(&["a"]);
257 assert!((mrr(&ranked, &relevant) - 0.5).abs() < 1e-9);
258 }
259
260 #[test]
261 fn mrr_third_position_is_one_third() {
262 let ranked = s(&["x", "y", "a"]);
263 let relevant = s(&["a"]);
264 assert!((mrr(&ranked, &relevant) - 1.0 / 3.0).abs() < 1e-9);
265 }
266
267 #[test]
268 fn mrr_absent_is_zero() {
269 let ranked = s(&["x", "y", "z"]);
270 let relevant = s(&["a"]);
271 assert_eq!(mrr(&ranked, &relevant), 0.0);
272 }
273
274 #[test]
275 fn mrr_empty_relevant_is_zero() {
276 let ranked = s(&["a", "b"]);
277 assert_eq!(mrr(&ranked, &[]), 0.0);
278 }
279
280 #[test]
281 fn mrr_empty_ranked_is_zero() {
282 assert_eq!(mrr(&[], &s(&["a"]),), 0.0);
283 }
284
285 #[test]
286 fn mrr_uses_first_hit_when_multiple_relevant() {
287 let ranked = s(&["x", "b", "a"]);
289 let relevant = s(&["a", "b"]);
290 assert!((mrr(&ranked, &relevant) - 0.5).abs() < 1e-9);
291 }
292
293 #[test]
296 fn mean_empty_is_zero() {
297 assert_eq!(mean(&[]), 0.0);
298 }
299
300 #[test]
301 fn mean_single() {
302 assert!((mean(&[0.75]) - 0.75).abs() < 1e-9);
303 }
304
305 #[test]
306 fn mean_normal() {
307 let v = [0.0, 0.5, 1.0];
308 assert!((mean(&v) - 0.5).abs() < 1e-9);
309 }
310
311 #[test]
312 fn mean_all_ones() {
313 assert!((mean(&[1.0, 1.0, 1.0]) - 1.0).abs() < 1e-9);
314 }
315
316 #[test]
319 fn stale_hit_rate_no_stale_is_zero() {
320 assert_eq!(stale_hit_rate(&s(&["a", "b", "c"]), &[], 4), 0.0);
322 assert_eq!(stale_hit_rate(&[], &[], 4), 0.0);
323 }
324
325 #[test]
326 fn stale_hit_rate_stale_in_top_k_is_one() {
327 let ranked = s(&["a", "b", "c", "d"]);
329 let stale = s(&["b"]);
330 assert_eq!(stale_hit_rate(&ranked, &stale, 4), 1.0);
331 }
332
333 #[test]
334 fn stale_hit_rate_stale_beyond_k_is_zero() {
335 let ranked = s(&["a", "b", "c", "d"]);
337 let stale = s(&["d"]);
338 assert_eq!(stale_hit_rate(&ranked, &stale, 2), 0.0);
339 }
340
341 #[test]
342 fn stale_hit_rate_stale_absent_is_zero() {
343 let ranked = s(&["a", "b", "c"]);
344 let stale = s(&["z"]);
345 assert_eq!(stale_hit_rate(&ranked, &stale, 4), 0.0);
346 }
347
348 #[test]
351 fn resolution_correct_relevant_above_stale_is_true() {
352 let ranked = s(&["new", "x", "old"]);
354 assert!(resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
355 }
356
357 #[test]
358 fn resolution_correct_stale_above_relevant_is_false() {
359 let ranked = s(&["old", "x", "new"]);
361 assert!(!resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
362 }
363
364 #[test]
365 fn resolution_correct_stale_absent_is_true() {
366 let ranked = s(&["new", "x", "y"]);
368 assert!(resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
369 }
370
371 #[test]
372 fn resolution_correct_relevant_absent_is_false() {
373 let ranked = s(&["old", "x", "y"]);
375 assert!(!resolution_correct(&ranked, &s(&["new"]), &s(&["old"])));
376 }
377
378 #[test]
379 fn resolution_correct_empty_relevant_is_false() {
380 let ranked = s(&["new", "old"]);
381 assert!(!resolution_correct(&ranked, &[], &s(&["old"])));
382 }
383}