1use 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#[derive(Debug, Clone)]
30pub struct RetrievalCaseResult {
31 pub name: &'static str,
33 pub returned: Vec<String>,
35 pub precision: f32,
37 pub recall: f32,
39 pub reciprocal_rank: f32,
41 pub skip_correct: bool,
44 pub leaked_forbidden: bool,
46 pub tokens: usize,
48}
49
50#[derive(Debug, Clone)]
52pub struct RetrievalReport {
53 pub cases: Vec<RetrievalCaseResult>,
55 pub precision: f32,
57 pub recall: f32,
59 pub mrr: f32,
61 pub skip_accuracy: f32,
63 pub mean_tokens: f32,
65 pub p95_tokens: f32,
67}
68
69impl RetrievalReport {
70 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 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
109fn 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
129pub 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 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#[derive(Debug, Clone)]
209pub struct IngestionCaseResult {
210 pub name: &'static str,
212 pub stored: bool,
214 pub correct: bool,
216 pub detail: Option<String>,
218}
219
220#[derive(Debug, Clone)]
222pub struct IngestionReport {
223 pub cases: Vec<IngestionCaseResult>,
225 pub accuracy: f32,
227 pub false_stores: usize,
229 pub missed_stores: usize,
231}
232
233impl IngestionReport {
234 pub fn failures(&self) -> Vec<&IngestionCaseResult> {
236 self.cases.iter().filter(|c| !c.correct).collect()
237 }
238}
239
240pub 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 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
332pub fn user() -> crate::core::UserId {
334 eval_user()
335}
336
337#[cfg(test)]
338mod tests {
339 use super::*;
340
341 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}