Skip to main content

urna_runtime/bm25/
index.rs

1//! `Bm25Index` struct + build/search. Encoding/decoding lives in
2//! `super::codec`.
3
4use std::collections::HashMap;
5
6use super::tokenize::tokenize;
7
8/// on-disk payload version for the bm25 section (`0x08`). v1 stored every
9/// posting as raw `(u32 doc, u32 tf)`; v2 delta-gaps the (sorted) doc ids
10/// and bitpacks the gaps, the term-frequencies, and the doc lengths with
11/// `intpack`. the decoded index is identical, so scores are unchanged. the
12/// reader still accepts v1. the section is optional and excluded from
13/// content_hash, so this bump is additive within v1.
14pub const BM25_PAYLOAD_VERSION: u32 = 2;
15pub const DEFAULT_K1: f32 = 1.5;
16pub const DEFAULT_B: f32 = 0.75;
17
18#[derive(Clone, Debug)]
19pub(super) struct Posting {
20    pub doc: u32,
21    pub tf: u32,
22}
23
24#[derive(Clone, Debug)]
25pub(super) struct TermEntry {
26    pub df: u32,
27    pub postings: Vec<Posting>,
28}
29
30pub struct Bm25Index {
31    pub k1: f32,
32    pub b: f32,
33    pub avgdl: f32,
34    pub n_docs: usize,
35    pub n_terms: usize,
36    pub(super) doc_lengths: Vec<u32>,
37    /// term -> entry. HashMap for O(1) query lookup.
38    pub(super) terms: HashMap<String, TermEntry>,
39}
40
41impl Bm25Index {
42    /// Build a BM25 index from the canonical chunk texts.
43    pub fn build(docs: &[String], k1: f32, b: f32) -> Self {
44        let n_docs = docs.len();
45        let mut doc_lengths = Vec::with_capacity(n_docs);
46        let mut term_postings: HashMap<String, Vec<Posting>> = HashMap::new();
47        for (doc_id, doc) in docs.iter().enumerate() {
48            let tokens = tokenize(doc);
49            doc_lengths.push(tokens.len() as u32);
50            let mut tf_map: HashMap<&str, u32> = HashMap::new();
51            for t in &tokens {
52                *tf_map.entry(t.as_str()).or_insert(0) += 1;
53            }
54            for (term, tf) in tf_map {
55                term_postings
56                    .entry(term.to_string())
57                    .or_default()
58                    .push(Posting {
59                        doc: doc_id as u32,
60                        tf,
61                    });
62            }
63        }
64        let total_dl: u64 = doc_lengths.iter().map(|x| *x as u64).sum();
65        let avgdl = if n_docs == 0 {
66            0.0
67        } else {
68            total_dl as f32 / n_docs as f32
69        };
70        // Sort terms alphabetically so the on-disk encoding is reproducible.
71        let mut entries: Vec<(String, Vec<Posting>)> = term_postings.into_iter().collect();
72        entries.sort_by(|a, b| a.0.cmp(&b.0));
73        let mut terms: HashMap<String, TermEntry> = HashMap::with_capacity(entries.len());
74        for (key, mut postings) in entries {
75            postings.sort_by_key(|p| p.doc);
76            terms.insert(
77                key,
78                TermEntry {
79                    df: postings.len() as u32,
80                    postings,
81                },
82            );
83        }
84        let n_terms = terms.len();
85        Self {
86            k1,
87            b,
88            avgdl,
89            n_docs,
90            n_terms,
91            doc_lengths,
92            terms,
93        }
94    }
95
96    /// Score `query_text` against the corpus, return the top-k `(doc, score)`
97    /// pairs. Empty query returns an empty vec.
98    pub fn search(&self, query_text: &str, k: usize) -> Vec<(usize, f32)> {
99        if k == 0 || self.n_docs == 0 {
100            return Vec::new();
101        }
102        let q_tokens = tokenize(query_text);
103        if q_tokens.is_empty() {
104            return Vec::new();
105        }
106        let mut scores: HashMap<u32, f32> = HashMap::new();
107        for term in q_tokens {
108            let Some(entry) = self.terms.get(&term) else {
109                continue;
110            };
111            let idf =
112                ((self.n_docs as f32 - entry.df as f32 + 0.5) / (entry.df as f32 + 0.5) + 1.0).ln();
113            for p in &entry.postings {
114                // the codec rejects out-of-range doc ids at open; `get` keeps
115                // this loop panic-free even so.
116                let dl = self.doc_lengths.get(p.doc as usize).copied().unwrap_or(0) as f32;
117                let denom =
118                    p.tf as f32 + self.k1 * (1.0 - self.b + self.b * dl / self.avgdl.max(1.0));
119                let term_score = idf * (p.tf as f32 * (self.k1 + 1.0)) / denom;
120                *scores.entry(p.doc).or_insert(0.0) += term_score;
121            }
122        }
123        let mut all: Vec<(usize, f32)> = scores.into_iter().map(|(d, s)| (d as usize, s)).collect();
124        // a corrupted payload (zero doc lengths, absurd df) can yield a
125        // non-finite score; drop those rather than rank them.
126        all.retain(|p| p.1.is_finite());
127        all.sort_by(|a, b| crate::order::cmp_score_desc(*a, *b));
128        all.truncate(k);
129        all
130    }
131}