pub mod analyzer;
pub mod bm25;
pub use analyzer::analyze;
use rustc_hash::FxHashMap;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
pub type TermId = u32;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct Posting {
pub slot: u32,
pub tf: u32,
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
struct Doc {
terms: Vec<(TermId, u32)>,
len: u32,
}
impl Doc {
#[inline]
fn term_freq(&self, term: TermId) -> u32 {
match self.terms.binary_search_by_key(&term, |&(id, _)| id) {
Ok(at) => self.terms[at].1,
Err(_) => 0,
}
}
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub struct TextIndex {
ids: FxHashMap<Arc<str>, TermId>,
names: Vec<Option<Arc<str>>>,
postings: Vec<Vec<Posting>>,
free_ids: Vec<TermId>,
docs: FxHashMap<u32, Doc>,
total_len: u64,
}
impl TextIndex {
pub fn new() -> Self {
Self::default()
}
pub fn build<I, S>(docs: I) -> Self
where
I: IntoIterator<Item = (u32, S)>,
S: AsRef<str>,
{
let mut index = Self::new();
for (slot, text) in docs {
index.remove_doc(slot);
let doc = index.intern_document(text.as_ref());
for &(term, tf) in &doc.terms {
index.postings[term as usize].push(Posting { slot, tf });
}
index.total_len += u64::from(doc.len);
index.docs.insert(slot, doc);
}
for list in &mut index.postings {
list.sort_unstable_by_key(|posting| posting.slot);
}
index
}
pub fn from_terms<S: AsRef<str>>(
terms: impl IntoIterator<Item = (S, Vec<Posting>)>,
empty_docs: &[u32],
) -> Self {
let mut index = Self::new();
for (name, postings) in terms {
if postings.is_empty() {
continue;
}
let name: Arc<str> = Arc::from(name.as_ref());
let id = index.names.len() as TermId;
index.names.push(Some(Arc::clone(&name)));
index.ids.insert(name, id);
index.postings.push(postings);
}
let mut docs: FxHashMap<u32, Doc> = FxHashMap::default();
for (id, list) in index.postings.iter().enumerate() {
for posting in list {
let doc = docs.entry(posting.slot).or_default();
doc.terms.push((id as TermId, posting.tf));
doc.len = doc.len.saturating_add(posting.tf);
}
}
for &slot in empty_docs {
docs.entry(slot).or_default();
}
index.total_len = docs.values().map(|doc| u64::from(doc.len)).sum();
index.docs = docs;
index
}
pub fn add_doc(&mut self, slot: u32, text: &str) {
self.remove_doc(slot);
let doc = self.intern_document(text);
for &(term, tf) in &doc.terms {
let list = &mut self.postings[term as usize];
let at = list.partition_point(|posting| posting.slot < slot);
list.insert(at, Posting { slot, tf });
}
self.total_len += u64::from(doc.len);
self.docs.insert(slot, doc);
}
pub fn remove_doc(&mut self, slot: u32) -> bool {
let Some(doc) = self.docs.remove(&slot) else {
return false;
};
self.total_len -= u64::from(doc.len);
for &(term, _) in &doc.terms {
let list = &mut self.postings[term as usize];
let at = list
.binary_search_by_key(&slot, |posting| posting.slot)
.ok()
.or_else(|| list.iter().position(|posting| posting.slot == slot));
if let Some(at) = at {
list.remove(at);
}
if list.is_empty() {
self.release_term(term);
}
}
true
}
fn intern_document(&mut self, text: &str) -> Doc {
let mut counts: FxHashMap<TermId, u32> = FxHashMap::default();
let mut len: u32 = 0;
for token in analyze(text) {
let term = self.intern_term(&token);
*counts.entry(term).or_insert(0) += 1;
len += 1;
}
let mut terms: Vec<(TermId, u32)> = counts.into_iter().collect();
terms.sort_unstable_by_key(|&(id, _)| id);
Doc { terms, len }
}
fn intern_term(&mut self, token: &str) -> TermId {
if let Some(&id) = self.ids.get(token) {
return id;
}
let name: Arc<str> = Arc::from(token);
let id = match self.free_ids.pop() {
Some(id) => {
self.names[id as usize] = Some(Arc::clone(&name));
id
}
None => {
self.names.push(Some(Arc::clone(&name)));
self.postings.push(Vec::new());
(self.postings.len() - 1) as TermId
}
};
self.ids.insert(name, id);
id
}
fn release_term(&mut self, term: TermId) {
if let Some(name) = self.names[term as usize].take() {
self.ids.remove(&name);
self.free_ids.push(term);
}
}
pub fn total_docs(&self) -> usize {
self.docs.len()
}
pub fn is_empty(&self) -> bool {
self.docs.is_empty()
}
pub fn vocabulary_len(&self) -> usize {
self.ids.len()
}
pub fn contains_doc(&self, slot: u32) -> bool {
self.docs.contains_key(&slot)
}
pub fn doc_len(&self, slot: u32) -> Option<u32> {
self.docs.get(&slot).map(|doc| doc.len)
}
pub fn avgdl(&self) -> f64 {
if self.docs.is_empty() {
return 0.0;
}
self.total_len as f64 / self.docs.len() as f64
}
#[cfg(test)]
pub fn df(&self, term: &str) -> usize {
self.postings_for(term).len()
}
pub fn term_id(&self, term: &str) -> Option<TermId> {
self.ids.get(term).copied()
}
#[cfg(test)]
pub fn postings_for(&self, term: &str) -> &[Posting] {
match self.ids.get(term) {
Some(&id) => &self.postings[id as usize],
None => &[],
}
}
pub fn postings_of(&self, term: TermId) -> &[Posting] {
self.postings.get(term as usize).map_or(&[], Vec::as_slice)
}
#[cfg(test)]
pub fn term_freq(&self, slot: u32, term: TermId) -> u32 {
self.docs.get(&slot).map_or(0, |doc| doc.term_freq(term))
}
pub fn doc_slots(&self) -> impl Iterator<Item = u32> + '_ {
self.docs.keys().copied()
}
pub fn iter_terms(&self) -> impl Iterator<Item = (&str, &[Posting])> + '_ {
self.names
.iter()
.enumerate()
.filter_map(|(id, name)| Some((name.as_deref()?, self.postings[id].as_slice())))
}
pub fn estimated_bytes(&self) -> usize {
const TABLE_SLACK: usize = 8;
let dictionary: usize = self
.ids
.keys()
.map(|name| {
name.len()
+ 2 * std::mem::size_of::<usize>()
+ 2 * std::mem::size_of::<Arc<str>>()
+ std::mem::size_of::<TermId>()
+ TABLE_SLACK
})
.sum();
let postings: usize = self
.postings
.iter()
.map(|list| {
list.capacity() * std::mem::size_of::<Posting>() + std::mem::size_of_val(list)
})
.sum();
let forward: usize = self
.docs
.values()
.map(|doc| {
doc.terms.capacity() * std::mem::size_of::<(TermId, u32)>()
+ std::mem::size_of::<(u32, Doc)>()
+ TABLE_SLACK
})
.sum();
dictionary + postings + forward + self.free_ids.capacity() * std::mem::size_of::<TermId>()
}
pub fn validate(&self) -> Result<(), String> {
let mut expected_total: u64 = 0;
let mut expected_postings: usize = 0;
for (&slot, doc) in &self.docs {
expected_total += u64::from(doc.len);
expected_postings += doc.terms.len();
self.validate_doc(slot, doc)?;
}
if expected_total != self.total_len {
return Err(format!(
"total_len {} != Σ doc lengths {expected_total}",
self.total_len
));
}
self.validate_postings(expected_postings)
}
fn validate_doc(&self, slot: u32, doc: &Doc) -> Result<(), String> {
let mut sum: u64 = 0;
let mut previous: Option<TermId> = None;
for &(term, tf) in &doc.terms {
if tf == 0 {
return Err(format!("doc {slot} keeps a zero-frequency term {term}"));
}
if previous.is_some_and(|last| last >= term) {
return Err(format!("doc {slot} terms are not strictly ascending"));
}
previous = Some(term);
sum += u64::from(tf);
if self.names[term as usize].is_none() {
return Err(format!("doc {slot} references freed term id {term}"));
}
let list = self.postings_of(term);
match list.binary_search_by_key(&slot, |posting| posting.slot) {
Ok(at) if list[at].tf == tf => {}
Ok(at) => {
return Err(format!(
"doc {slot} term {term}: forward tf {tf} != posting tf {}",
list[at].tf
))
}
Err(_) => return Err(format!("doc {slot} term {term} has no posting")),
}
}
if sum != u64::from(doc.len) {
return Err(format!("doc {slot} length {} != Σtf {sum}", doc.len));
}
Ok(())
}
fn validate_postings(&self, expected_postings: usize) -> Result<(), String> {
let mut seen: usize = 0;
for (id, name) in self.names.iter().enumerate() {
let Some(name) = name else { continue };
let list = &self.postings[id];
if list.is_empty() {
return Err(format!("term '{name}' is interned with no postings"));
}
if self.ids.get(name.as_ref()) != Some(&(id as TermId)) {
return Err(format!("term '{name}' does not resolve back to id {id}"));
}
seen += list.len();
let mut previous: Option<u32> = None;
for posting in list {
if previous.is_some_and(|last| last >= posting.slot) {
return Err(format!("term '{name}' postings are not strictly ascending"));
}
previous = Some(posting.slot);
match self.docs.get(&posting.slot) {
Some(doc) if doc.term_freq(id as TermId) == posting.tf => {}
Some(_) => {
return Err(format!(
"term '{name}' posting for slot {} disagrees with the forward view",
posting.slot
))
}
None => {
return Err(format!(
"term '{name}' keeps a posting for unindexed slot {}",
posting.slot
))
}
}
}
}
if seen != expected_postings {
return Err(format!(
"postings hold {seen} entries, the forward view {expected_postings}"
));
}
self.validate_id_space()
}
fn validate_id_space(&self) -> Result<(), String> {
if self.names.len() != self.postings.len() {
return Err(format!(
"{} term names but {} posting lists",
self.names.len(),
self.postings.len()
));
}
for &id in &self.free_ids {
if self.names[id as usize].is_some() {
return Err(format!("term id {id} is both live and free"));
}
if !self.postings[id as usize].is_empty() {
return Err(format!("freed term id {id} still has postings"));
}
}
let live = self.names.iter().filter(|name| name.is_some()).count();
if live != self.ids.len() {
return Err(format!(
"{live} named ids but {} dictionary entries",
self.ids.len()
));
}
if live + self.free_ids.len() != self.names.len() {
return Err(format!(
"id space leak: {live} live + {} free != {} slots",
self.free_ids.len(),
self.names.len()
));
}
Ok(())
}
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;