Skip to main content

kimetsu_brain/
eval.rs

1//! Retrieval quality metrics for `kimetsu brain eval`.
2//!
3//! Pure, no-I/O module. All functions operate on slices of `String`
4//! (memory keys / ranked result keys) so they are trivially unit-testable.
5
6use serde::{Deserialize, Serialize};
7
8// ─── Fixture types ────────────────────────────────────────────────────────────
9
10/// A single corpus memory: stable key (referenced by [`EvalCase::relevant`]) and text.
11#[derive(Debug, Clone, Serialize, Deserialize)]
12pub struct EvalMemory {
13    /// Stable short key used to cross-reference from [`EvalCase::relevant`].
14    pub key: String,
15    /// Full text of the memory to add to the corpus.
16    pub text: String,
17    /// Flagship 1 Pass A: optional RFC 3339 timestamp. When present and in the
18    /// PAST, the bench seeder stamps this memory with `valid_to` (expired) so
19    /// validity-aware retrieval excludes it. Omitting this field leaves the
20    /// memory valid indefinitely — existing fixtures are unchanged.
21    #[serde(default)]
22    pub valid_to: Option<String>,
23    /// Flagship 1 Pass A: optional key of another `EvalMemory` that supersedes
24    /// this one. When present, the bench seeder stamps `superseded_by` on this
25    /// memory (pointing to the survivor's DB id) so retrieval excludes it via
26    /// the existing `superseded_by IS NULL` guard.
27    /// Omitting this field leaves the memory active — existing fixtures unchanged.
28    #[serde(default)]
29    pub superseded_by_key: Option<String>,
30}
31
32/// Classification of an eval case for correctness measurement.
33///
34/// `Recall` is the default (and the only kind used by existing fixtures).
35/// The new kinds are used by `bench/dataset-correctness.json` to measure
36/// temporal correctness, contradiction resolution, and knowledge-update quality.
37#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
38#[serde(rename_all = "snake_case")]
39pub enum CaseKind {
40    /// Plain retrieval: the relevant memory should appear in the top-k.
41    #[default]
42    Recall,
43    /// A newer memory supersedes an older one; the query asks for current state.
44    /// The current/correct memory should win; the stale one should NOT appear.
45    KnowledgeUpdate,
46    /// Two memories make contradictory claims; the authoritative one should win.
47    Contradiction,
48    /// A fact is qualified by an as-of date; the most recent should win.
49    Temporal,
50    /// Multiple sessions produced overlapping memories; the canonical one wins.
51    MultiSession,
52}
53
54/// One eval case: a query plus the set of corpus keys that are relevant to it.
55#[derive(Debug, Clone, Serialize, Deserialize)]
56pub struct EvalCase {
57    pub query: String,
58    /// Keys from [`EvalMemory::key`] that are relevant to this query.
59    /// Explicit empty = known negative. Missing is invalid fixture data.
60    pub relevant: Vec<String>,
61    /// Optional task/fact family. Shared IDs and normalized queries also group cases.
62    #[serde(default)]
63    pub family: String,
64    /// Classification of this case. Defaults to [`CaseKind::Recall`].
65    /// Existing fixtures omit this field; `#[serde(default)]` keeps them valid.
66    #[serde(default)]
67    pub kind: CaseKind,
68    /// Keys of memories that should NOT appear in the top-k for this case
69    /// (superseded / contradicted / losing memories).
70    /// Empty by default — existing fixtures unchanged.
71    #[serde(default)]
72    pub stale: Vec<String>,
73}
74
75/// A committed eval fixture: a corpus of memories and a set of query cases.
76#[derive(Debug, Clone, Serialize, Deserialize)]
77pub struct EvalFixture {
78    pub memories: Vec<EvalMemory>,
79    pub cases: Vec<EvalCase>,
80}
81
82// ─── Metric math ─────────────────────────────────────────────────────────────
83
84/// Fraction of `relevant` items found in the **first `k`** positions of `ranked`.
85///
86/// Each relevant key is counted at most once even if it appears multiple times
87/// in `ranked`. Empty relevance returns zero; exclude known negatives from
88/// recall denominators and score their actual abstention separately.
89pub 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/// Metrics describe delivered results, never the pre-budget candidate pool.
106/// Missing class denominators serialize as null, not a perfect score.
107#[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    /// Equal weight for positive MRR and negative abstention accuracy when both
119    /// exist; otherwise the observed component only. Not factual probability.
120    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
181/// Mean Reciprocal Rank of the **first** relevant item in `ranked` (1-based).
182///
183/// Returns `1/rank` where `rank` is the 1-based position of the first relevant
184/// item. Returns `0.0` when no relevant item appears in `ranked`.
185pub 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
197/// Arithmetic mean of a slice of metric values.
198///
199/// Returns `0.0` for an empty slice.
200pub 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
207/// Per-case stale-hit rate: returns `1.0` if any `stale` key is present in the
208/// first `k` positions of `ranked`, else `0.0`.
209///
210/// Lower is better. Averaged across cases → mean stale-hit rate.
211/// Returns `0.0` when `stale` is empty (no stale keys defined for this case).
212pub 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
224/// Returns `true` when the case is "resolved correctly": every `relevant` key
225/// that appears in `ranked` outranks every `stale` key that appears in `ranked`.
226///
227/// More precisely: the rank of the **best** (lowest-index) relevant key must be
228/// strictly less than the rank of the **best** stale key. If no stale key
229/// appears in `ranked` at all, the case is resolved (stale is absent — ideal).
230/// If no relevant key appears, the case is unresolved.
231///
232/// Used for contradiction / knowledge-update cases. Averaged → resolution accuracy.
233pub fn resolution_correct(ranked: &[String], relevant: &[String], stale: &[String]) -> bool {
234    if relevant.is_empty() {
235        return false;
236    }
237    // Position of the first relevant key in ranked (best = lowest index).
238    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, // no relevant in ranked → unresolved
247    };
248
249    // Position of the first (best) stale key in ranked.
250    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, // stale absent from ranked → ideal, resolved
258        Some(stale_pos) => best_relevant < stale_pos,
259    }
260}
261
262// ─── Unit tests ───────────────────────────────────────────────────────────────
263
264#[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    // ── recall_at_k ──────────────────────────────────────────────────────────
273
274    #[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        // k > len(ranked): should still count everything in ranked.
318        let ranked = s(&["a", "b"]);
319        let relevant = s(&["a", "b", "c"]);
320        // 2 of 3 found in first 100 positions → 2/3.
321        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        // k=2: only "b" in first 2 → 0.5
330        assert!((recall_at_k(&ranked, &relevant, 2) - 0.5).abs() < 1e-9);
331        // k=4: both found → 1.0
332        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        // "a" appears twice in ranked, but should only count as 1 hit.
338        let ranked = s(&["a", "a", "b"]);
339        let relevant = s(&["a", "b"]);
340        // Both are in first 3 positions → 2/2 = 1.0 (not 3/2).
341        assert_eq!(recall_at_k(&ranked, &relevant, 3), 1.0);
342        // k=1: "a" appears → 1 of 2 relevant found = 0.5
343        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    // ── mrr ──────────────────────────────────────────────────────────────────
354
355    #[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        // "b" is at rank 2, "a" is at rank 3 — MRR should be 1/2.
397        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    // ── mean ─────────────────────────────────────────────────────────────────
403
404    #[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    // ── stale_hit_rate ────────────────────────────────────────────────────────
426
427    #[test]
428    fn stale_hit_rate_no_stale_is_zero() {
429        // No stale keys defined → always 0.0 regardless of ranked.
430        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        // "b" is stale and is at rank 2 (within k=4 window) → 1.0.
437        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        // "d" is stale but is at rank 4; window k=2 → 0.0.
445        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    // ── resolution_correct ────────────────────────────────────────────────────
458
459    #[test]
460    fn resolution_correct_relevant_above_stale_is_true() {
461        // relevant "new" at rank 1, stale "old" at rank 3 → resolved.
462        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        // stale "old" at rank 1, relevant "new" at rank 3 → NOT resolved.
469        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        // relevant present, stale absent from ranked → ideal resolution.
476        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        // relevant missing from ranked entirely → cannot be resolved.
483        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}