diff --git a/rust/lance-index/src/scalar/inverted.rs b/rust/lance-index/src/scalar/inverted.rs index 6a95876d174..ec7c22021cc 100644 --- a/rust/lance-index/src/scalar/inverted.rs +++ b/rust/lance-index/src/scalar/inverted.rs @@ -27,6 +27,7 @@ pub use compound::{ compound_search, compound_search_prepared_match, compound_search_prepared_match_with_score_floor, compound_search_with_base_scorer, compound_search_with_base_scorer_and_score_floor, exclusive_scaled_score_floor, + materialized_compound_top_k, }; #[doc(hidden)] pub use cross_column::cross_column_compound_search; diff --git a/rust/lance-index/src/scalar/inverted/compound.rs b/rust/lance-index/src/scalar/inverted/compound.rs index ed774c6337f..7bbba53ee52 100644 --- a/rust/lance-index/src/scalar/inverted/compound.rs +++ b/rust/lance-index/src/scalar/inverted/compound.rs @@ -2162,6 +2162,40 @@ impl TopKCollector { } } +/// Evaluate a compound query over exact, materialized leaf result sets. +/// +/// This is the bridge used by query-local residual postings: it keeps Boolean, +/// Boost, and MultiMatch semantics in the same scorer tree as the on-disk +/// compound path while allowing a different posting source. +#[doc(hidden)] +pub fn materialized_compound_top_k( + query: &FtsQuery, + leaves: Vec>, + limit: usize, +) -> Result<(Vec, Vec)> { + let mut leaf_count = 0; + let plan = CompoundScorerPlan::from_query(query, &mut leaf_count)?; + if leaf_count != leaves.len() { + return Err(Error::internal(format!( + "compound FTS planned {leaf_count} leaves but received {} materialized leaves", + leaves.len() + ))); + } + let mut scorers = leaves + .into_iter() + .map(|rows| { + let rows = rows + .into_iter() + .map(|(row_id, score)| ScoredRow { row_id, score }) + .collect(); + MaterializedScorer::try_new(rows).map(|scorer| Some(Box::new(scorer) as BoxScorer<'_>)) + }) + .collect::>>()?; + let mut scorer = plan.build(&mut scorers)?; + let rows = TopKCollector::new(limit).collect(scorer.as_mut())?; + Ok(rows.into_iter().map(|row| (row.row_id, row.score)).unzip()) +} + #[derive(Debug, Clone, Copy)] pub(super) enum DisjunctionScore { Sum, @@ -4414,6 +4448,7 @@ mod tests { use super::super::scorer::Scorer; use super::*; use crate::metrics::NoOpMetricsCollector; + use crate::scalar::inverted::query::MultiMatchQuery; fn rows(values: &[(u64, f32)]) -> Vec { values @@ -4426,6 +4461,25 @@ mod tests { Box::new(MaterializedScorer::try_new(rows(values)).unwrap()) } + #[test] + fn materialized_compound_top_k_preserves_multimatch_and_tie_order() { + let query = FtsQuery::MultiMatch(MultiMatchQuery { + match_queries: vec![ + MatchQuery::new("alpha".to_string()).with_column(Some("text".to_string())), + MatchQuery::new("alpha".to_string()).with_column(Some("text".to_string())), + ], + }); + let (row_ids, scores) = materialized_compound_top_k( + &query, + vec![vec![(7, 1.0), (3, 2.0)], vec![(7, 3.0), (5, 3.0)]], + 2, + ) + .unwrap(); + + assert_eq!(row_ids, vec![5, 7]); + assert_eq!(scores, vec![3.0, 3.0]); + } + fn zero_weight_wand<'a>( documents: &'a DocSet, scorer: Arc, diff --git a/rust/lance/src/dataset/mem_wal/index.rs b/rust/lance/src/dataset/mem_wal/index.rs index 3daf3e1274a..2735511faba 100644 --- a/rust/lance/src/dataset/mem_wal/index.rs +++ b/rust/lance/src/dataset/mem_wal/index.rs @@ -50,6 +50,7 @@ pub type RowPosition = u64; // Re-export public types used externally pub use btree::{BTreeIndexConfig, BTreeMemIndex}; pub use fts::{FtsIndexConfig, FtsMemIndex, FtsQueryExpr, SearchOptions}; +pub(crate) use fts::{QueryLocalFtsIndex, QueryLocalFtsStats}; pub use hnsw::{HnswIndexConfig, HnswMemIndex}; pub use pk_key::encode_pk_tuple; diff --git a/rust/lance/src/dataset/mem_wal/index/fts.rs b/rust/lance/src/dataset/mem_wal/index/fts.rs index fb34b9a2f79..50ab8cc3b36 100644 --- a/rust/lance/src/dataset/mem_wal/index/fts.rs +++ b/rust/lance/src/dataset/mem_wal/index/fts.rs @@ -51,19 +51,19 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex, OnceLock}; use arc_swap::ArcSwap; -use arrow_array::RecordBatch; +use arrow_array::{Array, RecordBatch, UInt64Array}; use crossbeam_skiplist::SkipMap; use fst::{Map, Streamer}; use lance_bitpacking::{BitPacker, BitPacker4x}; use lance_core::datatypes::Schema as LanceSchema; use lance_core::{Error, Result}; use lance_index::scalar::InvertedIndexParams; -use lance_index::scalar::inverted::query::{Operator, Tokens}; +use lance_index::scalar::inverted::query::{FtsQuery, Operator, Tokens}; use lance_index::scalar::inverted::tokenizer::document_tokenizer::{DocType, LanceTokenizer}; use lance_index::scalar::inverted::{DocSet, MemBM25Scorer, Scorer, TokenSet}; use lance_tokenizer::TokenStream; use rayon::prelude::*; -use rustc_hash::FxHashMap; +use rustc_hash::{FxHashMap, FxHashSet}; use super::RowPosition; use crate::index::scalar::inverted::{ResolvedFtsField, resolve_fts_field}; @@ -770,12 +770,15 @@ impl std::fmt::Debug for TokenizerPool { impl TokenizerPool { fn new(params: &InvertedIndexParams, cap: usize) -> Result { - let template = params.build()?; - Ok(Self { + Ok(Self::from_template(params.build()?, cap)) + } + + fn from_template(template: Box, cap: usize) -> Self { + Self { template, free: Mutex::new(Vec::new()), cap: cap.max(1), - }) + } } /// Acquire a tokenizer. Pops from the free list, otherwise clones the @@ -1003,6 +1006,11 @@ pub struct FtsMemIndex { /// The tail freezes into a partition once it reaches this many docs. freeze_threshold_rows: usize, + /// Query-local materializations disable freezes and tiered merges. Their + /// lifetime is bounded by one query, so background maintenance would only + /// outlive cancellation without providing reuse. + background_maintenance: bool, + /// Background tiered-merge slot. `None` = idle; `Some` with `result: None` /// = a merge is running on a worker thread; `Some` with `result: Some` = /// the merged partition is ready for the writer to install. Only the @@ -1011,6 +1019,154 @@ pub struct FtsMemIndex { merge: Arc>>, } +/// Query-owned term-only postings for one residual scan. +/// +/// This deliberately exposes only the immutable feature-materialization API +/// needed by hybrid execution. Unlike [`FtsMemIndex`], it never freezes or +/// starts a detached tiered merge; dropping the query drops all residual +/// postings. +#[derive(Debug)] +pub struct QueryLocalFtsIndex { + inner: FtsMemIndex, +} + +#[derive(Debug, Default)] +pub struct QueryLocalFtsStats { + doc_count: usize, + total_tokens: u64, + token_docs: FxHashMap, +} + +impl QueryLocalFtsStats { + pub(crate) fn checked_add_assign(&mut self, other: Self) -> Result<()> { + self.doc_count = self + .doc_count + .checked_add(other.doc_count) + .ok_or_else(|| Error::internal("query-local FTS document count overflow"))?; + self.total_tokens = self + .total_tokens + .checked_add(other.total_tokens) + .ok_or_else(|| Error::internal("query-local FTS total token count overflow"))?; + for (token, df) in other.token_docs { + let current = self.token_docs.entry(token).or_default(); + *current = current + .checked_add(df) + .ok_or_else(|| Error::internal("query-local FTS term document count overflow"))?; + } + Ok(()) + } + + pub(crate) fn add_to_scorer(&self, scorer: &mut MemBM25Scorer) -> Result<()> { + scorer.num_docs = scorer + .num_docs + .checked_add(self.doc_count) + .ok_or_else(|| Error::internal("residual BM25 document count overflow"))?; + scorer.total_tokens = scorer + .total_tokens + .checked_add(self.total_tokens) + .ok_or_else(|| Error::internal("residual BM25 total token count overflow"))?; + for (token, df) in &self.token_docs { + let current = scorer.token_docs.entry(token.clone()).or_default(); + *current = current + .checked_add(*df) + .ok_or_else(|| Error::internal("residual BM25 term document count overflow"))?; + } + Ok(()) + } +} + +impl QueryLocalFtsIndex { + #[cfg(test)] + pub(crate) fn try_with_params( + field_id: i32, + column_name: String, + params: InvertedIndexParams, + ) -> Result { + Ok(Self { + inner: FtsMemIndex::try_with_params_and_maintenance( + field_id, + column_name, + params, + false, + )?, + }) + } + + pub(crate) fn try_with_loaded_tokenizer( + field_id: i32, + column_name: String, + params: InvertedIndexParams, + tokenizer: Box, + ) -> Result { + params.validate_format_version()?; + let pool = TokenizerPool::from_template(tokenizer, FtsMemIndex::DEFAULT_TOKENIZER_POOL_CAP); + Ok(Self { + inner: FtsMemIndex::with_tokenizer_pool_and_maintenance( + field_id, + column_name, + params, + pool, + false, + ), + }) + } + + /// Create an empty query-local shard without rebuilding tokenizer assets. + /// + /// The tokenizer pool and its loaded template are shared with the seed; + /// each shard only clones a writer tokenizer from that in-memory template. + pub(crate) fn empty_sibling(&self) -> Self { + let resolved_field = OnceLock::new(); + if let Some(resolved) = self.inner.resolved_field.get() { + resolved_field + .set(resolved.clone()) + .expect("new query-local shard traversal is empty"); + } + + Self { + inner: FtsMemIndex { + field_id: self.inner.field_id, + source_column_name: self.inner.source_column_name.clone(), + params: self.inner.params.clone(), + resolved_field, + tokenizer_pool: self.inner.tokenizer_pool.clone(), + writer_tokenizer: Mutex::new(self.inner.tokenizer_pool.acquire()), + state: ArcSwap::from(IndexState::empty()), + freeze_threshold_rows: self.inner.freeze_threshold_rows, + background_maintenance: false, + merge: Arc::new(Mutex::new(None)), + }, + } + } + + pub(crate) fn exact_query_terms(&self, query: &FtsQuery) -> Result> { + self.inner.exact_query_terms(query) + } + + pub(crate) fn insert_with_row_ids_for_terms( + &self, + batch: &RecordBatch, + row_ids: &UInt64Array, + terms: &FxHashSet, + ) -> Result { + self.inner + .insert_with_row_ids_for_terms(batch, row_ids, terms) + } + + #[cfg(test)] + fn doc_count(&self) -> usize { + self.inner.doc_count() + } + + pub(crate) fn exact_leaf_results( + &self, + query: &FtsQuery, + scorer: &MemBM25Scorer, + ) -> Result>> { + self.inner.exact_leaf_results(query, scorer) + } +} + /// A tiered merge dispatched to a background worker. struct PendingMerge { /// `Arc::as_ptr` of each source partition, for identity-matching the @@ -1090,11 +1246,36 @@ impl FtsMemIndex { field_id: i32, column_name: String, params: InvertedIndexParams, + ) -> Result { + Self::try_with_params_and_maintenance(field_id, column_name, params, true) + } + + fn try_with_params_and_maintenance( + field_id: i32, + column_name: String, + params: InvertedIndexParams, + background_maintenance: bool, ) -> Result { params.validate_format_version()?; let pool = TokenizerPool::new(¶ms, Self::DEFAULT_TOKENIZER_POOL_CAP)?; + Ok(Self::with_tokenizer_pool_and_maintenance( + field_id, + column_name, + params, + pool, + background_maintenance, + )) + } + + fn with_tokenizer_pool_and_maintenance( + field_id: i32, + column_name: String, + params: InvertedIndexParams, + pool: TokenizerPool, + background_maintenance: bool, + ) -> Self { let writer_tokenizer = pool.template.box_clone(); - Ok(Self { + Self { field_id, source_column_name: column_name, params, @@ -1103,8 +1284,9 @@ impl FtsMemIndex { writer_tokenizer: Mutex::new(writer_tokenizer), state: ArcSwap::from(IndexState::empty()), freeze_threshold_rows: Self::DEFAULT_FREEZE_THRESHOLD_ROWS, + background_maintenance, merge: Arc::new(Mutex::new(None)), - }) + } } pub(crate) fn try_with_resolved_field( @@ -1256,7 +1438,46 @@ impl FtsMemIndex { self.insert_batch(batch, row_offset) } + /// Insert explicit, potentially non-contiguous rows while retaining + /// postings only for query terms. + /// The tokenizer still visits the complete document so BM25 document + /// lengths remain accurate when scoring with committed-index statistics. + pub(crate) fn insert_with_row_ids_for_terms( + &self, + batch: &RecordBatch, + row_ids: &UInt64Array, + terms: &FxHashSet, + ) -> Result { + if row_ids.len() != batch.num_rows() || row_ids.null_count() != 0 { + return Err(Error::invalid_input(format!( + "MemWAL FTS explicit row ids require {} non-null values, got len={} nulls={}", + batch.num_rows(), + row_ids.len(), + row_ids.null_count() + ))); + } + self.insert_batch_with_keys(batch, |row_index| Ok(row_ids.value(row_index)), Some(terms)) + } + fn insert_batch(&self, batch: &RecordBatch, row_offset: u64) -> Result<()> { + self.insert_batch_with_keys( + batch, + |row_index| { + row_offset + .checked_add(row_index as u64) + .ok_or_else(|| Error::invalid_input("MemWAL FTS row position overflow")) + }, + None, + ) + .map(|_| ()) + } + + fn insert_batch_with_keys( + &self, + batch: &RecordBatch, + row_position: impl Fn(usize) -> Result, + allowed_terms: Option<&FxHashSet>, + ) -> Result { let st = self.state.load_full(); let document_position_start = st.tail.doc_count(); if self.resolved_field.get().is_none() { @@ -1284,14 +1505,46 @@ impl FtsMemIndex { // per-document map and per-`(term, doc)` `Vec` allocation that // dominated insert cost. `FxHashMap` skips SipHash on the hot lookup. let mut term_builders: FxHashMap, BatchTermBuilder> = FxHashMap::default(); - let mut documents: Vec = Vec::with_capacity(batch.num_rows()); + let mut documents: Vec = if allowed_terms.is_some() { + Vec::new() + } else { + Vec::with_capacity(batch.num_rows()) + }; let mut total_tokens: u64 = 0; + let mut query_local_corpus_doc_count = 0usize; + let mut query_local_corpus_total_tokens = 0u64; let preserve_zero_token_documents = self.params.get_document_granularity().is_list_element(); let mut index_document = |key: DocumentKey, text: &str| -> Result<()> { let document_position = document_position_start + documents.len() as u64; - let num_tokens = index_text(text, document_position, tokenizer, &mut term_builders)?; - if preserve_zero_token_documents || num_tokens > 0 { + let (num_tokens, retained_term) = match allowed_terms { + Some(allowed_terms) => index_text_filtered( + text, + document_position, + tokenizer, + &mut term_builders, + allowed_terms, + )?, + None => ( + index_text(text, document_position, tokenizer, &mut term_builders)?, + false, + ), + }; + let belongs_in_corpus = preserve_zero_token_documents || num_tokens > 0; + if allowed_terms.is_some() && belongs_in_corpus { + query_local_corpus_doc_count = query_local_corpus_doc_count + .checked_add(1) + .ok_or_else(|| Error::internal("query-local FTS document count overflow"))?; + query_local_corpus_total_tokens = query_local_corpus_total_tokens + .checked_add(num_tokens as u64) + .ok_or_else(|| Error::internal("query-local FTS total token count overflow"))?; + } + let retain_document = if allowed_terms.is_some() { + retained_term + } else { + belongs_in_corpus + }; + if retain_document { documents.push(DocumentMetadata { key, num_tokens }); total_tokens += num_tokens as u64; } @@ -1301,15 +1554,28 @@ impl FtsMemIndex { for document in extracted_documents { index_document( DocumentKey { - row_position: row_offset + document.row_index as u64, + row_position: row_position(document.row_index)?, doc_index: document.doc_index, }, &document.text, )?; } + let query_local_stats = if allowed_terms.is_some() { + QueryLocalFtsStats { + doc_count: query_local_corpus_doc_count, + total_tokens: query_local_corpus_total_tokens, + token_docs: term_builders + .iter() + .map(|(term, builder)| (term.to_string(), builder.row_positions.len())) + .collect(), + } + } else { + QueryLocalFtsStats::default() + }; + if documents.is_empty() { - return Ok(()); + return Ok(query_local_stats); } // Drop the tokenizer guard before publishing so we don't hold it @@ -1326,10 +1592,124 @@ impl FtsMemIndex { self.params.has_positions(), ); - if st.tail.doc_count() >= self.freeze_threshold_rows as u64 { + if self.background_maintenance && st.tail.doc_count() >= self.freeze_threshold_rows as u64 { self.freeze(&st)?; } - Ok(()) + Ok(query_local_stats) + } + + /// Analyze every exact leaf and return the deduplicated query terms in + /// canonical leaf traversal order. + pub(crate) fn exact_query_terms(&self, query: &FtsQuery) -> Result> { + fn visit(index: &FtsMemIndex, query: &FtsQuery, terms: &mut Vec) -> Result<()> { + match query { + FtsQuery::Match(query) => { + if query.fuzziness != Some(0) { + return Err(Error::invalid_input( + "residual compound FTS only supports exact Match leaves", + )); + } + terms.extend(index.analyze_for_search(&query.terms)); + } + FtsQuery::Phrase(query) => { + terms.extend(index.analyze_for_search(&query.terms)); + } + FtsQuery::Boost(query) => { + visit(index, &query.positive, terms)?; + visit(index, &query.negative, terms)?; + } + FtsQuery::MultiMatch(query) => { + for query in &query.match_queries { + visit(index, &FtsQuery::Match(query.clone()), terms)?; + } + } + FtsQuery::Boolean(query) => { + for query in query + .should + .iter() + .chain(&query.must) + .chain(&query.must_not) + { + visit(index, query, terms)?; + } + } + } + Ok(()) + } + + let mut terms = Vec::new(); + visit(self, query, &mut terms)?; + let mut seen = HashSet::with_capacity(terms.len()); + terms.retain(|term| seen.insert(term.clone())); + Ok(terms) + } + + /// Materialize each exact leaf with a caller-supplied scorer. Compound + /// semantics are deliberately evaluated by the canonical lance-index + /// scorer instead of being duplicated here. + pub(crate) fn exact_leaf_results( + &self, + query: &FtsQuery, + scorer: &MemBM25Scorer, + ) -> Result>> { + fn visit( + index: &FtsMemIndex, + query: &FtsQuery, + scorer: &MemBM25Scorer, + leaves: &mut Vec>, + ) -> Result<()> { + match query { + FtsQuery::Match(query) => { + if query.fuzziness != Some(0) { + return Err(Error::invalid_input( + "residual compound FTS only supports exact Match leaves", + )); + } + let st = index.state.load_full(); + let tokens = index.analyze_for_search(&query.terms); + let rows = index + .search_match_with_scorer(&st, &tokens, query.operator, scorer) + .into_iter() + .map(|entry| (entry.row_position, entry.score)) + .collect(); + leaves.push(rows); + } + FtsQuery::Phrase(query) => { + let st = index.state.load_full(); + let tokens = index.analyze_for_search(&query.terms); + let rows = index + .search_phrase_with_scorer(&st, &tokens, query.slop, scorer) + .into_iter() + .map(|entry| (entry.row_position, entry.score)) + .collect(); + leaves.push(rows); + } + FtsQuery::Boost(query) => { + visit(index, &query.positive, scorer, leaves)?; + visit(index, &query.negative, scorer, leaves)?; + } + FtsQuery::MultiMatch(query) => { + for query in &query.match_queries { + visit(index, &FtsQuery::Match(query.clone()), scorer, leaves)?; + } + } + FtsQuery::Boolean(query) => { + for query in query + .should + .iter() + .chain(&query.must) + .chain(&query.must_not) + { + visit(index, query, scorer, leaves)?; + } + } + } + Ok(()) + } + + let mut leaves = Vec::new(); + visit(self, query, scorer, &mut leaves)?; + Ok(leaves) } /// Freeze the current tail into a new immutable partition and publish a @@ -1543,6 +1923,7 @@ impl FtsMemIndex { Operator::Or, &scorer, theta, + false, ) { topk.offer(e.score, e.key()); } @@ -1562,6 +1943,7 @@ impl FtsMemIndex { operator, &scorer, f32::NEG_INFINITY, + false, )); } results @@ -1569,6 +1951,76 @@ impl FtsMemIndex { } } + fn search_match_with_scorer( + &self, + st: &IndexState, + query_tokens: &Tokens, + operator: Operator, + scorer: &MemBM25Scorer, + ) -> Vec { + if operator == Operator::And && has_grouped_positions(query_tokens) { + let mut result_map: Option> = None; + for group in query_position_groups(query_tokens) { + let group_results = + self.search_match_strings_with_scorer(st, &group, Operator::Or, scorer); + let group_map = group_results + .into_iter() + .map(|entry| (entry.key(), entry.score)) + .collect::>(); + let Some(current) = result_map.as_mut() else { + result_map = Some(group_map); + continue; + }; + current.retain(|key, score| { + if let Some(group_score) = group_map.get(key) { + *score += group_score; + true + } else { + false + } + }); + } + return result_map + .unwrap_or_default() + .into_iter() + .map(|(key, score)| FtsEntry { + row_position: key.row_position, + doc_index: public_doc_index(&key.doc_index), + score, + }) + .collect(); + } + let tokens = query_tokens_to_vec(query_tokens); + self.search_match_strings_with_scorer(st, &tokens, operator, scorer) + } + + fn search_match_strings_with_scorer( + &self, + st: &IndexState, + tokens: &[String], + operator: Operator, + scorer: &MemBM25Scorer, + ) -> Vec { + if tokens.is_empty() { + return Vec::new(); + } + let tail = st.tail.snapshot(); + let mut results = Vec::new(); + for partition in st.partitions.iter() { + results.extend(partition.search_match(tokens, operator, scorer)); + } + results.extend(score_terms( + &tail, + &st.tail.terms, + tokens, + operator, + scorer, + f32::NEG_INFINITY, + true, + )); + results + } + fn search_grouped_and( &self, st: &IndexState, @@ -1691,6 +2143,57 @@ impl FtsMemIndex { results } + fn search_phrase_with_scorer( + &self, + st: &IndexState, + query_tokens: &Tokens, + slop: u32, + scorer: &MemBM25Scorer, + ) -> Vec { + if query_tokens.is_empty() || scorer.num_docs() == 0 { + return Vec::new(); + } + let groups = query_position_groups(query_tokens); + if groups.is_empty() { + return Vec::new(); + } + if groups.len() == 1 { + return self.search_match_strings_with_scorer(st, &groups[0], Operator::Or, scorer); + } + if !self.params.has_positions() { + return Vec::new(); + } + let has_grouped_terms = groups.iter().any(|group| group.len() > 1); + let tokens = position_groups_to_tokens(&groups); + let tail = st.tail.snapshot(); + let mut results = Vec::new(); + for partition in st.partitions.iter() { + if has_grouped_terms { + results.extend(partition.search_phrase_groups(&groups, slop, scorer)); + } else { + results.extend(partition.search_phrase(&tokens, slop, scorer)); + } + } + if has_grouped_terms { + results.extend(phrase_search_tail_groups( + &tail, + &st.tail.terms, + &groups, + slop, + scorer, + )); + } else { + results.extend(phrase_search_tail( + &tail, + &st.tail.terms, + &tokens, + slop, + scorer, + )); + } + results + } + fn search_fuzzy_tokens( &self, st: &IndexState, @@ -2264,8 +2767,33 @@ fn index_text( tokenizer: &mut dyn LanceTokenizer, term_builders: &mut FxHashMap, BatchTermBuilder>, ) -> Result { + index_text_with_predicate(text, document_position, tokenizer, term_builders, |_| true) + .map(|(num_tokens, _)| num_tokens) +} + +fn index_text_filtered( + text: &str, + document_position: u64, + tokenizer: &mut dyn LanceTokenizer, + term_builders: &mut FxHashMap, BatchTermBuilder>, + allowed_terms: &FxHashSet, +) -> Result<(u32, bool)> { + index_text_with_predicate(text, document_position, tokenizer, term_builders, |term| { + allowed_terms.contains(term) + }) +} + +#[inline] +fn index_text_with_predicate( + text: &str, + document_position: u64, + tokenizer: &mut dyn LanceTokenizer, + term_builders: &mut FxHashMap, BatchTermBuilder>, + mut retain_term: impl FnMut(&str) -> bool, +) -> Result<(u32, bool)> { let mut stream = tokenizer.token_stream_for_doc(text); let mut num_tokens = 0u32; + let mut retained_term = false; while let Some(token) = stream.next() { let position = u32::try_from(token.position).map_err(|_| { Error::invalid_input(format!( @@ -2274,13 +2802,16 @@ fn index_text( )) })?; let term = token.text.as_str(); - if let Some(builder) = term_builders.get_mut(term) { - builder.observe(document_position, position); - } else { - term_builders.insert( - Arc::::from(term), - BatchTermBuilder::with_first(document_position, position), - ); + if retain_term(term) { + retained_term = true; + if let Some(builder) = term_builders.get_mut(term) { + builder.observe(document_position, position); + } else { + term_builders.insert( + Arc::::from(term), + BatchTermBuilder::with_first(document_position, position), + ); + } } num_tokens = num_tokens.checked_add(1).ok_or_else(|| { Error::invalid_input(format!( @@ -2288,7 +2819,7 @@ fn index_text( )) })?; } - Ok(num_tokens) + Ok((num_tokens, retained_term)) } fn has_visible_chunk(slice: &TermSlice, visible_count: usize) -> bool { @@ -2379,6 +2910,11 @@ fn tail_token_df( /// Score `tokens` against the visible tail, summing each token's BM25 /// contribution per document. Uses the shared corpus-wide `scorer`. +/// +/// `retain_zero_weight_matches` is reserved for query-local residual postings +/// scored with committed-index statistics. A term absent from the committed +/// corpus has zero BM25 weight, but its fresh matching rows must remain visible +/// to compound membership and MUST_NOT evaluation. fn score_terms( snap: &Snapshot, terms: &SkipMap, Arc>>, @@ -2386,6 +2922,7 @@ fn score_terms( operator: Operator, scorer: &MemBM25Scorer, theta: f32, + retain_zero_weight_matches: bool, ) -> Vec { // Per-token tail data + its score upper bound (max freq over visible chunks, // scored at the most generous doc length of 1). If even the sum of those @@ -2401,7 +2938,7 @@ fn score_terms( continue; }; let qw = scorer.query_weight(token); - if qw == 0.0 { + if qw == 0.0 && !retain_zero_weight_matches { continue; } let slice = entry.value().load_full(); @@ -2411,7 +2948,9 @@ fn score_terms( .map(|c| c.max_freq) .max() .unwrap_or(0); - tail_ub += qw * scorer.doc_weight(max_freq, 1); + if qw != 0.0 { + tail_ub += qw * scorer.doc_weight(max_freq, 1); + } tail_terms.push((qw, slice)); } if tail_ub <= theta { @@ -2428,8 +2967,12 @@ fn score_terms( continue; }; for (i, &document_position) in chunk.row_positions.iter().enumerate() { - let dl = meta.dl(document_position).unwrap_or(1); - let score = qw * scorer.doc_weight(chunk.frequencies[i], dl); + let score = if qw == 0.0 { + 0.0 + } else { + let dl = meta.dl(document_position).unwrap_or(1); + qw * scorer.doc_weight(chunk.frequencies[i], dl) + }; *doc_scores.entry(document_position).or_default() += score; if let Some(doc_hits) = &mut doc_hits { *doc_hits.entry(document_position).or_default() += 1; @@ -4142,6 +4685,139 @@ mod tests { .unwrap() } + #[test] + fn query_term_allowlist_preserves_document_lengths_with_external_scorer() { + let schema = create_test_schema(); + let batch = create_test_batch(schema.as_ref()); + let row_ids = UInt64Array::from(vec![900, 42, 777]); + let terms = FxHashSet::from_iter(["hello".to_string()]); + let index = QueryLocalFtsIndex::try_with_params( + 1, + "description".to_string(), + InvertedIndexParams::default(), + ) + .unwrap(); + let full_index = FtsMemIndex::new(1, "description".to_string()); + + let stats = index + .insert_with_row_ids_for_terms(&batch, &row_ids, &terms) + .unwrap(); + full_index.insert(&batch, 0).unwrap(); + + // The unmatched nonempty row (row id 42) contributes no postings or + // metadata, but remains part of the approximate residual BM25 corpus. + assert_eq!(index.doc_count(), 2); + assert_eq!(index.inner.entry_count(), 2); + assert_eq!(stats.doc_count, 3); + assert_eq!(stats.total_tokens, 5); + assert_eq!(stats.token_docs.get("hello"), Some(&2)); + let committed_scorer = MemBM25Scorer::new(6, 3, HashMap::from([("hello".to_string(), 2)])); + let mut residual_scorer = committed_scorer.clone(); + stats.add_to_scorer(&mut residual_scorer).unwrap(); + assert_eq!(residual_scorer.num_docs, 6); + assert_eq!(residual_scorer.total_tokens, 11); + assert_eq!(residual_scorer.token_docs.get("hello"), Some(&4)); + + let query = FtsQuery::Match( + lance_index::scalar::inverted::query::MatchQuery::new("hello".to_string()) + .with_column(Some("description".to_string())), + ); + let leaves = index.exact_leaf_results(&query, &committed_scorer).unwrap(); + let full_leaves = full_index + .exact_leaf_results(&query, &committed_scorer) + .unwrap(); + let mut actual = leaves[0] + .iter() + .map(|(row_id, _)| *row_id) + .collect::>(); + actual.sort_unstable(); + assert_eq!(actual, vec![777, 900]); + let mut actual_scores = leaves[0] + .iter() + .map(|(_, score)| score.to_bits()) + .collect::>(); + let mut full_scores = full_leaves[0] + .iter() + .map(|(_, score)| score.to_bits()) + .collect::>(); + actual_scores.sort_unstable(); + full_scores.sort_unstable(); + assert_eq!(actual_scores, full_scores); + } + + #[test] + fn query_local_external_empty_scorer_retains_zero_score_membership() { + let schema = create_test_schema(); + let batch = create_test_batch(schema.as_ref()); + let row_ids = UInt64Array::from(vec![900, 42, 777]); + let terms = FxHashSet::from_iter(["hello".to_string()]); + let index = QueryLocalFtsIndex::try_with_params( + 1, + "description".to_string(), + InvertedIndexParams::default(), + ) + .unwrap(); + index + .insert_with_row_ids_for_terms(&batch, &row_ids, &terms) + .unwrap(); + + let committed_scorer = MemBM25Scorer::new(0, 0, HashMap::from([("hello".to_string(), 0)])); + let query = FtsQuery::Match( + lance_index::scalar::inverted::query::MatchQuery::new("hello".to_string()) + .with_column(Some("description".to_string())), + ); + let leaves = index.exact_leaf_results(&query, &committed_scorer).unwrap(); + + let mut actual = leaves[0].clone(); + actual.sort_unstable_by_key(|(row_id, _)| *row_id); + assert_eq!(actual.len(), 2); + assert_eq!(actual[0].0, 777); + assert_eq!(actual[1].0, 900); + assert!(actual.iter().all(|(_, score)| score.to_bits() == 0)); + } + + #[test] + fn query_local_materialization_never_starts_background_maintenance() { + let schema = create_test_schema(); + let batch = create_test_batch(schema.as_ref()); + let row_ids = UInt64Array::from(vec![900, 42, 777]); + let terms = FxHashSet::from_iter(["hello".to_string()]); + let params = InvertedIndexParams::default(); + let tokenizer = params.build().unwrap(); + let mut index = QueryLocalFtsIndex::try_with_loaded_tokenizer( + 1, + "description".to_string(), + params, + tokenizer, + ) + .unwrap(); + // Crossing the normal freeze threshold would create a partition and + // may launch a detached tiered merge. Query-local materialization must + // remain entirely in its query-owned tail instead. + index.inner.freeze_threshold_rows = 1; + index + .insert_with_row_ids_for_terms(&batch, &row_ids, &terms) + .unwrap(); + + assert!(index.inner.state.load().partitions.is_empty()); + assert!(index.inner.merge.lock().unwrap().is_none()); + assert_eq!(index.doc_count(), 2); + + let sibling = index.empty_sibling(); + assert!(Arc::ptr_eq( + &index.inner.tokenizer_pool, + &sibling.inner.tokenizer_pool + )); + assert_eq!(sibling.doc_count(), 0); + sibling + .insert_with_row_ids_for_terms(&batch, &UInt64Array::from(vec![901, 43, 778]), &terms) + .unwrap(); + assert_eq!(index.doc_count(), 2); + assert_eq!(sibling.doc_count(), 2); + assert!(sibling.inner.state.load().partitions.is_empty()); + assert!(sibling.inner.merge.lock().unwrap().is_none()); + } + fn create_element_test_batch() -> RecordBatch { let mut tags = ListBuilder::new(StringBuilder::new()); tags.values().append_value("alpha beta"); diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index 8a4fb75a399..fddd98b3339 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -114,7 +114,8 @@ use crate::io::exec::filtered_read::{ }; use crate::io::exec::fts::{ BoostQueryExec, CompoundQueryExec, CrossColumnCompoundQueryExec, FlatMatchFilterExec, - FlatMatchQueryExec, FtsDocumentExec, MatchQueryExec, PhraseQueryExec, SharedFtsScorer, + FlatMatchQueryExec, FtsDocumentExec, HybridCompoundQueryExec, MatchQueryExec, PhraseQueryExec, + SharedFtsScorer, }; use crate::io::exec::knn::MultivectorScoringExec; use crate::io::exec::scalar_index::{MaterializeIndexExec, ScalarIndexExec}; @@ -283,6 +284,76 @@ fn supports_compound_scorer(query: &FtsQuery) -> bool { !columns.is_empty() && (!matches!(query, FtsQuery::MultiMatch(_)) || columns.len() == 1) } +fn supports_indexed_stats_residual_compound(query: &FtsQuery) -> bool { + match query { + FtsQuery::Match(query) => query.fuzziness == Some(0), + // MemWAL phrase matching currently collapses tokenizer position gaps. + // Keep phrase queries on the established fallback until it can retain + // those gaps exactly (notably when stop words are configured). + FtsQuery::Phrase(_) => false, + FtsQuery::Boost(query) => { + supports_indexed_stats_residual_compound(&query.positive) + && supports_indexed_stats_residual_compound(&query.negative) + } + FtsQuery::MultiMatch(query) => query + .match_queries + .iter() + .all(|query| query.fuzziness == Some(0)), + FtsQuery::Boolean(query) => query + .should + .iter() + .chain(&query.must) + .chain(&query.must_not) + .all(supports_indexed_stats_residual_compound), + } +} + +const MAX_QUERY_LOCAL_RESIDUAL_ROWS: usize = 100_000; + +fn has_bounded_query_local_residual_rows(fragments: &[Fragment]) -> bool { + fragments + .iter() + .try_fold(0usize, |total, fragment| { + total.checked_add(fragment.physical_rows?) + }) + .is_some_and(|total| total <= MAX_QUERY_LOCAL_RESIDUAL_ROWS) +} + +fn has_complete_hybrid_fts_coverage( + segments: &[IndexMetadata], + residual_fragments: &[Fragment], + target_fragments: &[Fragment], +) -> bool { + let Some(target) = target_fragments + .iter() + .map(|fragment| u32::try_from(fragment.id).ok()) + .collect::>() + else { + return false; + }; + let Some(residual) = residual_fragments + .iter() + .map(|fragment| u32::try_from(fragment.id).ok()) + .collect::>() + else { + return false; + }; + let mut indexed = RoaringBitmap::new(); + for segment in segments { + let Some(coverage) = segment.fragment_bitmap.as_ref() else { + return false; + }; + if !indexed.is_disjoint(coverage) { + return false; + } + indexed |= coverage; + } + if !indexed.is_subset(&target) || !indexed.is_disjoint(&residual) { + return false; + } + indexed | residual == target +} + fn validate_fts_query_contract(query: &FtsQuery) -> Result<()> { fn validate_multiplier(name: &str, value: f32) -> Result<()> { if value.is_finite() && value >= 0.0 { @@ -4196,6 +4267,7 @@ impl Scanner { &self, query: &FtsQuery, params: &FtsSearchParams, + filter_plan: &ExprFilterPlan, prefilter_source: &PreFilterSource, document_granularity: DocumentGranularity, ) -> Result>> { @@ -4220,6 +4292,21 @@ impl Scanner { } let mut phrase_columns = HashSet::new(); collect_phrase_columns(query, &mut phrase_columns); + // Query-local residual scoring intentionally reuses committed-index + // BM25 statistics. Matching remains exact for the supported leaf + // shapes, but ranking is approximate until the appended rows are + // incorporated into a persistent index. + let allow_indexed_stats_residual = !cross_column + && !self.fast_search + && self.fragments.is_none() + && filter_plan.is_empty() + && self.external_row_mask.is_none() + && params.limit.is_some() + && document_granularity == DocumentGranularity::Row + && target_fragments + .iter() + .all(|fragment| fragment.deletion_file.is_none()) + && supports_indexed_stats_residual_compound(query); let segment_groups = futures::future::try_join_all(columns.into_iter().map(|column| { let phrase_columns = &phrase_columns; @@ -4243,8 +4330,12 @@ impl Scanner { ) .await?; let unindexed_fragments = self.retain_target_fragments(unindexed_fragments); + let has_bounded_residual = allow_indexed_stats_residual + && has_bounded_query_local_residual_rows(&unindexed_fragments); if !unindexed_fragments.is_empty() && (!self.fast_search || unindexed_fragments.len() == target_fragments.len()) + && !(has_bounded_residual + && unindexed_fragments.len() < target_fragments.len()) { // Flat and posting-backed leaves do not share a document // domain, so preserve the exact fallback for partial index @@ -4254,10 +4345,6 @@ impl Scanner { // indexed. return Ok(None); } - let unindexed_fragment_ids = unindexed_fragments - .iter() - .map(|fragment| fragment.id as u32) - .collect::(); let segments = match overlay_plan { FtsOverlayPlan::Unchanged(Some(segments)) => segments, FtsOverlayPlan::Unchanged(None) => { @@ -4271,6 +4358,23 @@ impl Scanner { } FtsOverlayPlan::RowLevel { .. } | FtsOverlayPlan::FullScan => return Ok(None), }; + if has_bounded_residual && !unindexed_fragments.is_empty() { + if !has_complete_hybrid_fts_coverage( + &segments, + &unindexed_fragments, + target_fragments, + ) { + return Ok(None); + } + if segments.is_empty() { + return Err(Error::internal( + "hybrid compound FTS requires one indexed segment", + )); + } + // Preserve the established semantic mismatch error before + // constructing query-local postings with the same tokenizer. + load_segment_details(&self.dataset, &column, &segments).await?; + } if cross_column { let details = futures::future::try_join_all( @@ -4306,7 +4410,7 @@ impl Scanner { } } - Ok(Some((column, segments, unindexed_fragment_ids))) + Ok(Some((column, segments, unindexed_fragments))) } })) .await?; @@ -4315,9 +4419,43 @@ impl Scanner { }; if !cross_column { - let (_, segments, _) = segment_groups.into_iter().next().ok_or_else(|| { - Error::internal("compound scorer requires one column".to_string()) - })?; + let (column, segments, unindexed_fragments) = + segment_groups.into_iter().next().ok_or_else(|| { + Error::internal("compound scorer requires one column".to_string()) + })?; + if allow_indexed_stats_residual && !unindexed_fragments.is_empty() { + let resolved = + resolve_fts_field(self.dataset.schema(), &column, document_granularity)?; + let scan_column = if resolved.has_lists() { + resolved.root_column.clone() + } else { + resolved.canonical_path.clone() + }; + let scan_projection = self + .dataset + .empty_projection() + .with_row_id() + .union_columns(&[scan_column], OnMissing::Error)?; + let PlannedFilteredScan { plan, .. } = self + .filtered_read( + &ExprFilterPlan::default(), + scan_projection, + /* make_deletions_null */ false, + Some(Arc::new(unindexed_fragments)), + None, + /* is_prefilter */ true, + None, + ) + .await?; + return Ok(Some(Arc::new(HybridCompoundQueryExec::new( + self.dataset.clone(), + query.clone(), + params.clone(), + column, + segments, + plan, + )))); + } return Ok(Some(Arc::new( CompoundQueryExec::new_with_segments( self.dataset.clone(), @@ -4334,9 +4472,17 @@ impl Scanner { let Some((_, _, first_unindexed_fragments)) = coverage_groups.next() else { return Ok(None); }; - if coverage_groups - .any(|(_, _, unindexed_fragments)| unindexed_fragments != first_unindexed_fragments) - { + let first_unindexed_fragment_ids = first_unindexed_fragments + .iter() + .map(|fragment| fragment.id as u32) + .collect::(); + if coverage_groups.any(|(_, _, unindexed_fragments)| { + unindexed_fragments + .iter() + .map(|fragment| fragment.id as u32) + .collect::() + != first_unindexed_fragment_ids + }) { // The cross-column scorer builds one shared prefilter. If column // coverage differs, that prefilter's union can re-admit stale // postings from a fragment invalidated only for another column. @@ -4348,7 +4494,6 @@ impl Scanner { .into_iter() .map(|(column, segments, _)| (column, segments)) .collect(); - let exec = CrossColumnCompoundQueryExec::new_with_segments( self.dataset.clone(), query.clone(), @@ -4371,7 +4516,13 @@ impl Scanner { if !document_granularity.is_list_element() && supports_compound_scorer(query) && let Some(plan) = self - .plan_compound_scorer(query, params, prefilter_source, document_granularity) + .plan_compound_scorer( + query, + params, + filter_plan, + prefilter_source, + document_granularity, + ) .await? { return Ok(plan); @@ -4441,6 +4592,7 @@ impl Scanner { .plan_compound_scorer( &child_query, params, + filter_plan, field_prefilter_source, document_granularity, ) @@ -7226,6 +7378,28 @@ mod test { assert!(error.to_string().contains("BoostQuery negative_boost")); } + #[test] + fn test_query_local_residual_row_bound() { + let fragment_with_rows = |id, physical_rows| { + let mut fragment = Fragment::new(id); + fragment.physical_rows = physical_rows; + fragment + }; + + assert!(has_bounded_query_local_residual_rows(&[ + fragment_with_rows(0, Some(40_000)), + fragment_with_rows(1, Some(60_000)), + ])); + assert!(!has_bounded_query_local_residual_rows(&[ + fragment_with_rows(0, Some(40_000)), + fragment_with_rows(1, Some(60_001)), + ])); + assert!(!has_bounded_query_local_residual_rows(&[ + fragment_with_rows(0, Some(1)), + fragment_with_rows(1, None), + ])); + } + #[test] fn test_normalize_fts_zero_boosts_recurses_and_preserves_nonzero_values() { fn boost_bits(query: &FtsQuery) -> Vec { diff --git a/rust/lance/src/dataset/tests/dataset_index.rs b/rust/lance/src/dataset/tests/dataset_index.rs index bdcd6e5719c..19673e83ae1 100644 --- a/rust/lance/src/dataset/tests/dataset_index.rs +++ b/rust/lance/src/dataset/tests/dataset_index.rs @@ -9,10 +9,11 @@ use std::sync::{Arc, Mutex}; use std::vec; use crate::dataset::ROW_ID; +use crate::dataset::WriteDestination; use crate::dataset::builder::DatasetBuilder; use crate::dataset::tests::dataset_migrations::scan_dataset; use crate::dataset::tests::dataset_transactions::{assert_results, execute_sql}; -use crate::dataset::transaction::{Operation, Transaction}; +use crate::dataset::transaction::{DataReplacementGroup, Operation, Transaction}; use crate::index::vector::VectorIndexParams; use crate::session::Session; use crate::utils::test::covering; @@ -61,6 +62,7 @@ use futures::{StreamExt, TryStreamExt}; use itertools::Itertools; use lance_arrow::json::ARROW_JSON_EXT_NAME; use lance_index::scalar::inverted::query::{FtsQuery, MultiMatchQuery}; +use lance_table::format::BasePath; use lance_testing::datagen::generate_random_array; use rand::Rng; use rstest::rstest; @@ -1322,6 +1324,12 @@ async fn compound_fts_results( .collect() } +fn scored_row_bits(rows: &[(u64, f32)]) -> Vec<(u64, u32)> { + rows.iter() + .map(|(row_id, score)| (*row_id, score.to_bits())) + .collect() +} + fn compound_fts_result_bits(batch: &RecordBatch) -> Vec<(u64, u32)> { let row_ids = batch[ROW_ID].as_primitive::().values(); let scores = batch[SCORE_COL].as_primitive::().values(); @@ -1817,7 +1825,7 @@ async fn test_top_level_cross_column_multimatch_uses_field_local_compound_scorer .await .unwrap(); // Index only the title after the append so it can retain a bounded plan - // while the partially covered body uses the exhaustive leaf fallback. + // while the partially covered body uses a query-local hybrid scorer. create_fragmented_fts_index(&mut partial_dataset, "title", true).await; partial_dataset .create_index( @@ -1829,13 +1837,8 @@ async fn test_top_level_cross_column_multimatch_uses_field_local_compound_scorer ) .await .unwrap(); - assert_compound_matches_independent_oracle( - &partial_dataset, - "partial_top_level_cross_column_multimatch", - &explicit_query, - LIMIT, - ) - .await; + let partial_results = + compound_fts_results(&partial_dataset, explicit_query.clone(), Some(LIMIT as i64)).await; let partial_plan = compound_fts_plan(&partial_dataset, explicit_query.clone(), LIMIT).await; assert!( !partial_plan.contains(CROSS_COLUMN_COMPOUND_FTS_SCORER), @@ -1843,18 +1846,23 @@ async fn test_top_level_cross_column_multimatch_uses_field_local_compound_scorer ); assert_eq!( partial_plan.matches("CompoundFtsScorer").count(), + 2, + "both fields should retain field-local bounded compound scorers:\n{partial_plan}" + ); + assert_eq!( + partial_plan.matches("HybridCompoundFtsScorer").count(), 1, - "the fully indexed title should retain its bounded compound scorer:\n{partial_plan}" + "only the partially covered body should use a query-local hybrid scorer:\n{partial_plan}" ); assert!( - partial_plan.contains("FlatMatchQuery"), - "the partially covered body should use the exact indexed-plus-flat fallback:\n{partial_plan}" + !partial_plan.contains("FlatMatchQuery"), + "the hybrid body scorer should replace the indexed-plus-flat fallback:\n{partial_plan}" ); let mut fast_scanner = partial_dataset.scan(); fast_scanner .with_row_id() - .full_text_search(FullTextSearchQuery::new_query(explicit_query)) + .full_text_search(FullTextSearchQuery::new_query(explicit_query.clone())) .unwrap() .fast_search(); fast_scanner.limit(Some(LIMIT as i64), None).unwrap(); @@ -1872,6 +1880,16 @@ async fn test_top_level_cross_column_multimatch_uses_field_local_compound_scorer !fast_plan.contains("FlatMatchQuery"), "fast search must skip the partially covered body's flat path:\n{fast_plan}" ); + + assert_eq!( + partial_results.len(), + LIMIT, + "the approximate residual path must still return a bounded top-k" + ); + assert!( + partial_results.iter().all(|(_, score)| score.is_finite()), + "committed-index statistics must produce finite residual scores" + ); } #[rstest] @@ -2612,18 +2630,71 @@ async fn test_same_column_compound_fast_search_excludes_unindexed_rows() { ]) .into(); - let mut exact_scanner = dataset.scan(); - exact_scanner + let mut hybrid_scanner = dataset.scan(); + hybrid_scanner .project(&["id"]) .unwrap() .full_text_search(FullTextSearchQuery::new_query(query.clone())) .unwrap(); - exact_scanner.limit(Some(2), None).unwrap(); - let exact = exact_scanner.try_into_batch().await.unwrap(); + hybrid_scanner.limit(Some(2), None).unwrap(); + let hybrid_plan = hybrid_scanner.explain_plan(false).await.unwrap(); + assert!( + hybrid_plan.contains("HybridCompoundFtsScorer"), + "partial coverage should build one indexed-statistics query-local residual index:\n{hybrid_plan}" + ); + assert!( + !hybrid_plan.contains("FlatMatchQuery"), + "hybrid compound scoring must not scan the residual once per leaf:\n{hybrid_plan}" + ); + let hybrid = hybrid_scanner.try_into_batch().await.unwrap(); assert_eq!( - exact["id"].as_primitive::().values(), + hybrid["id"].as_primitive::().values(), &[0, 2], - "exact search should include the appended hit" + "approximate residual search should include the appended hit" + ); + + let empty_terms_query: FtsQuery = BooleanQuery::new([ + (Occur::Must, compound_match_query("", "text", 1.0)), + (Occur::Should, compound_match_query(" ", "text", 1.0)), + ]) + .into(); + let empty_terms_plan = compound_fts_plan(&dataset, empty_terms_query.clone(), 2).await; + assert!( + empty_terms_plan.contains("HybridCompoundFtsScorer"), + "the empty analyzed-term case must exercise the hybrid short circuit:\n{empty_terms_plan}" + ); + let empty_results = compound_fts_results(&dataset, empty_terms_query, Some(2)).await; + assert!(empty_results.is_empty()); + + let mut filtered_scanner = dataset.scan(); + filtered_scanner + .with_row_id() + .filter("id >= 0") + .unwrap() + .full_text_search(FullTextSearchQuery::new_query(query.clone())) + .unwrap(); + filtered_scanner.prefilter(true); + filtered_scanner.limit(Some(2), None).unwrap(); + let filtered_plan = filtered_scanner.explain_plan(false).await.unwrap(); + assert!( + !filtered_plan.contains("HybridCompoundFtsScorer"), + "prefiltered residual scoring must retain the exact fallback:\n{filtered_plan}" + ); + + let phrase_query: FtsQuery = BooleanQuery::new([ + ( + Occur::Must, + PhraseQuery::new("fresh alpha".to_string()) + .with_column(Some("text".to_string())) + .into(), + ), + (Occur::Must, compound_match_query("fresh", "text", 1.0)), + ]) + .into(); + let phrase_plan = compound_fts_plan(&dataset, phrase_query, 2).await; + assert!( + !phrase_plan.contains("HybridCompoundFtsScorer"), + "phrase position gaps are not yet supported by the residual index:\n{phrase_plan}" ); let mut fast_scanner = dataset.scan(); @@ -2673,6 +2744,311 @@ async fn test_same_column_compound_fast_search_excludes_unindexed_rows() { ); } +#[tokio::test] +async fn test_partial_compound_hybrid_prunes_same_path_different_base_rewrite() { + let primary = TempStrDir::default(); + let base_one = TempStrDir::default(); + let base_two = TempStrDir::default(); + let initial = arrow_array::record_batch!( + ("text", Utf8, ["stable alpha", "stale alpha"]), + ("id", Int32, [0, 1]) + ) + .unwrap(); + let schema = initial.schema(); + let mut dataset = Dataset::write( + RecordBatchIterator::new(vec![initial].into_iter().map(Ok), schema.clone()), + &primary, + Some(WriteParams { + max_rows_per_file: 1, + initial_bases: Some(vec![ + BasePath::new(1, base_one.to_string(), None, false), + BasePath::new(2, base_two.to_string(), None, false), + ]), + target_bases: Some(vec![1]), + ..Default::default() + }), + ) + .await + .unwrap(); + assert_eq!(dataset.get_fragments().len(), 2); + assert!( + dataset + .get_fragments() + .iter() + .all(|fragment| { fragment.metadata().files[0].base_id == Some(1) }) + ); + let segment = dataset + .create_index_builder( + &["text"], + IndexType::Inverted, + &InvertedIndexParams::default().with_position(true), + ) + .name("text_idx".to_string()) + .execute_uncommitted() + .await + .unwrap(); + + let relative_path = dataset.get_fragment(1).unwrap().metadata().files[0] + .path + .clone(); + let replacement = + arrow_array::record_batch!(("text", Utf8, ["current beta"]), ("id", Int32, [1])).unwrap(); + let replacement_path = dataset + .data_file_dir_for_base(Some(2)) + .unwrap() + .join(relative_path.as_str()); + let object_writer = dataset + .object_store(Some(2)) + .await + .unwrap() + .create(&replacement_path) + .await + .unwrap(); + let mut writer = lance_file::versions::v2_1::create_writer( + object_writer, + schema.as_ref().try_into().unwrap(), + Default::default(), + ) + .unwrap(); + writer.write_batch(&replacement).await.unwrap(); + writer.finish().await.unwrap(); + let replacement_file = dataset + .create_data_file(&relative_path, Some(2)) + .await + .unwrap(); + assert_eq!(replacement_file.path, relative_path); + assert_eq!(replacement_file.base_id, Some(2)); + + let read_version = dataset.manifest.version; + let mut dataset = Dataset::commit( + WriteDestination::Dataset(Arc::new(dataset)), + Operation::DataReplacement { + replacements: vec![DataReplacementGroup(1, replacement_file)], + }, + Some(read_version), + None, + None, + Arc::new(Default::default()), + false, + ) + .await + .unwrap(); + dataset + .commit_existing_index_segments("text_idx", "text", vec![segment]) + .await + .unwrap(); + let committed = dataset + .load_index_by_name("text_idx") + .await + .unwrap() + .unwrap(); + let coverage = committed.fragment_bitmap.as_ref().unwrap(); + assert!( + coverage.contains(0), + "the unchanged physical file must remain covered" + ); + assert!( + !coverage.contains(1), + "the same path on a different registered base must be pruned" + ); + + let query: FtsQuery = BooleanQuery::new([ + (Occur::Must, compound_match_query("beta", "text", 1.0)), + (Occur::MustNot, compound_match_query("alpha", "text", 1.0)), + ]) + .into(); + let mut scanner = dataset.scan(); + scanner + .project(&["id"]) + .unwrap() + .full_text_search(FullTextSearchQuery::new_query(query)) + .unwrap(); + scanner.limit(Some(2), None).unwrap(); + let plan = scanner.explain_plan(false).await.unwrap(); + assert!( + plan.contains("HybridCompoundFtsScorer"), + "the physically pruned fragment should use hybrid residual scoring:\n{plan}" + ); + let results = scanner.try_into_batch().await.unwrap(); + assert_eq!( + results["id"].as_primitive::().values(), + &[1], + "the current beta row must be visible without leaking stale alpha membership" + ); + assert!( + results[SCORE_COL] + .as_primitive::() + .values() + .iter() + .all(|score| score.is_finite()) + ); +} + +#[tokio::test] +async fn test_partial_compound_hybrid_uses_mixed_approximate_statistics() { + let initial = arrow_array::record_batch!( + ("text", Utf8, ["fresh alpha", "blocked fresh alpha"]), + ("id", Int32, [0, 1]) + ) + .unwrap(); + let schema = initial.schema(); + let mut dataset = Dataset::write( + RecordBatchIterator::new(vec![initial].into_iter().map(Ok), schema), + "memory://", + None, + ) + .await + .unwrap(); + create_fragmented_fts_index(&mut dataset, "text", true).await; + + let appended = arrow_array::record_batch!( + ( + "text", + Utf8, + [ + "fresh alpha", + "fresh beta", + "fresh alpha", + "fresh beta beta", + "fresh beta blocked" + ] + ), + ("id", Int32, [2, 3, 4, 5, 6]) + ) + .unwrap(); + let schema = appended.schema(); + dataset + .append( + RecordBatchIterator::new(vec![appended].into_iter().map(Ok), schema), + Some(WriteParams { + // Keep the residual rows in separate fragments; execution may + // rechunk their scan batches before query-local indexing. + max_rows_per_file: 1, + ..Default::default() + }), + ) + .await + .unwrap(); + + let positive: FtsQuery = BooleanQuery::new([ + (Occur::Must, compound_match_query("fresh", "text", 1.0)), + (Occur::Should, compound_match_query("alpha", "text", 1.0)), + (Occur::MustNot, compound_match_query("blocked", "text", 1.0)), + ]) + .into(); + let boost_query: FtsQuery = BoostQuery::new( + positive, + compound_match_query("alpha", "text", 1.0), + Some(0.25), + ) + .into(); + let partial_boost = compound_fts_results(&dataset, boost_query.clone(), Some(10)).await; + assert_eq!( + partial_boost.len(), + 5, + "MUST_NOT must exclude the blocked row" + ); + let multimatch_query: FtsQuery = MultiMatchQuery::try_new( + "fresh alpha".to_string(), + vec!["text".to_string(), "text".to_string()], + ) + .unwrap() + .try_with_boosts(vec![1.0, 2.0]) + .unwrap() + .into(); + for (query_name, query) in [ + ("Boost", boost_query.clone()), + ("MultiMatch", multimatch_query.clone()), + ] { + let mut scanner = dataset.scan(); + scanner + .project(&["id"]) + .unwrap() + .full_text_search(FullTextSearchQuery::new_query(query)) + .unwrap(); + scanner.limit(Some(10), None).unwrap(); + let batch = scanner.try_into_batch().await.unwrap(); + let ids = batch["id"].as_primitive::().values(); + let scores = batch[SCORE_COL].as_primitive::().values(); + let score_bits = ids + .iter() + .copied() + .zip(scores.iter().map(|score| score.to_bits())) + .collect::>(); + let positions = ids + .iter() + .enumerate() + .map(|(position, row_id)| (*row_id, position)) + .collect::>(); + assert_eq!( + score_bits.get(&2), + score_bits.get(&4), + "{query_name} must preserve equal scores within the residual arm" + ); + assert!( + positions[&2] < positions[&4], + "{query_name} must preserve the row-id tie break within the residual arm" + ); + } + + let mut indexed_only_scanner = dataset.scan(); + indexed_only_scanner + .with_row_id() + .full_text_search(FullTextSearchQuery::new_query(boost_query.clone())) + .unwrap() + .fast_search(); + indexed_only_scanner.limit(Some(10), None).unwrap(); + let indexed_only = + compound_fts_result_bits(&indexed_only_scanner.try_into_batch().await.unwrap()) + .into_iter() + .collect::>(); + let partial_boost_bits = scored_row_bits(&partial_boost) + .into_iter() + .collect::>(); + for (row_id, score) in indexed_only { + assert_eq!( + partial_boost_bits.get(&row_id), + Some(&score), + "hybrid scoring must preserve committed-index scores for indexed row {row_id}" + ); + } + // The residual arm intentionally uses committed + query-local statistics, + // so its scores are not expected to equal either indexed-arm scores or a + // fully rebuilt index's exact global scores. + + let residual_only_query: FtsQuery = BooleanQuery::new([ + (Occur::Must, compound_match_query("beta", "text", 1.0)), + (Occur::MustNot, compound_match_query("blocked", "text", 1.0)), + ]) + .into(); + let mut residual_only_scanner = dataset.scan(); + residual_only_scanner + .project(&["id"]) + .unwrap() + .full_text_search(FullTextSearchQuery::new_query(residual_only_query)) + .unwrap(); + residual_only_scanner.limit(Some(10), None).unwrap(); + let residual_only_plan = residual_only_scanner.explain_plan(false).await.unwrap(); + assert!( + residual_only_plan.contains("HybridCompoundFtsScorer"), + "residual-only term membership must use the indexed-statistics hybrid path:\n{residual_only_plan}" + ); + let residual_only = residual_only_scanner.try_into_batch().await.unwrap(); + assert_eq!( + residual_only["id"].as_primitive::().values(), + &[5, 3], + "residual beta TF must rank id=5 first while MUST_NOT excludes id=6" + ); + let residual_only_scores = residual_only[SCORE_COL] + .as_primitive::() + .values(); + assert!( + residual_only_scores.iter().all(|score| score.is_finite()) + && residual_only_scores[0] > residual_only_scores[1], + "residual-only terms must retain membership and use query-local TF/DF scoring" + ); +} + #[tokio::test] async fn test_boolean_must_scores_sum_across_execution_paths() { let batch = arrow_array::record_batch!( diff --git a/rust/lance/src/index.rs b/rust/lance/src/index.rs index d10ec624d77..ac57366c46b 100644 --- a/rust/lance/src/index.rs +++ b/rust/lance/src/index.rs @@ -58,7 +58,7 @@ use lance_io::utils::{ CachedFileSize, read_last_block, read_message, read_message_from_buf, read_metadata_offset, read_version, }; -use lance_table::format::{Fragment, SelfDescribingFileReader}; +use lance_table::format::{DataFile, Fragment, SelfDescribingFileReader}; use lance_table::format::{IndexFile, IndexMetadata, list_index_files_with_sizes}; use lance_table::io::manifest::read_manifest_indexes; use roaring::RoaringBitmap; @@ -142,10 +142,60 @@ fn collect_subtree_field_ids(field: &Field, field_ids: &mut HashSet) { } } -fn fragment_field_paths<'a>( +/// Stable identity fields for a physical data file. +/// +/// This mirrors transaction rewrite validation, additionally resolves a +/// registered base to its physical binding, and deliberately excludes +/// `file_size_bytes`, which is a mutable cache rather than file identity. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum PhysicalBaseBinding<'a> { + Primary, + Registered { + path: &'a str, + is_dataset_root: bool, + }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct PhysicalDataFileIdentity<'a> { + base_id: Option, + base_binding: PhysicalBaseBinding<'a>, + path: &'a str, + fields: &'a [i32], + column_indices: &'a [i32], + file_major_version: u32, + file_minor_version: u32, +} + +impl<'a> PhysicalDataFileIdentity<'a> { + fn try_new(dataset: &'a Dataset, file: &'a DataFile) -> Option { + let base_binding = match file.base_id { + Some(base_id) => { + let base = dataset.manifest.base_paths.get(&base_id)?; + PhysicalBaseBinding::Registered { + path: &base.path, + is_dataset_root: base.is_dataset_root, + } + } + None => PhysicalBaseBinding::Primary, + }; + Some(Self { + base_id: file.base_id, + base_binding, + path: &file.path, + fields: file.fields.as_ref(), + column_indices: file.column_indices.as_ref(), + file_major_version: file.file_major_version, + file_minor_version: file.file_minor_version, + }) + } +} + +fn fragment_field_files<'a>( + dataset: &'a Dataset, fragment: &'a Fragment, indexed_field_ids: &HashSet, -) -> HashMap { +) -> Option>> { fragment .files .iter() @@ -153,7 +203,10 @@ fn fragment_field_paths<'a>( file.fields .iter() .filter(|field_id| indexed_field_ids.contains(field_id)) - .map(|field_id| (*field_id, file.path.as_str())) + .map(|field_id| { + PhysicalDataFileIdentity::try_new(dataset, file) + .map(|identity| (*field_id, identity)) + }) }) .collect() } @@ -221,9 +274,12 @@ async fn prune_stale_segment_coverage( let Some(current_fragment) = current_fragments.get(fragment_id) else { return true; }; + let historical_files = + fragment_field_files(&historical, historical_fragment, &indexed_field_ids); + let current_files = + fragment_field_files(dataset, current_fragment, &indexed_field_ids); let changed_files = - fragment_field_paths(historical_fragment, &indexed_field_ids) - != fragment_field_paths(current_fragment, &indexed_field_ids); + historical_files.is_none() || historical_files != current_files; let changed_overlays = prune_newer_overlays && current_fragment.overlays.iter().any(|overlay| { overlay.committed_version > version diff --git a/rust/lance/src/io/exec/fts.rs b/rust/lance/src/io/exec/fts.rs index a9bd8938184..3483a476d52 100644 --- a/rust/lance/src/io/exec/fts.rs +++ b/rust/lance/src/io/exec/fts.rs @@ -25,19 +25,24 @@ use datafusion_physical_plan::ExecutionPlanProperties; use datafusion_physical_plan::joins::{HashJoinExec, PartitionMode}; use datafusion_physical_plan::metrics::{BaselineMetrics, Count, Time}; use futures::future::try_join_all; -use futures::stream::{self}; +use futures::stream::{self, FuturesUnordered}; use futures::{FutureExt, StreamExt, TryStreamExt}; use itertools::Itertools; use lance_core::{ Error, ROW_ID, Result, - utils::{tokio::get_num_compute_intensive_cpus, tracing::StreamTracingExt}, + utils::{ + tokio::{get_num_compute_intensive_cpus, spawn_cpu}, + tracing::StreamTracingExt, + }, }; use lance_datafusion::utils::{ExecutionPlanMetricsSetExt, MetricsExt, PARTITIONS_SEARCHED_METRIC}; use lance_select::RowAddrMask; use lance_table::format::IndexMetadata; +use rustc_hash::FxHashSet; use super::PreFilterSource; use super::utils::{IndexMetrics, PreFilterMasks, build_prefilter}; +use crate::dataset::mem_wal::index::{QueryLocalFtsIndex, QueryLocalFtsStats}; use crate::index::scalar::inverted::{ ResolvedFtsField, fts_document_schema, load_segment_details, load_segments, transform_fts_document_stream, @@ -57,7 +62,8 @@ use lance_index::scalar::inverted::{ MemBM25Scorer, PreparedBm25Query, SCORE_COL, Scorer, build_global_bm25_scorer, compound_search, compound_search_prepared_match, compound_search_prepared_match_with_score_floor, compound_search_with_base_scorer, cross_column_compound_search, exclusive_scaled_score_floor, - flat_bm25_search_stream_with_options_and_scorer, fts_schema, prepare_bm25_query, + flat_bm25_search_stream_with_options_and_scorer, fts_schema, materialized_compound_top_k, + prepare_bm25_query, }; use lance_index::{prefilter::PreFilter, scalar::inverted::query::BooleanQuery}; use lance_tokenizer::{SimpleTokenizer, TextAnalyzer}; @@ -790,6 +796,398 @@ impl CompoundQueryExec { } } +#[derive(Debug)] +struct QueryLocalResidualShard { + index: QueryLocalFtsIndex, + stats: QueryLocalFtsStats, +} + +async fn index_query_local_residual_batch( + mut residual: QueryLocalResidualShard, + batch: RecordBatch, + allowed_terms: Arc>, +) -> Result { + spawn_cpu(move || { + let row_ids = batch + .column_by_name(ROW_ID) + .ok_or_else(|| { + Error::invalid_input( + "hybrid compound FTS residual input is missing _rowid".to_string(), + ) + })? + .as_primitive::(); + let stats = residual.index.insert_with_row_ids_for_terms( + &batch, + row_ids, + allowed_terms.as_ref(), + )?; + residual.stats.checked_add_assign(stats)?; + Ok(residual) + }) + .await +} + +/// Build a bounded set of independent residual posting shards. +/// +/// A single [`QueryLocalFtsIndex`] intentionally has one writer. Reusing one +/// index per CPU worker preserves that contract while allowing different scan +/// batches to tokenize in parallel. Completed workers immediately take the +/// next batch, so the entire stream is never collected in memory and the +/// number of live tokenizers/posting maps is bounded by the CPU pool size. +async fn index_query_local_residual( + residual_input: SendableRecordBatchStream, + seed: QueryLocalFtsIndex, + allowed_terms: Arc>, +) -> DataFusionResult> { + // Match flat FTS's CPU-task sizing. Dataset scan batches are normally + // row-bounded (often 8,192 rows), which can leave a small residual with + // only one or two tokenizer tasks. Byte rechunking keeps tasks substantial + // while exposing enough parallelism for variable-width text. + const ACCUMULATE_BYTES: usize = 256 * 1024; + const SLICE_BYTES: usize = 512 * 1024; + let input_schema = residual_input.schema(); + let mut residual_input = Box::pin(lance_arrow::stream::rechunk_stream_by_size( + residual_input, + input_schema, + ACCUMULATE_BYTES, + SLICE_BYTES, + )); + let parallelism = get_num_compute_intensive_cpus().max(1); + let mut initial_batches = Vec::with_capacity(parallelism); + let mut is_input_exhausted = false; + + while initial_batches.len() < parallelism { + let Some(batch) = residual_input.try_next().await? else { + is_input_exhausted = true; + break; + }; + initial_batches.push(batch); + } + + if initial_batches.is_empty() { + return Ok(vec![QueryLocalResidualShard { + index: seed, + stats: QueryLocalFtsStats::default(), + }]); + } + + // Construct every shard from the already-loaded seed before dispatching + // CPU work. This keeps tokenizer model I/O out of `spawn_cpu` closures. + let mut initial_shards = Vec::with_capacity(initial_batches.len()); + for _ in 1..initial_batches.len() { + initial_shards.push(QueryLocalResidualShard { + index: seed.empty_sibling(), + stats: QueryLocalFtsStats::default(), + }); + } + initial_shards.push(QueryLocalResidualShard { + index: seed, + stats: QueryLocalFtsStats::default(), + }); + + let mut in_flight = FuturesUnordered::new(); + for (shard, batch) in initial_shards.into_iter().zip(initial_batches) { + in_flight.push(index_query_local_residual_batch( + shard, + batch, + allowed_terms.clone(), + )); + } + + let mut shards = Vec::with_capacity(parallelism.min(in_flight.len())); + while let Some(shard) = in_flight.try_next().await? { + if is_input_exhausted { + shards.push(shard); + continue; + } + match residual_input.try_next().await? { + Some(batch) => in_flight.push(index_query_local_residual_batch( + shard, + batch, + allowed_terms.clone(), + )), + None => { + is_input_exhausted = true; + shards.push(shard); + } + } + } + Ok(shards) +} + +async fn query_local_residual_leaves( + shards: Vec, + query: FtsQuery, + scorer: Arc, +) -> Result>> { + let shard_leaves = stream::iter(shards.into_iter().map(|shard| { + let query = query.clone(); + let scorer = scorer.clone(); + spawn_cpu(move || shard.index.exact_leaf_results(&query, scorer.as_ref())) + })) + .buffered(get_num_compute_intensive_cpus().max(1)) + .try_collect::>() + .await?; + + let leaf_count = shard_leaves.first().map_or(0, Vec::len); + let mut merged = vec![Vec::new(); leaf_count]; + for leaves in shard_leaves { + if leaves.len() != leaf_count { + return Err(Error::internal(format!( + "hybrid compound FTS residual shards produced inconsistent leaf counts: expected {leaf_count}, got {}", + leaves.len() + ))); + } + for (merged, rows) in merged.iter_mut().zip(leaves) { + merged.extend(rows); + } + } + Ok(merged) +} + +fn residual_bm25_scorer( + committed_scorer: &MemBM25Scorer, + shards: &[QueryLocalResidualShard], +) -> Result { + let mut scorer = committed_scorer.clone(); + for shard in shards { + shard.stats.add_to_scorer(&mut scorer)?; + } + Ok(scorer) +} + +/// Compound FTS over committed postings plus a small append-only residual scan. +/// +/// The residual documents are tokenized once into query-local postings, rather +/// than once for every compound leaf. The indexed arm uses committed-index +/// BM25 statistics. The residual arm extends those statistics with the +/// query-local materialized documents, which matches the established mixed +/// flat-search approximation without rescanning the residual input or rebuilding +/// exact corpus statistics. +#[derive(Debug)] +pub(crate) struct HybridCompoundQueryExec { + dataset: Arc, + query: FtsQuery, + params: FtsSearchParams, + column: String, + segments: Arc<[IndexMetadata]>, + residual_input: Arc, + properties: Arc, + metrics: ExecutionPlanMetricsSet, +} + +impl HybridCompoundQueryExec { + pub(crate) fn new( + dataset: Arc, + query: FtsQuery, + params: FtsSearchParams, + column: String, + segments: Vec, + residual_input: Arc, + ) -> Self { + Self { + dataset, + query, + params, + column, + segments: Arc::from(segments), + residual_input, + properties: Arc::new(PlanProperties::new( + EquivalenceProperties::new(FTS_SCHEMA.clone()), + Partitioning::RoundRobinBatch(1), + EmissionType::Final, + Boundedness::Bounded, + )), + metrics: ExecutionPlanMetricsSet::new(), + } + } +} + +impl DisplayAs for HybridCompoundQueryExec { + fn fmt_as(&self, _t: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result { + write!( + f, + "HybridCompoundFtsScorer: column={}, query={}", + self.column, self.query + ) + } +} + +impl ExecutionPlan for HybridCompoundQueryExec { + fn name(&self) -> &str { + "HybridCompoundQueryExec" + } + + fn children(&self) -> Vec<&Arc> { + vec![&self.residual_input] + } + + fn required_input_distribution(&self) -> Vec { + vec![Distribution::SinglePartition] + } + + fn with_new_children( + self: Arc, + mut children: Vec>, + ) -> DataFusionResult> { + if children.len() != 1 { + return Err(DataFusionError::Internal(format!( + "hybrid compound FTS expected one residual child, got {}", + children.len() + ))); + } + let residual_input = children.pop().ok_or_else(|| { + DataFusionError::Internal("hybrid compound FTS lost its residual child".to_string()) + })?; + Ok(Arc::new(Self::new( + self.dataset.clone(), + self.query.clone(), + self.params.clone(), + self.column.clone(), + self.segments.to_vec(), + residual_input, + ))) + } + + #[instrument(name = "hybrid_compound_fts_exec", level = "debug", skip_all)] + fn execute( + &self, + partition: usize, + context: Arc, + ) -> DataFusionResult { + let dataset = self.dataset.clone(); + let query = self.query.clone(); + let params = self.params.clone(); + let column = self.column.clone(); + let segments = self.segments.clone(); + let residual_input = self.residual_input.clone(); + let metrics = Arc::new(FtsIndexMetrics::new(&self.metrics, partition)); + let schema = self.schema(); + + let stream = stream::once(async move { + let _timer = metrics.baseline_metrics.elapsed_compute().timer(); + let indices = + open_fts_segments(&dataset, &column, &segments, &metrics.index_metrics).await?; + let first_index = indices.first().ok_or_else(|| { + DataFusionError::Execution(format!( + "FTS index for column {column} has no committed segments" + )) + })?; + let field_id = dataset.schema().field_id(&column)?; + let tokenizer = first_index.tokenizer(); + let doc_type = tokenizer.doc_type(); + let residual_seed = QueryLocalFtsIndex::try_with_loaded_tokenizer( + field_id, + column.clone(), + first_index.params().clone(), + tokenizer, + )?; + let terms = residual_seed.exact_query_terms(&query)?; + if terms.is_empty() { + metrics.baseline_metrics.record_output(0); + return scored_documents_batch(schema, Vec::new()).map_err(DataFusionError::from); + } + let allowed_terms = Arc::new(terms.iter().cloned().collect::>()); + let query_tokens = Tokens::new(terms.clone(), doc_type); + let exact_params = params + .clone() + .with_fuzziness(Some(0)) + .with_phrase_slop(None); + + let residual_context = context.clone(); + let residual_indexing = async move { + let residual_input = residual_input.execute(partition, residual_context)?; + index_query_local_residual(residual_input, residual_seed, allowed_terms).await + }; + let scorer_build = async { + let scorer = build_global_bm25_scorer( + &indices, + &query_tokens, + &exact_params, + Some(metrics.as_ref()), + ) + .await?; + DataFusionResult::>::Ok(Arc::new(scorer)) + }; + let (residual_shards, committed_scorer) = + futures::future::try_join(residual_indexing, scorer_build).await?; + let residual_scorer = Arc::new(residual_bm25_scorer( + committed_scorer.as_ref(), + &residual_shards, + )?); + let limit = params.limit.ok_or_else(|| { + DataFusionError::Execution( + "hybrid compound FTS requires a bounded result limit".to_string(), + ) + })?; + + let prefilter = build_prefilter( + context, + partition, + &PreFilterSource::None, + dataset, + &segments, + PreFilterMasks { + overlay_block: None, + external_mask: None, + }, + )?; + let indexed_search = compound_search_with_base_scorer( + &indices, + &query, + ¶ms, + prefilter, + metrics.clone(), + committed_scorer, + ); + let residual_query = query.clone(); + let residual_search = async move { + let residual_leaves = query_local_residual_leaves( + residual_shards, + residual_query.clone(), + residual_scorer, + ) + .await?; + spawn_cpu(move || { + materialized_compound_top_k(&residual_query, residual_leaves, limit) + }) + .await + }; + let ((indexed_row_ids, indexed_scores), (residual_row_ids, residual_scores)) = + futures::future::try_join(indexed_search, residual_search).await?; + + let mut documents = indexed_row_ids + .into_iter() + .zip(indexed_scores) + .chain(residual_row_ids.into_iter().zip(residual_scores)) + .map(|(row_id, score)| ScoredDoc::new(row_id, score)) + .collect::>(); + documents.sort_unstable_by(|left, right| { + right + .score + .0 + .total_cmp(&left.score.0) + .then_with(|| left.row_id.cmp(&right.row_id)) + }); + documents.truncate(limit); + metrics.baseline_metrics.record_output(documents.len()); + scored_documents_batch(schema, documents).map_err(DataFusionError::from) + }); + Ok(Box::pin(RecordBatchStreamAdapter::new( + self.schema(), + stream.stream_in_current_span().boxed(), + ))) + } + + fn metrics(&self) -> Option { + Some(self.metrics.clone_inner()) + } + + fn properties(&self) -> &Arc { + &self.properties + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum WandExactnessCertificate { Exhaustive,