use std::{
cmp::Ordering,
collections::{BinaryHeap, HashMap},
ops::Range,
sync::Arc,
};
use bytes::Bytes;
use serde::Deserialize;
use crate::superfile::{
ReadError,
error::FtsError,
format::{
self, FST_SEPARATOR,
checksum::crc32c,
fts::{
HEADER_SIZE_V1_LEGACY as FTS_HEADER_SIZE, MAGIC_BYTES, U32_BYTES, U64_BYTES, hdr,
skip_entry, term_meta,
},
},
fts::{
bm25,
builder::{
DOC_LENGTHS_ENTRY_SIZE, SKIP_ENTRY_SIZE, TERM_META_POSITIONAL_SIZE, TERM_META_SIZE,
},
dict::{DictReader, make_key},
fst_value::FstValue,
positions::{decode_run, skip_run},
posting::{BLOCK_LEN, decode_block},
tokenize::{AsciiLowerTokenizer, Tokenizer as _},
},
lazy_source::{LazyByteSource, PrefetchedSource, RangeCoalescePlan, Source},
};
const TERM_RANGE_COALESCE_MAX_GAP: usize = 64 * 1024;
const TERM_RANGE_COALESCE_MAX_OVERFETCH: usize = 512 * 1024;
#[derive(Debug, Copy, Clone, Eq, PartialEq)]
pub enum BoolMode {
And,
Or,
}
#[derive(Default)]
pub(crate) struct ClauseLists<'a> {
pub musts: &'a [&'a str],
pub shoulds: &'a [&'a str],
pub negatives: &'a [&'a str],
pub must_phrases: &'a [Vec<String>],
pub should_phrases: &'a [Vec<String>],
pub negative_phrases: &'a [Vec<String>],
}
impl ClauseLists<'_> {
fn has_phrases(&self) -> bool {
!self.must_phrases.is_empty()
|| !self.should_phrases.is_empty()
|| !self.negative_phrases.is_empty()
}
fn no_positive_atoms(&self) -> bool {
self.musts.is_empty()
&& self.shoulds.is_empty()
&& self.must_phrases.is_empty()
&& self.should_phrases.is_empty()
}
fn no_negative_atoms(&self) -> bool {
self.negatives.is_empty() && self.negative_phrases.is_empty()
}
}
impl From<&str> for BoolMode {
fn from(s: &str) -> Self {
match s {
"and" => BoolMode::And,
"or" => BoolMode::Or,
_ => BoolMode::Or,
}
}
}
#[doc(hidden)]
#[derive(Debug, Copy, Clone, Eq, PartialEq)]
pub enum OrAlgo {
Bmm,
WandBmw,
Exhaustive,
Windowed,
}
const OR_WINDOW: u32 = 4096;
const OR_WINDOW_WORDS: usize = (OR_WINDOW as usize).div_ceil(64);
const OR_WINDOW_MIN_TERMS: usize = 3;
const OR_WINDOW_DOMINANCE_MULT: f32 = 1.5;
const WAND_BMW_2TERM_MAX_K: usize = 128;
const WAND_BMW_2TERM_DF_RATIO: u64 = 16;
fn no_dominant_term_ub(cursors: &[TermCursor]) -> bool {
let total: f32 = cursors.iter().map(|c| c.term_max_bm25).sum();
if total <= 0.0 {
return false;
}
let max = cursors
.iter()
.map(|c| c.term_max_bm25)
.fold(0.0f32, f32::max);
let avg = total / cursors.len() as f32;
max <= OR_WINDOW_DOMINANCE_MULT * avg
}
fn prefer_windowed_union(cursors: &[TermCursor]) -> bool {
cursors.len() >= OR_WINDOW_MIN_TERMS && no_dominant_term_ub(cursors)
}
fn two_term_has_rare_anchor(cursors: &[TermCursor]) -> bool {
if cursors.len() != 2 {
return false;
}
let lo = cursors[0].df.min(cursors[1].df);
let hi = cursors[0].df.max(cursors[1].df);
lo > 0 && hi >= lo.saturating_mul(WAND_BMW_2TERM_DF_RATIO)
}
#[derive(Debug, Clone)]
pub struct NormTable {
bytes: Vec<u8>,
lut: Box<[f32; 256]>,
}
impl NormTable {
fn new(doc_lengths: impl Iterator<Item = u32>, n_docs: usize, avgdl: f32) -> Self {
if avgdl <= 0.0 {
return Self::empty();
}
let inv_avgdl = 1.0_f32 / avgdl;
let mut lut = Box::new([0.0_f32; 256]);
for (b, slot) in lut.iter_mut().enumerate() {
let dl = bm25::dequantize_len(b as u8) as f32;
*slot = bm25::K1 * (1.0 - bm25::B + bm25::B * dl * inv_avgdl);
}
let mut bytes = Vec::with_capacity(n_docs);
for dl in doc_lengths {
bytes.push(bm25::quantize_len(dl));
}
Self { bytes, lut }
}
#[inline(always)]
fn get(&self, doc: u32) -> f32 {
self.lut[self.bytes[doc as usize] as usize]
}
#[cfg(test)]
fn len(&self) -> usize {
self.bytes.len()
}
fn empty() -> Self {
Self {
bytes: Vec::new(),
lut: Box::new([0.0; 256]),
}
}
}
#[derive(Debug, Clone)]
pub struct ColumnMeta {
pub name: String,
pub doc_lengths_range: Range<usize>,
pub avgdl: f32,
pub dl_norm_k1: NormTable,
pub positions: bool,
}
#[derive(Debug, Clone, Deserialize)]
pub struct FtsColumnConfig {
pub name: String,
#[serde(default = "default_tokenizer")]
pub tokenizer: String,
#[serde(default)]
pub positions: bool,
}
fn default_tokenizer() -> String {
"ascii_lower".to_string()
}
#[derive(Debug, Clone, Copy)]
pub struct OpenOptions {
pub verify_crc: bool,
}
impl Default for OpenOptions {
fn default() -> Self {
Self { verify_crc: true }
}
}
impl OpenOptions {
pub fn for_object_store() -> Self {
Self { verify_crc: false }
}
}
#[derive(Debug)]
pub struct FtsReader {
source: Source,
n_docs: u32,
n_terms_total: u32,
fst_range: Range<usize>,
postings_range: Range<usize>,
positions_range: Option<Range<usize>>,
columns: Vec<ColumnMeta>,
column_id_by_name: HashMap<String, u32>,
}
impl FtsReader {
pub fn open(blob: Bytes, columns_json: &str) -> Result<Self, FtsError> {
Self::open_with(blob, columns_json, OpenOptions::default())
}
pub fn open_with(blob: Bytes, columns_json: &str, opts: OpenOptions) -> Result<Self, FtsError> {
Self::open_with_source(Source::InMemory(blob), columns_json, opts)
}
pub async fn open_lazy(
source: Arc<dyn LazyByteSource>,
columns_json: &str,
opts: OpenOptions,
) -> Result<Self, FtsError> {
let fts_blob_len = source.size() as usize;
let header_fetch = format::fts::HEADER_SIZE_V2.min(fts_blob_len);
let header = fetch_lazy_range(source.as_ref(), 0..header_fetch, "fts header").await?;
if header.len() < FTS_HEADER_SIZE {
return Err(FtsError::Read(ReadError::MissingKv("fts header")));
}
if &header[0..MAGIC_BYTES] != format::fts::MAGIC {
return Err(FtsError::Read(ReadError::BadMagic {
section: "fts",
expected: format::fts::MAGIC,
actual: header[0..MAGIC_BYTES].to_vec(),
}));
}
let version = read_u32_le(&header[hdr::VERSION_OFF..hdr::VERSION_OFF + U32_BYTES]);
if version != format::fts::VERSION_V1_LEGACY && version != format::fts::VERSION_V2 {
return Err(FtsError::Read(ReadError::UnsupportedVersion(format!(
"fts section version {version}"
))));
}
let header_size = match version {
v if v == format::fts::VERSION_V2 => format::fts::HEADER_SIZE_V2,
_ => FTS_HEADER_SIZE,
};
if header.len() < header_size {
return Err(FtsError::Read(ReadError::MissingKv("fts header")));
}
let postings_offset =
read_u64_le(&header[hdr::POSTINGS_OFFSET_OFF..hdr::POSTINGS_OFFSET_OFF + U64_BYTES])
as usize;
let doc_lengths_table_offset =
read_u64_le(&header[hdr::DOC_LENGTHS_DIR_OFF..hdr::DOC_LENGTHS_DIR_OFF + U64_BYTES])
as usize;
let (fst_region, doc_lengths_tail) = futures::try_join!(
fetch_lazy_range(source.as_ref(), header_size..postings_offset, "fts/dict"),
fetch_lazy_range(
source.as_ref(),
doc_lengths_table_offset..fts_blob_len,
"fts/doc_lengths_tail",
),
)?;
let mut overlay = PrefetchedSource::new(source);
overlay.install(0, header);
overlay.install(header_size as u64, fst_region);
overlay.install(doc_lengths_table_offset as u64, doc_lengths_tail);
Self::open_with_source(Source::Lazy(Arc::new(overlay)), columns_json, opts)
}
pub(crate) fn open_with_source(
source: Source,
columns_json: &str,
opts: OpenOptions,
) -> Result<Self, FtsError> {
let source_len = source.len();
if source_len < FTS_HEADER_SIZE {
return Err(FtsError::Read(ReadError::MissingKv("fts header")));
}
let header = fetch_source_range(&source, 0..FTS_HEADER_SIZE, "fts header")?;
if &header[0..MAGIC_BYTES] != format::fts::MAGIC {
return Err(FtsError::Read(ReadError::BadMagic {
section: "fts",
expected: format::fts::MAGIC,
actual: header[0..MAGIC_BYTES].to_vec(),
}));
}
let version = read_u32_le(&header[hdr::VERSION_OFF..hdr::VERSION_OFF + U32_BYTES]);
let positional_blob = match version {
v if v == format::fts::VERSION_V1_LEGACY => false,
v if v == format::fts::VERSION_V2 => true,
_ => {
return Err(FtsError::Read(ReadError::UnsupportedVersion(format!(
"fts section version {version}"
))));
}
};
let header_size = match positional_blob {
true => format::fts::HEADER_SIZE_V2,
false => FTS_HEADER_SIZE,
};
if source_len < header_size {
return Err(FtsError::Read(ReadError::MissingKv("fts header")));
}
let n_columns =
read_u32_le(&header[hdr::N_COLUMNS_OFF..hdr::N_COLUMNS_OFF + U32_BYTES]) as usize;
let n_docs = read_u32_le(&header[hdr::N_DOCS_OFF..hdr::N_DOCS_OFF + U32_BYTES]);
let n_terms_total = read_u32_le(&header[hdr::N_TERMS_OFF..hdr::N_TERMS_OFF + U32_BYTES]);
let fst_offset =
read_u64_le(&header[hdr::FST_OFFSET_OFF..hdr::FST_OFFSET_OFF + U64_BYTES]) as usize;
let postings_offset =
read_u64_le(&header[hdr::POSTINGS_OFFSET_OFF..hdr::POSTINGS_OFFSET_OFF + U64_BYTES])
as usize;
let doc_lengths_table_offset =
read_u64_le(&header[hdr::DOC_LENGTHS_DIR_OFF..hdr::DOC_LENGTHS_DIR_OFF + U64_BYTES])
as usize;
let positions_offset: Option<usize> = match positional_blob {
true => {
let ext = fetch_source_range(
&source,
FTS_HEADER_SIZE..format::fts::HEADER_SIZE_V2,
"fts header ext",
)?;
Some(read_u64_le(&ext[0..U64_BYTES]) as usize)
}
false => None,
};
let postings_end = positions_offset.unwrap_or(doc_lengths_table_offset);
if fst_offset < header_size
|| postings_offset < fst_offset + 4
|| postings_end < postings_offset + 4
|| doc_lengths_table_offset < postings_end
|| doc_lengths_table_offset > source_len
|| positions_offset.is_some_and(|po| doc_lengths_table_offset < po + 4)
{
return Err(FtsError::Read(ReadError::MalformedVersion(format!(
"fts header offsets out of range: fst={fst_offset}, postings={postings_offset}, \
positions={positions_offset:?}, doc_lengths={doc_lengths_table_offset}, \
blob_len={}",
source_len
))));
}
let fst_range = fst_offset..postings_offset.saturating_sub(4); let postings_range = postings_offset..postings_end.saturating_sub(4); let positions_range: Option<Range<usize>> =
positions_offset.map(|po| po..doc_lengths_table_offset.saturating_sub(4));
if opts.verify_crc {
let fst_crc_bytes = fetch_source_range(
&source,
postings_offset.saturating_sub(4)..postings_offset,
"fts/dict crc",
)?;
let fst_crc_expected = read_u32_le(&fst_crc_bytes);
let fst_bytes = fetch_source_range(&source, fst_range.clone(), "fts/dict")?;
let fst_crc_actual = crc32c(&fst_bytes);
if fst_crc_expected != fst_crc_actual {
return Err(FtsError::Read(ReadError::ChecksumMismatch {
section: "fts/dict",
column: String::new(),
}));
}
}
if opts.verify_crc {
let postings_crc_pos = postings_end.saturating_sub(4);
let postings_crc_bytes =
fetch_source_range(&source, postings_crc_pos..postings_end, "fts/postings crc")?;
let postings_crc_expected = read_u32_le(&postings_crc_bytes);
let postings_bytes =
fetch_source_range(&source, postings_range.clone(), "fts/postings")?;
let postings_crc_actual = crc32c(&postings_bytes);
if postings_crc_expected != postings_crc_actual {
return Err(FtsError::Read(ReadError::ChecksumMismatch {
section: "fts/postings",
column: String::new(),
}));
}
}
if opts.verify_crc
&& let Some(pos_range) = &positions_range
{
let crc_pos = doc_lengths_table_offset.saturating_sub(4);
let crc_bytes = fetch_source_range(
&source,
crc_pos..doc_lengths_table_offset,
"fts/positions crc",
)?;
let crc_expected = read_u32_le(&crc_bytes);
let pos_bytes = fetch_source_range(&source, pos_range.clone(), "fts/positions")?;
let crc_actual = crc32c(&pos_bytes);
if crc_expected != crc_actual {
return Err(FtsError::Read(ReadError::ChecksumMismatch {
section: "fts/positions",
column: String::new(),
}));
}
}
let cols: Vec<FtsColumnConfig> = serde_json::from_str(columns_json).map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"inf.fts.columns JSON: {e}"
)))
})?;
if cols.len() != n_columns {
return Err(FtsError::Read(ReadError::MalformedVersion(format!(
"inf.fts.columns has {} entries, header says {}",
cols.len(),
n_columns
))));
}
let dir_size = n_columns * DOC_LENGTHS_ENTRY_SIZE;
let dir_end = doc_lengths_table_offset + dir_size;
if dir_end + 4 > source_len {
return Err(FtsError::Read(ReadError::MalformedVersion(
"doc-lengths directory runs past blob end".into(),
)));
}
let dir_region = fetch_source_range(
&source,
doc_lengths_table_offset..dir_end + 4,
"fts/doc_lengths_dir",
)?;
let dir_bytes = &dir_region[..dir_size];
if opts.verify_crc {
let dir_crc_expected = read_u32_le(&dir_region[dir_size..dir_size + 4]);
let dir_crc_actual = crc32c(dir_bytes);
if dir_crc_expected != dir_crc_actual {
return Err(FtsError::Read(ReadError::ChecksumMismatch {
section: "fts/doc_lengths_dir",
column: String::new(),
}));
}
}
let mut columns = Vec::with_capacity(n_columns);
let mut column_id_by_name = HashMap::with_capacity(n_columns);
for (i, col_cfg) in cols.iter().enumerate() {
let entry_off = i * DOC_LENGTHS_ENTRY_SIZE;
let column_id = u32::from_le_bytes([
dir_bytes[entry_off],
dir_bytes[entry_off + 1],
dir_bytes[entry_off + 2],
dir_bytes[entry_off + 3],
]);
let doc_lengths_offset =
read_u64_le(&dir_bytes[entry_off + 4..entry_off + 12]) as usize;
let avgdl_x1000 = read_u32_le(&dir_bytes[entry_off + 12..entry_off + 16]) as u64;
if column_id != i as u32 {
return Err(FtsError::Read(ReadError::MalformedVersion(format!(
"doc-lengths directory entry {i} has column_id {column_id}"
))));
}
let array_byte_len = 4 * n_docs as usize;
let array_end = doc_lengths_offset + array_byte_len;
if array_end + 4 > source_len {
return Err(FtsError::Read(ReadError::MalformedVersion(format!(
"doc-lengths array {i} runs past blob end"
))));
}
let array_region = fetch_source_range(
&source,
doc_lengths_offset..array_end + 4,
"fts/doc_lengths_array",
)?;
if opts.verify_crc {
let array_crc_expected =
read_u32_le(&array_region[array_byte_len..array_byte_len + 4]);
let array_crc_actual = crc32c(&array_region[..array_byte_len]);
if array_crc_expected != array_crc_actual {
return Err(FtsError::Read(ReadError::ChecksumMismatch {
section: "fts/doc_lengths_array",
column: format!(" (column '{}')", col_cfg.name),
}));
}
}
let avgdl = (avgdl_x1000 as f32) / format::fts::AVGDL_FIXED_POINT_SCALE;
let n = n_docs as usize;
let dl_norm_k1 = NormTable::new(
(0..n).map(|d| read_u32_le(&array_region[d * 4..d * 4 + 4])),
n,
avgdl,
);
columns.push(ColumnMeta {
name: col_cfg.name.clone(),
doc_lengths_range: doc_lengths_offset..array_end,
avgdl,
dl_norm_k1,
positions: col_cfg.positions,
});
column_id_by_name.insert(col_cfg.name.clone(), i as u32);
}
Ok(FtsReader {
source,
n_docs,
n_terms_total,
fst_range,
postings_range,
positions_range,
columns,
column_id_by_name,
})
}
pub fn n_docs(&self) -> u32 {
self.n_docs
}
pub fn n_terms(&self) -> u32 {
self.n_terms_total
}
pub fn fts_columns(&self) -> impl Iterator<Item = &str> {
self.columns.iter().map(|c| c.name.as_str())
}
pub fn fts_columns_config(&self) -> impl Iterator<Item = &ColumnMeta> {
self.columns.iter()
}
fn dict_bytes(&self) -> Result<Bytes, FtsError> {
fetch_source_range(&self.source, self.fst_range.clone(), "fts/dict")
}
async fn dict_bytes_async(&self) -> Result<Bytes, FtsError> {
self.source
.range_async(self.fst_range.clone())
.await
.map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"fts/dict range fetch failed: {e}"
)))
})
}
async fn fetch_term_postings(&self, terms: &[(usize, usize)]) -> Result<Vec<Bytes>, FtsError> {
if terms.is_empty() {
return Ok(Vec::new());
}
let base = self.postings_range.start;
let region_len = self.postings_range.len();
let mut ranges: Vec<Range<usize>> = Vec::with_capacity(terms.len());
for &(m, postings_length) in terms {
if postings_length < TERM_META_SIZE || m + postings_length > region_len {
return Err(FtsError::Read(ReadError::MalformedVersion(
"term postings range runs past postings region".into(),
)));
}
ranges.push(base + m..base + m + postings_length);
}
let plan = RangeCoalescePlan::new(
&ranges,
TERM_RANGE_COALESCE_MAX_GAP,
TERM_RANGE_COALESCE_MAX_OVERFETCH,
);
let fetched = self
.source
.get_ranges_parallel_async(plan.fetch_ranges())
.await
.map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"fts/postings term body range fetch failed: {e}"
)))
})?;
Ok(plan.restore(&fetched))
}
async fn fetch_term_positions(&self, terms: &[(u64, u32)]) -> Result<Vec<Bytes>, FtsError> {
if terms.iter().all(|&(_, len)| len == 0) {
return Ok(vec![Bytes::new(); terms.len()]);
}
let region = self.positions_range.as_ref().ok_or_else(|| {
FtsError::Read(ReadError::MalformedVersion(
"positional term in a blob with no positions region".into(),
))
})?;
let base = region.start;
let region_len = region.len();
let mut ranges: Vec<Range<usize>> = Vec::with_capacity(terms.len());
for &(off, len) in terms {
let off = off as usize;
let len = len as usize;
if off + len > region_len {
return Err(FtsError::Read(ReadError::MalformedVersion(
"term positions range runs past positions region".into(),
)));
}
ranges.push(base + off..base + off + len);
}
self.source
.get_ranges_parallel_async(&ranges)
.await
.map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"fts/positions term range fetch failed: {e}"
)))
})
}
async fn build_atom_cursors(
&self,
column_id: u32,
terms: &[&str],
phrases: &[Vec<String>],
) -> Result<Vec<Option<AnyCursor>>, FtsError> {
let col_meta = &self.columns[column_id as usize];
if !phrases.is_empty() && !col_meta.positions {
return Err(FtsError::PositionsUnavailable {
column: col_meta.name.clone(),
});
}
let mut out: Vec<Option<AnyCursor>> = Vec::with_capacity(terms.len() + phrases.len());
for term in terms {
let mut cursors = self.build_term_cursors(column_id, &[term]).await?;
out.push(cursors.pop().map(AnyCursor::Term));
}
for phrase in phrases {
let member_refs: Vec<&str> = phrase.iter().map(|t| t.as_str()).collect();
let cursors = self.build_term_cursors(column_id, &member_refs).await?;
if cursors.len() != member_refs.len() {
out.push(None);
continue;
}
let mut positional: Vec<(Option<TermMeta>, Option<u32>)> =
Vec::with_capacity(cursors.len());
for (cursor, term) in cursors.iter().zip(&member_refs) {
match cursor.bytes.is_empty() {
false => {
let term_meta = TermMeta::parse(cursor.bytes.as_ref(), 0, true)?;
positional.push((Some(term_meta), None));
}
true => {
let fst_bytes = self.dict_bytes_async().await?;
let dict = DictReader::open(&fst_bytes).map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"FST parse failed: {e}"
)))
})?;
let key = make_key(&col_meta.name, term);
let packed = dict
.lookup(&key)
.expect("inline member cursor was built from this dict");
let position = match FstValue::unpack(packed) {
FstValue::Inline { tf: slot, .. } => slot,
FstValue::Pfor { .. } => {
unreachable!("inline cursor from a PFOR FST value")
}
};
positional.push((None, Some(position)));
}
}
}
let pos_ranges: Vec<(u64, u32)> = positional
.iter()
.map(|(term_meta, _)| {
term_meta
.map(|tm| (tm.positions_offset, tm.positions_length))
.unwrap_or((0, 0))
})
.collect();
let positions = self.fetch_term_positions(&pos_ranges).await?;
out.push(Some(AnyCursor::Phrase(PhraseCursor::new(
cursors, positions, positional,
)?)));
}
Ok(out)
}
fn run_atoms_search(
&self,
column_id: u32,
mut musts: Vec<AnyCursor>,
mut shoulds: Vec<AnyCursor>,
k: usize,
mut filter: Option<AtomExcludeFilter>,
floor_eff: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
let dl_norm_k1 = &self.columns[column_id as usize].dl_norm_k1;
let initial_cap = k.min(self.n_docs as usize).max(1);
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::with_capacity(initial_cap);
let atom_slack = |atoms: &[AnyCursor], extra_ub: f32| -> Vec<f32> {
let total: f32 = atoms.iter().map(AnyCursor::term_max_bm25).sum();
atoms
.iter()
.map(|a| total - a.term_max_bm25() + extra_ub)
.collect()
};
if musts.is_empty() {
let others_ub = atom_slack(&shoulds, 0.0);
while let Some(doc) = shoulds
.iter()
.filter(|a| !a.is_exhausted())
.map(AnyCursor::current_doc_id)
.min()
{
let admitted = match filter.as_mut() {
Some(f) => f.admits(doc)?,
None => true,
};
if admitted {
let norm = dl_norm_k1.get(doc);
let score: f32 = shoulds
.iter()
.filter(|a| !a.is_exhausted() && a.current_doc_id() == doc)
.map(|a| a.score_current(norm))
.sum();
if score > floor_eff {
and_heap_push(&mut heap, k, None, score, doc);
}
}
let Some(next) = doc.checked_add(1) else {
break;
};
let bar = match heap.len() >= k {
true => heap.peek().expect("heap len == k").0.max(floor_eff),
false => floor_eff,
};
for (a, &others) in shoulds.iter_mut().zip(&others_ub) {
if !a.is_exhausted() && a.current_doc_id() == doc {
a.skip_to_pruned(next, bar - others, dl_norm_k1)?;
}
}
}
return Ok(drain_top_k_desc(heap));
}
let should_ub: f32 = shoulds.iter().map(AnyCursor::term_max_bm25).sum();
let must_others_ub = atom_slack(&musts, should_ub);
let should_others_ub: Vec<f32> = {
let must_ub_total: f32 = musts.iter().map(AnyCursor::term_max_bm25).sum();
atom_slack(&shoulds, must_ub_total)
};
let mut target = 0u32;
'docs: loop {
let bar = match heap.len() >= k {
true => heap.peek().expect("heap len == k").0.max(floor_eff),
false => floor_eff,
};
let mut aligned = target;
let mut i = 0usize;
while i < musts.len() {
let a = &mut musts[i];
a.skip_to_pruned(aligned, bar - must_others_ub[i], dl_norm_k1)?;
if a.is_exhausted() {
break 'docs;
}
let here = a.current_doc_id();
if here > aligned {
aligned = here;
i = 0;
continue;
}
i += 1;
}
let scoring_needed = match bar > f32::NEG_INFINITY {
true => {
let must_ub: f32 = musts
.iter_mut()
.map(|a| a.block_max_in_range(aligned, aligned))
.sum();
must_ub + should_ub >= bar
}
false => true,
};
let admitted = scoring_needed
&& match filter.as_mut() {
Some(f) => f.admits(aligned)?,
None => true,
};
if admitted {
let norm = dl_norm_k1.get(aligned);
let mut score: f32 = musts.iter().map(|a| a.score_current(norm)).sum();
for (sh, &others) in shoulds.iter_mut().zip(&should_others_ub) {
sh.skip_to_pruned(aligned, bar - others, dl_norm_k1)?;
if !sh.is_exhausted() && sh.current_doc_id() == aligned {
score += sh.score_current(norm);
}
}
if score > floor_eff {
and_heap_push(&mut heap, k, None, score, aligned);
}
}
let Some(next) = aligned.checked_add(1) else {
break;
};
target = next;
}
Ok(drain_top_k_desc(heap))
}
fn walk_atoms_match(
&self,
mut atoms: Vec<AnyCursor>,
mode: BoolMode,
mut filter: Option<AtomExcludeFilter>,
mut on_doc: impl FnMut(u32),
) -> Result<(), FtsError> {
match mode {
BoolMode::Or => {
while let Some(doc) = atoms
.iter()
.filter(|a| !a.is_exhausted())
.map(AnyCursor::current_doc_id)
.min()
{
let admitted = match filter.as_mut() {
Some(f) => f.admits(doc)?,
None => true,
};
if admitted {
on_doc(doc);
}
let Some(next) = doc.checked_add(1) else {
break;
};
for a in atoms.iter_mut() {
if !a.is_exhausted() && a.current_doc_id() == doc {
a.skip_to(next)?;
}
}
}
Ok(())
}
BoolMode::And => {
let mut target = 0u32;
'docs: loop {
let mut aligned = target;
let mut i = 0usize;
while i < atoms.len() {
let a = &mut atoms[i];
a.skip_to(aligned)?;
if a.is_exhausted() {
break 'docs;
}
let here = a.current_doc_id();
if here > aligned {
aligned = here;
i = 0;
continue;
}
i += 1;
}
let admitted = match filter.as_mut() {
Some(f) => f.admits(aligned)?,
None => true,
};
if admitted {
on_doc(aligned);
}
let Some(next) = aligned.checked_add(1) else {
break;
};
target = next;
}
Ok(())
}
}
}
pub(crate) async fn atoms_match_ids(
&self,
column: &str,
terms: &[&str],
phrases: &[Vec<String>],
mode: BoolMode,
) -> Result<Vec<u32>, FtsError> {
let column_id = self.resolve_column_id(column)?;
let built = self.build_atom_cursors(column_id, terms, phrases).await?;
let atoms: Vec<AnyCursor> = match mode {
BoolMode::And => {
if built.iter().any(Option::is_none) {
return Ok(Vec::new());
}
built.into_iter().flatten().collect()
}
BoolMode::Or => built.into_iter().flatten().collect(),
};
if atoms.is_empty() {
return Ok(Vec::new());
}
let mut out = Vec::new();
self.walk_atoms_match(atoms, mode, None, |d| out.push(d))?;
Ok(out)
}
pub(crate) async fn atoms_match_count(
&self,
column: &str,
terms: &[&str],
phrases: &[Vec<String>],
mode: BoolMode,
) -> Result<u64, FtsError> {
let column_id = self.resolve_column_id(column)?;
let built = self.build_atom_cursors(column_id, terms, phrases).await?;
let atoms: Vec<AnyCursor> = match mode {
BoolMode::And => {
if built.iter().any(Option::is_none) {
return Ok(0);
}
built.into_iter().flatten().collect()
}
BoolMode::Or => built.into_iter().flatten().collect(),
};
if atoms.is_empty() {
return Ok(0);
}
let mut n = 0u64;
self.walk_atoms_match(atoms, mode, None, |_| n += 1)?;
Ok(n)
}
fn resolve_column_id(&self, column: &str) -> Result<u32, FtsError> {
self.column_id_by_name
.get(column)
.copied()
.ok_or_else(|| FtsError::UnknownColumn(column.to_string()))
}
pub fn iter_column_terms(&self, column: &str) -> Result<Vec<Vec<u8>>, FtsError> {
self.iter_terms_with_prefix(column, b"")
}
pub fn iter_terms_with_prefix(
&self,
column: &str,
term_prefix: &[u8],
) -> Result<Vec<Vec<u8>>, FtsError> {
if !self.column_id_by_name.contains_key(column) {
return Ok(Vec::new());
}
let mut full_prefix = column.as_bytes().to_vec();
full_prefix.push(FST_SEPARATOR);
let column_prefix_len = full_prefix.len();
full_prefix.extend_from_slice(term_prefix);
let fst_bytes = self
.dict_bytes()
.expect("FST bytes must be available for term iteration");
let dict = DictReader::open(&fst_bytes).map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"FST parse failed: {e}"
)))
})?;
let pairs = dict.iter_prefix(&full_prefix);
Ok(pairs
.into_iter()
.map(|(key, _)| key[column_prefix_len..].to_vec())
.collect())
}
pub async fn search(
&self,
column: &str,
terms: &[&str],
k: usize,
mode: BoolMode,
) -> Result<Vec<(u32, f32)>, FtsError> {
self.search_with_floor(column, terms, k, mode, f32::NEG_INFINITY)
.await
}
pub async fn search_with_floor(
&self,
column: &str,
terms: &[&str],
k: usize,
mode: BoolMode,
floor: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
let column_id = self.resolve_column_id(column)?;
if terms.is_empty() || k == 0 {
return Ok(Vec::new());
}
let floor_eff = floor.next_down();
let (musts, shoulds): (&[&str], &[&str]) = match mode {
BoolMode::And => (terms, &[]),
BoolMode::Or => (&[], terms),
};
self.search_clauses(column_id, musts, shoulds, k, None, floor_eff)
.await
}
pub(crate) async fn search_excluding(
&self,
column: &str,
lists: ClauseLists<'_>,
k: usize,
floor: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
let column_id = self.resolve_column_id(column)?;
if k == 0 {
return Ok(Vec::new());
}
if lists.no_positive_atoms() {
if lists.no_negative_atoms() {
return Ok(Vec::new());
}
return Err(FtsError::NegationOnly);
}
let floor_eff = floor.next_down();
if lists.has_phrases() {
let must_atoms = self
.build_atom_cursors(column_id, lists.musts, lists.must_phrases)
.await?;
if must_atoms.iter().any(Option::is_none) {
return Ok(Vec::new());
}
let must_atoms: Vec<AnyCursor> = must_atoms.into_iter().flatten().collect();
let should_atoms: Vec<AnyCursor> = self
.build_atom_cursors(column_id, lists.shoulds, lists.should_phrases)
.await?
.into_iter()
.flatten()
.collect();
let negative_atoms: Vec<AnyCursor> = self
.build_atom_cursors(column_id, lists.negatives, lists.negative_phrases)
.await?
.into_iter()
.flatten()
.collect();
let filter = match negative_atoms.is_empty() {
true => None,
false => Some(AtomExcludeFilter::new(negative_atoms)),
};
return self.run_atoms_search(
column_id,
must_atoms,
should_atoms,
k,
filter,
floor_eff,
);
}
let mut filter = match lists.negatives {
[] => None,
_ => Some(ExcludeFilter::new(
self.build_term_cursors(column_id, lists.negatives).await?,
)),
};
self.search_clauses(
column_id,
lists.musts,
lists.shoulds,
k,
filter.as_mut(),
floor_eff,
)
.await
}
async fn search_clauses(
&self,
column_id: u32,
musts: &[&str],
shoulds: &[&str],
k: usize,
filter: Option<&mut ExcludeFilter>,
floor_eff: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
if musts.len() + shoulds.len() == 1 {
let term = musts.iter().chain(shoulds).next().expect("one atom");
return self
.search_single_term_bmw(column_id, term, k, filter, floor_eff)
.await;
}
if musts.is_empty() {
return self
.dispatch_multi_term_or(column_id, shoulds, k, filter, floor_eff)
.await;
}
let must_cursors = self.build_term_cursors(column_id, musts).await?;
if must_cursors.len() != musts.len() {
return Ok(Vec::new());
}
if shoulds.is_empty() {
return self.run_and_intersect(column_id, must_cursors, k, filter, floor_eff);
}
let should_cursors = self.build_term_cursors(column_id, shoulds).await?;
if should_cursors.is_empty() {
return self.run_and_intersect(column_id, must_cursors, k, filter, floor_eff);
}
self.run_must_should(
column_id,
must_cursors,
should_cursors,
k,
filter,
floor_eff,
)
}
pub async fn token_match(
&self,
column: &str,
tokens: &[&str],
mode: BoolMode,
) -> Result<Vec<u32>, FtsError> {
let column_id = self.resolve_column_id(column)?;
if tokens.is_empty() {
return Ok(Vec::new());
}
let cursors = self.build_term_cursors(column_id, tokens).await?;
Ok(match mode {
BoolMode::And => {
if cursors.len() != tokens.len() {
return Ok(Vec::new());
}
self.collect_and_intersect(column_id, cursors)
}
BoolMode::Or => or_merge_unranked(cursors),
})
}
pub async fn token_match_count(
&self,
column: &str,
tokens: &[&str],
mode: BoolMode,
) -> Result<u64, FtsError> {
let column_id = self.resolve_column_id(column)?;
if tokens.is_empty() {
return Ok(0);
}
let cursors = self.build_term_cursors(column_id, tokens).await?;
Ok(match mode {
BoolMode::And => {
if cursors.len() != tokens.len() {
return Ok(0);
}
self.count_and_intersect(column_id, cursors)
}
BoolMode::Or => or_count_unranked(cursors),
})
}
pub async fn term_df(&self, column: &str, token: &str) -> Result<u64, FtsError> {
let column_id = self.resolve_column_id(column)?;
let fst_bytes = self.dict_bytes_async().await?;
let dict = DictReader::open(&fst_bytes).map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"FST parse failed: {e}"
)))
})?;
let col_meta = &self.columns[column_id as usize];
let key = make_key(&col_meta.name, token);
Ok(match dict.lookup(&key) {
None => 0,
Some(packed) => match FstValue::unpack(packed) {
FstValue::Inline { .. } => 1,
FstValue::Pfor {
metadata_offset, ..
} => {
let fetched = self
.fetch_term_postings(&[(metadata_offset as usize, TERM_META_SIZE)])
.await?;
let header = fetched.first().expect("one fetched header range");
read_u32_le(&header.as_ref()[0..4]) as u64
}
},
})
}
pub async fn search_or_range_pretokenized(
&self,
column: &str,
terms: &[&str],
k: usize,
doc_id_start: u32,
doc_id_end: u32,
) -> Result<Vec<(u32, f32)>, FtsError> {
self.search_or_range_pretokenized_with_floor(
column,
terms,
k,
doc_id_start,
doc_id_end,
f32::NEG_INFINITY,
)
.await
}
pub async fn search_or_range_pretokenized_with_floor(
&self,
column: &str,
terms: &[&str],
k: usize,
doc_id_start: u32,
doc_id_end: u32,
floor: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
let column_id = self.resolve_column_id(column)?;
if terms.is_empty() || k == 0 || doc_id_start >= doc_id_end {
return Ok(Vec::new());
}
let cursors = self.build_term_cursors(column_id, terms).await?;
if cursors.is_empty() {
return Ok(Vec::new());
}
self.run_max_score_bmm_range(
column_id,
cursors,
k,
doc_id_start,
doc_id_end,
None,
floor.next_down(),
)
}
pub async fn search_multi(
&self,
columns: &[(&str, f32)],
query: &str,
k: usize,
mode: BoolMode,
) -> Result<Vec<(u32, f32)>, FtsError> {
let tok = AsciiLowerTokenizer;
let term_strings: Vec<String> = tok.tokenize(query).collect();
let term_refs: Vec<&str> = term_strings.iter().map(|s| s.as_str()).collect();
let mut combined: HashMap<u32, f32> = HashMap::new();
for (col_name, weight) in columns {
let per_col = self.search(col_name, &term_refs, usize::MAX, mode).await?;
for (doc_id, s) in per_col {
*combined.entry(doc_id).or_insert(0.0) += s * weight;
}
}
Ok(top_k(combined, k))
}
async fn search_single_term_bmw(
&self,
column_id: u32,
term: &str,
k: usize,
mut filter: Option<&mut ExcludeFilter>,
floor_eff: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
let fst_bytes = self.dict_bytes_async().await?;
let dict = DictReader::open(&fst_bytes).map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"FST parse failed: {e}"
)))
})?;
let col_meta = &self.columns[column_id as usize];
let key = make_key(&col_meta.name, term);
let Some(packed) = dict.lookup(&key) else {
return Ok(Vec::new());
};
let (metadata_offset, postings_length) = match FstValue::unpack(packed) {
FstValue::Inline { doc_id, tf } => {
let tf = match col_meta.positions {
true => 1,
false => tf,
};
let idf_t = bm25::idf(self.n_docs as u64, 1);
let idf_x_k1p1 = idf_t * (bm25::K1 + 1.0);
if let Some(f) = filter.as_deref_mut()
&& !f.admits(doc_id)
{
return Ok(Vec::new());
}
let dl_norm_k1 = col_meta.dl_norm_k1.get(doc_id);
let score = bm25::score_with_dl_norm_k1(idf_x_k1p1, tf, dl_norm_k1);
if score <= floor_eff {
return Ok(Vec::new());
}
return Ok(vec![(doc_id, score)]);
}
FstValue::Pfor {
metadata_offset,
postings_length,
} => (metadata_offset as usize, postings_length as usize),
};
let term_bytes = {
let mut fetched = self
.fetch_term_postings(&[(metadata_offset, postings_length)])
.await?;
fetched.pop().expect("one fetched range for one PFOR term")
};
let postings = term_bytes.as_ref();
let metadata_offset = 0usize;
let term_meta = TermMeta::parse(postings, metadata_offset, col_meta.positions)?;
let idf_t = bm25::idf(self.n_docs as u64, term_meta.df);
let idf_x_k1p1 = idf_t * (bm25::K1 + 1.0);
let dl_norm_k1 = &col_meta.dl_norm_k1;
let mut heap: BinaryHeap<TopKEntry> =
BinaryHeap::with_capacity(k.min(term_meta.num_blocks * BLOCK_LEN).max(1));
let mut buf_d = vec![0u32; BLOCK_LEN];
let mut buf_t = vec![0u32; BLOCK_LEN];
for i in 0..term_meta.num_blocks {
let (_, block_offset_in_term, block_max_bm25) = term_meta.skip_entry(postings, i);
if block_max_bm25 <= floor_eff {
continue;
}
if heap.len() >= k
&& let Some(TopKEntry(min_score, _)) = heap.peek()
&& block_max_bm25 <= *min_score
{
continue;
}
let block_end_in_term = term_meta.block_end_in_term(postings, i);
let block_bytes = &postings
[metadata_offset + block_offset_in_term..metadata_offset + block_end_in_term];
let n = decode_block(block_bytes, &mut buf_d, &mut buf_t);
for j in 0..n {
let doc_id = buf_d[j];
if let Some(f) = filter.as_deref_mut()
&& !f.admits(doc_id)
{
continue;
}
let tf = buf_t[j];
let score = bm25::score_with_dl_norm_k1(idf_x_k1p1, tf, dl_norm_k1.get(doc_id));
if score <= floor_eff {
continue;
}
if heap.len() < k {
heap.push(TopKEntry(score, doc_id));
} else if let Some(TopKEntry(min_score, _)) = heap.peek()
&& score > *min_score
{
heap.pop();
heap.push(TopKEntry(score, doc_id));
}
}
}
Ok(drain_top_k_desc(heap))
}
async fn build_term_cursors(
&self,
column_id: u32,
terms: &[&str],
) -> Result<Vec<TermCursor>, FtsError> {
let fst_bytes = self.dict_bytes_async().await?;
let dict = DictReader::open(&fst_bytes).map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"FST parse failed: {e}"
)))
})?;
let col_meta = &self.columns[column_id as usize];
enum Resolved {
Inline { doc_id: u32, tf: u32 },
Pfor,
}
let mut resolved: Vec<Resolved> = Vec::with_capacity(terms.len());
let mut pfor_offsets: Vec<(usize, usize)> = Vec::new();
for term in terms {
let key = make_key(&col_meta.name, term);
let Some(packed) = dict.lookup(&key) else {
continue;
};
match FstValue::unpack(packed) {
FstValue::Inline { doc_id, tf } => {
resolved.push(Resolved::Inline { doc_id, tf });
}
FstValue::Pfor {
metadata_offset,
postings_length,
} => {
pfor_offsets.push((metadata_offset as usize, postings_length as usize));
resolved.push(Resolved::Pfor);
}
}
}
let pfor_bytes = self.fetch_term_postings(&pfor_offsets).await?;
let mut pfor_iter = pfor_bytes.into_iter();
let mut cursors: Vec<TermCursor> = Vec::with_capacity(resolved.len());
for r in resolved {
match r {
Resolved::Inline { doc_id, tf } => {
let tf = match col_meta.positions {
true => 1,
false => tf,
};
let dl_norm_k1 = col_meta.dl_norm_k1.get(doc_id);
cursors.push(TermCursor::new_inline(
doc_id,
tf,
self.n_docs as u64,
dl_norm_k1,
));
}
Resolved::Pfor => {
let term_bytes = pfor_iter.next().expect("one fetched range per PFOR term");
cursors.push(TermCursor::new(
term_bytes,
self.n_docs as u64,
col_meta.positions,
)?);
}
}
}
Ok(cursors)
}
fn run_wand_bmw(
&self,
column_id: u32,
mut cursors: Vec<TermCursor>,
k: usize,
) -> Result<Vec<(u32, f32)>, FtsError> {
let col_meta = &self.columns[column_id as usize];
let dl_norm_k1 = &col_meta.dl_norm_k1;
let initial_cap = k.min(self.n_docs as usize).max(1);
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::with_capacity(initial_cap);
let mut threshold: f32 = 0.0;
let mut idx: Vec<usize> = Vec::with_capacity(cursors.len());
loop {
cursors.retain(|c| !c.is_exhausted());
if cursors.is_empty() {
break;
}
idx.clear();
idx.extend(0..cursors.len());
idx.sort_unstable_by_key(|&i| cursors[i].current_doc_id());
let mut accum_term_ub: f32 = 0.0;
let mut pivot_j: Option<usize> = None;
for (j, &ci) in idx.iter().enumerate() {
accum_term_ub += cursors[ci].term_max_bm25;
if accum_term_ub > threshold {
pivot_j = Some(j);
break;
}
}
let Some(mut pivot_j) = pivot_j else {
break;
};
let pivot_doc = cursors[idx[pivot_j]].current_doc_id();
while pivot_j + 1 < idx.len() && cursors[idx[pivot_j + 1]].current_doc_id() == pivot_doc
{
pivot_j += 1;
}
let mut accum_block_ub: f32 = 0.0;
for &ci in &idx[..=pivot_j] {
cursors[ci].shallow_advance_block_to(pivot_doc);
accum_block_ub += cursors[ci].inspect_block_max_bm25();
}
if accum_block_ub <= threshold {
let mut target = u32::MAX;
for &ci in &idx[..=pivot_j] {
let last = cursors[ci].inspect_block_last_doc_id();
if last < target {
target = last;
}
}
let mut effective_target = target.saturating_add(1);
for &ci in &idx[pivot_j + 1..] {
let d = cursors[ci].current_doc_id();
if d < effective_target {
effective_target = d;
}
}
cursors[idx[0]].skip_to(effective_target);
continue;
}
let mut aligned = true;
for &ci in &idx[..=pivot_j] {
if cursors[ci].current_doc_id() < pivot_doc {
cursors[ci].skip_to(pivot_doc);
if cursors[ci].current_doc_id() != pivot_doc {
aligned = false;
break;
}
}
}
if !aligned {
continue;
}
let norm = dl_norm_k1.get(pivot_doc);
let mut score: f32 = 0.0;
let mut idfs = [0.0_f32; 4];
let mut tfs = [0.0_f32; 4];
let mut packed = 0;
for cursor in &cursors {
if cursor.current_doc_id() == pivot_doc {
idfs[packed] = cursor.idf_x_k1p1;
tfs[packed] = cursor.current_tf() as f32;
packed += 1;
if packed == 4 {
score += bm25::score_simd_x4(idfs, tfs, norm);
idfs = [0.0; 4];
tfs = [0.0; 4];
packed = 0;
}
}
}
if packed > 0 {
score += bm25::score_simd_x4(idfs, tfs, norm);
}
if heap.len() < k {
heap.push(TopKEntry(score, pivot_doc));
if heap.len() == k {
threshold = heap.peek().expect("non-empty").0;
}
} else if let Some(TopKEntry(min_score, _)) = heap.peek()
&& score > *min_score
{
heap.pop();
heap.push(TopKEntry(score, pivot_doc));
threshold = heap.peek().expect("non-empty").0;
}
for cursor in cursors.iter_mut() {
if cursor.current_doc_id() == pivot_doc {
cursor.next();
}
}
}
Ok(drain_top_k_desc(heap))
}
fn run_max_score_bmm(
&self,
column_id: u32,
cursors: Vec<TermCursor>,
k: usize,
filter: Option<&mut ExcludeFilter>,
floor_eff: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
self.run_max_score_bmm_range(column_id, cursors, k, 0, u32::MAX, filter, floor_eff)
}
fn run_and_intersect(
&self,
column_id: u32,
mut cursors: Vec<TermCursor>,
k: usize,
filter: Option<&mut ExcludeFilter>,
floor_eff: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
if cursors.is_empty() {
return Ok(Vec::new());
}
let col_meta = &self.columns[column_id as usize];
let dl_norm_k1 = &col_meta.dl_norm_k1;
cursors.sort_by_key(|c| c.block_count());
let initial_cap = k.min(self.n_docs as usize).max(1);
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::with_capacity(initial_cap);
let mut sink = ScoreSink {
heap: &mut heap,
k,
filter,
floor_eff,
};
self.and_flat_merge(&mut cursors, dl_norm_k1, &mut sink);
Ok(drain_top_k_desc(heap))
}
fn run_must_should(
&self,
column_id: u32,
mut must_cursors: Vec<TermCursor>,
should_cursors: Vec<TermCursor>,
k: usize,
filter: Option<&mut ExcludeFilter>,
floor_eff: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
debug_assert!(
!must_cursors.is_empty() && !should_cursors.is_empty(),
"dispatch routes empty-side shapes to the AND/OR kernels"
);
let col_meta = &self.columns[column_id as usize];
let dl_norm_k1 = &col_meta.dl_norm_k1;
must_cursors.sort_by_key(|c| c.block_count());
let initial_cap = k.min(self.n_docs as usize).max(1);
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::with_capacity(initial_cap);
let should_ub = should_cursors.iter().map(|c| c.term_max_bm25).sum();
let mut sink = MustShouldSink {
heap: &mut heap,
k,
filter,
floor_eff,
shoulds: should_cursors,
should_ub,
dl_norm_k1,
};
self.and_flat_merge(&mut must_cursors, dl_norm_k1, &mut sink);
Ok(drain_top_k_desc(heap))
}
fn collect_and_intersect(&self, column_id: u32, mut cursors: Vec<TermCursor>) -> Vec<u32> {
if cursors.is_empty() {
return Vec::new();
}
let col_meta = &self.columns[column_id as usize];
let dl_norm_k1 = &col_meta.dl_norm_k1;
cursors.sort_by_key(|c| c.block_count());
let mut sink = CollectSink { out: Vec::new() };
self.and_flat_merge(&mut cursors, dl_norm_k1, &mut sink);
sink.out
}
fn count_and_intersect(&self, column_id: u32, mut cursors: Vec<TermCursor>) -> u64 {
if cursors.is_empty() {
return 0;
}
let col_meta = &self.columns[column_id as usize];
let dl_norm_k1 = &col_meta.dl_norm_k1;
cursors.sort_by_key(|c| c.block_count());
let mut sink = CountSink { n: 0 };
self.and_flat_merge(&mut cursors, dl_norm_k1, &mut sink);
sink.n
}
fn and_flat_merge<S: AndSink>(
&self,
cursors: &mut [TermCursor],
dl_norm_k1: &NormTable,
sink: &mut S,
) {
if cursors.len() == 2 {
self.and_flat_merge_2term(cursors, dl_norm_k1, sink);
} else {
self.and_flat_merge_general(cursors, dl_norm_k1, sink);
}
}
fn and_flat_merge_general<S: AndSink>(
&self,
cursors: &mut [TermCursor],
dl_norm_k1: &NormTable,
sink: &mut S,
) {
'outer: loop {
if cursors[0].is_exhausted() {
break;
}
let bar = sink.bar();
if bar > f32::NEG_INFINITY {
let range_start = cursors[0].current_doc_id();
let range_end = cursors[0].current_block_last_doc_id();
let leader_block_max = cursors[0].current_block_max_bm25();
let mut other_ub = 0.0_f32;
for c in cursors[1..].iter_mut() {
other_ub += c.block_max_in_range(range_start, range_end);
}
if leader_block_max + other_ub <= bar {
cursors[0].skip_to(range_end.saturating_add(1));
continue;
}
}
let leader_doc = cursors[0].current_doc_id();
let leader_block_end = cursors[0].current_block_last_doc_id();
let mut max_other = leader_doc;
let mut crossed_block = false;
for c in cursors[1..].iter_mut() {
c.skip_to(leader_doc);
if c.is_exhausted() {
break 'outer;
}
let here = c.current_doc_id();
if here > leader_block_end {
crossed_block = true;
}
if here > max_other {
max_other = here;
}
}
if max_other > leader_doc {
cursors[0].skip_to(max_other);
if cursors[0].is_exhausted() {
break 'outer;
}
if crossed_block {
continue;
}
}
let (leader_slice, others) = cursors.split_at_mut(1);
let c0 = &mut leader_slice[0];
let lb_n = c0.block_n;
let mut i = c0.pos;
while i < lb_n {
let a = c0.block_doc_ids[i];
let mut block_exhausted = false;
let mut all_match = true;
for o in others.iter_mut() {
while o.pos < o.block_n && o.block_doc_ids[o.pos] < a {
o.pos += 1;
}
if o.pos >= o.block_n {
block_exhausted = true;
break;
}
if o.block_doc_ids[o.pos] != a {
all_match = false;
break;
}
}
if block_exhausted {
break;
}
if all_match {
let score = if sink.needs_score() {
let norm = dl_norm_k1.get(a);
let mut score =
bm25::score_with_dl_norm_k1(c0.idf_x_k1p1, c0.block_tfs[i], norm);
for o in others.iter() {
score +=
bm25::score_with_dl_norm_k1(o.idf_x_k1p1, o.block_tfs[o.pos], norm);
}
score
} else {
0.0
};
sink.emit(a, score);
i += 1;
for o in others.iter_mut() {
o.pos += 1;
}
} else {
i += 1;
}
}
c0.pos = i;
if c0.pos >= c0.block_n {
c0.next();
}
for o in others.iter_mut() {
if o.pos >= o.block_n {
o.next();
}
}
}
}
fn and_flat_merge_2term<S: AndSink>(
&self,
cursors: &mut [TermCursor],
dl_norm_k1: &NormTable,
sink: &mut S,
) {
debug_assert_eq!(cursors.len(), 2);
let (left, right) = cursors.split_at_mut(1);
let c0 = &mut left[0];
let c1 = &mut right[0];
'outer: loop {
if c0.is_exhausted() || c1.is_exhausted() {
break;
}
let bar = sink.bar();
if bar > f32::NEG_INFINITY {
let range_start = c0.current_doc_id();
let range_end = c0.current_block_last_doc_id();
let ub =
c0.current_block_max_bm25() + c1.block_max_in_range(range_start, range_end);
if ub <= bar {
c0.skip_to(range_end.saturating_add(1));
continue;
}
}
c1.skip_to(c0.current_doc_id());
if c1.is_exhausted() {
break 'outer;
}
if c1.current_doc_id() > c0.current_doc_id() {
let crossed_block = c1.current_doc_id() > c0.current_block_last_doc_id();
c0.skip_to(c1.current_doc_id());
if c0.is_exhausted() {
break 'outer;
}
if crossed_block {
continue;
}
}
let lb_n = c0.block_n;
let rb_n = c1.block_n;
let mut i = c0.pos;
let mut j = c1.pos;
let c0_idf = c0.idf_x_k1p1;
let c1_idf = c1.idf_x_k1p1;
while i < lb_n && j < rb_n {
let a = c0.block_doc_ids[i];
let b = c1.block_doc_ids[j];
if a < b {
i += 1;
} else if a > b {
j += 1;
} else {
let score = if sink.needs_score() {
let norm = dl_norm_k1.get(a);
bm25::score_with_dl_norm_k1(c0_idf, c0.block_tfs[i], norm)
+ bm25::score_with_dl_norm_k1(c1_idf, c1.block_tfs[j], norm)
} else {
0.0
};
sink.emit(a, score);
i += 1;
j += 1;
}
}
c0.pos = i;
c1.pos = j;
if i >= lb_n {
c0.next();
}
if j >= rb_n {
c1.next();
}
}
}
fn run_max_score_bmm_range(
&self,
column_id: u32,
mut cursors: Vec<TermCursor>,
k: usize,
doc_id_start: u32,
doc_id_end: u32,
mut filter: Option<&mut ExcludeFilter>,
floor_eff: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
let col_meta = &self.columns[column_id as usize];
let dl_norm_k1 = &col_meta.dl_norm_k1;
if doc_id_start > 0 {
for cursor in &mut cursors {
cursor.skip_to(doc_id_start);
}
}
cursors.sort_unstable_by(|a, b| {
b.term_max_bm25
.partial_cmp(&a.term_max_bm25)
.unwrap_or(Ordering::Equal)
});
let n = cursors.len();
let mut partial_max = vec![0.0_f32; n + 1];
for i in (0..n).rev() {
partial_max[i] = partial_max[i + 1] + cursors[i].term_max_bm25;
}
let initial_cap = k.min(self.n_docs as usize).max(1);
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::with_capacity(initial_cap);
let mut threshold: f32 = floor_eff.max(0.0);
let recompute_f = |partial_max: &[f32], threshold: f32| -> usize {
let mut f = 0;
while f < partial_max.len() - 1 && partial_max[f] > threshold {
f += 1;
}
f
};
let mut f_essential: usize = recompute_f(&partial_max, threshold);
let total_term_ub = partial_max[0];
loop {
if f_essential == 1 {
if cursors[0].is_exhausted() || cursors[0].current_doc_id() >= doc_id_end {
break;
}
let block_ub = cursors[0].current_block_max_bm25()
+ (total_term_ub - cursors[0].term_max_bm25);
if block_ub <= threshold {
let end = cursors[0].current_block_last_doc_id();
cursors[0].skip_to(end.saturating_add(1));
continue;
}
let block_end = cursors[0].current_block_last_doc_id();
let mut f_changed = false;
let others_term_ub = total_term_ub - cursors[0].term_max_bm25;
while !cursors[0].is_exhausted()
&& cursors[0].current_doc_id() <= block_end
&& cursors[0].current_doc_id() < doc_id_end
{
let candidate = cursors[0].current_doc_id();
if let Some(f) = filter.as_deref_mut()
&& !f.admits(candidate)
{
cursors[0].next();
continue;
}
let norm = dl_norm_k1.get(candidate);
let essential_score = bm25::score_with_dl_norm_k1(
cursors[0].idf_x_k1p1,
cursors[0].current_tf(),
norm,
);
if essential_score + others_term_ub <= threshold {
cursors[0].next();
continue;
}
let mut idfs = [cursors[0].idf_x_k1p1, 0.0, 0.0, 0.0];
let mut tfs = [cursors[0].current_tf() as f32, 0.0, 0.0, 0.0];
let mut packed = 1;
let mut score: f32 = 0.0;
for cursor in cursors.iter_mut().skip(1) {
cursor.skip_to(candidate);
if cursor.current_doc_id() == candidate {
idfs[packed] = cursor.idf_x_k1p1;
tfs[packed] = cursor.current_tf() as f32;
packed += 1;
if packed == 4 {
score += bm25::score_simd_x4(idfs, tfs, norm);
idfs = [0.0; 4];
tfs = [0.0; 4];
packed = 0;
}
}
}
if packed > 0 {
score += bm25::score_simd_x4(idfs, tfs, norm);
}
if heap.len() < k {
heap.push(TopKEntry(score, candidate));
if heap.len() == k {
threshold = heap.peek().expect("non-empty").0.max(threshold);
let new_f = recompute_f(&partial_max, threshold);
if new_f != f_essential {
f_essential = new_f;
f_changed = true;
}
}
} else if score > threshold {
heap.pop();
heap.push(TopKEntry(score, candidate));
threshold = heap.peek().expect("non-empty").0.max(threshold);
let new_f = recompute_f(&partial_max, threshold);
if new_f != f_essential {
f_essential = new_f;
f_changed = true;
}
}
cursors[0].next();
if f_changed {
break;
}
}
continue;
}
let (candidate, leftmost_essential) = if f_essential == 2 {
let d0 = cursors[0].current_doc_id();
let d1 = cursors[1].current_doc_id();
if d0 == u32::MAX && d1 == u32::MAX {
break;
}
if d0 <= d1 { (d0, 0) } else { (d1, 1) }
} else {
let mut candidate = u32::MAX;
let mut leftmost_essential: usize = 0;
for (i, cursor) in cursors.iter().take(f_essential).enumerate() {
let d = cursor.current_doc_id();
if d < candidate {
candidate = d;
leftmost_essential = i;
}
}
if candidate == u32::MAX {
break;
}
(candidate, leftmost_essential)
};
if candidate >= doc_id_end {
break;
}
let leftmost_term_ub = cursors[leftmost_essential].term_max_bm25;
let leftmost_block_ub = cursors[leftmost_essential].current_block_max_bm25();
let others_ub = total_term_ub - leftmost_term_ub;
if leftmost_block_ub + others_ub <= threshold {
let last_in_block = cursors[leftmost_essential].current_block_last_doc_id();
cursors[leftmost_essential].skip_to(last_in_block.saturating_add(1));
continue;
}
let admitted = match filter.as_deref_mut() {
Some(f) => f.admits(candidate),
None => true,
};
if admitted {
let norm = dl_norm_k1.get(candidate);
let mut score: f32 = 0.0;
let mut idfs = [0.0_f32; 4];
let mut tfs = [0.0_f32; 4];
let mut packed = 0;
for cursor in cursors.iter().take(f_essential) {
if cursor.current_doc_id() == candidate {
idfs[packed] = cursor.idf_x_k1p1;
tfs[packed] = cursor.current_tf() as f32;
packed += 1;
if packed == 4 {
score += bm25::score_simd_x4(idfs, tfs, norm);
idfs = [0.0; 4];
tfs = [0.0; 4];
packed = 0;
}
}
}
if packed > 0 {
score += bm25::score_simd_x4(idfs, tfs, norm);
}
let non_essentials_term_ub = partial_max[f_essential];
if score + non_essentials_term_ub > threshold {
let mut remaining_block_ub: f32 = 0.0;
for cursor in cursors.iter_mut().skip(f_essential) {
cursor.shallow_advance_block_to(candidate);
remaining_block_ub += cursor.inspect_block_max_bm25();
}
if score + remaining_block_ub > threshold {
for cursor in cursors.iter_mut().skip(f_essential) {
let block_ub = cursor.inspect_block_max_bm25();
if score + remaining_block_ub <= threshold {
break;
}
cursor.skip_to(candidate);
if cursor.current_doc_id() == candidate {
score += bm25::score_with_dl_norm_k1(
cursor.idf_x_k1p1,
cursor.current_tf(),
norm,
);
}
remaining_block_ub -= block_ub;
}
}
}
if heap.len() < k {
heap.push(TopKEntry(score, candidate));
if heap.len() == k {
threshold = heap.peek().expect("non-empty").0.max(threshold);
f_essential = recompute_f(&partial_max, threshold);
}
} else if score > threshold {
heap.pop();
heap.push(TopKEntry(score, candidate));
threshold = heap.peek().expect("non-empty").0.max(threshold);
f_essential = recompute_f(&partial_max, threshold);
}
}
for cursor in cursors.iter_mut().take(f_essential) {
if cursor.current_doc_id() == candidate {
cursor.next();
}
}
}
Ok(drain_top_k_desc(heap))
}
fn run_windowed_union(
&self,
column_id: u32,
mut cursors: Vec<TermCursor>,
k: usize,
mut filter: Option<&mut ExcludeFilter>,
floor_eff: f32,
doc_id_start: u32,
doc_id_end: u32,
) -> Result<Vec<(u32, f32)>, FtsError> {
if k == 0 {
return Ok(Vec::new());
}
let col_meta = &self.columns[column_id as usize];
let dl_norm_k1 = &col_meta.dl_norm_k1;
if doc_id_start > 0 {
for c in &mut cursors {
c.skip_to(doc_id_start);
}
}
let initial_cap = k.min(self.n_docs as usize).max(1);
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::with_capacity(initial_cap);
let mut threshold: f32 = floor_eff.max(0.0);
let mut scores = vec![0.0f32; OR_WINDOW as usize];
let mut present = [0u64; OR_WINDOW_WORDS];
loop {
let mut min_doc = u32::MAX;
for c in &cursors {
if !c.is_exhausted() {
min_doc = min_doc.min(c.current_doc_id());
}
}
if min_doc == u32::MAX || min_doc >= doc_id_end {
break;
}
let base = min_doc & !(OR_WINDOW - 1);
let window_end = base.saturating_add(OR_WINDOW).min(doc_id_end);
for c in &mut cursors {
while !c.is_exhausted() {
let d = c.current_doc_id();
if d >= window_end {
break;
}
let pos = c.pos;
if pos + bm25::SCORE_SIMD_LANES <= c.block_n {
let doc_ids = [
c.block_doc_ids[pos],
c.block_doc_ids[pos + 1],
c.block_doc_ids[pos + 2],
c.block_doc_ids[pos + 3],
];
if doc_ids[bm25::SCORE_SIMD_LANES - 1] < window_end {
let contributions = bm25::score_one_term_x4(
c.idf_x_k1p1,
[
c.block_tfs[pos],
c.block_tfs[pos + 1],
c.block_tfs[pos + 2],
c.block_tfs[pos + 3],
],
[
dl_norm_k1.get(doc_ids[0]),
dl_norm_k1.get(doc_ids[1]),
dl_norm_k1.get(doc_ids[2]),
dl_norm_k1.get(doc_ids[3]),
],
);
for lane in 0..bm25::SCORE_SIMD_LANES {
let local = (doc_ids[lane] - base) as usize;
scores[local] += contributions[lane];
present[local >> 6] |= 1u64 << (local & 63);
}
c.advance_by(bm25::SCORE_SIMD_LANES);
continue;
}
}
let local = (d - base) as usize;
scores[local] += bm25::score_with_dl_norm_k1(
c.idf_x_k1p1,
c.current_tf(),
dl_norm_k1.get(d),
);
present[local >> 6] |= 1u64 << (local & 63);
c.next();
}
}
for (word_idx, word) in present.iter_mut().enumerate() {
let mut bits = *word;
*word = 0;
while bits != 0 {
let b = bits.trailing_zeros() as usize;
bits &= bits - 1;
let local = (word_idx << 6) | b;
let score = scores[local];
scores[local] = 0.0;
let doc = base + local as u32;
if let Some(f) = filter.as_deref_mut()
&& !f.admits(doc)
{
continue;
}
if heap.len() < k {
heap.push(TopKEntry(score, doc));
if heap.len() == k {
threshold = heap.peek().expect("non-empty").0.max(threshold);
}
} else if score > threshold {
heap.pop();
heap.push(TopKEntry(score, doc));
threshold = heap.peek().expect("non-empty").0.max(threshold);
}
}
}
}
Ok(drain_top_k_desc(heap))
}
fn run_exhaustive_union(
&self,
column_id: u32,
mut cursors: Vec<TermCursor>,
k: usize,
) -> Result<Vec<(u32, f32)>, FtsError> {
let col_meta = &self.columns[column_id as usize];
let dl_norm_k1 = &col_meta.dl_norm_k1;
let initial_cap = k.min(self.n_docs as usize).max(1);
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::with_capacity(initial_cap);
let mut threshold: f32 = 0.0;
loop {
let mut candidate = u32::MAX;
for cursor in &cursors {
let d = cursor.current_doc_id();
if d < candidate {
candidate = d;
}
}
if candidate == u32::MAX {
break;
}
let norm = dl_norm_k1.get(candidate);
let mut score: f32 = 0.0;
let mut idfs = [0.0_f32; 4];
let mut tfs = [0.0_f32; 4];
let mut packed = 0;
for cursor in cursors.iter_mut() {
if cursor.current_doc_id() == candidate {
idfs[packed] = cursor.idf_x_k1p1;
tfs[packed] = cursor.current_tf() as f32;
packed += 1;
if packed == 4 {
score += bm25::score_simd_x4(idfs, tfs, norm);
idfs = [0.0; 4];
tfs = [0.0; 4];
packed = 0;
}
cursor.next();
}
}
if packed > 0 {
score += bm25::score_simd_x4(idfs, tfs, norm);
}
if heap.len() < k {
heap.push(TopKEntry(score, candidate));
if heap.len() == k {
threshold = heap.peek().expect("non-empty").0;
}
} else if score > threshold {
heap.pop();
heap.push(TopKEntry(score, candidate));
threshold = heap.peek().expect("non-empty").0;
}
}
Ok(drain_top_k_desc(heap))
}
async fn dispatch_multi_term_or(
&self,
column_id: u32,
terms: &[&str],
k: usize,
filter: Option<&mut ExcludeFilter>,
floor_eff: f32,
) -> Result<Vec<(u32, f32)>, FtsError> {
let cursors = self.build_term_cursors(column_id, terms).await?;
if cursors.is_empty() {
return Ok(Vec::new());
}
let no_floor = floor_eff == f32::NEG_INFINITY;
if cursors.len() == 2
&& k <= WAND_BMW_2TERM_MAX_K
&& filter.is_none()
&& no_floor
&& two_term_has_rare_anchor(&cursors)
{
self.run_wand_bmw(column_id, cursors, k)
} else if prefer_windowed_union(&cursors) {
self.run_windowed_union(column_id, cursors, k, filter, floor_eff, 0, u32::MAX)
} else {
self.run_max_score_bmm(column_id, cursors, k, filter, floor_eff)
}
}
#[doc(hidden)]
pub async fn search_with_algo_for_bench(
&self,
column: &str,
terms: &[&str],
k: usize,
algo: OrAlgo,
) -> Result<Vec<(u32, f32)>, FtsError> {
let column_id = self.resolve_column_id(column)?;
if terms.is_empty() || k == 0 {
return Ok(Vec::new());
}
let cursors = self.build_term_cursors(column_id, terms).await?;
if cursors.is_empty() {
return Ok(Vec::new());
}
match algo {
OrAlgo::Bmm => self.run_max_score_bmm(column_id, cursors, k, None, f32::NEG_INFINITY),
OrAlgo::WandBmw => self.run_wand_bmw(column_id, cursors, k),
OrAlgo::Exhaustive => self.run_exhaustive_union(column_id, cursors, k),
OrAlgo::Windowed => {
self.run_windowed_union(column_id, cursors, k, None, f32::NEG_INFINITY, 0, u32::MAX)
}
}
}
}
#[derive(Debug, Copy, Clone)]
struct TopKEntry(f32, u32);
impl PartialEq for TopKEntry {
fn eq(&self, other: &Self) -> bool {
self.0 == other.0 && self.1 == other.1
}
}
impl Eq for TopKEntry {}
impl PartialOrd for TopKEntry {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for TopKEntry {
fn cmp(&self, other: &Self) -> Ordering {
other
.0
.partial_cmp(&self.0)
.unwrap_or(Ordering::Equal)
.then_with(|| self.1.cmp(&other.1))
}
}
fn drain_top_k_desc(heap: BinaryHeap<TopKEntry>) -> Vec<(u32, f32)> {
let mut out: Vec<(u32, f32)> = heap.into_iter().map(|TopKEntry(s, d)| (d, s)).collect();
out.sort_unstable_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(Ordering::Equal)
.then(a.0.cmp(&b.0))
});
out
}
struct PhraseMember {
cursor: TermCursor,
positions: Bytes,
term_meta: Option<TermMeta>,
inline_position: Option<u32>,
idf: f32,
run_offsets: Vec<u32>,
run_offsets_block: usize,
pos_scratch: Vec<u32>,
}
const NO_BLOCK_CACHED: usize = usize::MAX;
impl PhraseMember {
fn decode_current_positions(&mut self) -> Result<(), FtsError> {
self.pos_scratch.clear();
if let Some(p) = self.inline_position {
self.pos_scratch.push(p);
return Ok(());
}
let block = self.cursor.current_block;
if self.run_offsets_block != block {
self.run_offsets.clear();
let term_meta = self.term_meta.as_ref().expect("PFOR member has term meta");
let mut at =
term_meta.positions_block_offset(self.cursor.bytes.as_ref(), block) as usize;
for i in 0..self.cursor.block_n {
self.run_offsets.push(at as u32);
skip_run(&self.positions, &mut at, self.cursor.block_tfs[i]).ok_or_else(|| {
FtsError::Read(ReadError::MalformedVersion(
"position runs truncated within block".into(),
))
})?;
}
self.run_offsets_block = block;
}
let pair = self.cursor.pos;
let mut at = self.run_offsets[pair] as usize;
decode_run(
&self.positions,
&mut at,
self.cursor.block_tfs[pair],
&mut self.pos_scratch,
)
.ok_or_else(|| {
FtsError::Read(ReadError::MalformedVersion(
"position run truncated or overflowing".into(),
))
})?;
Ok(())
}
}
struct PhraseCursor {
members: Vec<PhraseMember>,
idf_x_k1p1: f32,
term_max_bm25: f32,
current_doc: u32,
current_tf: u32,
}
impl PhraseCursor {
fn new(
cursors: Vec<TermCursor>,
positions: Vec<Bytes>,
positional: Vec<(Option<TermMeta>, Option<u32>)>,
) -> Result<Self, FtsError> {
debug_assert!(cursors.len() >= 2, "single-token phrases degrade to terms");
debug_assert_eq!(cursors.len(), positions.len());
debug_assert_eq!(cursors.len(), positional.len());
let mut idf_sum = 0.0f32;
let mut min_scaled_bound = f32::INFINITY;
let members: Vec<PhraseMember> = cursors
.into_iter()
.zip(positions)
.zip(positional)
.map(|((cursor, positions), (term_meta, inline_position))| {
let idf = cursor.idf_x_k1p1 / (bm25::K1 + 1.0);
min_scaled_bound = min_scaled_bound.min(cursor.term_max_bm25 / idf);
idf_sum += idf;
PhraseMember {
cursor,
positions,
term_meta,
inline_position,
idf,
run_offsets: Vec::new(),
run_offsets_block: NO_BLOCK_CACHED,
pos_scratch: Vec::new(),
}
})
.collect();
let mut cursor = Self {
idf_x_k1p1: idf_sum * (bm25::K1 + 1.0),
term_max_bm25: idf_sum * min_scaled_bound,
members,
current_doc: 0,
current_tf: 0,
};
cursor.seek_match(0, f32::NEG_INFINITY, &NormTable::empty())?;
Ok(cursor)
}
#[inline]
fn is_exhausted(&self) -> bool {
self.current_doc == u32::MAX
}
#[inline]
fn current_doc_id(&self) -> u32 {
self.current_doc
}
fn skip_to(&mut self, target: u32) -> Result<(), FtsError> {
if self.is_exhausted() || self.current_doc >= target {
return Ok(());
}
self.seek_match(target, f32::NEG_INFINITY, &NormTable::empty())
}
fn skip_to_pruned(
&mut self,
target: u32,
bar: f32,
dl_norm_k1: &NormTable,
) -> Result<(), FtsError> {
if self.is_exhausted() || self.current_doc >= target {
return Ok(());
}
self.seek_match(target, bar, dl_norm_k1)
}
fn seek_match(
&mut self,
mut from: u32,
bar: f32,
dl_norm_k1: &NormTable,
) -> Result<(), FtsError> {
'docs: loop {
let mut aligned = from;
let mut i = 0usize;
while i < self.members.len() {
let c = &mut self.members[i].cursor;
c.skip_to(aligned);
if c.is_exhausted() {
self.current_doc = u32::MAX;
self.current_tf = 0;
return Ok(());
}
let here = c.current_doc_id();
if here > aligned {
aligned = here;
i = 0;
continue;
}
i += 1;
}
if bar > f32::NEG_INFINITY {
let min_tf = self
.members
.iter()
.map(|m| m.cursor.current_tf())
.min()
.expect("members >= 2");
let ub =
bm25::score_with_dl_norm_k1(self.idf_x_k1p1, min_tf, dl_norm_k1.get(aligned));
if ub < bar {
from = match aligned.checked_add(1) {
Some(next) => next,
None => {
self.current_doc = u32::MAX;
self.current_tf = 0;
return Ok(());
}
};
continue 'docs;
}
}
let tf = self.verify_at_aligned()?;
if tf > 0 {
self.current_doc = aligned;
self.current_tf = tf;
return Ok(());
}
from = match aligned.checked_add(1) {
Some(next) => next,
None => {
self.current_doc = u32::MAX;
self.current_tf = 0;
return Ok(());
}
};
continue 'docs;
}
}
fn verify_at_aligned(&mut self) -> Result<u32, FtsError> {
for m in self.members.iter_mut() {
m.decode_current_positions()?;
}
let (anchor, rest) = self.members.split_first_mut().expect("members >= 2");
let mut tf = 0u32;
'anchors: for &p in &anchor.pos_scratch {
for (i, m) in rest.iter().enumerate() {
let want = match p.checked_add(i as u32 + 1) {
Some(w) => w,
None => continue 'anchors,
};
if m.pos_scratch.binary_search(&want).is_err() {
continue 'anchors;
}
}
tf += 1;
}
Ok(tf)
}
#[inline]
fn score_current(&self, dl_norm_k1: f32) -> f32 {
bm25::score_with_dl_norm_k1(self.idf_x_k1p1, self.current_tf, dl_norm_k1)
}
fn block_max_in_range(&mut self, range_start: u32, range_end: u32) -> f32 {
let mut min_scaled = f32::INFINITY;
for m in self.members.iter_mut() {
let b = m.cursor.block_max_in_range(range_start, range_end);
min_scaled = min_scaled.min(b / m.idf);
}
let idf_sum = self.idf_x_k1p1 / (bm25::K1 + 1.0);
idf_sum * min_scaled
}
}
enum AnyCursor {
Term(TermCursor),
Phrase(PhraseCursor),
}
impl AnyCursor {
#[inline]
fn is_exhausted(&self) -> bool {
match self {
AnyCursor::Term(c) => c.is_exhausted(),
AnyCursor::Phrase(c) => c.is_exhausted(),
}
}
#[inline]
fn current_doc_id(&self) -> u32 {
match self {
AnyCursor::Term(c) => c.current_doc_id(),
AnyCursor::Phrase(c) => c.current_doc_id(),
}
}
fn skip_to(&mut self, target: u32) -> Result<(), FtsError> {
match self {
AnyCursor::Term(c) => {
c.skip_to(target);
Ok(())
}
AnyCursor::Phrase(c) => c.skip_to(target),
}
}
fn skip_to_pruned(
&mut self,
target: u32,
bar: f32,
dl_norm_k1: &NormTable,
) -> Result<(), FtsError> {
match self {
AnyCursor::Term(c) => {
c.skip_to(target);
Ok(())
}
AnyCursor::Phrase(c) => c.skip_to_pruned(target, bar, dl_norm_k1),
}
}
#[inline]
fn score_current(&self, dl_norm_k1: f32) -> f32 {
match self {
AnyCursor::Term(c) => {
bm25::score_with_dl_norm_k1(c.idf_x_k1p1, c.current_tf(), dl_norm_k1)
}
AnyCursor::Phrase(c) => c.score_current(dl_norm_k1),
}
}
#[inline]
fn term_max_bm25(&self) -> f32 {
match self {
AnyCursor::Term(c) => c.term_max_bm25,
AnyCursor::Phrase(c) => c.term_max_bm25,
}
}
#[inline]
fn block_max_in_range(&mut self, range_start: u32, range_end: u32) -> f32 {
match self {
AnyCursor::Term(c) => c.block_max_in_range(range_start, range_end),
AnyCursor::Phrase(c) => c.block_max_in_range(range_start, range_end),
}
}
}
struct AtomExcludeFilter {
atoms: Vec<AnyCursor>,
last_doc: u32,
}
impl AtomExcludeFilter {
fn new(atoms: Vec<AnyCursor>) -> Self {
Self { atoms, last_doc: 0 }
}
fn admits(&mut self, doc: u32) -> Result<bool, FtsError> {
debug_assert!(
doc >= self.last_doc,
"AtomExcludeFilter fed non-monotonic doc: {doc} < {}",
self.last_doc
);
self.last_doc = doc;
for a in &mut self.atoms {
a.skip_to(doc)?;
if !a.is_exhausted() && a.current_doc_id() == doc {
return Ok(false);
}
}
Ok(true)
}
}
struct ExcludeFilter {
cursors: Vec<TermCursor>,
last_doc: u32,
}
impl ExcludeFilter {
fn new(cursors: Vec<TermCursor>) -> Self {
Self {
cursors,
last_doc: 0,
}
}
}
impl ExcludeFilter {
#[inline]
fn admits(&mut self, doc: u32) -> bool {
debug_assert!(
doc >= self.last_doc,
"ExcludeFilter fed non-monotonic doc: {doc} < {}",
self.last_doc
);
self.last_doc = doc;
for c in &mut self.cursors {
c.skip_to(doc);
if !c.is_exhausted() && c.current_doc_id() == doc {
return false;
}
}
true
}
}
trait AndSink {
fn bar(&self) -> f32 {
f32::NEG_INFINITY
}
fn needs_score(&self) -> bool;
fn emit(&mut self, doc: u32, score: f32);
}
struct ScoreSink<'a> {
heap: &'a mut BinaryHeap<TopKEntry>,
k: usize,
filter: Option<&'a mut ExcludeFilter>,
floor_eff: f32,
}
impl AndSink for ScoreSink<'_> {
fn bar(&self) -> f32 {
if self.heap.len() >= self.k {
self.heap
.peek()
.expect("heap len == k")
.0
.max(self.floor_eff)
} else {
self.floor_eff
}
}
fn needs_score(&self) -> bool {
true
}
fn emit(&mut self, doc: u32, score: f32) {
if score > self.floor_eff {
and_heap_push(self.heap, self.k, self.filter.as_deref_mut(), score, doc);
}
}
}
struct MustShouldSink<'a> {
heap: &'a mut BinaryHeap<TopKEntry>,
k: usize,
filter: Option<&'a mut ExcludeFilter>,
floor_eff: f32,
shoulds: Vec<TermCursor>,
should_ub: f32,
dl_norm_k1: &'a NormTable,
}
impl AndSink for MustShouldSink<'_> {
fn bar(&self) -> f32 {
let full_bar = if self.heap.len() >= self.k {
self.heap
.peek()
.expect("heap len == k")
.0
.max(self.floor_eff)
} else {
self.floor_eff
};
full_bar - self.should_ub
}
fn needs_score(&self) -> bool {
true
}
fn emit(&mut self, doc: u32, must_score: f32) {
let norm = self.dl_norm_k1.get(doc);
let mut score = must_score;
for c in &mut self.shoulds {
c.skip_to(doc);
if !c.is_exhausted() && c.current_doc_id() == doc {
score += bm25::score_with_dl_norm_k1(c.idf_x_k1p1, c.current_tf(), norm);
}
}
if score > self.floor_eff {
and_heap_push(self.heap, self.k, self.filter.as_deref_mut(), score, doc);
}
}
}
struct CollectSink {
out: Vec<u32>,
}
impl AndSink for CollectSink {
fn needs_score(&self) -> bool {
false
}
fn emit(&mut self, doc: u32, _score: f32) {
self.out.push(doc);
}
}
struct CountSink {
n: u64,
}
impl AndSink for CountSink {
fn needs_score(&self) -> bool {
false
}
fn emit(&mut self, _doc: u32, _score: f32) {
self.n += 1;
}
}
#[inline]
fn and_heap_push(
heap: &mut BinaryHeap<TopKEntry>,
k: usize,
filter: Option<&mut ExcludeFilter>,
score: f32,
doc_id: u32,
) {
if let Some(f) = filter
&& !f.admits(doc_id)
{
return;
}
if heap.len() < k {
heap.push(TopKEntry(score, doc_id));
} else if let Some(&worst) = heap.peek()
&& (score > worst.0 || (score == worst.0 && doc_id < worst.1))
{
heap.pop();
heap.push(TopKEntry(score, doc_id));
}
}
fn top_k(scores: HashMap<u32, f32>, k: usize) -> Vec<(u32, f32)> {
let mut sorted: Vec<(u32, f32)> = scores.into_iter().collect();
sorted.sort_unstable_by_key(|(d, _)| *d);
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::with_capacity(k.min(sorted.len()).max(1));
for (doc_id, score) in sorted {
if heap.len() < k {
heap.push(TopKEntry(score, doc_id));
} else if let Some(TopKEntry(top_score, _)) = heap.peek()
&& score > *top_score
{
heap.pop();
heap.push(TopKEntry(score, doc_id));
}
}
drain_top_k_desc(heap)
}
fn fetch_source_range(source: &Source, range: Range<usize>, what: &str) -> Result<Bytes, FtsError> {
source.get_range(range).map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"{what} lazy source range fetch failed: {e}"
)))
})
}
async fn fetch_lazy_range(
source: &dyn LazyByteSource,
range: Range<usize>,
what: &str,
) -> Result<Bytes, FtsError> {
source
.range(range.start as u64, range.len() as u64)
.await
.map_err(|e| {
FtsError::Read(ReadError::MalformedVersion(format!(
"{what} lazy source range fetch failed: {e}"
)))
})
}
#[inline]
fn read_u32_le(b: &[u8]) -> u32 {
u32::from_le_bytes([b[0], b[1], b[2], b[3]])
}
#[inline]
fn read_u64_le(b: &[u8]) -> u64 {
let mut buf = [0u8; 8];
buf.copy_from_slice(&b[0..8]);
u64::from_le_bytes(buf)
}
fn or_walk_unranked(mut cursors: Vec<TermCursor>, mut emit: impl FnMut(u32)) {
loop {
let min_doc = cursors
.iter()
.filter(|c| !c.is_exhausted())
.map(TermCursor::current_doc_id)
.min();
let Some(min_doc) = min_doc else { break };
emit(min_doc);
for c in cursors.iter_mut() {
if !c.is_exhausted() && c.current_doc_id() == min_doc {
c.next();
}
}
}
}
fn or_merge_unranked(cursors: Vec<TermCursor>) -> Vec<u32> {
let mut out = Vec::new();
or_walk_unranked(cursors, |doc| out.push(doc));
out
}
fn or_count_unranked(mut cursors: Vec<TermCursor>) -> u64 {
let mut present = [0u64; OR_WINDOW_WORDS];
let mut n = 0u64;
loop {
let mut min_doc = u32::MAX;
for c in &cursors {
if !c.is_exhausted() {
min_doc = min_doc.min(c.current_doc_id());
}
}
if min_doc == u32::MAX {
break;
}
let base = min_doc & !(OR_WINDOW - 1);
let window_end = base.saturating_add(OR_WINDOW);
for c in &mut cursors {
while !c.is_exhausted() {
let d = c.current_doc_id();
if d >= window_end {
break;
}
let local = (d - base) as usize;
present[local >> 6] |= 1u64 << (local & 63);
c.next();
}
}
for word in present.iter_mut() {
n += word.count_ones() as u64;
*word = 0;
}
}
n
}
#[derive(Debug, Copy, Clone)]
struct TermMeta {
df: u64,
postings_length: usize,
num_blocks: usize,
skip_start: usize,
positions_offset: u64,
positions_length: u32,
}
impl TermMeta {
fn parse(postings: &[u8], metadata_offset: usize, positional: bool) -> Result<Self, FtsError> {
let term_meta_size = match positional {
true => TERM_META_POSITIONAL_SIZE,
false => TERM_META_SIZE,
};
if metadata_offset + term_meta_size > postings.len() {
return Err(FtsError::Read(ReadError::MalformedVersion(
"term metadata offset out of postings region".into(),
)));
}
let df = read_u32_le(
&postings[metadata_offset + term_meta::DF_OFF
..metadata_offset + term_meta::DF_OFF + U32_BYTES],
) as u64;
let postings_length = read_u32_le(
&postings[metadata_offset + term_meta::POSTINGS_LENGTH_OFF
..metadata_offset + term_meta::POSTINGS_LENGTH_OFF + U32_BYTES],
) as usize;
let num_blocks = read_u32_le(
&postings[metadata_offset + term_meta::NUM_BLOCKS_OFF
..metadata_offset + term_meta::NUM_BLOCKS_OFF + U32_BYTES],
) as usize;
let (positions_offset, positions_length) = match positional {
true => (
read_u64_le(
&postings[metadata_offset + term_meta::POSITIONS_OFFSET_OFF
..metadata_offset + term_meta::POSITIONS_OFFSET_OFF + U64_BYTES],
),
read_u32_le(
&postings[metadata_offset + term_meta::POSITIONS_LENGTH_OFF
..metadata_offset + term_meta::POSITIONS_LENGTH_OFF + U32_BYTES],
),
),
false => (0, 0),
};
let skip_start = metadata_offset + term_meta_size;
let skip_end = skip_start + num_blocks * SKIP_ENTRY_SIZE;
if skip_end > postings.len() {
return Err(FtsError::Read(ReadError::MalformedVersion(
"skip table runs past postings region".into(),
)));
}
Ok(Self {
df,
postings_length,
num_blocks,
skip_start,
positions_offset,
positions_length,
})
}
#[inline]
fn skip_entry(&self, postings: &[u8], i: usize) -> (u32, usize, f32) {
debug_assert!(i < self.num_blocks, "skip entry {i} >= {}", self.num_blocks);
let entry_off = self.skip_start + i * SKIP_ENTRY_SIZE;
let last_doc_id = read_u32_le(
&postings[entry_off + skip_entry::LAST_DOC_ID_OFF
..entry_off + skip_entry::LAST_DOC_ID_OFF + U32_BYTES],
);
let block_offset = read_u32_le(
&postings[entry_off + skip_entry::BLOCK_OFFSET_OFF
..entry_off + skip_entry::BLOCK_OFFSET_OFF + U32_BYTES],
) as usize;
let max_bm25_x1000 = read_u32_le(
&postings[entry_off + skip_entry::MAX_BM25_OFF
..entry_off + skip_entry::MAX_BM25_OFF + U32_BYTES],
);
(
last_doc_id,
block_offset,
max_bm25_x1000.saturating_add(1) as f32 / format::fts::BLOCK_MAX_BM25_FIXED_POINT_SCALE,
)
}
#[inline]
fn positions_block_offset(&self, postings: &[u8], i: usize) -> u32 {
debug_assert!(i < self.num_blocks, "skip entry {i} >= {}", self.num_blocks);
let entry_off = self.skip_start + i * SKIP_ENTRY_SIZE;
read_u32_le(
&postings[entry_off + skip_entry::POSITIONS_BLOCK_OFFSET_OFF
..entry_off + skip_entry::POSITIONS_BLOCK_OFFSET_OFF + U32_BYTES],
)
}
#[inline]
fn block_end_in_term(&self, postings: &[u8], i: usize) -> usize {
if i + 1 < self.num_blocks {
let next_off = self.skip_start + (i + 1) * SKIP_ENTRY_SIZE;
read_u32_le(&postings[next_off + 4..next_off + 8]) as usize
} else {
self.postings_length
}
}
}
#[derive(Debug, Clone, Copy)]
struct BlockMeta {
last_doc_id: u32,
block_byte_offset: usize,
block_byte_end: usize,
block_max_bm25: f32,
}
struct TermCursor {
idf_x_k1p1: f32,
term_max_bm25: f32,
df: u64,
blocks: Vec<BlockMeta>,
block_doc_ids: Vec<u32>,
block_tfs: Vec<u32>,
block_n: usize,
current_block: usize,
pos: usize,
inspect_block: usize,
bytes: Bytes,
}
impl TermCursor {
fn new(term_bytes: Bytes, n_docs: u64, positional: bool) -> Result<Self, FtsError> {
let postings: &[u8] = term_bytes.as_ref();
let metadata_offset = 0usize;
let term_meta = TermMeta::parse(postings, metadata_offset, positional)?;
let idf = bm25::idf(n_docs, term_meta.df);
let mut blocks: Vec<BlockMeta> = Vec::with_capacity(term_meta.num_blocks);
let mut term_max_bm25: f32 = 0.0;
for i in 0..term_meta.num_blocks {
let (last_doc_id, block_offset_in_term, block_max_bm25) =
term_meta.skip_entry(postings, i);
term_max_bm25 = term_max_bm25.max(block_max_bm25);
blocks.push(BlockMeta {
last_doc_id,
block_byte_offset: metadata_offset + block_offset_in_term,
block_byte_end: metadata_offset + term_meta.block_end_in_term(postings, i),
block_max_bm25,
});
}
let mut cursor = Self {
idf_x_k1p1: idf * (bm25::K1 + 1.0),
term_max_bm25,
df: term_meta.df,
blocks,
block_doc_ids: vec![0u32; BLOCK_LEN],
block_tfs: vec![0u32; BLOCK_LEN],
block_n: 0,
current_block: 0,
pos: 0,
inspect_block: 0,
bytes: term_bytes,
};
if !cursor.blocks.is_empty() {
cursor.decode_current_block();
}
Ok(cursor)
}
fn new_inline(doc_id: u32, tf: u32, n_docs: u64, dl_norm_k1: f32) -> Self {
let idf = bm25::idf(n_docs, 1);
let idf_x_k1p1 = idf * (bm25::K1 + 1.0);
let block_max_bm25 = bm25::score_with_dl_norm_k1(idf_x_k1p1, tf, dl_norm_k1);
let blocks = vec![BlockMeta {
last_doc_id: doc_id,
block_byte_offset: 0,
block_byte_end: 0,
block_max_bm25,
}];
let mut block_doc_ids = vec![0u32; BLOCK_LEN];
let mut block_tfs = vec![0u32; BLOCK_LEN];
block_doc_ids[0] = doc_id;
block_tfs[0] = tf;
Self {
idf_x_k1p1,
term_max_bm25: block_max_bm25,
df: 1,
blocks,
block_doc_ids,
block_tfs,
block_n: 1,
current_block: 0,
pos: 0,
inspect_block: 0,
bytes: Bytes::new(),
}
}
fn decode_current_block(&mut self) {
let block = self.blocks[self.current_block];
let bytes = self
.bytes
.slice(block.block_byte_offset..block.block_byte_end);
self.block_n = decode_block(&bytes, &mut self.block_doc_ids, &mut self.block_tfs);
self.pos = 0;
}
fn is_exhausted(&self) -> bool {
self.current_block >= self.blocks.len()
}
#[inline(always)]
fn block_count(&self) -> usize {
self.blocks.len()
}
#[inline(always)]
fn current_doc_id(&self) -> u32 {
if self.is_exhausted() || self.pos >= self.block_n {
u32::MAX
} else {
self.block_doc_ids[self.pos]
}
}
#[inline(always)]
fn current_tf(&self) -> u32 {
debug_assert!(!self.is_exhausted() && self.pos < self.block_n);
self.block_tfs[self.pos]
}
#[inline(always)]
fn current_block_max_bm25(&self) -> f32 {
if self.is_exhausted() {
0.0
} else {
self.blocks[self.current_block].block_max_bm25
}
}
#[inline(always)]
fn current_block_last_doc_id(&self) -> u32 {
if self.is_exhausted() {
u32::MAX
} else {
self.blocks[self.current_block].last_doc_id
}
}
fn shallow_advance_block_to(&mut self, target: u32) {
if self.inspect_block < self.current_block {
self.inspect_block = self.current_block;
}
while self.inspect_block < self.blocks.len()
&& self.blocks[self.inspect_block].last_doc_id < target
{
self.inspect_block += 1;
}
}
fn block_max_in_range(&mut self, range_start: u32, range_end: u32) -> f32 {
self.shallow_advance_block_to(range_start);
let mut max: f32 = 0.0;
let mut i = self.inspect_block;
while i < self.blocks.len() {
let block_start = if i == 0 {
0u32
} else {
self.blocks[i - 1].last_doc_id.saturating_add(1)
};
if block_start > range_end {
break;
}
let m = self.blocks[i].block_max_bm25;
if m > max {
max = m;
}
i += 1;
}
max
}
fn inspect_block_max_bm25(&self) -> f32 {
if self.inspect_block >= self.blocks.len() {
0.0
} else {
self.blocks[self.inspect_block].block_max_bm25
}
}
fn inspect_block_last_doc_id(&self) -> u32 {
if self.inspect_block >= self.blocks.len() {
u32::MAX
} else {
self.blocks[self.inspect_block].last_doc_id
}
}
#[inline(always)]
fn next(&mut self) {
if self.is_exhausted() {
return;
}
self.pos += 1;
if self.pos >= self.block_n {
self.advance_block();
}
}
#[inline(always)]
fn advance_by(&mut self, count: usize) {
debug_assert!(!self.is_exhausted());
debug_assert!(count > 0 && self.pos + count <= self.block_n);
self.pos += count;
if self.pos == self.block_n {
self.advance_block();
}
}
#[inline(always)]
fn advance_block(&mut self) {
self.current_block += 1;
if self.current_block > self.inspect_block {
self.inspect_block = self.current_block;
}
if self.current_block < self.blocks.len() {
self.decode_current_block();
}
}
#[inline(always)]
fn skip_to(&mut self, target: u32) {
if self.is_exhausted() {
return;
}
let cur_block = self.current_block;
let cur_block_last = self.blocks[cur_block].last_doc_id;
if cur_block_last >= target {
let n = self.block_n;
while self.pos < n && self.block_doc_ids[self.pos] < target {
self.pos += 1;
}
if self.pos < n {
return;
}
}
self.skip_to_cross_block(target);
}
#[cold]
fn skip_to_cross_block(&mut self, target: u32) {
while self.current_block < self.blocks.len()
&& self.blocks[self.current_block].last_doc_id < target
{
self.current_block += 1;
}
if self.current_block > self.inspect_block {
self.inspect_block = self.current_block;
}
if self.is_exhausted() {
return;
}
self.decode_current_block();
while self.pos < self.block_n && self.block_doc_ids[self.pos] < target {
self.pos += 1;
}
if self.pos >= self.block_n {
self.current_block += 1;
if self.current_block > self.inspect_block {
self.inspect_block = self.current_block;
}
if self.current_block < self.blocks.len() {
self.decode_current_block();
}
}
}
}
#[cfg(test)]
mod tests {
use std::{collections::HashSet, sync::Arc};
use super::*;
use crate::superfile::{BytesLazyByteSource, fts::builder::FtsBuilder};
fn build_blob() -> (Bytes, String) {
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false)
.expect("register column");
b.add_doc(0, 0, "rust async runtime").expect("add doc");
b.add_doc(0, 1, "tokio is a rust runtime").expect("add doc");
b.add_doc(0, 2, "java spring boot").expect("add doc");
let bytes = b.finish().expect("finish");
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
(Bytes::from(bytes), json.to_string())
}
#[test]
fn open_accepts_valid_blob() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open should succeed");
assert_eq!(r.n_docs(), 3);
assert!(r.n_terms() > 0);
assert_eq!(r.fts_columns().collect::<Vec<_>>(), vec!["body"]);
}
#[test]
fn open_rejects_bad_magic() {
let (mut blob_vec, json) = build_blob();
let mut bytes = blob_vec.to_vec();
bytes[0] = b'X';
blob_vec = Bytes::from(bytes);
let err = FtsReader::open(blob_vec, &json).expect_err("expected error");
assert!(matches!(err, FtsError::Read(ReadError::BadMagic { .. })));
}
#[test]
fn open_rejects_short_blob() {
let err = FtsReader::open(Bytes::from(vec![0u8; 8]), "[]").expect_err("expected error");
assert!(matches!(err, FtsError::Read(_)));
}
#[test]
fn open_rejects_columns_json_mismatch() {
let (blob, _) = build_blob();
let bad_json = r#"[{"name":"body","tokenizer":"ascii_lower"},{"name":"title","tokenizer":"ascii_lower"}]"#;
let err = FtsReader::open(blob, bad_json).expect_err("expected error");
assert!(matches!(
err,
FtsError::Read(ReadError::MalformedVersion(_))
));
}
#[tokio::test]
async fn search_returns_exact_doc_ids_for_known_term() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["rust"], 10, BoolMode::Or)
.await
.expect("FTS search");
let ids: Vec<u32> = hits.iter().map(|(d, _)| *d).collect();
assert!(ids.contains(&0), "doc 0 should match");
assert!(ids.contains(&1), "doc 1 should match");
assert!(!ids.contains(&2), "doc 2 should not match");
}
#[tokio::test]
async fn token_match_or_unions_and_intersects_unranked() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
assert_eq!(
r.token_match("body", &["rust"], BoolMode::Or)
.await
.expect("single"),
vec![0, 1]
);
assert_eq!(
r.token_match("body", &["rust", "java"], BoolMode::Or)
.await
.expect("or"),
vec![0, 1, 2]
);
assert_eq!(
r.token_match("body", &["rust", "runtime"], BoolMode::And)
.await
.expect("and"),
vec![0, 1]
);
assert!(
r.token_match("body", &["rust", "zzz"], BoolMode::And)
.await
.expect("and absent")
.is_empty()
);
assert_eq!(
r.token_match("body", &["java", "zzz"], BoolMode::Or)
.await
.expect("or absent"),
vec![2]
);
assert!(
r.token_match("body", &[], BoolMode::And)
.await
.expect("empty")
.is_empty()
);
}
#[tokio::test]
async fn token_match_count_matches_token_match_len() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let cases: &[(&[&str], BoolMode)] = &[
(&["rust"], BoolMode::Or),
(&["rust", "java"], BoolMode::Or),
(&["rust", "runtime"], BoolMode::And),
(&["rust", "zzz"], BoolMode::And),
(&["java", "zzz"], BoolMode::Or),
(&[], BoolMode::And),
];
for (tokens, mode) in cases {
let len = r
.token_match("body", tokens, *mode)
.await
.expect("token_match")
.len() as u64;
let count = r
.token_match_count("body", tokens, *mode)
.await
.expect("token_match_count");
assert_eq!(count, len, "count vs len for {tokens:?} {mode:?}");
}
}
#[tokio::test]
async fn or_count_spans_multiple_windows() {
const N_DOCS: u32 = OR_WINDOW * 2 + 500;
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
for i in 0..N_DOCS {
let mut text = String::from("alpha "); if i % 2 == 0 {
text.push_str("beta ");
}
if i % 3 == 0 {
text.push_str("gamma ");
}
if i % 5 == 0 {
text.push_str("delta ");
}
b.add_doc(0, i, text.trim()).expect("add doc");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let shapes: &[&[&str]] = &[
&["alpha"], &["beta", "gamma"], &["alpha", "beta", "gamma", "delta"], &["gamma", "zzz_absent"], ];
for terms in shapes {
let merge_len = r
.token_match("body", terms, BoolMode::Or)
.await
.expect("token_match")
.len() as u64;
let count = r
.token_match_count("body", terms, BoolMode::Or)
.await
.expect("token_match_count");
assert_eq!(
count, merge_len,
"windowed count vs merge len for {terms:?}"
);
}
assert_eq!(
r.token_match_count("body", &["alpha"], BoolMode::Or)
.await
.expect("count"),
N_DOCS as u64
);
}
#[tokio::test]
async fn token_match_doc_set_matches_bm25_for_same_terms() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let mut bm25: Vec<u32> = r
.search("body", &["rust", "java"], 10, BoolMode::Or)
.await
.expect("search")
.into_iter()
.map(|(d, _)| d)
.collect();
bm25.sort_unstable();
let boolean = r
.token_match("body", &["rust", "java"], BoolMode::Or)
.await
.expect("boolean");
assert_eq!(bm25, boolean, "boolean Or doc set == bm25 doc set");
}
#[tokio::test]
async fn exhaustive_and_bmm_agree_on_top_k() {
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false)
.expect("register column");
let docs = [
"alpha",
"beta",
"gamma",
"alpha beta",
"alpha gamma",
"beta gamma",
"alpha beta gamma",
"delta",
"epsilon",
"alpha delta",
"beta epsilon",
"gamma delta",
"alpha beta delta",
"alpha epsilon gamma",
"delta epsilon",
"alpha alpha alpha",
"beta beta beta",
"gamma gamma",
"alpha beta gamma delta epsilon",
"epsilon",
];
for (i, text) in docs.iter().enumerate() {
b.add_doc(0, i as u32, text).expect("add doc");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let terms: &[&str] = &["alpha", "beta", "gamma"];
let bmm = r
.search_with_algo_for_bench("body", terms, 5, OrAlgo::Bmm)
.await
.expect("bmm");
let exh = r
.search_with_algo_for_bench("body", terms, 5, OrAlgo::Exhaustive)
.await
.expect("exhaustive");
assert_eq!(bmm.len(), exh.len(), "result length mismatch");
for ((d_bmm, s_bmm), (d_exh, s_exh)) in bmm.iter().zip(exh.iter()) {
assert_eq!(d_bmm, d_exh, "doc_id mismatch");
assert!(
(s_bmm - s_exh).abs() < 1e-4,
"score mismatch: bmm={s_bmm} exhaustive={s_exh}"
);
}
}
#[tokio::test]
async fn search_missing_term_or_returns_empty() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["nonexistent"], 10, BoolMode::Or)
.await
.expect("search");
assert!(hits.is_empty());
}
#[tokio::test]
async fn search_and_short_circuits_on_missing_term() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["rust", "nonexistent"], 10, BoolMode::And)
.await
.expect("search");
assert!(hits.is_empty());
}
#[tokio::test]
async fn search_and_intersects_term_postings() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["rust", "runtime"], 10, BoolMode::And)
.await
.expect("search");
let ids: Vec<u32> = hits.iter().map(|(d, _)| *d).collect();
assert!(ids.contains(&0));
assert!(ids.contains(&1));
assert!(!ids.contains(&2));
}
#[tokio::test]
async fn search_unknown_column_errors() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let err = r
.search("title", &["rust"], 10, BoolMode::Or)
.await
.expect_err("expected error");
assert!(matches!(err, FtsError::UnknownColumn(_)));
}
#[tokio::test]
async fn search_empty_terms_returns_empty() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &[], 10, BoolMode::Or)
.await
.expect("FTS search");
assert!(hits.is_empty());
}
#[tokio::test]
async fn search_zero_k_returns_empty() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["rust"], 0, BoolMode::Or)
.await
.expect("FTS search");
assert!(hits.is_empty());
}
#[tokio::test]
async fn search_results_sorted_by_score_desc() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["rust"], 10, BoolMode::Or)
.await
.expect("FTS search");
for w in hits.windows(2) {
assert!(w[0].1 >= w[1].1, "scores should be descending");
}
}
#[tokio::test]
async fn search_limits_to_k() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["rust"], 1, BoolMode::Or)
.await
.expect("FTS search");
assert_eq!(hits.len(), 1);
}
fn build_mixed_df_blob() -> (Bytes, String) {
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false)
.expect("register column");
b.add_doc(0, 0, "common rust uniqzero").expect("add doc");
b.add_doc(0, 1, "common rust").expect("add doc");
b.add_doc(0, 2, "common uniqtwo").expect("add doc");
let bytes = b.finish().expect("finish");
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
(Bytes::from(bytes), json.to_string())
}
#[test]
fn df1_inline_form_flag_set_on_fst_value() {
let (blob, _json) = build_mixed_df_blob();
let header_size = 48usize;
let fst_off =
u64::from_le_bytes(blob[24..32].try_into().expect("fst_off slice is 8 bytes")) as usize;
let postings_off = u64::from_le_bytes(
blob[32..40]
.try_into()
.expect("postings_off slice is 8 bytes"),
) as usize;
let fst_bytes = &blob[fst_off..postings_off - 4];
let dict = DictReader::open(fst_bytes).expect("open dict");
assert_eq!(header_size, 48);
let val_common = dict.lookup(b"body\x1Fcommon").expect("common in FST");
let val_rust = dict.lookup(b"body\x1Frust").expect("rust in FST");
let val_uniq_d0 = dict.lookup(b"body\x1Funiqzero").expect("uniqzero in FST");
let val_uniq_d2 = dict.lookup(b"body\x1Funiqtwo").expect("uniqtwo in FST");
assert_eq!(val_common & 1, 0, "df=3 common term must use PFOR form");
assert_eq!(val_rust & 1, 0, "df=2 rust term must use PFOR form");
assert_eq!(val_uniq_d0 & 1, 1, "df=1 uniqzero must use inline form");
assert_eq!(val_uniq_d2 & 1, 1, "df=1 uniqtwo must use inline form");
match FstValue::unpack(val_uniq_d0) {
FstValue::Inline { doc_id, tf } => {
assert_eq!(doc_id, 0);
assert_eq!(tf, 1);
}
FstValue::Pfor { .. } => panic!("expected inline form"),
}
match FstValue::unpack(val_uniq_d2) {
FstValue::Inline { doc_id, tf } => {
assert_eq!(doc_id, 2);
assert_eq!(tf, 1);
}
FstValue::Pfor { .. } => panic!("expected inline form"),
}
}
#[tokio::test]
async fn df1_single_term_search_returns_one_doc() {
let (blob, json) = build_mixed_df_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["uniqzero"], 10, BoolMode::Or)
.await
.expect("FTS search");
assert_eq!(hits.len(), 1, "df=1 term should return exactly one hit");
assert_eq!(hits[0].0, 0, "uniqzero lives in doc 0");
assert!(hits[0].1 > 0.0, "score must be positive");
}
#[tokio::test]
async fn df1_in_or_query_combines_with_df_ge_2() {
let (blob, json) = build_mixed_df_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["uniqtwo", "rust"], 10, BoolMode::Or)
.await
.expect("FTS search");
let ids: Vec<u32> = hits.iter().map(|(d, _)| *d).collect();
assert!(ids.contains(&0));
assert!(ids.contains(&1));
assert!(ids.contains(&2));
}
#[tokio::test]
async fn df1_in_and_query_intersects_correctly() {
let (blob, json) = build_mixed_df_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["uniqzero", "rust"], 10, BoolMode::And)
.await
.expect("FTS search");
let ids: Vec<u32> = hits.iter().map(|(d, _)| *d).collect();
assert_eq!(ids, vec![0]);
let hits = r
.search("body", &["uniqzero", "uniqtwo"], 10, BoolMode::And)
.await
.expect("FTS search");
assert!(hits.is_empty());
}
#[tokio::test]
async fn df1_missing_term_returns_empty() {
let (blob, json) = build_mixed_df_blob();
let r = FtsReader::open(blob, &json).expect("open FtsReader");
let hits = r
.search("body", &["nonexistentunique"], 10, BoolMode::Or)
.await
.expect("FTS search");
assert!(hits.is_empty());
}
#[test]
fn df1_inline_path_skips_postings_region_writes() {
let tok = Arc::new(AsciiLowerTokenizer);
let mut b_inline = FtsBuilder::new(tok.clone());
b_inline
.register_column("body".into(), false)
.expect("register column");
for i in 0..20 {
b_inline
.add_doc(0, i, &format!("uniq{i:03}"))
.expect("add doc");
}
let blob_inline = b_inline.finish().expect("finish inline");
let mut b_pfor = FtsBuilder::new(tok);
b_pfor
.register_column("body".into(), false)
.expect("register column");
for i in 0..20 {
let text = (0..20)
.map(|j| format!("uniq{j:03}"))
.collect::<Vec<_>>()
.join(" ");
b_pfor.add_doc(0, i, &text).expect("add doc");
}
let blob_pfor = b_pfor.finish().expect("finish pfor");
let postings_off_i = u64::from_le_bytes(
blob_inline[32..40]
.try_into()
.expect("postings_off_i slice is 8 bytes"),
) as usize;
let positions_off_i = u64::from_le_bytes(
blob_inline[48..56]
.try_into()
.expect("positions_off_i slice is 8 bytes"),
) as usize;
let postings_size_inline = positions_off_i - postings_off_i;
let postings_off_p = u64::from_le_bytes(
blob_pfor[32..40]
.try_into()
.expect("postings_off_p slice is 8 bytes"),
) as usize;
let positions_off_p = u64::from_le_bytes(
blob_pfor[48..56]
.try_into()
.expect("positions_off_p slice is 8 bytes"),
) as usize;
let postings_size_pfor = positions_off_p - postings_off_p;
assert_eq!(
postings_size_inline, 4,
"all-df=1 postings region should hold only the trailing CRC32; \
got {postings_size_inline} bytes"
);
assert!(
postings_size_pfor > 20 * 36,
"PFOR postings region should be hundreds of bytes; got {postings_size_pfor}"
);
}
async fn exclude_filter_for(reader: &FtsReader, terms: &[&str]) -> ExcludeFilter {
let column_id = reader.resolve_column_id("body").expect("column exists");
let cursors = reader
.build_term_cursors(column_id, terms)
.await
.expect("build cursors");
ExcludeFilter::new(cursors)
}
#[tokio::test]
async fn exclude_filter_rejects_docs_in_negated_list() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let mut f = exclude_filter_for(&r, &["rust"]).await;
assert!(!f.admits(0));
assert!(!f.admits(1));
assert!(f.admits(2));
}
#[tokio::test]
async fn exclude_filter_missing_term_excludes_nothing() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let mut f = exclude_filter_for(&r, &["nonexistent"]).await;
assert!(f.admits(0));
assert!(f.admits(1));
assert!(f.admits(2));
}
#[tokio::test]
async fn exclude_filter_multiple_negated_terms() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let mut f = exclude_filter_for(&r, &["rust", "java"]).await;
assert!(!f.admits(0));
assert!(!f.admits(1));
assert!(!f.admits(2));
}
#[tokio::test]
#[cfg(debug_assertions)]
#[should_panic(expected = "non-monotonic")]
async fn exclude_filter_panics_on_non_monotonic_feed() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let mut f = exclude_filter_for(&r, &["rust"]).await;
let _ = f.admits(1);
let _ = f.admits(0);
}
#[test]
fn open_with_verify_crc_off_succeeds() {
let (blob, json) = build_blob();
let r = FtsReader::open_with(blob, &json, OpenOptions { verify_crc: false })
.expect("open with crc off");
assert_eq!(r.n_docs(), 3);
assert_eq!(r.fts_columns().collect::<Vec<_>>(), vec!["body"]);
}
#[test]
fn open_with_object_store_options_matches_crc_off() {
let opts = OpenOptions::for_object_store();
assert!(!opts.verify_crc);
let (blob, json) = build_blob();
FtsReader::open_with(blob, &json, opts).expect("open object-store options");
}
#[test]
fn default_open_options_verifies_crc() {
assert!(OpenOptions::default().verify_crc);
}
#[test]
fn default_tokenizer_helper_is_ascii_lower() {
assert_eq!(default_tokenizer(), "ascii_lower");
}
#[test]
fn fts_column_config_missing_tokenizer_defaults() {
let (blob, _) = build_blob();
let json = r#"[{"name":"body"}]"#;
let r = FtsReader::open(blob, json).expect("open with terse json");
let cfg = r.fts_columns_config().next().expect("one column");
assert_eq!(cfg.name, "body");
}
#[test]
fn fts_columns_config_exposes_per_column_metadata() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let cols: Vec<&ColumnMeta> = r.fts_columns_config().collect();
assert_eq!(cols.len(), 1);
assert_eq!(cols[0].name, "body");
assert!(cols[0].avgdl > 0.0);
assert_eq!(cols[0].dl_norm_k1.len(), 3);
}
#[test]
fn norm_table_footprint_is_one_byte_per_doc() {
const N: u32 = 5_000;
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false)
.expect("register column");
for d in 0..N {
let words = (d % 40) + 1;
let text: String = (0..words).map(|w| format!("t{}x{w} ", d % 97)).collect();
b.add_doc(0, d, text.trim()).expect("add doc");
}
let bytes = b.finish().expect("finish");
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(Bytes::from(bytes), json).expect("open");
let nt = &r.columns[0].dl_norm_k1;
let per_doc = nt.bytes.capacity(); let lut = std::mem::size_of_val(&*nt.lut); let m2_bytes = per_doc + lut;
let f32_baseline = N as usize * std::mem::size_of::<f32>();
assert_eq!(nt.bytes.len(), N as usize, "one bucket byte per doc");
assert_eq!(nt.lut.len(), 256, "fixed 256-entry decode table");
assert!(
m2_bytes < f32_baseline,
"norm table {m2_bytes} B not smaller than f32 baseline {f32_baseline} B"
);
assert_eq!(
per_doc * 4,
f32_baseline,
"per-doc term is exactly 4× smaller"
);
}
#[test]
fn iter_column_terms_lists_every_term_in_lex_order() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let terms: Vec<String> = r
.iter_column_terms("body")
.expect("iter terms")
.into_iter()
.map(|b| String::from_utf8(b).expect("utf8"))
.collect();
let mut sorted = terms.clone();
sorted.sort();
assert_eq!(terms, sorted, "terms must be in lex order");
for expected in [
"rust", "async", "runtime", "tokio", "java", "spring", "boot",
] {
assert!(terms.contains(&expected.to_string()), "missing {expected}");
}
}
#[test]
fn iter_column_terms_unknown_column_is_empty() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
assert!(r.iter_column_terms("nope").expect("ok").is_empty());
}
#[test]
fn iter_terms_with_prefix_bounds_the_walk() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let terms: Vec<String> = r
.iter_terms_with_prefix("body", b"run")
.expect("prefix walk")
.into_iter()
.map(|b| String::from_utf8(b).expect("utf8"))
.collect();
assert_eq!(terms, vec!["runtime".to_string()]);
assert!(
r.iter_terms_with_prefix("body", b"zzz")
.expect("prefix walk")
.is_empty()
);
}
#[tokio::test]
async fn term_df_reports_document_frequency() {
let (blob, json) = build_mixed_df_blob();
let r = FtsReader::open(blob, &json).expect("open");
assert_eq!(r.term_df("body", "common").await.expect("df"), 3);
assert_eq!(r.term_df("body", "rust").await.expect("df"), 2);
assert_eq!(r.term_df("body", "uniqzero").await.expect("df"), 1);
assert_eq!(r.term_df("body", "missing").await.expect("df"), 0);
}
#[tokio::test]
async fn term_df_unknown_column_errors() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let err = r.term_df("nope", "rust").await.expect_err("error");
assert!(matches!(err, FtsError::UnknownColumn(_)));
}
fn build_phrase_blob() -> (Bytes, &'static str) {
use crate::superfile::fts::builder::FtsBuilder;
let mut b = FtsBuilder::new(crate::test_helpers::default_tokenizer());
b.register_column("title".into(), true).expect("register");
let docs = [
"new york city",
"york new haven",
"the new york times",
"new haven york",
"new york new york",
];
for (i, d) in docs.iter().enumerate() {
b.add_doc(0, i as u32, d).expect("add doc");
}
(
Bytes::from(b.finish().expect("finish")),
r#"[{"name":"title","tokenizer":"ascii_lower","positions":true}]"#,
)
}
fn phrase(terms: &[&str]) -> Vec<Vec<String>> {
vec![terms.iter().map(|t| t.to_string()).collect()]
}
#[tokio::test]
async fn phrase_matches_adjacent_in_order_only() {
let (blob, json) = build_phrase_blob();
let r = FtsReader::open(blob, json).expect("open");
let phrases = phrase(&["new", "york"]);
let hits = r
.search_excluding(
"title",
ClauseLists {
should_phrases: &phrases,
..ClauseLists::default()
},
10,
f32::NEG_INFINITY,
)
.await
.expect("phrase search");
let ids: Vec<u32> = hits.iter().map(|(d, _)| *d).collect();
let mut sorted = ids.clone();
sorted.sort_unstable();
assert_eq!(sorted, vec![0, 2, 4], "adjacency in order only");
assert_eq!(hits[0].0, 4, "double occurrence ranks first");
}
#[tokio::test]
async fn phrase_composes_with_clauses() {
let (blob, json) = build_phrase_blob();
let r = FtsReader::open(blob, json).expect("open");
let ny = phrase(&["new", "york"]);
let hits = r
.search_excluding(
"title",
ClauseLists {
musts: &["the"],
must_phrases: &ny,
..ClauseLists::default()
},
10,
f32::NEG_INFINITY,
)
.await
.expect("must phrase + term");
assert_eq!(
hits.iter().map(|(d, _)| *d).collect::<Vec<_>>(),
vec![2],
"+\"new york\" +the"
);
let hits = r
.search_excluding(
"title",
ClauseLists {
shoulds: &["haven"],
negative_phrases: &ny,
..ClauseLists::default()
},
10,
f32::NEG_INFINITY,
)
.await
.expect("negated phrase");
let mut ids: Vec<u32> = hits.iter().map(|(d, _)| *d).collect();
ids.sort_unstable();
assert_eq!(ids, vec![1, 3], "haven docs don't contain the phrase");
}
#[tokio::test]
async fn phrase_with_absent_member_matches_nothing() {
let (blob, json) = build_phrase_blob();
let r = FtsReader::open(blob, json).expect("open");
let ghost = phrase(&["new", "zealand"]);
let hits = r
.search_excluding(
"title",
ClauseLists {
must_phrases: &ghost,
..ClauseLists::default()
},
10,
f32::NEG_INFINITY,
)
.await
.expect("ghost phrase");
assert!(hits.is_empty());
}
#[tokio::test]
async fn phrase_on_positionless_column_is_typed_error() {
use crate::superfile::fts::builder::FtsBuilder;
let mut b = FtsBuilder::new(crate::test_helpers::default_tokenizer());
b.register_column("title".into(), false).expect("register");
b.add_doc(0, 0, "new york").expect("add doc");
let blob = Bytes::from(b.finish().expect("finish"));
let r =
FtsReader::open(blob, r#"[{"name":"title","tokenizer":"ascii_lower"}]"#).expect("open");
let phrases = phrase(&["new", "york"]);
let err = r
.search_excluding(
"title",
ClauseLists {
should_phrases: &phrases,
..ClauseLists::default()
},
10,
f32::NEG_INFINITY,
)
.await
.expect_err("must be a typed error");
assert!(matches!(err, FtsError::PositionsUnavailable { .. }));
}
#[tokio::test]
async fn search_excluding_drops_negated_docs() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let hits = r
.search_excluding(
"body",
ClauseLists {
shoulds: &["runtime"],
negatives: &["async"],
..ClauseLists::default()
},
10,
f32::NEG_INFINITY,
)
.await
.expect("search excluding");
let ids: Vec<u32> = hits.iter().map(|(d, _)| *d).collect();
assert_eq!(ids, vec![1], "doc 0 excluded by negated 'async'");
}
#[tokio::test]
async fn search_excluding_negation_only_errors() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let err = r
.search_excluding(
"body",
ClauseLists {
negatives: &["rust"],
..ClauseLists::default()
},
10,
f32::NEG_INFINITY,
)
.await
.expect_err("negation-only");
assert!(matches!(err, FtsError::NegationOnly));
}
#[tokio::test]
async fn search_excluding_no_terms_at_all_is_empty() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let hits = r
.search_excluding("body", ClauseLists::default(), 10, f32::NEG_INFINITY)
.await
.expect("empty");
assert!(hits.is_empty());
}
#[tokio::test]
async fn search_with_floor_prunes_below_floor() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
let hits = r
.search_with_floor("body", &["rust"], 10, BoolMode::Or, 1e9)
.await
.expect("floored search");
assert!(hits.is_empty(), "floor above all scores prunes everything");
}
#[tokio::test]
async fn search_multi_weights_and_combines_columns() {
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("title".into(), false).expect("register");
b.register_column("body".into(), false).expect("register");
b.add_doc(0, 0, "rust").expect("add");
b.add_doc(1, 0, "systems").expect("add");
b.add_doc(0, 1, "python").expect("add");
b.add_doc(1, 1, "rust ml").expect("add");
b.add_doc(0, 2, "go").expect("add");
b.add_doc(1, 2, "concurrency").expect("add");
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"title","tokenizer":"ascii_lower"},{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let hits = r
.search_multi(&[("title", 1.0), ("body", 1.0)], "rust", 10, BoolMode::Or)
.await
.expect("multi");
let ids: HashSet<u32> = hits.iter().map(|(d, _)| *d).collect();
assert!(ids.contains(&0));
assert!(ids.contains(&1));
assert!(!ids.contains(&2));
}
#[tokio::test]
async fn search_or_range_restricts_to_doc_id_window() {
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
for i in 0..8u32 {
b.add_doc(0, i, "alpha beta").expect("add");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let hits = r
.search_or_range_pretokenized("body", &["alpha", "beta"], 100, 2, 5)
.await
.expect("ranged search");
let ids: HashSet<u32> = hits.iter().map(|(d, _)| *d).collect();
assert_eq!(
ids,
[2u32, 3, 4].into_iter().collect(),
"only docs in [2,5) returned"
);
}
#[tokio::test]
async fn search_or_range_degenerate_inputs_are_empty() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
assert!(
r.search_or_range_pretokenized("body", &[], 10, 0, 3)
.await
.expect("empty terms")
.is_empty()
);
assert!(
r.search_or_range_pretokenized("body", &["rust"], 0, 0, 3)
.await
.expect("zero k")
.is_empty()
);
assert!(
r.search_or_range_pretokenized("body", &["rust"], 10, 3, 3)
.await
.expect("empty range")
.is_empty()
);
}
#[tokio::test]
async fn search_or_range_with_floor_prunes() {
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
for i in 0..8u32 {
b.add_doc(0, i, "alpha beta").expect("add");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let hits = r
.search_or_range_pretokenized_with_floor("body", &["alpha", "beta"], 100, 0, 8, 1e9)
.await
.expect("floored ranged search");
assert!(hits.is_empty(), "floor above all scores prunes everything");
}
#[tokio::test]
async fn search_with_algo_wand_bmw_agrees_with_bmm() {
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
let docs = [
"alpha beta",
"alpha",
"beta gamma",
"alpha beta gamma",
"gamma",
"alpha gamma",
"beta",
"alpha beta gamma",
];
for (i, t) in docs.iter().enumerate() {
b.add_doc(0, i as u32, t).expect("add");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let terms: &[&str] = &["alpha", "beta", "gamma"];
let bmm = r
.search_with_algo_for_bench("body", terms, 5, OrAlgo::Bmm)
.await
.expect("bmm");
let wand = r
.search_with_algo_for_bench("body", terms, 5, OrAlgo::WandBmw)
.await
.expect("wand");
assert_eq!(bmm.len(), wand.len());
for ((db, sb), (dw, sw)) in bmm.iter().zip(wand.iter()) {
assert_eq!(db, dw, "doc_id mismatch");
assert!((sb - sw).abs() < 1e-4, "score mismatch {sb} vs {sw}");
}
}
#[tokio::test]
async fn wand_bmw_exercises_block_skips_on_multi_block_lists() {
const N_DOCS: u32 = 400;
const K: usize = 5;
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
for i in 0..N_DOCS {
let mut text = String::new();
text.push_str("alpha ");
if i % 2 == 0 {
text.push_str("beta ");
}
if i % 5 == 0 {
text.push_str("gamma ");
}
if i % 13 == 0 {
text.push_str("delta ");
}
if i % 29 == 0 {
text.push_str("epsilon ");
}
b.add_doc(0, i, text.trim()).expect("add doc");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let terms: &[&str] = &["alpha", "beta", "gamma", "delta", "epsilon"];
let wand = r
.search_with_algo_for_bench("body", terms, K, OrAlgo::WandBmw)
.await
.expect("wand");
let bmm = r
.search_with_algo_for_bench("body", terms, K, OrAlgo::Bmm)
.await
.expect("bmm");
assert_eq!(wand.len(), bmm.len(), "result length mismatch");
assert_eq!(wand.len(), K, "expected a full top-K");
for ((dw, sw), (db, sb)) in wand.iter().zip(bmm.iter()) {
assert_eq!(dw, db, "doc_id mismatch wand={dw} bmm={db}");
assert!((sw - sb).abs() < 1e-4, "score mismatch {sw} vs {sb}");
}
}
#[tokio::test]
async fn windowed_union_agrees_with_bmm() {
const N_DOCS: u32 = OR_WINDOW * 2 + 500;
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
for i in 0..N_DOCS {
let mut text = String::from("alpha zeta eta theta "); if i % 2 == 0 {
text.push_str("beta ");
}
if i % 3 == 0 {
text.push_str("gamma ");
}
if i % 5 == 0 {
text.push_str("delta ");
}
if i % 7 == 0 {
text.push_str("epsilon ");
}
b.add_doc(0, i, text.trim()).expect("add doc");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let col = r.resolve_column_id("body").expect("col");
let uniform_terms: &[&str] = &["zeta", "eta", "theta"];
let uniform_cursors = r
.build_term_cursors(col, uniform_terms)
.await
.expect("uniform cursors");
assert!(
prefer_windowed_union(&uniform_cursors),
"production router should select windowed union for equal upper bounds"
);
let shapes: &[&[&str]] = &[
&["alpha", "beta"],
&["alpha", "beta", "gamma"],
&["beta", "gamma", "delta"], &["alpha", "beta", "gamma", "delta", "epsilon"],
uniform_terms,
];
for terms in shapes {
for k in [1usize, 5, 50, 1000] {
let bmm = r
.search_with_algo_for_bench("body", terms, k, OrAlgo::Bmm)
.await
.expect("bmm");
let win = r
.search_with_algo_for_bench("body", terms, k, OrAlgo::Windowed)
.await
.expect("windowed");
assert_eq!(bmm.len(), win.len(), "len mismatch {terms:?} k={k}");
for ((db, sb), (dw, sw)) in bmm.iter().zip(win.iter()) {
assert_eq!(db, dw, "doc_id mismatch {terms:?} k={k}: bmm={db} win={dw}");
assert!(
(sb - sw).abs() < 1e-4,
"score mismatch {terms:?} k={k}: {sb} vs {sw}"
);
}
}
}
}
#[tokio::test]
async fn wand_bmw_2term_no_floor_agrees_with_bmm() {
const N_DOCS: u32 = OR_WINDOW * 2 + 500;
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
for i in 0..N_DOCS {
let mut text = String::from("alpha ");
if i % 2 == 0 {
text.push_str("beta ");
}
if i % 3 == 0 {
text.push_str("gamma ");
}
b.add_doc(0, i, text.trim()).expect("add doc");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let col = r.resolve_column_id("body").expect("col");
let shapes: &[&[&str]] = &[&["alpha", "beta"], &["beta", "gamma"], &["alpha", "gamma"]];
for terms in shapes {
for k in [1usize, 5, 50, 128] {
let cw = r.build_term_cursors(col, terms).await.expect("cursors");
let cb = r.build_term_cursors(col, terms).await.expect("cursors");
let wand = r.run_wand_bmw(col, cw, k).expect("wand");
let bmm = r
.run_max_score_bmm(col, cb, k, None, f32::NEG_INFINITY)
.expect("bmm");
assert_eq!(wand.len(), bmm.len(), "len mismatch {terms:?} k={k}");
for ((dw, sw), (db, sb)) in wand.iter().zip(bmm.iter()) {
assert_eq!(dw, db, "doc mismatch {terms:?} k={k}: {dw} vs {db}");
assert!(
(sw - sb).abs() < 1e-4,
"score mismatch {terms:?} k={k}: {sw} vs {sb}"
);
}
}
}
}
#[tokio::test]
async fn two_term_rare_anchor_gates_on_df_ratio() {
const N_DOCS: u32 = 4000;
let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
for i in 0..N_DOCS {
let mut text = String::from("common "); if i % 2 == 0 {
text.push_str("frequent "); }
if i % 200 == 0 {
text.push_str("rare "); }
b.add_doc(0, i, text.trim()).expect("add doc");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let col = r.resolve_column_id("body").expect("col");
let anchored = r
.build_term_cursors(col, &["common", "rare"])
.await
.expect("cursors");
assert!(
two_term_has_rare_anchor(&anchored),
"rare+common should have a rare anchor"
);
let uniform = r
.build_term_cursors(col, &["common", "frequent"])
.await
.expect("cursors");
assert!(
!two_term_has_rare_anchor(&uniform),
"two common terms should not anchor"
);
}
#[tokio::test]
async fn windowed_union_negation_agrees_with_bmm() {
const N_DOCS: u32 = OR_WINDOW + 1000; let tok = Arc::new(AsciiLowerTokenizer);
let mut b = FtsBuilder::new(tok);
b.register_column("body".into(), false).expect("register");
for i in 0..N_DOCS {
let mut text = String::from("alpha ");
if i % 2 == 0 {
text.push_str("beta ");
}
if i % 3 == 0 {
text.push_str("gamma ");
}
if i % 5 == 0 {
text.push_str("delta ");
}
if i % 7 == 0 {
text.push_str("epsilon ");
}
b.add_doc(0, i, text.trim()).expect("add doc");
}
let blob = Bytes::from(b.finish().expect("finish"));
let json = r#"[{"name":"body","tokenizer":"ascii_lower"}]"#;
let r = FtsReader::open(blob, json).expect("open");
let col = r.resolve_column_id("body").expect("col");
let cases: &[(&[&str], &[&str])] = &[
(&["alpha", "beta", "gamma"], &["delta"]),
(&["beta", "gamma", "delta"], &["epsilon"]),
(&["alpha", "beta", "gamma", "delta"], &["epsilon", "gamma"]),
];
for (pos, neg) in cases {
for k in [1usize, 5, 50] {
let mut wf =
ExcludeFilter::new(r.build_term_cursors(col, neg).await.expect("neg cursors"));
let win = r
.run_windowed_union(
col,
r.build_term_cursors(col, pos).await.expect("pos cursors"),
k,
Some(&mut wf),
f32::NEG_INFINITY,
0,
u32::MAX,
)
.expect("windowed");
let mut bf =
ExcludeFilter::new(r.build_term_cursors(col, neg).await.expect("neg cursors"));
let bmm = r
.run_max_score_bmm(
col,
r.build_term_cursors(col, pos).await.expect("pos cursors"),
k,
Some(&mut bf),
f32::NEG_INFINITY,
)
.expect("bmm");
assert_eq!(win.len(), bmm.len(), "len {pos:?} -{neg:?} k={k}");
for ((dw, sw), (db, sb)) in win.iter().zip(bmm.iter()) {
assert_eq!(
dw, db,
"doc mismatch {pos:?} -{neg:?} k={k}: win={dw} bmm={db}"
);
assert!(
(sw - sb).abs() < 1e-4,
"score mismatch {pos:?} -{neg:?} k={k}: {sw} vs {sb}"
);
}
}
}
let pos: &[&str] = &["alpha", "beta", "gamma"];
let neg: &[&str] = &["delta"];
let unfiltered = r
.run_windowed_union(
col,
r.build_term_cursors(col, pos).await.expect("pos"),
N_DOCS as usize,
None,
f32::NEG_INFINITY,
0,
u32::MAX,
)
.expect("unfiltered");
let mut f = ExcludeFilter::new(r.build_term_cursors(col, neg).await.expect("neg"));
let filtered = r
.run_windowed_union(
col,
r.build_term_cursors(col, pos).await.expect("pos"),
N_DOCS as usize,
Some(&mut f),
f32::NEG_INFINITY,
0,
u32::MAX,
)
.expect("filtered");
assert!(
filtered.len() < unfiltered.len(),
"negation should drop docs: filtered={} unfiltered={}",
filtered.len(),
unfiltered.len()
);
}
#[tokio::test]
async fn search_with_algo_empty_and_zero_k_short_circuit() {
let (blob, json) = build_blob();
let r = FtsReader::open(blob, &json).expect("open");
assert!(
r.search_with_algo_for_bench("body", &[], 5, OrAlgo::Bmm)
.await
.expect("empty")
.is_empty()
);
assert!(
r.search_with_algo_for_bench("body", &["rust"], 0, OrAlgo::Exhaustive)
.await
.expect("zero k")
.is_empty()
);
}
#[test]
fn read_u32_le_and_u64_le_decode_little_endian() {
let b32 = [0x78, 0x56, 0x34, 0x12];
assert_eq!(read_u32_le(&b32), 0x1234_5678);
let b64 = [0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00];
assert_eq!(read_u64_le(&b64), 1);
}
#[test]
fn top_k_keeps_highest_scores_with_doc_id_tiebreak() {
let mut scores: HashMap<u32, f32> = HashMap::new();
scores.insert(0, 1.0);
scores.insert(1, 3.0);
scores.insert(2, 2.0);
scores.insert(3, 3.0); let out = top_k(scores, 2);
assert_eq!(out, vec![(1, 3.0), (3, 3.0)]);
}
#[test]
fn top_k_smaller_than_k_returns_all_sorted() {
let mut scores: HashMap<u32, f32> = HashMap::new();
scores.insert(5, 2.0);
scores.insert(9, 5.0);
let out = top_k(scores, 10);
assert_eq!(out, vec![(9, 5.0), (5, 2.0)]);
}
#[test]
fn drain_top_k_desc_orders_descending_with_tiebreak() {
let mut heap: BinaryHeap<TopKEntry> = BinaryHeap::new();
heap.push(TopKEntry(1.0, 4));
heap.push(TopKEntry(2.0, 1));
heap.push(TopKEntry(2.0, 0)); let out = drain_top_k_desc(heap);
assert_eq!(out, vec![(0, 2.0), (1, 2.0), (4, 1.0)]);
}
#[tokio::test]
async fn open_lazy_round_trips_a_search() {
let (blob, json) = build_blob();
let src: Arc<dyn LazyByteSource> = Arc::new(BytesLazyByteSource::new(blob));
let r = FtsReader::open_lazy(src, &json, OpenOptions::for_object_store())
.await
.expect("open_lazy");
assert_eq!(r.n_docs(), 3);
let hits = r
.search("body", &["rust"], 10, BoolMode::Or)
.await
.expect("search over lazy reader");
let ids: HashSet<u32> = hits.iter().map(|(d, _)| *d).collect();
assert!(ids.contains(&0) && ids.contains(&1));
}
}