gemini_memory_rs/evals/
harness.rs

1//! Running the evaluation corpus against the real engine.
2//!
3//! The harness exercises the same code path a live session would: a plan is
4//! derived from the utterance, queries are fused, and the assembler produces a
5//! budgeted snapshot. Nothing is stubbed, so a regression in ranking or budget
6//! shows up here rather than in production.
7
8use chrono::Utc;
9use std::sync::Arc;
10
11use super::fixtures::{
12    IngestionCase, RetrievalCase, corpus, eval_user, ingestion_cases, retrieval_cases,
13};
14use super::metrics;
15use crate::bm25::{IndexedMemory, MemoryIndex};
16use crate::core::{
17    AdmissionVerdict, IngestionConfig, MemoryError, MemoryStatus, RetrievalConfig, SessionId,
18    TurnId, admit_observation,
19};
20use crate::ingestion::{
21    MemoryObservationExtractor, ObservationExtractionContext, RuleBasedObservationExtractor,
22};
23use crate::retrieval::{
24    DeterministicPlanner, IndexHandle, KnownEntities, LocalMemoryRetriever, MemoryRetriever,
25    RetrievalRequest,
26};
27
28/// Per-case retrieval outcome.
29#[derive(Debug, Clone)]
30pub struct RetrievalCaseResult {
31    /// Case name.
32    pub name: &'static str,
33    /// Records returned, in rank order.
34    pub returned: Vec<String>,
35    /// Precision over the case's relevant set.
36    pub precision: f32,
37    /// Recall over the case's relevant set.
38    pub recall: f32,
39    /// Reciprocal rank of the first relevant record.
40    pub reciprocal_rank: f32,
41    /// Whether the context the case produced matched the expectation: facts
42    /// for a case that needs memory, nothing for one that does not.
43    pub skip_correct: bool,
44    /// Whether any forbidden record was returned.
45    pub leaked_forbidden: bool,
46    /// Tokens the assembled context cost.
47    pub tokens: usize,
48}
49
50/// Aggregate retrieval report.
51#[derive(Debug, Clone)]
52pub struct RetrievalReport {
53    /// Per-case detail.
54    pub cases: Vec<RetrievalCaseResult>,
55    /// Mean precision across cases that expected memory.
56    pub precision: f32,
57    /// Mean recall across cases that expected memory.
58    pub recall: f32,
59    /// Mean reciprocal rank.
60    pub mrr: f32,
61    /// Fraction of cases whose skip decision was right.
62    pub skip_accuracy: f32,
63    /// Mean context size in tokens.
64    pub mean_tokens: f32,
65    /// 95th-percentile context size in tokens.
66    pub p95_tokens: f32,
67}
68
69impl RetrievalReport {
70    /// Cases where a forbidden record surfaced.
71    pub fn leaks(&self) -> Vec<&'static str> {
72        self.cases
73            .iter()
74            .filter(|c| c.leaked_forbidden)
75            .map(|c| c.name)
76            .collect()
77    }
78
79    /// A one-line-per-case rendering, for eyeballing a run.
80    pub fn render(&self) -> String {
81        use std::fmt::Write as _;
82        let mut out = String::new();
83        for case in &self.cases {
84            let _ = writeln!(
85                out,
86                "{:<44} p={:.2} r={:.2} rr={:.2} tokens={:<4} {}",
87                case.name,
88                case.precision,
89                case.recall,
90                case.reciprocal_rank,
91                case.tokens,
92                if case.skip_correct { "" } else { "SKIP-WRONG" }
93            );
94        }
95        let _ = writeln!(
96            out,
97            "\nprecision={:.3} recall={:.3} mrr={:.3} skip={:.3} tokens(mean)={:.1} tokens(p95)={:.1}",
98            self.precision,
99            self.recall,
100            self.mrr,
101            self.skip_accuracy,
102            self.mean_tokens,
103            self.p95_tokens
104        );
105        out
106    }
107}
108
109/// Build a retriever over the evaluation corpus.
110fn eval_retriever() -> (LocalMemoryRetriever, Arc<DeterministicPlanner>) {
111    let records = corpus();
112    let index = MemoryIndex::build(
113        records
114            .iter()
115            .filter(|m| m.status == MemoryStatus::Active)
116            .map(IndexedMemory::from_canonical),
117    );
118    let planner = Arc::new(DeterministicPlanner::with_entities(
119        KnownEntities::from_index(&index),
120    ));
121    let canonical = Arc::new(IndexHandle::with_index(index));
122    let overlay = Arc::new(IndexHandle::new());
123    (
124        LocalMemoryRetriever::new(canonical, overlay, RetrievalConfig::default()),
125        planner,
126    )
127}
128
129/// Run the retrieval evaluation.
130pub async fn run_retrieval_eval() -> Result<RetrievalReport, MemoryError> {
131    let (retriever, planner) = eval_retriever();
132    let mut results = Vec::new();
133
134    for case in retrieval_cases() {
135        results.push(run_retrieval_case(&retriever, &planner, &case).await?);
136    }
137
138    let recall_cases: Vec<&RetrievalCaseResult> = results
139        .iter()
140        .zip(retrieval_cases())
141        .filter(|(_, case)| case.expects_memory)
142        .map(|(result, _)| result)
143        .collect();
144
145    let precision = metrics::mean(&recall_cases.iter().map(|c| c.precision).collect::<Vec<_>>());
146    let recall = metrics::mean(&recall_cases.iter().map(|c| c.recall).collect::<Vec<_>>());
147    let mrr = metrics::mean(
148        &recall_cases
149            .iter()
150            .map(|c| c.reciprocal_rank)
151            .collect::<Vec<_>>(),
152    );
153    let skip_accuracy = metrics::mean(
154        &results
155            .iter()
156            .map(|c| if c.skip_correct { 1.0 } else { 0.0 })
157            .collect::<Vec<_>>(),
158    );
159    let mut tokens: Vec<f32> = results.iter().map(|c| c.tokens as f32).collect();
160    let mean_tokens = metrics::mean(&tokens);
161    let p95_tokens = metrics::percentile(&mut tokens, 95.0);
162
163    Ok(RetrievalReport {
164        cases: results,
165        precision,
166        recall,
167        mrr,
168        skip_accuracy,
169        mean_tokens,
170        p95_tokens,
171    })
172}
173
174async fn run_retrieval_case(
175    retriever: &LocalMemoryRetriever,
176    planner: &DeterministicPlanner,
177    case: &RetrievalCase,
178) -> Result<RetrievalCaseResult, MemoryError> {
179    let now = Utc::now();
180    let plan = planner.plan(case.query, TurnId(1), 1, now);
181    let snapshot = retriever.prepare(RetrievalRequest { plan, now }).await?;
182
183    let returned: Vec<String> = snapshot
184        .facts
185        .iter()
186        .map(|f| f.memory_id.to_string())
187        .collect();
188    let relevant: Vec<String> = case.relevant.iter().map(|r| (*r).to_string()).collect();
189    let forbidden: Vec<String> = case.forbidden.iter().map(|r| (*r).to_string()).collect();
190
191    Ok(RetrievalCaseResult {
192        name: case.name,
193        precision: metrics::precision(&returned, &relevant),
194        recall: metrics::recall(&returned, &relevant),
195        reciprocal_rank: metrics::reciprocal_rank(&returned, &relevant),
196        // The observable contract, not the internal flag. A question that
197        // needs no memory must surface no facts; whether the planner declined
198        // to search or searched and scored nothing is invisible from outside,
199        // and only one of those two is a decision the planner can get right.
200        skip_correct: returned.is_empty() != case.expects_memory,
201        leaked_forbidden: returned.iter().any(|r| forbidden.contains(r)),
202        tokens: usize::from(snapshot.token_count),
203        returned,
204    })
205}
206
207/// Per-case ingestion outcome.
208#[derive(Debug, Clone)]
209pub struct IngestionCaseResult {
210    /// Case name.
211    pub name: &'static str,
212    /// Whether anything was admitted.
213    pub stored: bool,
214    /// Whether that matched the expectation.
215    pub correct: bool,
216    /// Why it differed, when it did.
217    pub detail: Option<String>,
218}
219
220/// Aggregate ingestion report.
221#[derive(Debug, Clone)]
222pub struct IngestionReport {
223    /// Per-case detail.
224    pub cases: Vec<IngestionCaseResult>,
225    /// Fraction of cases that behaved as specified.
226    pub accuracy: f32,
227    /// Utterances stored that should not have been.
228    pub false_stores: usize,
229    /// Utterances not stored that should have been.
230    pub missed_stores: usize,
231}
232
233impl IngestionReport {
234    /// Cases that did not behave as specified.
235    pub fn failures(&self) -> Vec<&IngestionCaseResult> {
236        self.cases.iter().filter(|c| !c.correct).collect()
237    }
238}
239
240/// Run the ingestion evaluation.
241pub async fn run_ingestion_eval() -> Result<IngestionReport, MemoryError> {
242    let extractor = RuleBasedObservationExtractor::new();
243    let config = IngestionConfig::default();
244    let mut results = Vec::new();
245    let mut false_stores = 0;
246    let mut missed_stores = 0;
247
248    for case in ingestion_cases() {
249        let observations = extractor
250            .extract(
251                ObservationExtractionContext::user_turn(
252                    case.utterance,
253                    SessionId::new("ses_eval"),
254                    TurnId(1),
255                    Utc::now(),
256                )
257                .attributed_to(case.speaker),
258            )
259            .await?;
260
261        // Admission is part of ingestion, so a candidate the policy refuses is
262        // not "stored" however confidently the extractor produced it.
263        let admitted: Vec<_> = observations
264            .into_iter()
265            .filter(|o| matches!(admit_observation(o, &config), AdmissionVerdict::Accept(_)))
266            .collect();
267
268        let stored = !admitted.is_empty();
269        let mut detail = None;
270        let mut correct = stored == case.stores;
271
272        if stored != case.stores {
273            if stored {
274                false_stores += 1;
275            } else {
276                missed_stores += 1;
277            }
278            detail = Some(format!("expected stored={}, got {stored}", case.stores));
279        } else if let Some(observation) = admitted.first() {
280            correct = check_expectations(&case, observation, &mut detail);
281        }
282
283        results.push(IngestionCaseResult {
284            name: case.name,
285            stored,
286            correct,
287            detail,
288        });
289    }
290
291    let accuracy = metrics::mean(
292        &results
293            .iter()
294            .map(|c| if c.correct { 1.0 } else { 0.0 })
295            .collect::<Vec<_>>(),
296    );
297
298    Ok(IngestionReport {
299        cases: results,
300        accuracy,
301        false_stores,
302        missed_stores,
303    })
304}
305
306fn check_expectations(
307    case: &IngestionCase,
308    observation: &crate::core::MemoryObservation,
309    detail: &mut Option<String>,
310) -> bool {
311    if let Some(expected) = case.kind
312        && observation.kind != expected
313    {
314        *detail = Some(format!(
315            "expected kind {expected:?}, got {:?}",
316            observation.kind
317        ));
318        return false;
319    }
320    if let Some(expected) = case.explicitness
321        && observation.explicitness != expected
322    {
323        *detail = Some(format!(
324            "expected explicitness {expected:?}, got {:?}",
325            observation.explicitness
326        ));
327        return false;
328    }
329    true
330}
331
332/// The evaluation user, re-exported for callers driving the harness directly.
333pub fn user() -> crate::core::UserId {
334    eval_user()
335}
336
337#[cfg(test)]
338mod tests {
339    use super::*;
340
341    /// Acceptance thresholds (ยง42). Lowering one is a product decision.
342    const MIN_PRECISION: f32 = 0.85;
343    const MIN_RECALL: f32 = 0.80;
344    const MAX_MEAN_TOKENS: f32 = 250.0;
345    const MAX_TOKENS: usize = 500;
346
347    #[tokio::test]
348    async fn retrieval_meets_its_acceptance_thresholds() {
349        let report = run_retrieval_eval().await.unwrap();
350        assert!(
351            report.precision >= MIN_PRECISION,
352            "precision {:.3} below {MIN_PRECISION}\n{}",
353            report.precision,
354            report.render()
355        );
356        assert!(
357            report.recall >= MIN_RECALL,
358            "recall {:.3} below {MIN_RECALL}\n{}",
359            report.recall,
360            report.render()
361        );
362    }
363
364    #[tokio::test]
365    async fn memory_is_skipped_for_questions_that_do_not_need_it() {
366        let report = run_retrieval_eval().await.unwrap();
367        assert_eq!(
368            report.skip_accuracy,
369            1.0,
370            "skip decisions were wrong somewhere\n{}",
371            report.render()
372        );
373    }
374
375    #[tokio::test]
376    async fn superseded_and_forbidden_records_never_surface() {
377        let report = run_retrieval_eval().await.unwrap();
378        assert!(
379            report.leaks().is_empty(),
380            "forbidden records leaked in: {:?}\n{}",
381            report.leaks(),
382            report.render()
383        );
384    }
385
386    #[tokio::test]
387    async fn context_stays_inside_its_budget() {
388        let report = run_retrieval_eval().await.unwrap();
389        assert!(
390            report.mean_tokens <= MAX_MEAN_TOKENS,
391            "mean context {:.1} tokens exceeds {MAX_MEAN_TOKENS}\n{}",
392            report.mean_tokens,
393            report.render()
394        );
395        for case in &report.cases {
396            assert!(
397                case.tokens <= MAX_TOKENS,
398                "case `{}` returned {} tokens",
399                case.name,
400                case.tokens
401            );
402        }
403    }
404
405    #[tokio::test]
406    async fn ingestion_behaves_as_specified() {
407        let report = run_ingestion_eval().await.unwrap();
408        assert_eq!(
409            report.accuracy,
410            1.0,
411            "ingestion failures: {:?}",
412            report.failures()
413        );
414    }
415
416    #[tokio::test]
417    async fn nothing_is_stored_that_should_not_be() {
418        let report = run_ingestion_eval().await.unwrap();
419        assert_eq!(report.false_stores, 0, "{:?}", report.failures());
420    }
421
422    #[tokio::test]
423    async fn the_report_renders_every_case() {
424        let report = run_retrieval_eval().await.unwrap();
425        let rendered = report.render();
426        for case in &report.cases {
427            assert!(rendered.contains(case.name));
428        }
429        assert!(rendered.contains("precision="));
430    }
431}