Skip to main content

mempal_runtime/
longmemeval.rs

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}