use std::collections::HashMap;
const K1: f32 = 1.2;
const B: f32 = 0.75;
#[derive(Debug, Clone, Copy)]
struct Posting {
doc: u32,
field: u8,
tf: u16,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Hit {
pub doc_idx: usize,
pub score: f32,
pub matched_terms: Vec<String>,
}
#[derive(Debug)]
pub struct Bm25Index<D> {
field_weights: Vec<f32>,
docs: Vec<D>,
field_lens: Vec<u32>,
avg_lens: Vec<f32>,
postings: HashMap<String, Vec<Posting>>,
}
#[derive(Debug)]
pub struct Bm25Builder<D> {
index: Bm25Index<D>,
}
impl<D> Bm25Builder<D> {
pub fn new(field_weights: &[f32]) -> Self {
assert!(field_weights.len() <= u8::MAX as usize, "too many fields");
Self {
index: Bm25Index {
field_weights: field_weights.to_vec(),
docs: Vec::new(),
field_lens: Vec::new(),
avg_lens: vec![0.0; field_weights.len()],
postings: HashMap::new(),
},
}
}
pub fn add(&mut self, payload: D, fields: &[Vec<String>]) {
let num_fields = self.index.field_weights.len();
let doc = self.index.docs.len() as u32;
self.index.docs.push(payload);
for field in 0..num_fields {
let tokens = fields.get(field).map(Vec::as_slice).unwrap_or(&[]);
self.index.field_lens.push(tokens.len() as u32);
let mut counts: Vec<(&str, u16)> = Vec::new();
for token in tokens {
match counts.iter_mut().find(|(t, _)| *t == token.as_str()) {
Some((_, c)) => *c = c.saturating_add(1),
None => counts.push((token.as_str(), 1)),
}
}
for (term, tf) in counts {
self.index
.postings
.entry(term.to_string())
.or_default()
.push(Posting {
doc,
field: field as u8,
tf,
});
}
}
}
pub fn build(mut self) -> Bm25Index<D> {
let num_fields = self.index.field_weights.len();
let num_docs = self.index.docs.len();
if num_docs > 0 {
for field in 0..num_fields {
let total: u64 = (0..num_docs)
.map(|d| u64::from(self.index.field_lens[d * num_fields + field]))
.sum();
self.index.avg_lens[field] = total as f32 / num_docs as f32;
}
}
self.index
}
}
impl<D> Bm25Index<D> {
pub fn len(&self) -> usize {
self.docs.len()
}
pub fn is_empty(&self) -> bool {
self.docs.is_empty()
}
pub fn doc(&self, idx: usize) -> &D {
&self.docs[idx]
}
pub fn docs(&self) -> &[D] {
&self.docs
}
pub fn search(&self, query: &[String], filter: impl Fn(&D) -> bool) -> Vec<Hit> {
let num_fields = self.field_weights.len();
let n = self.docs.len() as f32;
let mut acc: HashMap<u32, (f32, Vec<usize>)> = HashMap::new();
let mut allowed: HashMap<u32, bool> = HashMap::new();
let mut seen_terms: Vec<&str> = Vec::new();
for (term_idx, term) in query.iter().enumerate() {
if seen_terms.contains(&term.as_str()) {
continue;
}
seen_terms.push(term);
let Some(postings) = self.postings.get(term) else {
continue;
};
let mut per_doc: Vec<(u32, f32)> = Vec::new();
for p in postings {
let len = self.field_lens[p.doc as usize * num_fields + p.field as usize] as f32;
let avg = self.avg_lens[p.field as usize].max(f32::EPSILON);
let norm = 1.0 - B + B * len / avg;
let weighted = self.field_weights[p.field as usize] * f32::from(p.tf) / norm;
match per_doc.last_mut() {
Some((d, tf)) if *d == p.doc => *tf += weighted,
_ => per_doc.push((p.doc, weighted)),
}
}
let df = per_doc.len() as f32;
let idf = (1.0 + (n - df + 0.5) / (df + 0.5)).ln();
for (doc, tf) in per_doc {
let ok = *allowed
.entry(doc)
.or_insert_with(|| filter(&self.docs[doc as usize]));
if !ok {
continue;
}
let score = idf * tf * (K1 + 1.0) / (tf + K1);
let entry = acc.entry(doc).or_insert((0.0, Vec::new()));
entry.0 += score;
entry.1.push(term_idx);
}
}
let mut hits: Vec<Hit> = acc
.into_iter()
.filter(|(_, (score, _))| *score > 0.0)
.map(|(doc, (score, terms))| Hit {
doc_idx: doc as usize,
score,
matched_terms: terms.into_iter().map(|i| query[i].clone()).collect(),
})
.collect();
hits.sort_by(|a, b| {
b.score
.total_cmp(&a.score)
.then_with(|| a.doc_idx.cmp(&b.doc_idx))
});
hits
}
}
#[cfg(test)]
mod tests {
use super::*;
fn toks(s: &str) -> Vec<String> {
s.split_whitespace().map(str::to_string).collect()
}
#[test]
fn field_weights_change_ranking() {
let mut b = Bm25Builder::new(&[3.0, 1.0]);
b.add("body", &[toks("alpha"), toks("retry")]);
b.add("name", &[toks("retry"), toks("alpha")]);
let idx = b.build();
let hits = idx.search(&toks("retry"), |_| true);
assert_eq!(hits.len(), 2);
assert_eq!(*idx.doc(hits[0].doc_idx), "name");
}
#[test]
fn rare_term_outranks_common_term() {
let mut b = Bm25Builder::new(&[1.0]);
b.add(0, &[toks("common rare")]);
b.add(1, &[toks("common other")]);
b.add(2, &[toks("common")]);
b.add(3, &[toks("common")]);
let idx = b.build();
let rare = idx.search(&toks("rare"), |_| true);
let common = idx.search(&toks("common"), |_| true);
assert!(rare[0].score > common[0].score);
let both = idx.search(&toks("common rare"), |_| true);
assert_eq!(both[0].doc_idx, 0);
}
#[test]
fn equal_scores_are_ordered_by_insertion() {
let mut b = Bm25Builder::new(&[1.0]);
for i in 0..5 {
b.add(i, &[toks("same tokens")]);
}
let idx = b.build();
let hits = idx.search(&toks("same"), |_| true);
let order: Vec<usize> = hits.iter().map(|h| h.doc_idx).collect();
assert_eq!(order, vec![0, 1, 2, 3, 4]);
}
#[test]
fn matched_terms_and_filter() {
let mut b = Bm25Builder::new(&[1.0]);
b.add(10, &[toks("retry backoff")]);
b.add(20, &[toks("backoff")]);
b.add(30, &[toks("parse config")]);
let idx = b.build();
let hits = idx.search(&toks("retry backoff missing"), |_| true);
assert_eq!(hits.len(), 2);
assert_eq!(hits[0].doc_idx, 0);
assert_eq!(hits[0].matched_terms, toks("retry backoff"));
assert_eq!(hits[1].matched_terms, toks("backoff"));
let filtered = idx.search(&toks("backoff"), |d| *d != 10);
assert_eq!(filtered.len(), 1);
assert_eq!(*idx.doc(filtered[0].doc_idx), 20);
}
#[test]
fn empty_index_and_query() {
let idx: Bm25Index<u8> = Bm25Builder::new(&[1.0]).build();
assert!(idx.search(&toks("x"), |_| true).is_empty());
let mut b = Bm25Builder::new(&[1.0]);
b.add(1u8, &[toks("x")]);
assert!(b.build().search(&[], |_| true).is_empty());
}
}