use std::collections::HashMap;
use crate::bm25::bm25_score;
use crate::token::tokenize;
#[derive(Debug, Clone, PartialEq)]
pub struct TextMatch {
pub key: Vec<u8>,
pub score: f64,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct TextStats {
pub docs: u64,
pub tokens: u64,
pub postings: u64,
pub approx_bytes: u64,
}
type ScoredList<'s> = (&'s Buckets, f64, f64);
#[derive(Debug, Clone, Copy)]
struct Loc {
tf: u32,
band: u8,
slot: u32,
}
fn band_of(dl: u32) -> u8 {
(31 - dl.max(1).leading_zeros()).min(15) as u8
}
const BAND_MIN_DL: [u32; 16] = [
1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192, 16384, 32768,
];
type Bands = Vec<(u8, Vec<u32>)>;
#[derive(Debug)]
pub enum Buckets {
One { id: u32, tf: u32, dl: u32 },
Many(Box<ManyBuckets>),
}
#[derive(Debug, Default)]
pub struct ManyBuckets {
buckets: Vec<(u32, Bands)>,
index: HashMap<u32, Loc>,
}
impl Default for Buckets {
fn default() -> Self {
Buckets::Many(Box::default())
}
}
impl ManyBuckets {
fn insert(&mut self, tf: u32, dl: u32, id: u32) {
let pos = self.buckets.iter().position(|(t, _)| *t <= tf);
let slot = match pos {
Some(i) if self.buckets[i].0 == tf => &mut self.buckets[i].1,
Some(i) => {
self.buckets.insert(i, (tf, Bands::new()));
&mut self.buckets[i].1
}
None => {
self.buckets.push((tf, Bands::new()));
&mut self.buckets.last_mut().expect("just pushed").1
}
};
let band = band_of(dl);
let bi = match slot.binary_search_by_key(&band, |(b, _)| *b) {
Ok(i) => i,
Err(i) => {
slot.insert(i, (band, Vec::new()));
i
}
};
let v = &mut slot[bi].1;
self.index.insert(id, Loc { tf, band, slot: v.len() as u32 });
v.push(id);
}
fn remove(&mut self, id: u32) {
let Some(loc) = self.index.remove(&id) else { return };
if let Some(i) = self.buckets.iter().position(|(t, _)| *t == loc.tf)
&& let Ok(bi) = self.buckets[i].1.binary_search_by_key(&loc.band, |(b, _)| *b)
{
let v = &mut self.buckets[i].1[bi].1;
let si = loc.slot as usize;
v.swap_remove(si);
if let Some(&moved) = v.get(si) {
self.index.get_mut(&moved).expect("indexed posting").slot = loc.slot;
}
if v.is_empty() {
self.buckets[i].1.remove(bi);
}
if self.buckets[i].1.is_empty() {
self.buckets.remove(i);
}
}
}
}
impl Buckets {
fn new_one(tf: u32, dl: u32, id: u32) -> Self {
Buckets::One { id, tf, dl }
}
fn insert(&mut self, tf: u32, dl: u32, id: u32) {
match self {
Buckets::One { id: id0, tf: tf0, dl: dl0 } => {
let (id0, tf0, dl0) = (*id0, *tf0, *dl0);
let mut many = ManyBuckets::default();
many.insert(tf0, dl0, id0);
many.insert(tf, dl, id);
*self = Buckets::Many(Box::new(many));
}
Buckets::Many(m) => m.insert(tf, dl, id),
}
}
fn remove(&mut self, _tf: u32, _dl: u32, id: u32) {
match self {
Buckets::One { id: id0, .. } => {
if *id0 == id {
*self = Buckets::Many(Box::default());
}
}
Buckets::Many(m) => m.remove(id),
}
}
fn len(&self) -> usize {
match self {
Buckets::One { .. } => 1,
Buckets::Many(m) => m.index.len(),
}
}
fn is_empty(&self) -> bool {
self.len() == 0
}
fn max_tf(&self) -> u32 {
match self {
Buckets::One { tf, .. } => *tf,
Buckets::Many(m) => m.buckets.first().map_or(1, |(t, _)| *t),
}
}
fn get(&self, id: u32) -> Option<u32> {
match self {
Buckets::One { id: id0, tf, .. } => (*id0 == id).then_some(*tf),
Buckets::Many(m) => m.index.get(&id).map(|l| l.tf),
}
}
fn tf_groups(&self) -> Vec<(u32, BandsView<'_>)> {
match self {
Buckets::One { id, tf, dl } => {
vec![(*tf, BandsView::One(band_of(*dl), std::slice::from_ref(id)))]
}
Buckets::Many(m) => m
.buckets
.iter()
.map(|(t, bands)| (*t, BandsView::Slice(bands.as_slice())))
.collect(),
}
}
}
enum BandsView<'a> {
One(u8, &'a [u32]),
Slice(&'a [(u8, Vec<u32>)]),
}
impl BandsView<'_> {
fn iter(&self) -> impl Iterator<Item = (u8, &[u32])> + '_ {
let (one, slice) = match self {
BandsView::One(b, ids) => (Some((*b, *ids)), [].as_slice()),
BandsView::Slice(s) => (None, *s),
};
one.into_iter().chain(slice.iter().map(|(b, v)| (*b, v.as_slice())))
}
}
#[derive(Debug, Default)]
pub struct TextSegment {
postings: HashMap<Vec<u8>, Buckets>,
docs: HashMap<Vec<u8>, (u32, u32, Vec<u8>)>,
id_key: Vec<Option<Vec<u8>>>,
id_dl: Vec<u32>,
free_ids: Vec<u32>,
total_len: u64,
}
impl TextSegment {
pub fn new() -> Self {
Self::default()
}
pub fn apply(&mut self, key: &[u8], text: Option<&[u8]>) {
if let Some((old_id, old_len, old_text)) = self.docs.remove(key) {
self.total_len -= u64::from(old_len);
for (t, tf) in tf_of(&tokenize(&old_text)) {
if let Some(list) = self.postings.get_mut(&t) {
list.remove(tf, old_len, old_id);
if list.is_empty() {
self.postings.remove(&t);
}
}
}
self.id_key[old_id as usize] = None;
self.free_ids.push(old_id);
}
let Some(text) = text else { return };
let toks = tokenize(text);
if toks.is_empty() {
return;
}
let dl = toks.len() as u32;
let id = match self.free_ids.pop() {
Some(id) => {
self.id_key[id as usize] = Some(key.to_vec());
self.id_dl[id as usize] = dl;
id
}
None => {
self.id_key.push(Some(key.to_vec()));
self.id_dl.push(dl);
(self.id_key.len() - 1) as u32
}
};
self.docs.insert(key.to_vec(), (id, dl, text.to_vec()));
self.total_len += u64::from(dl);
for (t, tf) in tf_of(&toks) {
match self.postings.entry(t) {
std::collections::hash_map::Entry::Occupied(mut e) => {
e.get_mut().insert(tf, dl, id);
}
std::collections::hash_map::Entry::Vacant(v) => {
v.insert(Buckets::new_one(tf, dl, id));
}
}
}
}
pub fn matches(&self, query: &[u8], limit: usize) -> Vec<TextMatch> {
let mut q_tokens = tokenize(query);
q_tokens.sort();
q_tokens.dedup();
if q_tokens.is_empty() || self.docs.is_empty() {
return Vec::new();
}
let n_docs = self.docs.len() as f64;
let avgdl = self.total_len as f64 / n_docs;
let mut lists: Vec<ScoredList<'_>> = Vec::new();
for t in &q_tokens {
let Some(list) = self.postings.get(t) else { continue };
let df = list.len() as f64;
let max_tf = f64::from(list.max_tf());
lists.push((list, df, crate::bm25::bm25_upper(max_tf, df, n_docs)));
}
if lists.is_empty() {
return Vec::new();
}
lists.sort_by(|a, b| b.2.total_cmp(&a.2));
let tail_ub: Vec<f64> = {
let mut acc = 0.0;
let mut v: Vec<f64> = lists.iter().rev().map(|l| { acc += l.2; acc }).collect();
v.reverse();
v
};
let mut scores: HashMap<u32, f64> = HashMap::new();
let mut kth_threshold = 0.0_f64;
let mut walked = 0usize;
for (i, (list, df, _ub)) in lists.iter().enumerate() {
if i > 0 && scores.len() >= limit && tail_ub[i] < kth_threshold {
break;
}
walked = i + 1;
let groups = list.tf_groups();
for (bi, (tf, bands)) in groups.iter().enumerate() {
if scores.len() >= limit {
let bound = crate::bm25::bm25_upper(f64::from(*tf), *df, n_docs);
if bound + tail_ub[i + 1..].first().copied().unwrap_or(0.0)
< kth_of(&scores, limit)
{
let ids: Vec<u32> = scores.keys().copied().collect();
let walked_tfs: Vec<u32> =
groups[..bi].iter().map(|(t, _)| *t).collect();
for &id in &ids {
if let Some(tf2) = list.get(id)
&& !walked_tfs.contains(&tf2)
{
let dl_u = self.id_dl[id as usize];
*scores.get_mut(&id).expect("accumulated") += bm25_score(
f64::from(tf2),
*df,
n_docs,
f64::from(dl_u),
avgdl,
);
}
}
break;
}
}
for (b, band) in bands.iter() {
if band.is_empty() {
continue;
}
let bound = bm25_score(
f64::from(*tf),
*df,
n_docs,
f64::from(BAND_MIN_DL[b as usize]),
avgdl,
);
if lists.len() == 1
&& scores.len() >= limit
&& bound < kth_of(&scores, limit)
{
break;
}
for &id in band {
let dl = f64::from(self.id_dl[id as usize]);
*scores.entry(id).or_insert(0.0) +=
bm25_score(f64::from(*tf), *df, n_docs, dl, avgdl);
}
}
}
if scores.len() >= limit && i + 1 < lists.len() {
kth_threshold = kth_of(&scores, limit);
}
}
if walked < lists.len() {
let ids: Vec<u32> = scores.keys().copied().collect();
for (list, df, _) in &lists[walked..] {
for &id in &ids {
if let Some(tf) = list.get(id) {
let dl_u = self.id_dl[id as usize];
*scores.get_mut(&id).expect("accumulated") += bm25_score(
f64::from(tf),
*df,
n_docs,
f64::from(dl_u),
avgdl,
);
}
}
}
}
let key_of = |id: u32| -> &[u8] {
self.id_key[id as usize].as_deref().expect("live posting id")
};
let mut top: Vec<(f64, &[u8])> = Vec::with_capacity(limit + 1);
for (id, score) in &scores {
let cand = (*score, key_of(*id));
if top.len() < limit {
top.push(cand);
if top.len() == limit {
top.sort_by(|a, b| b.0.total_cmp(&a.0).then_with(|| a.1.cmp(b.1)));
}
} else if better(cand, top[limit - 1]) {
let pos = top
.partition_point(|e| better(*e, cand));
top.insert(pos, cand);
top.pop();
}
}
if top.len() < limit {
top.sort_by(|a, b| b.0.total_cmp(&a.0).then_with(|| a.1.cmp(b.1)));
}
top.into_iter()
.map(|(score, k)| TextMatch { key: k.to_vec(), score })
.collect()
}
pub fn stats(&self) -> TextStats {
let postings: u64 = self.postings.values().map(|l| l.len() as u64).sum();
let many_postings: u64 = self
.postings
.values()
.map(|l| match l {
Buckets::One { .. } => 0,
Buckets::Many(m) => m.index.len() as u64,
})
.sum();
let token_bytes: u64 = self.postings.keys().map(|t| (t.len() + 48) as u64).sum();
let doc_bytes: u64 = self
.docs
.iter()
.map(|(k, (_, _, text))| (2 * k.len() + text.len() + 110) as u64)
.sum();
TextStats {
docs: self.docs.len() as u64,
tokens: self.postings.len() as u64,
postings,
approx_bytes: token_bytes + many_postings * 30 + doc_bytes,
}
}
pub fn contains(&self, key: &[u8]) -> bool {
self.docs.contains_key(key)
}
}
fn tf_of(toks: &[Vec<u8>]) -> HashMap<Vec<u8>, u32> {
let mut tf = HashMap::new();
for t in toks {
*tf.entry(t.clone()).or_insert(0) += 1;
}
tf
}
fn better(a: (f64, &[u8]), b: (f64, &[u8])) -> bool {
a.0 > b.0 || (a.0 == b.0 && a.1 < b.1)
}
fn kth_of(scores: &HashMap<u32, f64>, limit: usize) -> f64 {
let mut v: Vec<f64> = scores.values().copied().collect();
let idx = limit - 1;
v.select_nth_unstable_by(idx, |a, b| b.total_cmp(a));
v[idx]
}
#[cfg(test)]
mod tests {
use super::*;
fn seg() -> TextSegment {
let mut s = TextSegment::new();
s.apply(b"d1", Some("rust full text search engine".as_bytes()));
s.apply(b"d2", Some("rust systems programming".as_bytes()));
s.apply(b"d3", Some("全文检索引擎 rust 実装".as_bytes()));
s
}
#[test]
fn ranked_or_semantics() {
let s = seg();
let hits = s.matches(b"rust search", 10);
assert_eq!(hits.len(), 3, "OR semantics: every rust doc matches");
assert_eq!(hits[0].key, b"d1".to_vec(), "d1 matches both terms → top");
let hits = s.matches(b"programming", 10);
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].key, b"d2".to_vec());
}
#[test]
fn cjk_query_bigrams() {
let s = seg();
let hits = s.matches("检索".as_bytes(), 10);
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].key, b"d3".to_vec());
assert!(s.matches("数据库".as_bytes(), 10).is_empty());
}
#[test]
fn update_and_remove() {
let mut s = seg();
s.apply(b"d1", Some(b"totally different now"));
assert!(s.matches(b"engine", 10).is_empty(), "old tokens gone");
assert_eq!(s.matches(b"different", 10)[0].key, b"d1".to_vec());
s.apply(b"d2", None);
assert!(!s.contains(b"d2"));
assert!(s.matches(b"programming", 10).is_empty());
let st = s.stats();
assert_eq!(st.docs, 2);
assert!(st.tokens > 0 && st.approx_bytes > 0);
}
#[test]
fn maxscore_pruning_matches_naive() {
let mut s = TextSegment::new();
for i in 0..500u32 {
let mut body = String::from("common filler words here");
if i % 5 == 0 {
body.push_str(" mid");
}
if i == 42 || i == 99 {
body.push_str(" rare");
}
for _ in 0..(i % 7) {
body.push_str(" pad");
}
s.apply(format!("k{i:03}").as_bytes(), Some(body.as_bytes()));
}
let naive = |query: &str, limit: usize| -> Vec<(Vec<u8>, f64)> {
let q = tokenize(query.as_bytes());
let n_docs = s.docs.len() as f64;
let avgdl = s.total_len as f64 / n_docs;
let mut sc: HashMap<Vec<u8>, f64> = HashMap::new();
for t in &q {
let Some(list) = s.postings.get(t) else { continue };
let df = list.len() as f64;
for (tf, bands) in list.tf_groups() {
for (_b, band) in bands.iter() {
for &id in band {
let k = s.id_key[id as usize].clone().expect("live id");
let dl = f64::from(s.id_dl[id as usize]);
*sc.entry(k).or_insert(0.0) +=
bm25_score(f64::from(tf), df, n_docs, dl, avgdl);
}
}
}
}
let mut v: Vec<(Vec<u8>, f64)> = sc.into_iter().collect();
v.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
v.truncate(limit);
v
};
for (q, limit) in [("rare common", 10), ("mid common", 5), ("rare mid common", 3), ("common", 7)] {
let got: Vec<(Vec<u8>, f64)> =
s.matches(q.as_bytes(), limit).into_iter().map(|m| (m.key, m.score)).collect();
let want = naive(q, limit);
assert_eq!(got, want, "query {q:?} limit {limit}");
}
}
#[test]
fn bucket_stop_keeps_walked_doc_contributions() {
let mut s = TextSegment::new();
for i in 0..2000u32 {
s.apply(format!("c{i:04}").as_bytes(), Some(b"common common"));
}
s.apply(b"special", Some(b"rare common pad pad pad"));
let naive_ok = {
let hits = s.matches(b"rare common", 5);
hits[0].key == b"special".to_vec()
};
assert!(naive_ok);
let mut s2 = TextSegment::new();
for i in 0..2000u32 {
s2.apply(format!("c{i:04}").as_bytes(), Some(b"common common"));
}
s2.apply(b"special", Some(b"rare only pad pad pad"));
let with_common = s.matches(b"rare common", 1)[0].score;
let without_common = s2.matches(b"rare common", 1)[0].score;
assert!(
with_common > without_common + 1e-9,
"skipped-bucket contribution lost: {with_common} vs {without_common}"
);
}
#[test]
fn limit_and_empty_query() {
let s = seg();
assert_eq!(s.matches(b"rust", 2).len(), 2);
assert!(s.matches(b"", 10).is_empty());
assert!(s.matches(b"!!!", 10).is_empty());
}
}