1use std::collections::{BTreeMap, BTreeSet};
2use std::fs;
3use std::path::{Path, PathBuf};
4use std::time::Instant;
5
6use crate::aaak::{AaakCodec, AaakMeta};
7use crate::core::{
8 config::Config,
9 db::Database,
10 types::{BootstrapEvidenceArgs, Drawer, SourceType, TaxonomyEntry},
11 utils::{build_drawer_id, route_room_from_taxonomy},
12};
13use crate::embed::{ConfiguredEmbedderFactory, Embedder};
14use crate::ingest::normalize::CURRENT_NORMALIZE_VERSION;
15use crate::search::search;
16use anyhow::{Context, Result, bail};
17use clap::ValueEnum;
18use serde::{Deserialize, Deserializer, Serialize};
19use tempfile::tempdir;
20
21const BENCH_WING: &str = "longmemeval";
22const DEFAULT_TOP_K: usize = 50;
23const METRIC_KS: [usize; 6] = [1, 3, 5, 10, 30, 50];
24const TECHNICAL_KEYWORDS: &[&str] = &[
25 "code", "python", "function", "bug", "error", "api", "database", "server", "deploy", "git",
26 "test", "debug", "refactor",
27];
28const PLANNING_KEYWORDS: &[&str] = &[
29 "plan",
30 "roadmap",
31 "milestone",
32 "deadline",
33 "priority",
34 "sprint",
35 "backlog",
36 "scope",
37 "requirement",
38 "spec",
39];
40const DECISION_KEYWORDS: &[&str] = &[
41 "decided",
42 "chose",
43 "picked",
44 "switched",
45 "migrated",
46 "replaced",
47 "trade-off",
48 "alternative",
49 "option",
50 "approach",
51];
52const PERSONAL_KEYWORDS: &[&str] = &[
53 "family", "friend", "birthday", "vacation", "hobby", "health", "feeling", "love", "home",
54 "weekend",
55];
56const KNOWLEDGE_KEYWORDS: &[&str] = &[
57 "learn",
58 "study",
59 "degree",
60 "school",
61 "university",
62 "course",
63 "research",
64 "paper",
65 "book",
66 "reading",
67];
68const ROOM_KEYWORDS: &[(&str, &[&str])] = &[
69 ("technical", TECHNICAL_KEYWORDS),
70 ("planning", PLANNING_KEYWORDS),
71 ("decisions", DECISION_KEYWORDS),
72 ("personal", PERSONAL_KEYWORDS),
73 ("knowledge", KNOWLEDGE_KEYWORDS),
74];
75
76#[derive(Debug, Clone, Copy, PartialEq, Eq, ValueEnum)]
77pub enum BenchMode {
78 Raw,
79 Aaak,
80 Rooms,
81}
82
83#[derive(Debug, Clone, Copy, PartialEq, Eq, ValueEnum)]
84pub enum LongMemEvalGranularity {
85 Session,
86 Turn,
87}
88
89#[derive(Debug, Clone)]
90pub struct LongMemEvalArgs {
91 pub data_file: PathBuf,
92 pub mode: BenchMode,
93 pub granularity: LongMemEvalGranularity,
94 pub limit: usize,
95 pub skip: usize,
96 pub top_k: usize,
97 pub out: Option<PathBuf>,
98}
99
100#[derive(Debug, Clone, Deserialize)]
101struct LongMemEvalTurn {
102 role: String,
103 content: String,
104}
105
106#[derive(Debug, Clone, Deserialize)]
107struct LongMemEvalEntry {
108 question_id: String,
109 question_type: String,
110 question: String,
111 #[serde(deserialize_with = "deserialize_answer")]
112 answer: String,
113 haystack_sessions: Vec<Vec<LongMemEvalTurn>>,
114 haystack_session_ids: Vec<String>,
115 haystack_dates: Vec<String>,
116 answer_session_ids: Vec<String>,
117}
118
119#[derive(Debug, Clone, Deserialize)]
120#[serde(untagged)]
121enum LongMemEvalAnswerValue {
122 Text(String),
123 Signed(i64),
124 Unsigned(u64),
125 Float(f64),
126 Bool(bool),
127}
128
129fn deserialize_answer<'de, D>(deserializer: D) -> std::result::Result<String, D::Error>
130where
131 D: Deserializer<'de>,
132{
133 let value = LongMemEvalAnswerValue::deserialize(deserializer)?;
134 Ok(match value {
135 LongMemEvalAnswerValue::Text(text) => text,
136 LongMemEvalAnswerValue::Signed(number) => number.to_string(),
137 LongMemEvalAnswerValue::Unsigned(number) => number.to_string(),
138 LongMemEvalAnswerValue::Float(number) => number.to_string(),
139 LongMemEvalAnswerValue::Bool(value) => value.to_string(),
140 })
141}
142
143#[derive(Debug, Clone)]
144struct CorpusItem {
145 corpus_id: String,
146 original_text: String,
147 retrieval_text: String,
148 timestamp: String,
149 drawer_id: String,
150}
151
152#[derive(Debug, Clone, Default, PartialEq)]
153struct AggregateMetrics {
154 recall_any: BTreeMap<usize, Vec<f64>>,
155 recall_all: BTreeMap<usize, Vec<f64>>,
156 ndcg_any: BTreeMap<usize, Vec<f64>>,
157}
158
159#[derive(Debug, Clone, Default, PartialEq)]
160struct MetricSnapshot {
161 recall_any: BTreeMap<usize, f64>,
162 ndcg_any: BTreeMap<usize, f64>,
163}
164
165#[derive(Debug, Clone, Default, PartialEq)]
166struct EntryMetricSnapshot {
167 session: MetricSnapshot,
168 turn: MetricSnapshot,
169}
170
171#[derive(Debug, Clone, PartialEq)]
172struct BenchmarkSummary {
173 mode: BenchMode,
174 granularity: LongMemEvalGranularity,
175 question_count: usize,
176 elapsed_secs: f64,
177 session: MetricSnapshot,
178 turn: MetricSnapshot,
179 per_type: BTreeMap<String, MetricSnapshot>,
180}
181
182#[derive(Debug, Clone, Serialize, PartialEq)]
183struct RankedItemLog {
184 corpus_id: String,
185 text: String,
186 timestamp: String,
187}
188
189#[derive(Debug, Clone, Serialize, PartialEq)]
190struct RetrievalLog {
191 query: String,
192 ranked_items: Vec<RankedItemLog>,
193 metrics: RetrievalMetricLog,
194}
195
196#[derive(Debug, Clone, Serialize, PartialEq)]
197struct RetrievalMetricLog {
198 session: BTreeMap<String, f64>,
199 turn: BTreeMap<String, f64>,
200}
201
202#[derive(Debug, Clone, Serialize, PartialEq)]
203struct BenchmarkLogEntry {
204 question_id: String,
205 question_type: String,
206 question: String,
207 answer: String,
208 retrieval_results: RetrievalLog,
209}
210
211pub async fn run_longmemeval_command(config: &Config, args: LongMemEvalArgs) -> Result<()> {
212 if args.top_k == 0 {
213 bail!("--top-k must be greater than 0");
214 }
215
216 use crate::embed::EmbedderFactory;
217
218 let entries = load_entries(&args.data_file, args.limit, args.skip)?;
219 let embedder = ConfiguredEmbedderFactory::new(config.clone())
220 .build()
221 .await
222 .context("failed to initialize embedder for LongMemEval benchmark")?;
223 let (summary, logs) = run_benchmark_with_embedder(&*embedder, &entries, &args).await?;
224
225 print_summary(&summary);
226
227 if let Some(out_path) = args.out.as_deref() {
228 write_results_log(out_path, &logs)?;
229 println!("results_log: {}", out_path.display());
230 }
231
232 Ok(())
233}
234
235fn load_entries(path: &Path, limit: usize, skip: usize) -> Result<Vec<LongMemEvalEntry>> {
236 let contents = fs::read_to_string(path)
237 .with_context(|| format!("failed to read LongMemEval data from {}", path.display()))?;
238 let mut entries = serde_json::from_str::<Vec<LongMemEvalEntry>>(&contents)
239 .with_context(|| format!("failed to parse LongMemEval JSON {}", path.display()))?;
240
241 if skip > 0 {
242 entries = entries.into_iter().skip(skip).collect();
243 }
244 if limit > 0 {
245 entries.truncate(limit);
246 }
247
248 Ok(entries)
249}
250
251async fn run_benchmark_with_embedder<E: Embedder + ?Sized>(
252 embedder: &E,
253 entries: &[LongMemEvalEntry],
254 args: &LongMemEvalArgs,
255) -> Result<(BenchmarkSummary, Vec<BenchmarkLogEntry>)> {
256 let ks = selected_ks(args.top_k);
257 let started = Instant::now();
258 let scratch = tempdir().context("failed to create benchmark scratch directory")?;
259
260 let mut session_metrics = AggregateMetrics::default();
261 let mut turn_metrics = AggregateMetrics::default();
262 let mut per_type = BTreeMap::<String, AggregateMetrics>::new();
263 let mut logs = Vec::with_capacity(entries.len());
264
265 for (index, entry) in entries.iter().enumerate() {
266 let db_path = scratch.path().join(format!("q-{index}.db"));
267 let db = Database::open(&db_path)
268 .with_context(|| format!("failed to open benchmark database {}", db_path.display()))?;
269
270 if args.mode == BenchMode::Rooms {
271 install_rooms_taxonomy(&db)?;
272 }
273
274 let corpus_items = build_corpus(entry, args.granularity, args.mode);
275 if corpus_items.is_empty() {
276 continue;
277 }
278
279 ingest_corpus(&db, embedder, &corpus_items, args.mode).await?;
280
281 let results = search(&db, embedder, &entry.question, None, None, args.top_k)
282 .await
283 .with_context(|| format!("search failed for question {}", entry.question_id))?;
284 let rankings = map_results_to_rankings(&results, &corpus_items);
285 let entry_metrics = score_entry(entry, &rankings, &corpus_items, &ks);
286
287 for &k in &ks {
288 push_metric_set(&mut session_metrics, k, &entry_metrics.session);
289 push_metric_set(&mut turn_metrics, k, &entry_metrics.turn);
290 let bucket = per_type.entry(entry.question_type.clone()).or_default();
291 push_metric_set(bucket, k, &entry_metrics.session);
292 }
293
294 logs.push(build_log_entry(
295 entry,
296 &corpus_items,
297 &rankings,
298 &entry_metrics,
299 args.top_k,
300 ));
301 }
302
303 let elapsed_secs = started.elapsed().as_secs_f64();
304 let summary = BenchmarkSummary {
305 mode: args.mode,
306 granularity: args.granularity,
307 question_count: logs.len(),
308 elapsed_secs,
309 session: summarize_metrics(&session_metrics),
310 turn: summarize_metrics(&turn_metrics),
311 per_type: per_type
312 .into_iter()
313 .map(|(name, metrics)| (name, summarize_metrics(&metrics)))
314 .collect(),
315 };
316
317 Ok((summary, logs))
318}
319
320fn build_corpus(
321 entry: &LongMemEvalEntry,
322 granularity: LongMemEvalGranularity,
323 mode: BenchMode,
324) -> Vec<CorpusItem> {
325 let codec = AaakCodec::default();
326 let mut items = Vec::new();
327
328 for ((session, session_id), date) in entry
329 .haystack_sessions
330 .iter()
331 .zip(entry.haystack_session_ids.iter())
332 .zip(entry.haystack_dates.iter())
333 {
334 match granularity {
335 LongMemEvalGranularity::Session => {
336 let Some(text) = join_user_turns(session) else {
337 continue;
338 };
339 let item_ordinal = items.len();
340 items.push(build_corpus_item(
341 &codec,
342 item_ordinal,
343 session_id.clone(),
344 text,
345 date.clone(),
346 mode,
347 ));
348 }
349 LongMemEvalGranularity::Turn => {
350 let mut turn_index = 0usize;
351 for turn in session {
352 if turn.role != "user" {
353 continue;
354 }
355 let item_ordinal = items.len();
356 items.push(build_corpus_item(
357 &codec,
358 item_ordinal,
359 format!("{session_id}_turn_{turn_index}"),
360 turn.content.clone(),
361 date.clone(),
362 mode,
363 ));
364 turn_index += 1;
365 }
366 }
367 }
368 }
369
370 items
371}
372
373fn build_corpus_item(
374 codec: &AaakCodec,
375 item_ordinal: usize,
376 corpus_id: String,
377 original_text: String,
378 timestamp: String,
379 mode: BenchMode,
380) -> CorpusItem {
381 let retrieval_text = match mode {
382 BenchMode::Raw | BenchMode::Rooms => original_text.clone(),
383 BenchMode::Aaak => codec
384 .encode(
385 &original_text,
386 &AaakMeta {
387 wing: BENCH_WING.to_string(),
388 room: "benchmark".to_string(),
389 date: timestamp.clone(),
390 source: corpus_id.clone(),
391 },
392 )
393 .document
394 .to_string(),
395 };
396 let drawer_id_seed = format!("{item_ordinal}\n{corpus_id}\n{retrieval_text}");
397 let drawer_id = build_drawer_id(BENCH_WING, None, &drawer_id_seed);
398
399 CorpusItem {
400 corpus_id,
401 original_text,
402 retrieval_text,
403 timestamp,
404 drawer_id,
405 }
406}
407
408fn join_user_turns(session: &[LongMemEvalTurn]) -> Option<String> {
409 let turns = session
410 .iter()
411 .filter(|turn| turn.role == "user")
412 .map(|turn| turn.content.as_str())
413 .collect::<Vec<_>>();
414 (!turns.is_empty()).then(|| turns.join("\n"))
415}
416
417async fn ingest_corpus<E: Embedder + ?Sized>(
418 db: &Database,
419 embedder: &E,
420 items: &[CorpusItem],
421 mode: BenchMode,
422) -> Result<()> {
423 let texts = items
424 .iter()
425 .map(|item| item.retrieval_text.as_str())
426 .collect::<Vec<_>>();
427 let vectors = embedder
428 .embed(&texts)
429 .await
430 .context("failed to embed benchmark corpus")?;
431 let taxonomy = if mode == BenchMode::Rooms {
432 db.taxonomy_entries()
433 .context("failed to load rooms taxonomy for benchmark ingest")?
434 } else {
435 Vec::new()
436 };
437
438 for (item, vector) in items.iter().zip(vectors.iter()) {
439 let room = match mode {
440 BenchMode::Rooms => Some(route_room_from_taxonomy(
441 &item.original_text,
442 BENCH_WING,
443 &taxonomy,
444 )),
445 BenchMode::Raw | BenchMode::Aaak => None,
446 };
447
448 let drawer = Drawer::new_bootstrap_evidence(BootstrapEvidenceArgs {
449 id: item.drawer_id.clone(),
450 content: item.retrieval_text.clone(),
451 wing: BENCH_WING.to_string(),
452 room: room.clone(),
453 source_file: Some(format!("longmemeval://{}", item.corpus_id)),
454 source_type: SourceType::Conversation,
455 added_at: item.timestamp.clone(),
456 chunk_index: Some(0),
457 importance: 0,
458 });
459 let drawer = Drawer {
460 normalize_version: CURRENT_NORMALIZE_VERSION,
461 ..drawer
462 };
463 db.insert_drawer(&drawer)
464 .with_context(|| format!("failed to insert drawer {}", item.drawer_id))?;
465 db.insert_vector(&item.drawer_id, vector)
466 .with_context(|| format!("failed to insert vector for {}", item.drawer_id))?;
467 }
468
469 Ok(())
470}
471
472fn install_rooms_taxonomy(db: &Database) -> Result<()> {
473 for (room, keywords) in ROOM_KEYWORDS {
474 db.upsert_taxonomy_entry(&TaxonomyEntry {
475 wing: BENCH_WING.to_string(),
476 room: (*room).to_string(),
477 display_name: Some((*room).to_string()),
478 keywords: keywords
479 .iter()
480 .map(|keyword| (*keyword).to_string())
481 .collect(),
482 })
483 .with_context(|| format!("failed to install benchmark taxonomy room {room}"))?;
484 }
485
486 Ok(())
487}
488
489fn map_results_to_rankings(
490 results: &[crate::core::types::SearchResult],
491 items: &[CorpusItem],
492) -> Vec<usize> {
493 let drawer_to_index = items
494 .iter()
495 .enumerate()
496 .map(|(index, item)| (item.drawer_id.as_str(), index))
497 .collect::<BTreeMap<_, _>>();
498 let mut rankings = Vec::with_capacity(items.len());
499 let mut seen = BTreeSet::new();
500
501 for result in results {
502 if let Some(&index) = drawer_to_index.get(result.drawer_id.as_str())
503 && seen.insert(index)
504 {
505 rankings.push(index);
506 }
507 }
508
509 for index in 0..items.len() {
510 if seen.insert(index) {
511 rankings.push(index);
512 }
513 }
514
515 rankings
516}
517
518fn score_entry(
519 entry: &LongMemEvalEntry,
520 rankings: &[usize],
521 items: &[CorpusItem],
522 ks: &[usize],
523) -> EntryMetricSnapshot {
524 let corpus_ids = items
525 .iter()
526 .map(|item| item.corpus_id.clone())
527 .collect::<Vec<_>>();
528 let session_level_ids = corpus_ids
529 .iter()
530 .map(|corpus_id| session_id_from_corpus_id(corpus_id))
531 .collect::<Vec<_>>();
532 let answer_session_ids = entry
533 .answer_session_ids
534 .iter()
535 .map(|id| id.as_str())
536 .collect::<BTreeSet<_>>();
537 let answer_turn_ids = corpus_ids
538 .iter()
539 .filter(|corpus_id| {
540 answer_session_ids.contains(session_id_from_corpus_id(corpus_id).as_str())
541 })
542 .map(String::as_str)
543 .collect::<BTreeSet<_>>();
544
545 let mut snapshot = EntryMetricSnapshot::default();
546
547 for &k in ks {
548 let (session_any, _session_all, session_ndcg) =
549 evaluate_retrieval(rankings, &answer_session_ids, &session_level_ids, k);
550 snapshot.session.recall_any.insert(k, session_any);
551 snapshot.session.ndcg_any.insert(k, session_ndcg);
552
553 let (turn_any, _turn_all, turn_ndcg) =
554 evaluate_retrieval(rankings, &answer_turn_ids, &corpus_ids, k);
555 snapshot.turn.recall_any.insert(k, turn_any);
556 snapshot.turn.ndcg_any.insert(k, turn_ndcg);
557 }
558
559 snapshot
560}
561
562fn push_metric_set(target: &mut AggregateMetrics, k: usize, snapshot: &MetricSnapshot) {
563 if let Some(value) = snapshot.recall_any.get(&k) {
564 target.recall_any.entry(k).or_default().push(*value);
565 }
566 if let Some(value) = snapshot.ndcg_any.get(&k) {
567 target.ndcg_any.entry(k).or_default().push(*value);
568 }
569 if let Some(value) = snapshot.recall_any.get(&k) {
570 target.recall_all.entry(k).or_default().push(*value);
571 }
572}
573
574fn summarize_metrics(metrics: &AggregateMetrics) -> MetricSnapshot {
575 MetricSnapshot {
576 recall_any: metrics
577 .recall_any
578 .iter()
579 .map(|(&k, values)| (k, average(values)))
580 .collect(),
581 ndcg_any: metrics
582 .ndcg_any
583 .iter()
584 .map(|(&k, values)| (k, average(values)))
585 .collect(),
586 }
587}
588
589fn average(values: &[f64]) -> f64 {
590 if values.is_empty() {
591 return 0.0;
592 }
593
594 values.iter().sum::<f64>() / values.len() as f64
595}
596
597fn selected_ks(top_k: usize) -> Vec<usize> {
598 METRIC_KS
599 .into_iter()
600 .filter(|k| *k <= top_k)
601 .collect::<Vec<_>>()
602}
603
604fn build_log_entry(
605 entry: &LongMemEvalEntry,
606 items: &[CorpusItem],
607 rankings: &[usize],
608 metrics: &EntryMetricSnapshot,
609 top_k: usize,
610) -> BenchmarkLogEntry {
611 let ranked_items = rankings
612 .iter()
613 .take(top_k.min(items.len()))
614 .map(|&index| RankedItemLog {
615 corpus_id: items[index].corpus_id.clone(),
616 text: truncate_log_text(&items[index].original_text),
617 timestamp: items[index].timestamp.clone(),
618 })
619 .collect();
620
621 BenchmarkLogEntry {
622 question_id: entry.question_id.clone(),
623 question_type: entry.question_type.clone(),
624 question: entry.question.clone(),
625 answer: entry.answer.clone(),
626 retrieval_results: RetrievalLog {
627 query: entry.question.clone(),
628 ranked_items,
629 metrics: RetrievalMetricLog {
630 session: metric_map_for_log(&metrics.session),
631 turn: metric_map_for_log(&metrics.turn),
632 },
633 },
634 }
635}
636
637fn metric_map_for_log(snapshot: &MetricSnapshot) -> BTreeMap<String, f64> {
638 let mut map = BTreeMap::new();
639 for (&k, &value) in &snapshot.recall_any {
640 map.insert(format!("recall_any@{k}"), value);
641 }
642 for (&k, &value) in &snapshot.ndcg_any {
643 map.insert(format!("ndcg_any@{k}"), value);
644 }
645 map
646}
647
648fn truncate_log_text(text: &str) -> String {
649 let compact = text.split_whitespace().collect::<Vec<_>>().join(" ");
650 if compact.chars().count() <= 500 {
651 return compact;
652 }
653
654 compact.chars().take(500).collect::<String>()
655}
656
657fn print_summary(summary: &BenchmarkSummary) {
658 println!();
659 println!("============================================================");
660 println!(" mempal × LongMemEval");
661 println!("============================================================");
662 println!(" Questions: {}", summary.question_count);
663 println!(" Mode: {}", summary.mode.as_str());
664 println!(" Granularity: {}", summary.granularity.as_str());
665 println!(" Time: {:.1}s", summary.elapsed_secs);
666 println!();
667 println!(" SESSION-LEVEL METRICS:");
668 for (&k, &recall) in &summary.session.recall_any {
669 let ndcg = summary
670 .session
671 .ndcg_any
672 .get(&k)
673 .copied()
674 .unwrap_or_default();
675 println!(" Recall@{k:2}: {recall:.3} NDCG@{k:2}: {ndcg:.3}");
676 }
677 println!();
678 println!(" TURN-LEVEL METRICS:");
679 for (&k, &recall) in &summary.turn.recall_any {
680 let ndcg = summary.turn.ndcg_any.get(&k).copied().unwrap_or_default();
681 println!(" Recall@{k:2}: {recall:.3} NDCG@{k:2}: {ndcg:.3}");
682 }
683 if !summary.per_type.is_empty() {
684 println!();
685 println!(" QUESTION TYPES (session Recall@5 / Recall@10 / NDCG@10):");
686 for (question_type, metrics) in &summary.per_type {
687 let r5 = metrics.recall_any.get(&5).copied().unwrap_or_default();
688 let r10 = metrics.recall_any.get(&10).copied().unwrap_or_default();
689 let nd10 = metrics.ndcg_any.get(&10).copied().unwrap_or_default();
690 println!(" {question_type}: R@5={r5:.3} R@10={r10:.3} NDCG@10={nd10:.3}");
691 }
692 }
693}
694
695fn write_results_log(path: &Path, entries: &[BenchmarkLogEntry]) -> Result<()> {
696 if let Some(parent) = path.parent()
697 && !parent.as_os_str().is_empty()
698 {
699 fs::create_dir_all(parent)
700 .with_context(|| format!("failed to create results directory {}", parent.display()))?;
701 }
702
703 let body = entries
704 .iter()
705 .map(|entry| {
706 serde_json::to_string(entry).context("failed to serialize LongMemEval log entry")
707 })
708 .collect::<Result<Vec<_>>>()?
709 .join("\n");
710 fs::write(path, body)
711 .with_context(|| format!("failed to write LongMemEval results to {}", path.display()))?;
712 Ok(())
713}
714
715fn dcg(relevances: &[f64], k: usize) -> f64 {
716 relevances
717 .iter()
718 .take(k)
719 .enumerate()
720 .map(|(index, relevance)| relevance / ((index + 2) as f64).log2())
721 .sum()
722}
723
724fn ndcg(rankings: &[usize], correct_ids: &BTreeSet<&str>, corpus_ids: &[String], k: usize) -> f64 {
725 let relevances = rankings
726 .iter()
727 .take(k)
728 .map(|&index| {
729 if correct_ids.contains(corpus_ids[index].as_str()) {
730 1.0
731 } else {
732 0.0
733 }
734 })
735 .collect::<Vec<_>>();
736 let mut ideal = relevances.clone();
737 ideal.sort_by(|left, right| right.partial_cmp(left).unwrap_or(std::cmp::Ordering::Equal));
738 let ideal_dcg = dcg(&ideal, k);
739 if ideal_dcg == 0.0 {
740 return 0.0;
741 }
742
743 dcg(&relevances, k) / ideal_dcg
744}
745
746fn evaluate_retrieval(
747 rankings: &[usize],
748 correct_ids: &BTreeSet<&str>,
749 corpus_ids: &[String],
750 k: usize,
751) -> (f64, f64, f64) {
752 let top_k_ids = rankings
753 .iter()
754 .take(k)
755 .map(|&index| corpus_ids[index].as_str())
756 .collect::<BTreeSet<_>>();
757 let recall_any = if correct_ids.iter().any(|id| top_k_ids.contains(id)) {
758 1.0
759 } else {
760 0.0
761 };
762 let recall_all = if correct_ids.iter().all(|id| top_k_ids.contains(id)) {
763 1.0
764 } else {
765 0.0
766 };
767 let ndcg = ndcg(rankings, correct_ids, corpus_ids, k);
768
769 (recall_any, recall_all, ndcg)
770}
771
772fn session_id_from_corpus_id(corpus_id: &str) -> String {
773 corpus_id
774 .split_once("_turn_")
775 .map(|(session_id, _)| session_id.to_string())
776 .unwrap_or_else(|| corpus_id.to_string())
777}
778
779impl BenchMode {
780 fn as_str(self) -> &'static str {
781 match self {
782 Self::Raw => "raw",
783 Self::Aaak => "aaak",
784 Self::Rooms => "rooms",
785 }
786 }
787}
788
789impl LongMemEvalGranularity {
790 fn as_str(self) -> &'static str {
791 match self {
792 Self::Session => "session",
793 Self::Turn => "turn",
794 }
795 }
796}
797
798pub fn default_top_k() -> usize {
799 DEFAULT_TOP_K
800}
801
802#[cfg(test)]
803mod tests {
804 use super::*;
805
806 #[derive(Default)]
807 struct TestEmbedder;
808
809 #[async_trait::async_trait]
810 impl Embedder for TestEmbedder {
811 async fn embed(
812 &self,
813 texts: &[&str],
814 ) -> std::result::Result<Vec<Vec<f32>>, crate::embed::EmbedError> {
815 Ok(texts.iter().map(|text| fake_embedding(text)).collect())
816 }
817
818 fn dimensions(&self) -> usize {
819 384
820 }
821
822 fn name(&self) -> &str {
823 "test"
824 }
825 }
826
827 fn fake_embedding(text: &str) -> Vec<f32> {
828 let mut embedding = vec![0.0_f32; 384];
829 for token in text
830 .split(|ch: char| !ch.is_alphanumeric())
831 .filter(|token| !token.is_empty())
832 {
833 let mut hash = 0usize;
834 for byte in token.to_ascii_lowercase().bytes() {
835 hash = hash.wrapping_mul(33).wrapping_add(usize::from(byte));
836 }
837 embedding[hash % 384] += 1.0;
838 }
839 embedding
840 }
841
842 fn sample_entry() -> LongMemEvalEntry {
843 LongMemEvalEntry {
844 question_id: "q-1".to_string(),
845 question_type: "decision".to_string(),
846 question: "Which auth provider did the user decide to use?".to_string(),
847 answer: "Clerk".to_string(),
848 haystack_sessions: vec![
849 vec![
850 LongMemEvalTurn {
851 role: "user".to_string(),
852 content: "We decided to use Clerk for auth because pricing was better."
853 .to_string(),
854 },
855 LongMemEvalTurn {
856 role: "assistant".to_string(),
857 content: "I will note the auth choice.".to_string(),
858 },
859 ],
860 vec![LongMemEvalTurn {
861 role: "user".to_string(),
862 content: "Deployment notes for Render and Postgres.".to_string(),
863 }],
864 ],
865 haystack_session_ids: vec!["sess_auth".to_string(), "sess_deploy".to_string()],
866 haystack_dates: vec!["2026-04-08".to_string(), "2026-04-09".to_string()],
867 answer_session_ids: vec!["sess_auth".to_string()],
868 }
869 }
870
871 fn duplicate_session_entry() -> LongMemEvalEntry {
872 LongMemEvalEntry {
873 question_id: "q-dup".to_string(),
874 question_type: "multi-session".to_string(),
875 question: "How many times did I repeat the note?".to_string(),
876 answer: "2".to_string(),
877 haystack_sessions: vec![
878 vec![LongMemEvalTurn {
879 role: "user".to_string(),
880 content: "Pick up three shirts from the store.".to_string(),
881 }],
882 vec![LongMemEvalTurn {
883 role: "user".to_string(),
884 content: "Pick up three shirts from the store.".to_string(),
885 }],
886 ],
887 haystack_session_ids: vec!["sess_1".to_string(), "sess_2".to_string()],
888 haystack_dates: vec!["2026-04-08".to_string(), "2026-04-09".to_string()],
889 answer_session_ids: vec!["sess_1".to_string(), "sess_2".to_string()],
890 }
891 }
892
893 fn duplicate_session_id_entry() -> LongMemEvalEntry {
894 LongMemEvalEntry {
895 question_id: "q-dup-id".to_string(),
896 question_type: "multi-session".to_string(),
897 question: "Which repeated session should be recalled?".to_string(),
898 answer: "2".to_string(),
899 haystack_sessions: vec![
900 vec![LongMemEvalTurn {
901 role: "user".to_string(),
902 content: "Remember the pickup is at the downtown store.".to_string(),
903 }],
904 vec![LongMemEvalTurn {
905 role: "user".to_string(),
906 content: "Remember the pickup is at the downtown store.".to_string(),
907 }],
908 ],
909 haystack_session_ids: vec!["sess_dup".to_string(), "sess_dup".to_string()],
910 haystack_dates: vec!["2026-04-08".to_string(), "2026-04-09".to_string()],
911 answer_session_ids: vec!["sess_dup".to_string()],
912 }
913 }
914
915 #[test]
916 fn test_load_entries_accepts_numeric_answer() {
917 let temp = tempdir().expect("tempdir should be created");
918 let path = temp.path().join("longmemeval.json");
919 fs::write(
920 &path,
921 r#"[{
922 "question_id":"q-1",
923 "question_type":"multi-session",
924 "question":"How many items do I need to pick up?",
925 "question_date":"2023/02/15 (Wed) 23:50",
926 "answer":3,
927 "answer_session_ids":["answer_1"],
928 "haystack_dates":["2023/02/15 (Wed) 01:41"],
929 "haystack_session_ids":["sess_1"],
930 "haystack_sessions":[[
931 {"role":"user","content":"Remember to pick up three shirts."}
932 ]]
933 }]"#,
934 )
935 .expect("fixture should be written");
936
937 let entries = load_entries(&path, 0, 0).expect("numeric answer should parse");
938
939 assert_eq!(entries.len(), 1);
940 assert_eq!(entries[0].answer, "3");
941 }
942
943 #[test]
944 fn test_build_corpus_assigns_unique_drawer_ids_for_duplicate_content() {
945 let items = build_corpus(
946 &duplicate_session_entry(),
947 LongMemEvalGranularity::Session,
948 BenchMode::Raw,
949 );
950
951 assert_eq!(items.len(), 2);
952 assert_ne!(items[0].drawer_id, items[1].drawer_id);
953 }
954
955 #[test]
956 fn test_build_corpus_assigns_unique_drawer_ids_for_duplicate_session_ids() {
957 let items = build_corpus(
958 &duplicate_session_id_entry(),
959 LongMemEvalGranularity::Session,
960 BenchMode::Raw,
961 );
962
963 assert_eq!(items.len(), 2);
964 assert_ne!(items[0].drawer_id, items[1].drawer_id);
965 }
966
967 #[test]
968 fn test_build_corpus_session_granularity() {
969 let items = build_corpus(
970 &sample_entry(),
971 LongMemEvalGranularity::Session,
972 BenchMode::Raw,
973 );
974
975 assert_eq!(items.len(), 2);
976 assert_eq!(items[0].corpus_id, "sess_auth");
977 assert!(items[0].original_text.contains("We decided to use Clerk"));
978 }
979
980 #[test]
981 fn test_build_corpus_turn_granularity() {
982 let items = build_corpus(
983 &sample_entry(),
984 LongMemEvalGranularity::Turn,
985 BenchMode::Raw,
986 );
987
988 assert_eq!(items.len(), 2);
989 assert_eq!(items[0].corpus_id, "sess_auth_turn_0");
990 assert_eq!(items[1].corpus_id, "sess_deploy_turn_0");
991 }
992
993 #[test]
994 fn test_build_corpus_aaak_mode_uses_encoded_text() {
995 let items = build_corpus(
996 &sample_entry(),
997 LongMemEvalGranularity::Session,
998 BenchMode::Aaak,
999 );
1000
1001 assert!(
1002 items[0]
1003 .retrieval_text
1004 .starts_with("V1|longmemeval|benchmark|")
1005 );
1006 assert!(items[0].retrieval_text.contains("Clerk"));
1007 }
1008
1009 #[test]
1010 fn test_evaluate_retrieval_matches_expected_metrics() {
1011 let rankings = vec![1, 0, 2];
1012 let correct_ids = BTreeSet::from(["sess_auth"]);
1013 let corpus_ids = vec![
1014 "sess_other".to_string(),
1015 "sess_auth".to_string(),
1016 "sess_third".to_string(),
1017 ];
1018
1019 let (recall_any, recall_all, ndcg_score) =
1020 evaluate_retrieval(&rankings, &correct_ids, &corpus_ids, 1);
1021
1022 assert_eq!(recall_any, 1.0);
1023 assert_eq!(recall_all, 1.0);
1024 assert!((ndcg_score - 1.0).abs() < f64::EPSILON);
1025 }
1026
1027 #[tokio::test]
1028 async fn test_run_benchmark_raw_small_sample() {
1029 let args = LongMemEvalArgs {
1030 data_file: PathBuf::from("unused.json"),
1031 mode: BenchMode::Raw,
1032 granularity: LongMemEvalGranularity::Session,
1033 limit: 0,
1034 skip: 0,
1035 top_k: 5,
1036 out: None,
1037 };
1038 let entry = sample_entry();
1039
1040 let (summary, logs) = run_benchmark_with_embedder(&TestEmbedder, &[entry], &args)
1041 .await
1042 .expect("benchmark should run");
1043
1044 assert_eq!(summary.question_count, 1);
1045 assert_eq!(summary.session.recall_any.get(&1), Some(&1.0));
1046 assert_eq!(logs.len(), 1);
1047 assert_eq!(
1048 logs[0].retrieval_results.metrics.session["recall_any@1"],
1049 1.0
1050 );
1051 }
1052
1053 #[tokio::test]
1054 async fn test_run_benchmark_rooms_small_sample() {
1055 let args = LongMemEvalArgs {
1056 data_file: PathBuf::from("unused.json"),
1057 mode: BenchMode::Rooms,
1058 granularity: LongMemEvalGranularity::Session,
1059 limit: 0,
1060 skip: 0,
1061 top_k: 5,
1062 out: None,
1063 };
1064 let entry = sample_entry();
1065
1066 let (summary, _logs) = run_benchmark_with_embedder(&TestEmbedder, &[entry], &args)
1067 .await
1068 .expect("rooms benchmark should run");
1069
1070 assert_eq!(summary.question_count, 1);
1071 assert_eq!(summary.session.recall_any.get(&1), Some(&1.0));
1072 }
1073}