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    /// Empty = off-domain query (exercises noise floor, recall trivially 1.0).
60    pub relevant: Vec<String>,
61    /// Classification of this case. Defaults to [`CaseKind::Recall`].
62    /// Existing fixtures omit this field; `#[serde(default)]` keeps them valid.
63    #[serde(default)]
64    pub kind: CaseKind,
65    /// Keys of memories that should NOT appear in the top-k for this case
66    /// (superseded / contradicted / losing memories).
67    /// Empty by default — existing fixtures unchanged.
68    #[serde(default)]
69    pub stale: Vec<String>,
70}
71
72/// A committed eval fixture: a corpus of memories and a set of query cases.
73#[derive(Debug, Clone, Serialize, Deserialize)]
74pub struct EvalFixture {
75    pub memories: Vec<EvalMemory>,
76    pub cases: Vec<EvalCase>,
77}
78
79// ─── Metric math ─────────────────────────────────────────────────────────────
80
81/// Fraction of `relevant` items found in the **first `k`** positions of `ranked`.
82///
83/// Each relevant key is counted at most once even if it appears multiple times
84/// in `ranked`. Returns `1.0` when `relevant` is empty (trivial recall for
85/// off-domain / noise queries). Returns `0.0` when `k == 0`.
86pub 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
101/// Mean Reciprocal Rank of the **first** relevant item in `ranked` (1-based).
102///
103/// Returns `1/rank` where `rank` is the 1-based position of the first relevant
104/// item. Returns `0.0` when no relevant item appears in `ranked`.
105pub 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
117/// Arithmetic mean of a slice of metric values.
118///
119/// Returns `0.0` for an empty slice.
120pub 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
127/// Per-case stale-hit rate: returns `1.0` if any `stale` key is present in the
128/// first `k` positions of `ranked`, else `0.0`.
129///
130/// Lower is better. Averaged across cases → mean stale-hit rate.
131/// Returns `0.0` when `stale` is empty (no stale keys defined for this case).
132pub 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
144/// Returns `true` when the case is "resolved correctly": every `relevant` key
145/// that appears in `ranked` outranks every `stale` key that appears in `ranked`.
146///
147/// More precisely: the rank of the **best** (lowest-index) relevant key must be
148/// strictly less than the rank of the **best** stale key. If no stale key
149/// appears in `ranked` at all, the case is resolved (stale is absent — ideal).
150/// If no relevant key appears, the case is unresolved.
151///
152/// Used for contradiction / knowledge-update cases. Averaged → resolution accuracy.
153pub fn resolution_correct(ranked: &[String], relevant: &[String], stale: &[String]) -> bool {
154    if relevant.is_empty() {
155        return false;
156    }
157    // Position of the first relevant key in ranked (best = lowest index).
158    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, // no relevant in ranked → unresolved
167    };
168
169    // Position of the first (best) stale key in ranked.
170    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, // stale absent from ranked → ideal, resolved
178        Some(stale_pos) => best_relevant < stale_pos,
179    }
180}
181
182// ─── Unit tests ───────────────────────────────────────────────────────────────
183
184#[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    // ── recall_at_k ──────────────────────────────────────────────────────────
193
194    #[test]
195    fn recall_at_k_empty_relevant_is_one() {
196        // Off-domain queries: no relevant items → trivially 1.0.
197        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        // k > len(ranked): should still count everything in ranked.
209        let ranked = s(&["a", "b"]);
210        let relevant = s(&["a", "b", "c"]);
211        // 2 of 3 found in first 100 positions → 2/3.
212        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        // k=2: only "b" in first 2 → 0.5
221        assert!((recall_at_k(&ranked, &relevant, 2) - 0.5).abs() < 1e-9);
222        // k=4: both found → 1.0
223        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        // "a" appears twice in ranked, but should only count as 1 hit.
229        let ranked = s(&["a", "a", "b"]);
230        let relevant = s(&["a", "b"]);
231        // Both are in first 3 positions → 2/2 = 1.0 (not 3/2).
232        assert_eq!(recall_at_k(&ranked, &relevant, 3), 1.0);
233        // k=1: "a" appears → 1 of 2 relevant found = 0.5
234        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    // ── mrr ──────────────────────────────────────────────────────────────────
245
246    #[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        // "b" is at rank 2, "a" is at rank 3 — MRR should be 1/2.
288        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    // ── mean ─────────────────────────────────────────────────────────────────
294
295    #[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    // ── stale_hit_rate ────────────────────────────────────────────────────────
317
318    #[test]
319    fn stale_hit_rate_no_stale_is_zero() {
320        // No stale keys defined → always 0.0 regardless of ranked.
321        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        // "b" is stale and is at rank 2 (within k=4 window) → 1.0.
328        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        // "d" is stale but is at rank 4; window k=2 → 0.0.
336        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    // ── resolution_correct ────────────────────────────────────────────────────
349
350    #[test]
351    fn resolution_correct_relevant_above_stale_is_true() {
352        // relevant "new" at rank 1, stale "old" at rank 3 → resolved.
353        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        // stale "old" at rank 1, relevant "new" at rank 3 → NOT resolved.
360        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        // relevant present, stale absent from ranked → ideal resolution.
367        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        // relevant missing from ranked entirely → cannot be resolved.
374        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}