diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index 8efa561bbe4..3002da8323f 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -116,7 +116,7 @@ use crate::io::exec::filtered_read::{ use crate::io::exec::fts::{ BoostQueryExec, CombinedFieldsQueryExec, CompoundQueryExec, CrossColumnCompoundQueryExec, FlatMatchFilterExec, FlatMatchQueryExec, FtsDocumentExec, HybridCompoundQueryExec, - MatchQueryExec, PhraseQueryExec, SharedFtsScorer, + MatchQueryExec, PhraseQueryExec, SharedFtsScorer, SharedFtsScorerExec, }; use crate::io::exec::knn::MultivectorScoringExec; use crate::io::exec::scalar_index::{MaterializeIndexExec, ScalarIndexExec}; @@ -164,6 +164,12 @@ pub(crate) fn validate_batch_size(batch_size: usize) -> Result { Ok(validated) } +/// A restricted corpus scan must count every document before selecting output rows. +enum FlatMatchFilter<'a> { + BeforeScoring(&'a ExprFilterPlan), + AtEmission(PreFilterSource), +} + enum FtsOverlayPlan { Unchanged(Option>), RowLevel { @@ -4825,6 +4831,15 @@ impl Scanner { self.fragments_covered_by_fts_query(query).await?, ) .await?; + // Choose corpus-aware scoring only for the resolved root Match. Recursive leaves + // retain their existing scoring and approximation contracts. + if let FtsQuery::Match(match_query) = query + && let Some(plan) = self + .plan_restricted_match_query(match_query, ¶ms, filter_plan, &prefilter_source) + .await? + { + return Ok(plan); + } // Data overlay masking blocks stale rows from indexed leaves and re-evaluates only those // rows from their current values on the flat-text path. let fts_exec = self @@ -5365,7 +5380,7 @@ impl Scanner { HashMap::new(), &flat_query, &flat_params, - filter_plan, + FlatMatchFilter::BeforeScoring(filter_plan), None, ) .await?; @@ -5391,7 +5406,7 @@ impl Scanner { HashMap::new(), &flat_query, &flat_params, - filter_plan, + FlatMatchFilter::BeforeScoring(filter_plan), None, ) .await?; @@ -5442,7 +5457,7 @@ impl Scanner { stale_rows, &flat_query, &flat_params, - filter_plan, + FlatMatchFilter::BeforeScoring(filter_plan), shared_scorer, ) .await?, @@ -5465,7 +5480,7 @@ impl Scanner { HashMap::new(), &flat_query, &flat_params, - filter_plan, + FlatMatchFilter::BeforeScoring(filter_plan), None, ) .await?; @@ -5476,6 +5491,158 @@ impl Scanner { Self::combine_fts_leaf_plans(phrase_plan, flat_phrase_plan, params) } + /// Use global append-only corpus statistics independently of the selected candidates. + /// This entry is called only for a root scalar-text Match with a user prefilter or + /// explicit fragment selection. Other query shapes retain their existing planner. + async fn plan_restricted_match_query( + &self, + query: &MatchQuery, + params: &FtsSearchParams, + filter_plan: &ExprFilterPlan, + prefilter_source: &PreFilterSource, + ) -> Result>> { + let is_primary_fts = self.full_text_query.is_some() + && self.nearest.is_none() + && self.filter.query_filter.is_none(); + let has_postfilter = !self.prefilter && self.filter.expr_filter.is_some(); + let has_restriction = + self.fragments.is_some() || (self.prefilter && !filter_plan.is_empty()); + if !is_primary_fts + || self.fast_search + || has_postfilter + || !has_restriction + || params.wand_factor != 1.0 + || params.phrase_slop.is_some() + || query.fuzziness != Some(0) + || query.document_granularity != Some(DocumentGranularity::Row) + || self + .dataset + .fragments() + .iter() + .any(|f| f.deletion_file.is_some()) + { + return Ok(None); + } + let Some(column) = query.column.as_ref() else { + return Ok(None); + }; + let resolved = resolve_fts_field(self.dataset.schema(), column, DocumentGranularity::Row)?; + if resolved.has_lists() + || !self + .dataset + .schema() + .field_by_id(resolved.final_field_id) + .is_some_and(|field| { + matches!( + field.data_type(), + DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View + ) + }) + { + return Ok(None); + } + if params.limit == Some(0) + || self.dataset.fragments().is_empty() + || self.fragments.as_ref().is_some_and(Vec::is_empty) + { + return Ok(Some(Arc::new(EmptyExec::new(fts_schema( + DocumentGranularity::Row, + ))))); + } + + let index = self + .dataset + .load_scalar_index( + IndexCriteria::default() + .for_column(column) + .supports_fts() + .with_fts_document_granularity(DocumentGranularity::Row), + ) + .await?; + let (corpus_fragments, segments) = if let Some(index) = index { + // Keep the full residual corpus, including fragments outside candidate selection. + let corpus_fragments = self.dataset.unindexed_fragments(&index.name).await?; + if corpus_fragments.is_empty() { + return Ok(None); + } + let FtsOverlayPlan::Unchanged(segments) = self + .fts_overlay_plan(column, DocumentGranularity::Row, self.dataset.fragments()) + .await? + else { + return Ok(None); + }; + let segments = match segments { + Some(segments) => segments, + None => load_segments(&self.dataset, column, DocumentGranularity::Row) + .await? + .ok_or_else(|| Error::internal("FTS index has no physical segments"))?, + }; + let details = futures::future::try_join_all( + segments + .iter() + .map(|segment| load_physical_fts_details(&self.dataset, column, segment)), + ) + .await?; + if details.iter().any(|details| { + !matches!( + details.posting_format_version, + Some(INVERTED_INDEX_VERSION_V2 | INVERTED_INDEX_VERSION_V3) + ) + }) { + return Ok(None); + } + (corpus_fragments, Some(segments)) + } else { + (self.dataset.fragments().to_vec(), None) + }; + let candidate_filter = self + .prefilter_source( + filter_plan, + corpus_fragments + .iter() + .map(|fragment| fragment.id as u32) + .collect(), + ) + .await?; + let (indexed_plan, shared_scorer) = match segments { + Some(segments) => { + let shared_scorer = Arc::new(SharedFtsScorer::new()); + let indexed_plan = MatchQueryExec::new_with_segments_and_document_granularity( + self.dataset.clone(), + query.clone(), + params.clone(), + prefilter_source.clone(), + segments, + DocumentGranularity::Row, + ) + .with_shared_scorer(shared_scorer.clone()) + .with_external_mask(self.external_row_mask.clone()); + ( + Some(Arc::new(indexed_plan) as Arc), + Some(shared_scorer), + ) + } + None => (None, None), + }; + // Retain this producer even when the selected residual domain is empty. Its + // scorer must be published before indexed candidates can be pruned to top-k. + let flat_plan = self + .plan_flat_match_query( + corpus_fragments, + HashMap::new(), + query, + params, + FlatMatchFilter::AtEmission(candidate_filter), + shared_scorer.clone(), + ) + .await?; + let plan = Self::combine_fts_leaf_plans(indexed_plan, Some(flat_plan), params)?; + Ok(Some(match shared_scorer { + Some(shared_scorer) => Arc::new(SharedFtsScorerExec::new(plan, shared_scorer)), + None => plan, + })) + } + async fn plan_match_query( &self, query: &MatchQuery, @@ -5531,7 +5698,7 @@ impl Scanner { HashMap::new(), query, params, - filter_plan, + FlatMatchFilter::BeforeScoring(filter_plan), None, ) .await?; @@ -5557,7 +5724,7 @@ impl Scanner { HashMap::new(), query, params, - filter_plan, + FlatMatchFilter::BeforeScoring(filter_plan), None, ) .await?; @@ -5601,7 +5768,7 @@ impl Scanner { stale_rows, query, params, - filter_plan, + FlatMatchFilter::BeforeScoring(filter_plan), shared_scorer, ) .await?, @@ -5625,7 +5792,7 @@ impl Scanner { HashMap::new(), query, params, - filter_plan, + FlatMatchFilter::BeforeScoring(filter_plan), None, ) .await?; @@ -5677,9 +5844,14 @@ impl Scanner { stale_rows: HashMap, query: &MatchQuery, params: &FtsSearchParams, - filter_plan: &ExprFilterPlan, + filter: FlatMatchFilter<'_>, shared_scorer: Option>, ) -> Result> { + let unfiltered = ExprFilterPlan::default(); + let (filter_plan, candidate_filter) = match filter { + FlatMatchFilter::BeforeScoring(filter_plan) => (filter_plan, None), + FlatMatchFilter::AtEmission(candidate_filter) => (&unfiltered, Some(candidate_filter)), + }; let column = query .column .as_ref() @@ -5720,6 +5892,9 @@ impl Scanner { if let Some(shared_scorer) = shared_scorer { flat_match_plan = flat_match_plan.with_shared_scorer(shared_scorer); } + if let Some(candidate_filter) = candidate_filter { + flat_match_plan = flat_match_plan.with_candidate_filter(candidate_filter); + } let flat_match_plan: Arc = Arc::new(flat_match_plan); // Unindexed fragments and stale rows never reach the index-side prefilter, // so apply the external row-address mask to the flat FTS results here @@ -17024,45 +17199,46 @@ full_filter=name LIKE Utf8(\"test%2\"), refine_filter=name LIKE Utf8(\"test%2\") .await?; log::info!("Test case: Full text search with unindexed rows and prefilter"); - // After routing flat FTS through `FilteredReadExec`, the BTree on `i` - // pushes into the unindexed-fragment scan too — no more `FilterExec` on - // top of an unfiltered `LanceScan`. Legacy uses the `MaterializeIndex` - // shape, v2 uses `LanceRead` with `full_filter` set. + // The shared scorer reads the full unindexed text corpus for BM25 statistics. + // Keep the BTree prefilter in separate candidate inputs so it restricts + // emitted matches without changing corpus statistics. let expected = if data_storage_version == LanceFileVersion::Legacy { r#"ProjectionExec: expr=[s@2 as s, _score@1 as _score, _rowid@0 as _rowid] Take: columns="_rowid, _score, (s)" CoalesceBatchesExec: target_batch_size=8192 + MatchCorpus + SortExec: expr=[_score@1 DESC NULLS LAST], preserve_partitioning=[false] + CoalescePartitionsExec + UnionExec + MatchQuery: column=s, query=[hello] + CoalescePartitionsExec + UnionExec + MaterializeIndex: query=[i > 10]@i_idx(BTree) + ProjectionExec: expr=[_rowid@1 as _rowid] + FilterExec: i@0 > 10 + LanceScan: uri=..., projection=[i], row_id=true, row_addr=false, ordered=false, range=None + FlatMatchQuery: column=s, query=hello + LanceScan: uri=..., projection=[s], row_id=true, row_addr=false, ordered=true, range=None + CoalescePartitionsExec + UnionExec + MaterializeIndex: query=[i > 10]@i_idx(BTree) + ProjectionExec: expr=[_rowid@1 as _rowid] + FilterExec: i@0 > 10 + LanceScan: uri=..., projection=[i], row_id=true, row_addr=false, ordered=false, range=None"# + } else { + r#"ProjectionExec: expr=[s@2 as s, _score@1 as _score, _rowid@0 as _rowid] + LanceRead: uri=..., projection=[s], source=stream(_rowid) + MatchCorpus SortExec: expr=[_score@1 DESC NULLS LAST], preserve_partitioning=[false] CoalescePartitionsExec UnionExec MatchQuery: column=s, query=[hello] - CoalescePartitionsExec - UnionExec - MaterializeIndex: query=[i > 10]@i_idx(BTree) - ProjectionExec: expr=[_rowid@1 as _rowid] - FilterExec: i@0 > 10 - LanceScan: uri=..., projection=[i], row_id=true, row_addr=false, ordered=false, range=None + LanceRead: uri=..., projection=[], num_fragments=5, range_before=None, range_after=None, row_id=true, row_addr=false, full_filter=i > Int32(10), refine_filter=-- + ScalarIndexQuery: query=[i > 10]@i_idx(BTree) FlatMatchQuery: column=s, query=hello - CoalescePartitionsExec - UnionExec - Take: columns="_rowid, (s)" - CoalesceBatchesExec: target_batch_size=8192 - MaterializeIndex: query=[i > 10]@i_idx(BTree) - ProjectionExec: expr=[_rowid@2 as _rowid, s@1 as s] - FilterExec: i@0 > 10 - LanceScan: uri=..., projection=[i, s], row_id=true, row_addr=false, ordered=false, range=None"# - } else { - r#"ProjectionExec: expr=[s@2 as s, _score@1 as _score, _rowid@0 as _rowid] - LanceRead: uri=..., projection=[s], source=stream(_rowid) - SortExec: expr=[_score@1 DESC NULLS LAST], preserve_partitioning=[false] - CoalescePartitionsExec - UnionExec - MatchQuery: column=s, query=[hello] - LanceRead: uri=..., projection=[], num_fragments=5, range_before=None, range_after=None, row_id=true, row_addr=false, full_filter=i > Int32(10), refine_filter=-- - ScalarIndexQuery: query=[i > 10]@i_idx(BTree) - FlatMatchQuery: column=s, query=hello - LanceRead: uri=..., projection=[s], num_fragments=1, range_before=None, range_after=None, row_id=true, row_addr=false, full_filter=i > Int32(10), refine_filter=-- - ScalarIndexQuery: query=[i > 10]@i_idx(BTree)"# + LanceRead: uri=..., projection=[s], num_fragments=1, range_before=None, range_after=None, row_id=true, row_addr=false, full_filter=--, refine_filter=-- + LanceRead: uri=..., projection=[], num_fragments=1, range_before=None, range_after=None, row_id=true, row_addr=false, full_filter=i > Int32(10), refine_filter=-- + ScalarIndexQuery: query=[i > 10]@i_idx(BTree)"# }; assert_plan_equals( &dataset.dataset, diff --git a/rust/lance/src/dataset/tests/dataset_index.rs b/rust/lance/src/dataset/tests/dataset_index.rs index cc44fced886..fd8f1b55438 100644 --- a/rust/lance/src/dataset/tests/dataset_index.rs +++ b/rust/lance/src/dataset/tests/dataset_index.rs @@ -6,6 +6,7 @@ use std::future::Future; use std::pin::Pin; use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; +use std::time::Duration; use std::vec; use crate::dataset::ROW_ID; @@ -39,7 +40,7 @@ use lance_core::cache::{ CacheBackend, CacheCodec, CacheEntry, InternalCacheKey, LanceCache, QuickCacheBackend, }; use lance_core::utils::tempfile::TempStrDir; -use lance_datafusion::exec::ExecutionSummaryCounts; +use lance_datafusion::exec::{ExecutionSummaryCounts, LanceExecutionOptions, get_session_context}; use lance_datafusion::utils::PARTITIONS_SEARCHED_METRIC; use lance_datagen::{BatchCount, Dimension, RowCount, array, gen_batch}; use lance_file::reader::{FileReader, FileReaderOptions}; @@ -57,8 +58,10 @@ use lance_index::{IndexType, scalar::ScalarIndexParams, vector::DIST_COL}; use lance_io::scheduler::{ScanScheduler, SchedulerConfig}; use lance_io::utils::CachedFileSize; use lance_linalg::distance::MetricType; +use lance_select::IndexExprResult; use datafusion::common::{assert_contains, assert_not_contains}; +use datafusion::physical_plan::display::DisplayableExecutionPlan; use futures::{StreamExt, TryStreamExt}; use itertools::Itertools; use lance_arrow::json::ARROW_JSON_EXT_NAME; @@ -4045,6 +4048,1113 @@ async fn test_fts_unindexed_data() { assert_eq!(results.num_rows(), 1); } +#[derive(Clone, Copy, Debug)] +enum RestrictedMatchDomain { + Prefilter, + IndexedFragment, + Filter(&'static str), + Fragments(&'static [usize]), +} + +#[derive(Debug, serde::Serialize)] +struct RawMatchCorpus { + documents: usize, + total_tokens: usize, + average_length: f64, + document_lengths: Vec, + document_frequencies: Vec, + query_terms: Vec, +} + +fn raw_match_oracle(corpus: &[(i32, &str)], query: &str) -> (RawMatchCorpus, Vec<(i32, f64)>) { + // Derive every statistic from raw strings, without calling a production + // tokenizer or scorer. Fixture words need no stemming or normalization. + let documents: Vec<_> = corpus + .iter() + .map(|(id, text)| (*id, text.split_whitespace().collect::>())) + // Empty analyzed documents do not contribute to BM25 corpus statistics. + .filter(|(_, tokens)| !tokens.is_empty()) + .collect(); + let query_terms: Vec<_> = query.split_whitespace().map(str::to_owned).collect(); + let document_frequencies: Vec<_> = query_terms + .iter() + .map(|term| { + documents + .iter() + .filter(|(_, document)| document.contains(&term.as_str())) + .count() + }) + .collect(); + let total_tokens: usize = documents.iter().map(|(_, tokens)| tokens.len()).sum(); + let average_length = if documents.is_empty() { + 0.0 + } else { + total_tokens as f64 / documents.len() as f64 + }; + let mut scores: Vec<(i32, f64)> = documents + .iter() + .filter_map(|(id, document)| { + let score = query_terms + .iter() + .zip(&document_frequencies) + .filter_map(|(term, document_frequency)| { + let frequency = document + .iter() + .filter(|token| **token == term.as_str()) + .count() as f64; + if frequency == 0.0 { + return None; + } + let inverse_document_frequency = + ((documents.len() as f64 - *document_frequency as f64 + 0.5) + / (*document_frequency as f64 + 0.5) + + 1.0) + .ln(); + Some( + inverse_document_frequency * frequency * (1.2 + 1.0) + / (frequency + + 1.2 + * (1.0 - 0.75 + 0.75 * document.len() as f64 / average_length)), + ) + }) + .sum::(); + (score > 0.0).then_some((*id, score)) + }) + .collect(); + scores.sort_by(|left, right| { + right + .1 + .total_cmp(&left.1) + .then_with(|| left.0.cmp(&right.0)) + }); + ( + RawMatchCorpus { + documents: documents.len(), + total_tokens, + average_length, + document_lengths: documents.iter().map(|(_, tokens)| tokens.len()).collect(), + document_frequencies, + query_terms, + }, + scores, + ) +} + +#[derive(Debug, serde::Serialize)] +struct RestrictedMatchObservation { + stage: &'static str, + limit: Option, + plan: std::result::Result, + rows: std::result::Result, String>, + scan_stats: Option, +} + +async fn restricted_match_observation( + dataset: &Dataset, + domain: RestrictedMatchDomain, + stage: &'static str, + limit: Option, + terms: &str, +) -> RestrictedMatchObservation { + let query = MatchQuery::new(terms.to_owned()) + .with_column(Some("text".to_owned())) + .with_fuzziness(Some(0)) + .with_operator(Operator::Or) + .with_document_granularity(DocumentGranularity::Row); + let collected_stats = Arc::new(Mutex::new(None::)); + let stats_setter = collected_stats.clone(); + let mut scanner = dataset.scan(); + scanner + .scan_stats_callback(Arc::new(move |stats| { + *stats_setter.lock().unwrap() = Some(stats.clone()); + })) + .project(&["id"]) + .unwrap() + .full_text_search(FullTextSearchQuery::new_query(query.into())) + .unwrap() + .limit(limit, None) + .unwrap(); + match domain { + RestrictedMatchDomain::Prefilter => { + scanner.prefilter(true).filter("id < 2").unwrap(); + } + RestrictedMatchDomain::IndexedFragment => { + scanner.with_fragments(vec![dataset.get_fragments()[0].clone().into()]); + } + RestrictedMatchDomain::Filter(filter) => { + scanner.prefilter(true).filter(filter).unwrap(); + } + RestrictedMatchDomain::Fragments(positions) => { + let fragments = dataset.get_fragments(); + scanner.with_fragments( + positions + .iter() + .map(|position| fragments[*position].clone().into()) + .collect(), + ); + } + } + let plan = tokio::time::timeout(Duration::from_secs(10), scanner.explain_plan(false)) + .await + .map_err(|_| "restricted Match planning timed out after 10 seconds".to_owned()) + .and_then(|result| result.map_err(|error| error.to_string())); + let batch = tokio::time::timeout(Duration::from_secs(10), scanner.try_into_batch()) + .await + .map_err(|_| "restricted Match execution timed out after 10 seconds".to_owned()) + .and_then(|result| result.map_err(|error| error.to_string())); + let rows = batch.and_then(|batch| { + let ids = batch + .column_by_name("id") + .and_then(|column| column.as_primitive_opt::()) + .ok_or_else(|| "restricted Match result is missing an Int32 id column".to_owned())?; + let scores = batch + .column_by_name(SCORE_COL) + .and_then(|column| column.as_primitive_opt::()) + .ok_or_else(|| { + "restricted Match result is missing a Float32 score column".to_owned() + })?; + Ok(ids + .values() + .iter() + .copied() + .zip(scores.values().iter().copied()) + .collect()) + }); + let scan_stats = collected_stats.lock().unwrap().take().map(|stats| { + serde_json::json!({ + "iops": stats.iops, + "requests": stats.requests, + "bytes_read": stats.bytes_read, + "all_counts": stats.all_counts, + "all_times": stats.all_times, + }) + }); + RestrictedMatchObservation { + stage, + limit, + plan, + rows, + scan_stats, + } +} + +const RESTRICTED_MATCH_CORPUS: [(i32, &str); 8] = [ + (0, "alpha"), + (1, "beta"), + (2, "beta gamma delta"), + (3, "alpha"), + (4, "alpha alpha"), + (5, "epsilon"), + (6, "alpha beta"), + (7, "alpha"), +]; + +fn restricted_match_fragments(layout: &[&[usize]]) -> Vec)>> { + layout + .iter() + .map(|positions| { + positions + .iter() + .map(|position| { + let (id, text) = RESTRICTED_MATCH_CORPUS[*position]; + (id, Some(text)) + }) + .collect() + }) + .collect() +} + +async fn create_restricted_match_index(dataset: &mut Dataset) { + let params = InvertedIndexParams::default() + .base_tokenizer("whitespace".to_owned()) + .stem(false) + .remove_stop_words(false); + dataset + .create_index(&["text"], IndexType::Inverted, None, ¶ms, true) + .await + .unwrap(); +} + +async fn restricted_match_dataset( + fragments: &[Vec<(i32, Option<&str>)>], + indexed_fragments: usize, + stable_row_ids: bool, +) -> Dataset { + assert!(!fragments.is_empty() && indexed_fragments <= fragments.len()); + let mut batches = fragments.iter().map(|rows| { + arrow_array::record_batch!( + ( + "id", + Int32, + rows.iter().map(|(id, _)| *id).collect::>() + ), + ( + "text", + Utf8, + rows.iter().map(|(_, text)| *text).collect::>() + ) + ) + .unwrap() + }); + let initial = batches.next().unwrap(); + let schema = initial.schema(); + let mut dataset = Dataset::write( + RecordBatchIterator::new(vec![initial].into_iter().map(Ok), schema), + "memory://", + Some(WriteParams { + enable_stable_row_ids: stable_row_ids, + ..Default::default() + }), + ) + .await + .unwrap(); + for (position, batch) in batches.enumerate() { + if position + 1 == indexed_fragments { + create_restricted_match_index(&mut dataset).await; + } + let schema = batch.schema(); + dataset + .append( + RecordBatchIterator::new(vec![batch].into_iter().map(Ok), schema), + None, + ) + .await + .unwrap(); + } + if indexed_fragments == fragments.len() { + create_restricted_match_index(&mut dataset).await; + } + assert_eq!(dataset.get_fragments().len(), fragments.len()); + dataset +} + +async fn restricted_match_maintenance_observations( + dataset: &mut Dataset, + domain: RestrictedMatchDomain, + terms: &str, + limits: &[Option], +) -> Vec { + let mut observations = Vec::with_capacity(limits.len() * 2); + for stage in ["before_optimize", "after_optimize"] { + if stage == "after_optimize" { + dataset + .optimize_indices(&OptimizeOptions::append()) + .await + .unwrap(); + } + for limit in limits { + observations + .push(restricted_match_observation(dataset, domain, stage, *limit, terms).await); + } + } + observations +} + +fn assert_restricted_match_oracle( + name: &str, + corpus: &[(i32, &str)], + domain: RestrictedMatchDomain, + candidate_ids: &[i32], + observations: &[RestrictedMatchObservation], +) { + let (corpus_stats, mut expected) = raw_match_oracle(corpus, "alpha beta"); + expected.retain(|(id, _)| candidate_ids.contains(id)); + let evidence = serde_json::json!({ + "case": name, + "raw_rows": corpus, + "domain": format!("{domain:?}"), + "candidate_ids": candidate_ids, + "corpus": corpus_stats, + "expected": expected, + "observations": observations, + }); + for observation in observations { + assert!(observation.plan.is_ok(), "planning failed: {evidence:#}"); + let mut actual = observation + .rows + .as_ref() + .unwrap_or_else(|error| panic!("execution failed: {error}; {evidence:#}")) + .clone(); + let mut expected = expected.clone(); + if let Some(limit) = observation.limit { + expected.truncate(limit as usize); + } + // Append ordering can change physical row IDs of tied residual rows. + // All-hit assertions compare the requested logical IDs and their scores. + actual.sort_by_key(|(id, _)| *id); + expected.sort_by_key(|(id, _)| *id); + assert_eq!(actual.len(), expected.len(), "wrong hits: {evidence:#}"); + for ((actual_id, actual_score), (expected_id, expected_score)) in + actual.iter().zip(&expected) + { + assert_eq!(actual_id, expected_id, "wrong candidate: {evidence:#}"); + assert!( + (f64::from(*actual_score) - expected_score).abs() < 1e-5, + "wrong global corpus score: {evidence:#}" + ); + } + } +} + +fn assert_restricted_match_membership( + name: &str, + expected_ids: &[i32], + observations: &[RestrictedMatchObservation], +) { + let evidence = serde_json::json!({ + "case": name, + "expected_ids": expected_ids, + "observations": observations, + }); + let mut expected_ids = expected_ids.to_vec(); + expected_ids.sort_unstable(); + for observation in observations { + assert!(observation.plan.is_ok(), "planning failed: {evidence:#}"); + let rows = observation + .rows + .as_ref() + .unwrap_or_else(|error| panic!("execution failed: {error}; {evidence:#}")); + let mut actual_ids: Vec<_> = rows.iter().map(|(id, _)| *id).collect(); + actual_ids.sort_unstable(); + assert_eq!(actual_ids, expected_ids, "wrong candidates: {evidence:#}"); + assert!( + rows.iter().all(|(_, score)| score.is_finite()), + "non-finite score: {evidence:#}" + ); + } +} + +#[rstest] +#[case::reported(&["alpha", "beta"])] +#[case::non_tie(&["alpha", "beta", "beta gamma delta"])] +#[tokio::test] +async fn test_fts_9058_restricted_match_corpus_stability( + #[case] indexed_texts: &[&str], + #[values( + RestrictedMatchDomain::Prefilter, + RestrictedMatchDomain::IndexedFragment + )] + domain: RestrictedMatchDomain, + #[values(false, true)] use_stable_row_id: bool, +) { + // Based on the filter and fragment cases reported by sbrunk and + // lance-gatekeeper[bot] in https://github.com/lance-format/lance/issues/9058. + let initial = arrow_array::record_batch!( + ("text", Utf8, indexed_texts.to_vec()), + ( + "id", + Int32, + (0..indexed_texts.len() as i32).collect::>() + ) + ) + .unwrap(); + let schema = initial.schema(); + let mut dataset = Dataset::write( + RecordBatchIterator::new(vec![initial].into_iter().map(Ok), schema), + "memory://", + Some(WriteParams { + enable_stable_row_ids: use_stable_row_id, + ..Default::default() + }), + ) + .await + .unwrap(); + let params = InvertedIndexParams::default() + .base_tokenizer("whitespace".to_owned()) + .stem(false) + .remove_stop_words(false); + dataset + .create_index(&["text"], IndexType::Inverted, None, ¶ms, true) + .await + .unwrap(); + + let first_appended_id = indexed_texts.len() as i32; + let appended = arrow_array::record_batch!( + ("text", Utf8, vec!["alpha"; 10]), + ( + "id", + Int32, + (first_appended_id..first_appended_id + 10).collect::>() + ) + ) + .unwrap(); + let schema = appended.schema(); + dataset + .append( + RecordBatchIterator::new(vec![appended].into_iter().map(Ok), schema), + None, + ) + .await + .unwrap(); + assert_eq!(dataset.get_fragments().len(), 2); + + let mut corpus = indexed_texts.to_vec(); + corpus.extend(["alpha"; 10]); + let raw_rows: Vec<_> = corpus + .iter() + .enumerate() + .map(|(id, text)| (id as i32, *text)) + .collect(); + let (corpus_stats, mut expected) = raw_match_oracle(&raw_rows, "alpha beta"); + let candidate_count = match domain { + RestrictedMatchDomain::Prefilter => 2, + RestrictedMatchDomain::IndexedFragment => indexed_texts.len(), + _ => unreachable!("initial fixtures only use the two reported restrictions"), + }; + expected.retain(|(id, _)| (*id as usize) < candidate_count); + assert_eq!(expected[0].0, 1, "the full corpus must prefer beta"); + + // Collect both sides of maintenance before comparing scores, so a failing + // baseline still records both plans and the actual change in results. + let mut observations = Vec::with_capacity(4); + for stage in ["before_optimize", "after_optimize"] { + if stage == "after_optimize" { + dataset + .optimize_indices(&OptimizeOptions::append()) + .await + .unwrap(); + } + for limit in [Some(1), None] { + observations.push( + restricted_match_observation(&dataset, domain, stage, limit, "alpha beta").await, + ); + } + } + let evidence = serde_json::json!({ + "indexed_texts": indexed_texts, + "appended_texts": vec!["alpha"; 10], + "domain": format!("{domain:?}"), + "stable_row_ids": use_stable_row_id, + "corpus": corpus_stats, + "expected": expected, + "observations": observations, + }); + for observation in &observations { + assert!( + observation.plan.is_ok(), + "restricted Match planning failed: {evidence:#}" + ); + let rows = observation.rows.as_ref().unwrap_or_else(|error| { + panic!("restricted Match execution failed: {error}; {evidence:#}") + }); + let expected = match observation.limit { + Some(limit) => &expected[..limit as usize], + None => &expected, + }; + assert_eq!( + rows.len(), + expected.len(), + "candidate membership changed: {evidence:#}" + ); + for ((actual_id, actual_score), (expected_id, expected_score)) in rows.iter().zip(expected) + { + assert_eq!( + actual_id, expected_id, + "wrong ranked candidate: {evidence:#}" + ); + assert!( + (f64::from(*actual_score) - expected_score).abs() < 1e-5, + "score differs from the global corpus oracle: {evidence:#}" + ); + } + if observation.stage == "before_optimize" && observation.limit == Some(1) { + // GlobalLimit can finish before its child reaches EOF. The scanner's + // completion callback must still see the residual corpus read metrics. + let rows_scanned = observation + .scan_stats + .as_ref() + .and_then(|stats| stats["all_counts"]["rows_scanned"].as_u64()) + .expect("restricted Match callback is missing row-read metrics"); + assert!( + rows_scanned >= 10, + "residual corpus reads were hidden: {evidence:#}" + ); + } + } +} + +#[rstest] +#[case::top(0)] +#[case::offset(1)] +#[tokio::test] +async fn test_fts_9058_default_match_uses_global_corpus(#[case] offset: i64) { + let raw_rows: Vec<_> = [(0, "alpha"), (1, "beta")] + .into_iter() + .chain((2..12).map(|id| (id, "alpha"))) + .collect(); + let fragments = [&raw_rows[..2], &raw_rows[2..]] + .map(|rows| rows.iter().map(|(id, text)| (*id, Some(*text))).collect()); + let mut dataset = restricted_match_dataset(&fragments, 1, false).await; + let (_, mut expected) = raw_match_oracle(&raw_rows, "alpha beta"); + expected.retain(|(id, _)| *id < 2); + let mut observations = Vec::with_capacity(2); + for stage in ["before_optimize", "after_optimize"] { + if stage == "after_optimize" { + dataset + .optimize_indices(&OptimizeOptions::append()) + .await + .unwrap(); + } + // Exercise the reported query without explicitly resolving fuzziness + // or document granularity, including the scanner's top-k offset. + let query = MatchQuery::new("alpha beta".to_owned()).with_column(Some("text".to_owned())); + let mut scanner = dataset.scan(); + scanner + .project(&["id"]) + .unwrap() + .full_text_search(FullTextSearchQuery::new_query(query.into())) + .unwrap() + .prefilter(true) + .filter("id < 2") + .unwrap() + .limit(Some(1), Some(offset)) + .unwrap(); + let plan = tokio::time::timeout(Duration::from_secs(10), scanner.explain_plan(false)) + .await + .unwrap() + .unwrap(); + let batch = tokio::time::timeout(Duration::from_secs(10), scanner.try_into_batch()) + .await + .unwrap() + .unwrap(); + observations.push((stage, plan, batch)); + } + for (stage, plan, batch) in observations { + if stage == "before_optimize" { + assert_contains!(&plan, "MatchCorpus"); + } + assert_eq!(batch.num_rows(), 1, "{stage}: {plan}"); + let (expected_id, expected_score) = expected[offset as usize]; + assert_eq!( + batch["id"].as_primitive::().value(0), + expected_id + ); + let score = batch[SCORE_COL].as_primitive::().value(0); + assert!( + (f64::from(score) - expected_score).abs() < 1e-5, + "{stage}: score {score} differs from {expected_score}; {plan}" + ); + } +} + +#[rstest] +#[case::mixed_fragments(RestrictedMatchDomain::Fragments(&[0, 1]), &[0, 1, 2, 3, 4])] +#[case::residual_fragment(RestrictedMatchDomain::Fragments(&[2]), &[5, 6, 7])] +#[case::empty_fragments(RestrictedMatchDomain::Fragments(&[]), &[])] +#[case::one_filtered_hit(RestrictedMatchDomain::Filter("id = 1"), &[1])] +#[case::no_filtered_hits(RestrictedMatchDomain::Filter("id < 0"), &[])] +#[case::residual_filtered_hits(RestrictedMatchDomain::Filter("id >= 3"), &[3, 4, 5, 6, 7])] +#[tokio::test] +async fn test_fts_9058_candidate_domains( + #[case] domain: RestrictedMatchDomain, + #[case] candidate_ids: &[i32], + #[values(false, true)] stable_row_ids: bool, +) { + let fragments = restricted_match_fragments(&[&[0, 1, 2], &[3, 4], &[5, 6, 7]]); + let mut dataset = restricted_match_dataset(&fragments, 1, stable_row_ids).await; + let observations = restricted_match_maintenance_observations( + &mut dataset, + domain, + "alpha beta", + &[Some(1), None], + ) + .await; + assert_restricted_match_oracle( + &format!("9058-domain-{domain:?}.json"), + &RESTRICTED_MATCH_CORPUS, + domain, + candidate_ids, + &observations, + ); +} + +#[rstest] +#[case::allow("id = 6", &[6])] +#[case::block("id != 1", &[0, 2, 3, 4, 5, 6, 7])] +#[tokio::test] +async fn test_fts_9058_scalar_index_candidates( + #[case] filter: &'static str, + #[case] candidate_ids: &[i32], + #[values(false, true)] stable_row_ids: bool, +) { + let fragments = restricted_match_fragments(&[&[0, 1, 2], &[3, 4], &[5, 6, 7]]); + let mut dataset = restricted_match_dataset(&fragments, 1, stable_row_ids).await; + // Build the candidate index after every append so it covers the two + // residual fragments as well as the fragment with the text index. + dataset + .create_index( + &["id"], + IndexType::BTree, + None, + &ScalarIndexParams::default(), + true, + ) + .await + .unwrap(); + + let query = MatchQuery::new("alpha beta".to_owned()).with_column(Some("text".to_owned())); + let mut scanner = dataset.scan(); + scanner + .full_text_search(FullTextSearchQuery::new_query(query.into())) + .unwrap() + .prefilter(true) + .filter(filter) + .unwrap(); + let plan = tokio::time::timeout(Duration::from_secs(10), scanner.create_plan()) + .await + .unwrap() + .unwrap(); + let mut pending = vec![plan]; + let mut candidate_indices = Vec::new(); + while let Some(node) = pending.pop() { + if node.name() == "FlatMatchQueryExec" { + let children = node.children(); + assert_eq!(children.len(), 2); + assert_eq!(children[1].name(), "ScalarIndexExec"); + candidate_indices.push(children[1].clone()); + } + pending.extend(node.children().into_iter().cloned()); + } + assert_eq!(candidate_indices.len(), 1); + let batches = tokio::time::timeout( + Duration::from_secs(10), + datafusion::physical_plan::collect( + candidate_indices.pop().unwrap(), + get_session_context(&LanceExecutionOptions::default()).task_ctx(), + ), + ) + .await + .unwrap() + .unwrap(); + assert_eq!(batches.len(), 1); + let (selection, covered_fragments) = IndexExprResult::deserialize(&batches[0]).unwrap(); + assert!(selection.is_exact()); + // Scalar prefilters are scoped to the flat branch's two residual fragments. + assert_eq!(covered_fragments.iter().collect::>(), vec![1, 2]); + assert_eq!(selection.upper.block_list().is_some(), filter == "id != 1"); + + let domain = RestrictedMatchDomain::Filter(filter); + let observations = restricted_match_maintenance_observations( + &mut dataset, + domain, + "alpha beta", + &[Some(1), None], + ) + .await; + for observation in &observations { + let plan = observation.plan.as_ref().unwrap(); + if observation.stage == "before_optimize" { + assert_contains!(plan, "MatchCorpus"); + } + assert_contains!(plan, "ScalarIndexQuery:"); + assert_contains!(plan, "@id_idx(BTree)"); + } + assert_restricted_match_oracle( + filter, + &RESTRICTED_MATCH_CORPUS, + domain, + candidate_ids, + &observations, + ); +} + +#[rstest] +#[case::forward("forward", &[&[0_usize, 1, 2] as &[usize], &[3, 4], &[5, 6, 7]])] +#[case::reversed("reversed", &[&[0_usize, 1, 2] as &[usize], &[5, 6, 7], &[3, 4]])] +#[case::three_residuals("three-residuals", &[&[0_usize, 1, 2] as &[usize], &[3], &[4, 5], &[6, 7]])] +#[tokio::test] +async fn test_fts_9058_residual_layout(#[case] name: &str, #[case] layout: &[&[usize]]) { + let fragments = restricted_match_fragments(layout); + let mut dataset = restricted_match_dataset(&fragments, 1, false).await; + let domain = RestrictedMatchDomain::Prefilter; + let observations = restricted_match_maintenance_observations( + &mut dataset, + domain, + "alpha beta", + &[Some(1), None], + ) + .await; + assert_restricted_match_oracle( + &format!("9058-layout-{name}.json"), + &RESTRICTED_MATCH_CORPUS, + domain, + &[0, 1], + &observations, + ); +} + +#[rstest] +#[tokio::test] +async fn test_fts_9058_no_index_matches_complete_index( + #[values( + RestrictedMatchDomain::Prefilter, + RestrictedMatchDomain::IndexedFragment + )] + domain: RestrictedMatchDomain, +) { + let fragments = restricted_match_fragments(&[&[0, 1, 2], &[3, 4], &[5, 6, 7]]); + let mut dataset = restricted_match_dataset(&fragments, 0, false).await; + let mut observations = Vec::with_capacity(4); + for stage in ["without_index", "complete_index"] { + if stage == "complete_index" { + create_restricted_match_index(&mut dataset).await; + } + for limit in [Some(1), None] { + observations.push( + restricted_match_observation(&dataset, domain, stage, limit, "alpha beta").await, + ); + } + } + let candidate_ids: &[i32] = match domain { + RestrictedMatchDomain::Prefilter => &[0, 1], + RestrictedMatchDomain::IndexedFragment => &[0, 1, 2], + _ => unreachable!("coverage comparison uses the two reported restrictions"), + }; + assert_restricted_match_oracle( + &format!("9058-no-index-{domain:?}.json"), + &RESTRICTED_MATCH_CORPUS, + domain, + candidate_ids, + &observations, + ); +} + +#[rstest] +#[case::zero_limit("zero-limit", "alpha beta", Some(0))] +#[case::no_matches("no-matches", "absent", Some(1))] +#[case::empty_query("empty-query", "", Some(1))] +#[tokio::test] +async fn test_fts_9058_empty_results_complete( + #[case] name: &str, + #[case] terms: &str, + #[case] limit: Option, + #[values( + RestrictedMatchDomain::Prefilter, + RestrictedMatchDomain::IndexedFragment + )] + domain: RestrictedMatchDomain, +) { + let fragments = restricted_match_fragments(&[&[0, 1, 2], &[3, 4], &[5, 6, 7]]); + let mut dataset = restricted_match_dataset(&fragments, 1, false).await; + let observations = + restricted_match_maintenance_observations(&mut dataset, domain, terms, &[limit]).await; + assert_restricted_match_membership( + &format!("9058-empty-{name}-{domain:?}.json"), + &[], + &observations, + ); +} + +#[rstest] +#[case::all_null("all-null", &[None, None], &[None, None], &[])] +#[case::all_empty("all-empty", &[Some(""), Some("")], &[Some(""), Some("")], &[])] +#[case::empty_index("empty-index", &[None, Some("")], &[Some("alpha"), None], &[2])] +#[case::empty_residual("empty-residual", &[Some("alpha"), None], &[None, Some("")], &[0])] +#[tokio::test] +async fn test_fts_9058_empty_text_corpus( + #[case] name: &str, + #[case] initial: &[Option<&str>], + #[case] residual: &[Option<&str>], + #[case] expected_ids: &[i32], +) { + let fragments = vec![ + initial + .iter() + .enumerate() + .map(|(id, text)| (id as i32, *text)) + .collect(), + residual + .iter() + .enumerate() + .map(|(id, text)| ((initial.len() + id) as i32, *text)) + .collect(), + ]; + let mut dataset = restricted_match_dataset(&fragments, 1, false).await; + let observations = restricted_match_maintenance_observations( + &mut dataset, + RestrictedMatchDomain::Filter("id >= 0"), + "alpha beta", + &[None], + ) + .await; + // These controls check empty/null handling without imposing a different + // null-document counting convention on existing index formats. + assert_restricted_match_membership( + &format!("9058-text-{name}.json"), + expected_ids, + &observations, + ); +} + +#[rstest] +#[case::indexed_row("indexed-row", "id = 1", &[0, 2, 3, 4, 6, 7])] +#[case::residual_row("residual-row", "id = 6", &[0, 1, 2, 3, 4, 7])] +#[case::all_rows("all-rows", "true", &[])] +#[tokio::test] +async fn test_fts_9058_deleted_candidates_stay_excluded( + #[case] name: &str, + #[case] deletion: &str, + #[case] expected_ids: &[i32], + #[values(false, true)] stable_row_ids: bool, +) { + let fragments = restricted_match_fragments(&[&[0, 1, 2], &[3, 4], &[5, 6, 7]]); + let mut dataset = restricted_match_dataset(&fragments, 1, stable_row_ids).await; + dataset.delete(deletion).await.unwrap(); + let observations = restricted_match_maintenance_observations( + &mut dataset, + RestrictedMatchDomain::Filter("id >= 0"), + "alpha beta", + &[None], + ) + .await; + // Deletions retain their established corpus-statistics semantics. The new + // restriction path must preserve live candidates and never revive postings. + assert_restricted_match_membership( + &format!("9058-deleted-{name}-stable-{stable_row_ids}.json"), + expected_ids, + &observations, + ); +} + +#[rstest] +#[tokio::test] +async fn test_fts_9058_modest_indexed_corpus_with_tail( + #[values( + RestrictedMatchDomain::Prefilter, + RestrictedMatchDomain::IndexedFragment + )] + domain: RestrictedMatchDomain, +) { + let raw_rows: Vec<_> = (0..1100) + .map(|id| { + let text = match id { + 0 | 1000.. => "alpha", + 1 => "beta", + 2 => "beta gamma delta", + _ => "gamma delta", + }; + (id, text) + }) + .collect(); + let fragments = [&raw_rows[..1000], &raw_rows[1000..]] + .map(|rows| rows.iter().map(|(id, text)| (*id, Some(*text))).collect()); + let mut dataset = restricted_match_dataset(&fragments, 1, false).await; + let observations = restricted_match_maintenance_observations( + &mut dataset, + domain, + "alpha beta", + &[Some(1), None], + ) + .await; + let candidate_ids: Vec<_> = match domain { + RestrictedMatchDomain::Prefilter => vec![0, 1], + RestrictedMatchDomain::IndexedFragment => (0..1000).collect(), + _ => unreachable!("modest corpus uses the two reported restrictions"), + }; + assert_restricted_match_oracle( + &format!("9058-modest-{domain:?}.json"), + &raw_rows, + domain, + &candidate_ids, + &observations, + ); +} + +#[rstest] +#[tokio::test] +async fn test_fts_9058_reuses_optimized_physical_plan( + #[values(false, true)] stable_row_ids: bool, + #[values(false, true)] select_fragment: bool, +) { + let fragments = restricted_match_fragments(&[&[0, 1, 2], &[3, 4], &[5, 6, 7]]); + let dataset = restricted_match_dataset(&fragments, 1, stable_row_ids).await; + let query = MatchQuery::new("alpha beta".to_owned()) + .with_column(Some("text".to_owned())) + .with_fuzziness(Some(0)) + .with_document_granularity(DocumentGranularity::Row); + let mut scanner = dataset.scan(); + scanner + .project(&["id"]) + .unwrap() + .full_text_search(FullTextSearchQuery::new_query(query.into())) + .unwrap() + .limit(Some(1), None) + .unwrap(); + if select_fragment { + scanner.with_fragments(vec![dataset.get_fragments()[0].clone().into()]); + } else { + scanner.prefilter(true).filter("id < 2").unwrap(); + } + // Reusing the Scanner alone would create a fresh plan and miss scorer state + // retained by the physical executors. Reuse the same plan and context here. + let plan = tokio::time::timeout(Duration::from_secs(10), scanner.create_plan()) + .await + .expect("restricted Match planning timed out") + .unwrap(); + let explanation = DisplayableExecutionPlan::new(plan.as_ref()) + .indent(false) + .to_string(); + assert!(explanation.contains("MatchCorpus"), "{explanation}"); + assert!( + explanation + .lines() + .any(|line| line.trim_start().starts_with("MatchQuery:")), + "{explanation}" + ); + assert!(explanation.contains("FlatMatchQuery:"), "{explanation}"); + let context = get_session_context(&LanceExecutionOptions::default()).task_ctx(); + let (_, mut expected) = raw_match_oracle(&RESTRICTED_MATCH_CORPUS, "alpha beta"); + expected.retain(|(id, _)| *id < if select_fragment { 3 } else { 2 }); + let expected = expected[0]; + let assert_top_result = |batches: Vec| { + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 1); + for batch in batches { + let ids = batch["id"].as_primitive::(); + let scores = batch[SCORE_COL].as_primitive::(); + for (id, score) in ids.values().iter().zip(scores.values()) { + assert_eq!(*id, expected.0); + assert!((f64::from(*score) - expected.1).abs() < 1e-5); + } + } + }; + for _ in 0..2 { + let batches = tokio::time::timeout( + Duration::from_secs(10), + datafusion::physical_plan::collect(plan.clone(), context.clone()), + ) + .await + .expect("reused restricted Match physical plan timed out") + .unwrap(); + assert_top_result(batches); + } + let (left, right) = tokio::time::timeout( + Duration::from_secs(10), + futures::future::try_join( + datafusion::physical_plan::collect(plan.clone(), context.clone()), + datafusion::physical_plan::collect(plan.clone(), context.clone()), + ), + ) + .await + .expect("concurrent restricted Match physical plan executions timed out") + .unwrap(); + assert_top_result(left); + assert_top_result(right); +} + +#[derive(Clone, Copy, Debug)] +enum RestrictedMatchControl { + Unfiltered, + Postfilter, + Fuzzy, + Boolean, +} + +#[rstest] +#[tokio::test] +async fn test_fts_9058_other_query_paths_keep_candidate_filtering( + #[values( + RestrictedMatchControl::Unfiltered, + RestrictedMatchControl::Postfilter, + RestrictedMatchControl::Fuzzy, + RestrictedMatchControl::Boolean + )] + control: RestrictedMatchControl, +) { + let fragments = restricted_match_fragments(&[&[0, 1, 2], &[3, 4], &[5, 6, 7]]); + let dataset = restricted_match_dataset(&fragments, 1, false).await; + let match_query = |terms: &str| { + MatchQuery::new(terms.to_owned()) + .with_column(Some("text".to_owned())) + .with_document_granularity(DocumentGranularity::Row) + }; + let query: FtsQuery = match control { + RestrictedMatchControl::Fuzzy => match_query("alpha beta").with_fuzziness(Some(1)).into(), + RestrictedMatchControl::Boolean => BooleanQuery::new([ + (Occur::Should, match_query("alpha").into()), + (Occur::Should, match_query("beta").into()), + ]) + .into(), + _ => match_query("alpha beta").into(), + }; + let mut scanner = dataset.scan(); + scanner + .project(&["id"]) + .unwrap() + .full_text_search(FullTextSearchQuery::new_query(query)) + .unwrap() + .limit(Some(1), None) + .unwrap(); + match control { + RestrictedMatchControl::Unfiltered => {} + RestrictedMatchControl::Postfilter => { + scanner.prefilter(false).filter("id < 2").unwrap(); + } + RestrictedMatchControl::Fuzzy | RestrictedMatchControl::Boolean => { + scanner.prefilter(true).filter("id < 2").unwrap(); + } + } + let explanation = tokio::time::timeout(Duration::from_secs(10), scanner.explain_plan(false)) + .await + .unwrap() + .unwrap(); + let plan = tokio::time::timeout(Duration::from_secs(10), scanner.create_plan()) + .await + .unwrap() + .unwrap(); + let mut pending = vec![plan]; + let mut flat_child_counts = Vec::new(); + while let Some(node) = pending.pop() { + if node.name() == "FlatMatchQueryExec" { + flat_child_counts.push(node.children().len()); + } + pending.extend(node.children().into_iter().cloned()); + } + let batch = tokio::time::timeout(Duration::from_secs(10), scanner.try_into_batch()) + .await + .unwrap() + .unwrap(); + let ids = batch["id"].as_primitive::().values(); + let scores = batch[SCORE_COL].as_primitive::().values(); + let evidence = serde_json::json!({ + "control": format!("{control:?}"), + "plan": explanation, + "flat_child_counts": flat_child_counts, + "ids": ids.as_ref(), + "scores": scores.as_ref(), + }); + assert!( + !explanation.contains("MatchCorpus"), + "unrelated query acquired restricted corpus execution: {evidence:#}" + ); + // The original flat paths have one input and apply any candidate predicate + // before scoring. A separate emission-selection child belongs only to the + // corrected restricted exact root Match path. + assert!( + !flat_child_counts.is_empty(), + "missing flat path: {evidence:#}" + ); + assert!( + flat_child_counts.iter().all(|count| *count == 1), + "unrelated query acquired corpus emission filtering: {evidence:#}" + ); + // Postfiltering happens after global top-k. In this unchanged mixed path, + // residual id 6 wins that top-k and is then excluded by id < 2. + let expected_rows = usize::from(!matches!(control, RestrictedMatchControl::Postfilter)); + assert_eq!( + batch.num_rows(), + expected_rows, + "wrong candidates: {evidence:#}" + ); + assert!(scores.iter().all(|score| score.is_finite()), "{evidence:#}"); + if !matches!(control, RestrictedMatchControl::Unfiltered) { + assert!(ids.iter().all(|id| *id < 2), "filter leaked: {evidence:#}"); + } + if matches!(control, RestrictedMatchControl::Unfiltered) { + assert!(explanation.contains("MatchQuery:"), "{evidence:#}"); + assert!(!explanation.contains("CompoundFtsScorer"), "{evidence:#}"); + } +} + #[tokio::test] async fn test_fts_v1_remains_queryable_after_append_optimize() { let params = InvertedIndexParams::default().format_version(InvertedListFormatVersion::V1); @@ -4310,10 +5420,8 @@ async fn test_fts_without_index() { #[tokio::test] async fn test_fts_without_index_uses_scalar_index_for_prefilter() { - // Verify that flat FTS (no inverted index on text) routes its prefilter - // through `FilteredReadExec` so a scalar index on the filter column is - // actually used. Six rows with two distinct ids: a prefilter of `id = 1` - // must match exactly the three text rows tagged with id=1. + // Flat FTS must use the scalar index for candidates while keeping all six + // rows in the scoring corpus. The id=1 candidates contain two alpha matches. let text = StringArray::from(vec![ "alpha bravo", "charlie delta", @@ -4361,11 +5469,21 @@ async fn test_fts_without_index_uses_scalar_index_for_prefilter() { .unwrap(); let plan = scan.analyze_plan().await.unwrap(); - // The flat-FTS path now reads via `FilteredReadExec` (prints as `LanceRead`) - // with the prefilter plumbed into it, so the scalar index on `id` is used. + // The text read supplies the global corpus. The scalar index independently + // supplies the candidate mask, which is applied when matches are emitted. assert_contains!(&plan, "FlatMatchQuery"); - assert_contains!(&plan, "LanceRead"); - assert_contains!(&plan, "full_filter=id = Int32(1)"); + let corpus_read = plan + .lines() + .find(|line| line.contains("LanceRead:") && line.contains("projection=[text]")) + .unwrap_or_else(|| panic!("missing text corpus read: {plan}")); + assert_contains!(corpus_read, "full_filter=--, refine_filter=--"); + assert_contains!(corpus_read, "rows_scanned=6"); + let candidate_index = plan + .lines() + .find(|line| line.contains("ScalarIndexQuery:")) + .unwrap_or_else(|| panic!("missing scalar candidate index: {plan}")); + assert_contains!(candidate_index, "query=[id = 1]@id_idx(BTree)"); + assert_contains!(candidate_index, "indices_loaded=1"); // The legacy plan ran a `LanceScan` wrapped in a manual `LanceFilterExec`; // make sure we did not regress to that shape. assert_not_contains!(&plan, "LanceScan:"); @@ -4378,6 +5496,10 @@ async fn test_fts_without_index_uses_scalar_index_for_prefilter() { 2, "expected the two id=1 rows that match `alpha`, got plan:\n{plan}" ); + assert_eq!(results["id"].as_primitive::().values(), &[1, 1]); + let mut matched_text = results["text"].as_string::().iter().collect_vec(); + matched_text.sort_unstable(); + assert_eq!(matched_text, [Some("alpha bravo"), Some("alpha echo")]); } #[tokio::test] diff --git a/rust/lance/src/io/exec/fts.rs b/rust/lance/src/io/exec/fts.rs index 38f40e9dc44..1bbc46d21a2 100644 --- a/rust/lance/src/io/exec/fts.rs +++ b/rust/lance/src/io/exec/fts.rs @@ -3,18 +3,19 @@ use std::cmp::Ordering; use std::collections::{HashMap, HashSet}; -use std::sync::{Arc, OnceLock}; +use std::sync::{Arc, Mutex, OnceLock, Weak}; use arrow::array::{AsArray, BooleanBuilder, ListBuilder, UInt32Builder}; use arrow::datatypes::{Float32Type, UInt64Type}; use arrow_array::{Array, BooleanArray, Float32Array, OffsetSizeTrait, RecordBatch, UInt64Array}; use arrow_schema::{DataType, Field, SchemaRef}; +use datafusion::common::tree_node::{Transformed, TreeNode}; use datafusion::common::{NullEquality, Statistics}; use datafusion::error::{DataFusionError, Result as DataFusionResult}; use datafusion::execution::SendableRecordBatchStream; use datafusion::physical_plan::empty::EmptyExec; -use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType}; -use datafusion::physical_plan::metrics::{ExecutionPlanMetricsSet, Gauge, MetricsSet}; +use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType, reset_plan_states}; +use datafusion::physical_plan::metrics::{ExecutionPlanMetricsSet, Gauge, MetricValue, MetricsSet}; use datafusion::physical_plan::repartition::RepartitionExec; use datafusion::physical_plan::stream::RecordBatchStreamAdapter; use datafusion::physical_plan::union::UnionExec; @@ -31,6 +32,7 @@ use itertools::Itertools; use lance_core::{ Error, ROW_ID, Result, utils::{ + futures::FinallyStreamExt, tokio::{get_num_compute_intensive_cpus, spawn_cpu}, tracing::StreamTracingExt, }, @@ -41,8 +43,10 @@ use lance_table::format::IndexMetadata; use rustc_hash::FxHashSet; use super::PreFilterSource; +use super::row_addr_mask::apply_mask; use super::utils::{ - IndexMetrics, PreFilterMasks, build_prefilter, build_prefilter_restricted_to_fragments, + FilteredRowIdsToPrefilter, IndexMetrics, PreFilterMasks, SelectionVectorToPrefilter, + build_prefilter, build_prefilter_restricted_to_fragments, }; use crate::dataset::mem_wal::index::{QueryLocalFtsIndex, QueryLocalFtsStats}; use crate::index::scalar::inverted::{ @@ -68,7 +72,10 @@ use lance_index::scalar::inverted::{ flat_bm25_search_stream_with_options_and_scorer, fts_schema, materialized_compound_top_k, prepare_bm25_query, validate_combined_tokenizers, }; -use lance_index::{prefilter::PreFilter, scalar::inverted::query::BooleanQuery}; +use lance_index::{ + prefilter::{FilterLoader, PreFilter}, + scalar::inverted::query::BooleanQuery, +}; use lance_tokenizer::{SimpleTokenizer, TextAnalyzer}; use tracing::instrument; use uuid::Uuid; @@ -2394,6 +2401,281 @@ impl Drop for SharedFtsScorerProducer { } } +/// Owns one restricted Match plan's scorer lifecycle. Each execution gets a +/// fresh producer/consumer pair, including after failure or child replacement. +/// Re-execution still requires replayable children: an overlay-stale scalar +/// prefilter can contain a consumable `OneShotExec` that this owner cannot replay. +#[derive(Debug)] +pub(crate) struct SharedFtsScorerExec { + input: Arc, + scorer: Arc, + /// Retain metric handles, never a finished execution's task-owning plan. + /// Concurrent executions have independent scorers. Rebound metrics describe + /// the latest execution; template-retained counters keep their existing semantics. + metrics: Arc>, +} + +#[derive(Debug)] +struct SharedFtsExecutionMetrics { + execution: Arc<()>, + active: Option>, + snapshot: MetricsSet, +} + +struct SharedFtsMetricsGuard { + input: Arc, + template: Arc, + metrics: Arc>, + execution: Arc<()>, +} + +impl Drop for SharedFtsMetricsGuard { + fn drop(&mut self) { + let snapshot = + SharedFtsScorerExec::execution_metrics(self.template.as_ref(), self.input.as_ref()); + match self.metrics.lock() { + Ok(mut metrics) if Arc::ptr_eq(&metrics.execution, &self.execution) => { + metrics.snapshot = snapshot; + metrics.active = None; + } + Ok(_) => {} + Err(error) => log::warn!("could not retain restricted Match runtime metrics: {error}"), + } + } +} + +impl SharedFtsScorerExec { + pub(crate) fn new(input: Arc, scorer: Arc) -> Self { + Self { + metrics: Arc::new(Mutex::new(SharedFtsExecutionMetrics { + execution: Arc::new(()), + active: None, + snapshot: MetricsSet::new(), + })), + input, + scorer, + } + } + + fn prepare_execution(&self) -> DataFusionResult> { + let scorer = Arc::new(SharedFtsScorer::new()); + let mut consumers = 0; + let mut producers = 0; + // Reset other execution state too, particularly a previous top-k sort's + // score threshold. This owner only encloses the restricted root Match. + let input = reset_plan_states(self.input.clone())?; + let input = input + .transform_up(|plan| { + if let Some(node) = plan.downcast_ref::() + && node + .shared_scorer + .as_ref() + .is_some_and(|shared| Arc::ptr_eq(shared, &self.scorer)) + { + consumers += 1; + return Ok(Transformed::yes(Arc::new(MatchQueryExec { + dataset: node.dataset.clone(), + query: node.query.clone(), + tokenized_query: node.tokenized_query.clone(), + params: node.params.clone(), + prefilter_source: node.prefilter_source.clone(), + base_scorer: node.base_scorer.clone(), + prepared_query: node.prepared_query.clone(), + shared_scorer: Some(scorer.clone()), + segment_selection: node.segment_selection.clone(), + overlay_block: node.overlay_block.clone(), + document_granularity: node.document_granularity, + schema: node.schema.clone(), + external_mask: node.external_mask.clone(), + properties: node.properties.clone(), + metrics: node.metrics.clone(), + }) + as Arc)); + } + if let Some(node) = plan.downcast_ref::() + && node + .shared_scorer + .as_ref() + .is_some_and(|shared| Arc::ptr_eq(shared, &self.scorer)) + { + producers += 1; + return Ok(Transformed::yes(Arc::new(FlatMatchQueryExec { + dataset: node.dataset.clone(), + query: node.query.clone(), + tokenized_query: node.tokenized_query.clone(), + params: node.params.clone(), + unindexed_input: node.unindexed_input.clone(), + candidate_filter: node.candidate_filter.clone(), + base_scorer: node.base_scorer.clone(), + shared_scorer: Some(scorer.clone()), + preset_segments: node.preset_segments.clone(), + document_granularity: node.document_granularity, + document_column: node.document_column.clone(), + schema: node.schema.clone(), + properties: node.properties.clone(), + metrics: node.metrics.clone(), + }) + as Arc)); + } + Ok(Transformed::no(plan)) + })? + .data; + if consumers != 1 || producers != 1 { + return Err(DataFusionError::Internal(format!( + "restricted Match corpus requires one scorer consumer and producer, got {consumers} and {producers}" + ))); + } + Ok(input) + } + + fn execution_metrics(template: &dyn ExecutionPlan, runtime: &dyn ExecutionPlan) -> MetricsSet { + fn collect(plan: &dyn ExecutionPlan, metrics: &mut MetricsSet, only_named: bool) { + if let Some(node_metrics) = plan.metrics() { + for metric in node_metrics.iter() { + if !only_named + || matches!( + metric.value(), + MetricValue::Count { .. } + | MetricValue::Gauge { .. } + | MetricValue::Time { .. } + | MetricValue::PruningMetrics { .. } + | MetricValue::Ratio { .. } + | MetricValue::Custom { .. } + ) + { + metrics.push(metric.clone()); + } + } + } + for child in plan.children() { + collect(child.as_ref(), metrics, only_named); + } + } + let mut template_metrics = MetricsSet::new(); + collect(template, &mut template_metrics, false); + let mut seen = template_metrics + .iter() + .map(|metric| Arc::as_ptr(metric) as usize) + .collect::>(); + // The wrapper has its input's output and compute metrics. Descendant + // baseline metrics describe different operators, so only their named + // I/O, scorer and other diagnostic metrics belong in this aggregation. + let mut runtime_metrics = runtime.metrics().unwrap_or_default(); + for child in runtime.children() { + collect(child.as_ref(), &mut runtime_metrics, true); + } + let mut metrics = MetricsSet::new(); + for metric in runtime_metrics.iter() { + // Some leaf resets intentionally retain metric handles. The visible + // template already exposes these, so do not count them twice. + if seen.insert(Arc::as_ptr(metric) as usize) { + metrics.push(metric.clone()); + } + } + metrics + } +} + +impl DisplayAs for SharedFtsScorerExec { + fn fmt_as(&self, _format: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result { + write!(f, "MatchCorpus") + } +} + +impl ExecutionPlan for SharedFtsScorerExec { + fn name(&self) -> &str { + "SharedFtsScorerExec" + } + + fn properties(&self) -> &Arc { + self.input.properties() + } + + fn children(&self) -> Vec<&Arc> { + vec![&self.input] + } + + fn required_input_distribution(&self) -> Vec { + vec![Distribution::SinglePartition] + } + + fn maintains_input_order(&self) -> Vec { + vec![true] + } + + fn with_new_children( + self: Arc, + mut children: Vec>, + ) -> DataFusionResult> { + if children.len() != 1 { + return Err(DataFusionError::Internal(format!( + "restricted Match corpus requires one input, got {}", + children.len() + ))); + } + let input = children.pop().ok_or_else(|| { + DataFusionError::Internal("restricted Match corpus input is missing".to_string()) + })?; + Ok(Arc::new(Self::new(input, self.scorer.clone()))) + } + + fn execute( + &self, + partition: usize, + context: Arc, + ) -> DataFusionResult { + if partition != 0 { + return Err(DataFusionError::Internal(format!( + "restricted Match corpus requires partition 0, got {partition}" + ))); + } + let input = self.prepare_execution()?; + if input.output_partitioning().partition_count() != 1 { + return Err(DataFusionError::Internal( + "restricted Match corpus input must have one output partition".to_string(), + )); + } + let execution = Arc::new(()); + let result = input.execute(partition, context); + let snapshot = Self::execution_metrics(self.input.as_ref(), input.as_ref()); + *self.metrics.lock().map_err(|_| { + DataFusionError::Internal("restricted Match runtime lock was poisoned".to_string()) + })? = SharedFtsExecutionMetrics { + execution: execution.clone(), + active: Some(Arc::downgrade(&input)), + snapshot, + }; + let stream = result?; + let schema = stream.schema(); + let guard = SharedFtsMetricsGuard { + input, + template: self.input.clone(), + metrics: self.metrics.clone(), + execution, + }; + // The captured guard also runs when this closure is dropped before EOF, + // preserving lazily registered metrics on error or cancellation. + let stream = stream.finally(move || drop(guard)); + Ok(Box::pin(RecordBatchStreamAdapter::new(schema, stream))) + } + + fn metrics(&self) -> Option { + let (active, snapshot) = { + let metrics = self.metrics.lock().ok()?; + ( + metrics.active.as_ref().and_then(Weak::upgrade), + metrics.snapshot.clone(), + ) + }; + // Expose active metrics before EOF without retaining the execution + // between inspections. + Some(match active { + Some(input) => Self::execution_metrics(self.input.as_ref(), input.as_ref()), + None => snapshot, + }) + } +} + /// Time spent resolving an exact ordered UUID selection to committed FTS segments. pub const FTS_SEGMENT_BIND_DURATION_METRIC: &str = "fts_segment_bind_duration"; @@ -4000,6 +4282,8 @@ pub struct FlatMatchQueryExec { tokenized_query: Arc>, params: FtsSearchParams, unindexed_input: Arc, + /// Selects emitted rows after the full input has contributed corpus statistics. + candidate_filter: PreFilterSource, /// Optional override for the BM25 scorer normally built locally inside /// `execute()`. See [`MatchQueryExec::with_base_scorer`]. base_scorer: Option>, @@ -4083,6 +4367,7 @@ impl FlatMatchQueryExec { tokenized_query: Arc::new(OnceLock::new()), params, unindexed_input, + candidate_filter: PreFilterSource::None, base_scorer: None, shared_scorer: None, preset_segments: None, @@ -4142,6 +4427,7 @@ impl FlatMatchQueryExec { base_scorer: None, shared_scorer: None, preset_segments: Some(segments), + candidate_filter: PreFilterSource::None, document_granularity, document_column, schema, @@ -4161,6 +4447,13 @@ impl FlatMatchQueryExec { self } + /// Keep corpus collection independent of the candidate domain, including + /// when no input rows are eligible to appear in the search results. + pub(crate) fn with_candidate_filter(mut self, candidate_filter: PreFilterSource) -> Self { + self.candidate_filter = candidate_filter; + self + } + pub fn query(&self) -> &MatchQuery { &self.query } @@ -4188,7 +4481,11 @@ impl ExecutionPlan for FlatMatchQueryExec { } fn children(&self) -> Vec<&Arc> { - vec![&self.unindexed_input] + let mut children = vec![&self.unindexed_input]; + if let Some(candidate_filter) = self.candidate_filter.execution_plan() { + children.push(candidate_filter); + } + children } fn required_input_distribution(&self) -> Vec { @@ -4196,25 +4493,35 @@ impl ExecutionPlan for FlatMatchQueryExec { // output partition, so the input must be coalesced to one partition. Without // this, EnforceDistribution may round-robin the scan across `target_partitions` // and only partition 0 is consumed, silently dropping the other fragments. - vec![Distribution::SinglePartition] + vec![Distribution::SinglePartition; self.children().len()] } fn with_new_children( self: Arc, - mut children: Vec>, + children: Vec>, ) -> DataFusionResult> { - if children.len() != 1 { - return Err(DataFusionError::Internal( - "Unexpected number of children".to_string(), - )); + let expected_children = self.children().len(); + if children.len() != expected_children { + return Err(DataFusionError::Internal(format!( + "flat Match expected {expected_children} children, got {}", + children.len() + ))); } - let unindexed_input = children.pop().unwrap(); + let mut children = children.into_iter(); + let unindexed_input = children.next().ok_or_else(|| { + DataFusionError::Internal("flat Match corpus input is missing".to_string()) + })?; + let candidate_filter = match children.next() { + Some(child) => self.candidate_filter.with_execution_plan(child)?, + None => PreFilterSource::None, + }; Ok(Arc::new(Self { dataset: self.dataset.clone(), query: self.query.clone(), tokenized_query: self.tokenized_query.clone(), params: self.params.clone(), unindexed_input, + candidate_filter, base_scorer: self.base_scorer.clone(), shared_scorer: self.shared_scorer.clone(), preset_segments: self.preset_segments.clone(), @@ -4235,6 +4542,7 @@ impl ExecutionPlan for FlatMatchQueryExec { let query = self.query.clone(); let tokenized_query = self.tokenized_query.clone(); let ds = self.dataset.clone(); + let candidate_filter = self.candidate_filter.clone(); let preset_base_scorer = self.base_scorer.clone(); let shared_scorer_producer = self.shared_scorer.clone().map(SharedFtsScorerProducer::new); let preset_segments = self.preset_segments.clone(); @@ -4255,10 +4563,19 @@ impl ExecutionPlan for FlatMatchQueryExec { "column not set for MatchQuery {}", query.terms )))?; - let unindexed_input = document_input( - self.unindexed_input.execute(partition, context)?, - &document_column, - )?; + let unindexed_input = self + .unindexed_input + .execute(partition, context.clone()) + .and_then(|stream| document_input(stream, &document_column).map_err(Into::into)); + let unindexed_input = match unindexed_input { + Ok(input) => input, + Err(error) => { + if let Some(producer) = shared_scorer_producer { + producer.publish_error(&error); + } + return Err(error); + } + }; let stream = stream::once(async move { let shared_scorer_producer = shared_scorer_producer; @@ -4332,7 +4649,32 @@ impl ExecutionPlan for FlatMatchQueryExec { if let Some(producer) = shared_scorer_producer { producer.publish(Arc::new(scorer)); } - Ok(stream) + // Candidate selection is deliberately inside this node: even + // an empty selection must not remove the corpus producer. + let mask = match candidate_filter { + PreFilterSource::FilteredRowIds(input) => { + Box::new(FilteredRowIdsToPrefilter::new( + input.execute(partition, context)?, + )) + .load() + .await? + } + PreFilterSource::ScalarIndexQuery(input) => { + Box::new(SelectionVectorToPrefilter( + input.execute(partition, context)?, + )) + .load() + .await? + } + PreFilterSource::None => return Ok(stream), + }; + let schema = stream.schema(); + let stream = stream.map(move |batch| { + let _timer = metrics.baseline_metrics.elapsed_compute().timer(); + apply_mask(&mask, batch?) + }); + Ok(Box::pin(RecordBatchStreamAdapter::new(schema, stream)) + as SendableRecordBatchStream) } Err(error) => { if let Some(producer) = shared_scorer_producer { @@ -5351,21 +5693,32 @@ impl ExecutionPlan for BooleanQueryExec { #[cfg(test)] mod tests { + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; + use std::time::Duration; use crate::index::DatasetIndexExt; use arrow_array::{ ArrayRef, Float32Array, Int32Array, RecordBatch, RecordBatchIterator, StringArray, - UInt64Array, + UInt64Array, record_batch, }; use arrow_schema::DataType; + use datafusion::datasource::memory::MemorySourceConfig; use datafusion::error::{DataFusionError, Result as DataFusionResult}; - use datafusion::physical_plan::metrics::ExecutionPlanMetricsSet; + use datafusion::physical_plan::coalesce_partitions::CoalescePartitionsExec; + use datafusion::physical_plan::metrics::{ExecutionPlanMetricsSet, MetricBuilder, MetricsSet}; + use datafusion::physical_plan::stream::RecordBatchStreamAdapter; + use datafusion::physical_plan::{ + DisplayAs, DisplayFormatType, PlanProperties, SendableRecordBatchStream, + }; use datafusion::{execution::TaskContext, physical_plan::ExecutionPlan}; - use futures::TryStreamExt; + use futures::{FutureExt, TryStreamExt, stream}; + use lance_core::utils::futures::StreamOnDropExt; use lance_core::{ROW_ID, utils::address::RowAddress}; use lance_datafusion::datagen::DatafusionDatagenExt; - use lance_datafusion::exec::{ExecutionStatsCallback, ExecutionSummaryCounts}; + use lance_datafusion::exec::{ + ExecutionStatsCallback, ExecutionSummaryCounts, collect_execution_metrics, + }; use lance_datafusion::utils::{INDEX_CACHE_HITS_METRIC, PARTITIONS_SEARCHED_METRIC}; use lance_datagen::{BatchCount, ByteCount, RowCount}; use lance_index::metrics::NoOpMetricsCollector; @@ -5381,7 +5734,12 @@ mod tests { }; use lance_index::scalar::{FullTextSearchQuery, InvertedIndexParams}; use lance_index::{IndexCriteria, IndexType}; + use lance_select::{ + RowAddrMask, RowAddrTreeMap, + result::{IndexExprResult, IndexExprResultWireFormat}, + }; use lance_table::format::IndexMetadata; + use rstest::rstest; use uuid::Uuid; use crate::{ @@ -5396,8 +5754,8 @@ mod tests { use super::{ BoolSlot, BoostQueryExec, CombinedFieldsQueryExec, CompoundQueryExec, CrossColumnCompoundQueryExec, FTS_SEGMENT_BIND_DURATION_METRIC, FlatMatchFilterExec, - FlatMatchQueryExec, MatchQueryExec, PhraseQueryExec, WAND_TIE_COMPLETION_BUDGET, - WandExactnessCertificate, build_boolean_query_children, + FlatMatchQueryExec, MatchQueryExec, PhraseQueryExec, SharedFtsScorer, SharedFtsScorerExec, + WAND_TIE_COMPLETION_BUDGET, WandExactnessCertificate, build_boolean_query_children, classify_wand_exactness_certificate, default_text_tokenizer, open_fts_segments, tokenizer_for_match_query, }; @@ -5749,6 +6107,462 @@ mod tests { ); } + fn memory_input(batch: RecordBatch) -> Arc { + MemorySourceConfig::try_new_exec(&[vec![batch.clone()]], batch.schema(), None).unwrap() + } + + fn candidate_input(row_ids: Vec, is_scalar_index_query: bool) -> PreFilterSource { + if is_scalar_index_query { + let mask = RowAddrMask::from_allowed(RowAddrTreeMap::from_iter(row_ids)); + let batch = IndexExprResult::exact(mask) + .serialize( + &roaring::RoaringBitmap::from_iter([0]), + IndexExprResultWireFormat::TwoMask, + ) + .unwrap(); + PreFilterSource::ScalarIndexQuery(memory_input(batch)) + } else { + PreFilterSource::FilteredRowIds(memory_input( + record_batch!((ROW_ID, UInt64, row_ids)).unwrap(), + )) + } + } + + async fn indexed_match_dataset() -> Arc { + let mut dataset = lance_datagen::gen_batch() + .col( + "text", + lance_datagen::array::cycle_utf8_literals(&["alpha", "beta"]), + ) + .into_ram_dataset(FragmentCount::from(1), FragmentRowCount::from(2)) + .await + .unwrap(); + dataset + .create_index( + &["text"], + IndexType::Inverted, + None, + &InvertedIndexParams::default() + .stem(false) + .remove_stop_words(false), + true, + ) + .await + .unwrap(); + Arc::new(dataset) + } + + #[rstest] + #[case::row_ids(false)] + #[case::selection_vector(true)] + #[tokio::test] + async fn test_9058_flat_candidates_preserve_corpus(#[case] is_scalar_index_query: bool) { + let dataset = indexed_match_dataset().await; + let query = MatchQuery::new("alpha beta".to_string()) + .with_column(Some("text".to_string())) + .with_document_granularity(DocumentGranularity::Row); + let corpus = memory_input( + record_batch!( + ("text", Utf8, ["alpha", "beta", "alpha"]), + (ROW_ID, UInt64, [100, 101, 102]) + ) + .unwrap(), + ); + let unfiltered = FlatMatchQueryExec::new( + dataset.clone(), + query.clone(), + FtsSearchParams::default(), + corpus.clone(), + ) + .unwrap(); + let expected = execute_results(&unfiltered).await.unwrap(); + let shared = Arc::new(SharedFtsScorer::new()); + let filtered = Arc::new( + FlatMatchQueryExec::new(dataset, query, FtsSearchParams::default(), corpus.clone()) + .unwrap() + .with_shared_scorer(shared.clone()) + .with_candidate_filter(candidate_input(vec![101], is_scalar_index_query)), + ); + + assert_eq!(filtered.children().len(), 2); + assert!( + filtered + .required_input_distribution() + .iter() + .all(|distribution| { + matches!( + distribution, + datafusion_physical_expr::Distribution::SinglePartition + ) + }) + ); + assert_eq!( + execute_results(filtered.as_ref()).await.unwrap(), + vec![expected[1]] + ); + assert_eq!(filtered.metrics().unwrap().output_rows(), Some(1)); + let scorer = shared.wait().await.unwrap(); + assert_eq!(scorer.num_docs(), 5); + assert_eq!(scorer.num_docs_containing_token("alpha"), 3); + assert_eq!(scorer.num_docs_containing_token("beta"), 2); + + // Replacing the candidate child must preserve its source kind. Empty + // candidates still require the producer to collect the complete corpus. + let empty = candidate_input(vec![], is_scalar_index_query); + let expanded_corpus = memory_input( + record_batch!( + ("text", Utf8, ["alpha", "beta", "alpha", "beta"]), + (ROW_ID, UInt64, [100, 101, 102, 103]) + ) + .unwrap(), + ); + let rebuilt = filtered + .clone() + .with_new_children(vec![ + expanded_corpus, + empty.execution_plan().unwrap().clone(), + ]) + .unwrap(); + assert!(execute_results(rebuilt.as_ref()).await.unwrap().is_empty()); + assert_eq!(rebuilt.metrics().unwrap().output_rows(), Some(0)); + assert_eq!(shared.wait().await.unwrap().num_docs(), 6); + let error = filtered.with_new_children(vec![]).unwrap_err(); + assert!(matches!(error, DataFusionError::Internal(_))); + assert!(error.to_string().contains("expected 2 children")); + } + + #[test] + fn test_9058_corpus_owner_rejects_missing_producer() { + let input = Arc::new(EmptyExec::new(FTS_SCHEMA.clone())); + let owner = Arc::new(SharedFtsScorerExec::new( + input, + Arc::new(SharedFtsScorer::new()), + )); + let Err(error) = owner.execute(0, Arc::new(TaskContext::default())) else { + panic!("a scorer owner without a producer must fail before execution"); + }; + assert!(matches!(error, DataFusionError::Internal(_))); + assert!(error.to_string().contains("got 0 and 0")); + let error = owner.with_new_children(vec![]).unwrap_err(); + assert!(matches!(error, DataFusionError::Internal(_))); + assert!(error.to_string().contains("requires one input")); + } + + #[derive(Debug, Clone, Copy)] + enum ResidualFailure { + Execute, + Stream, + Cancel, + } + + /// A replayable source whose first execution fails or remains pending. Its + /// retained metric handles also exercise the owner's metric de-duplication. + #[derive(Debug, Clone)] + struct FailingResidualExec { + input: Arc, + failure: ResidualFailure, + attempts: Arc, + started: Arc, + dropped: Arc, + metrics: ExecutionPlanMetricsSet, + } + + impl FailingResidualExec { + fn new(input: Arc, failure: ResidualFailure) -> Self { + Self { + input, + failure, + attempts: Arc::new(AtomicUsize::new(0)), + started: Arc::new(tokio::sync::Notify::new()), + dropped: Arc::new(AtomicBool::new(false)), + metrics: ExecutionPlanMetricsSet::new(), + } + } + } + + impl DisplayAs for FailingResidualExec { + fn fmt_as(&self, _: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result { + write!(f, "FailingResidual") + } + } + + impl ExecutionPlan for FailingResidualExec { + fn name(&self) -> &str { + "FailingResidualExec" + } + fn properties(&self) -> &Arc { + self.input.properties() + } + fn children(&self) -> Vec<&Arc> { + vec![&self.input] + } + fn with_new_children( + self: Arc, + children: Vec>, + ) -> DataFusionResult> { + assert_eq!(children.len(), 1); + Ok(Arc::new(Self { + input: children[0].clone(), + ..self.as_ref().clone() + })) + } + fn metrics(&self) -> Option { + Some(self.metrics.clone_inner()) + } + fn execute( + &self, + partition: usize, + context: Arc, + ) -> DataFusionResult { + MetricBuilder::new(&self.metrics) + .counter("residual_executions", partition) + .add(1); + if self.attempts.fetch_add(1, Ordering::SeqCst) != 0 { + return self.input.execute(partition, context); + } + match self.failure { + ResidualFailure::Execute => Err(DataFusionError::Execution( + "injected residual execute failure".to_string(), + )), + ResidualFailure::Stream => Ok(Box::pin(RecordBatchStreamAdapter::new( + self.schema(), + stream::once(async { + Err(DataFusionError::Execution( + "injected residual stream failure".to_string(), + )) + }), + ))), + ResidualFailure::Cancel => { + let started = self.started.clone(); + let dropped = self.dropped.clone(); + let stream = stream::once(async move { + started.notify_one(); + futures::future::pending::>().await + }) + .on_drop(move || dropped.store(true, Ordering::SeqCst)); + Ok(Box::pin(RecordBatchStreamAdapter::new( + self.schema(), + stream, + ))) + } + } + } + } + + #[rstest] + #[case::execute(ResidualFailure::Execute, "injected residual execute failure")] + #[case::stream(ResidualFailure::Stream, "injected residual stream failure")] + #[tokio::test] + async fn test_9058_flat_publishes_input_error( + #[case] failure: ResidualFailure, + #[case] message: &str, + ) { + let dataset = indexed_match_dataset().await; + let input = Arc::new(FailingResidualExec::new( + memory_input( + record_batch!(("text", Utf8, ["alpha"]), (ROW_ID, UInt64, [100])).unwrap(), + ), + failure, + )); + let shared = Arc::new(SharedFtsScorer::new()); + let flat = FlatMatchQueryExec::new( + dataset, + MatchQuery::new("alpha".to_string()) + .with_column(Some("text".to_string())) + .with_document_granularity(DocumentGranularity::Row), + FtsSearchParams::default(), + input, + ) + .unwrap() + .with_shared_scorer(shared.clone()); + assert_execution_error(execute_results(&flat).await.unwrap_err(), message); + let error = tokio::time::timeout(Duration::from_secs(1), shared.wait()) + .await + .unwrap() + .unwrap_err(); + assert_execution_error(error, message); + } + + #[rstest] + #[case::execute(ResidualFailure::Execute)] + #[case::stream(ResidualFailure::Stream)] + #[case::cancel(ResidualFailure::Cancel)] + #[tokio::test] + async fn test_9058_shared_corpus_retries_after_failure(#[case] failure: ResidualFailure) { + let dataset = indexed_match_dataset().await; + let query = MatchQuery::new("alpha beta".to_string()) + .with_column(Some("text".to_string())) + .with_document_granularity(DocumentGranularity::Row); + let residual = Arc::new(FailingResidualExec::new( + memory_input( + record_batch!( + ("text", Utf8, vec!["alpha"; 10]), + (ROW_ID, UInt64, (100..110).collect::>()) + ) + .unwrap(), + ), + failure, + )); + let shared = Arc::new(SharedFtsScorer::new()); + let indexed = Arc::new( + MatchQueryExec::new( + dataset.clone(), + query.clone(), + FtsSearchParams::default().with_limit(Some(1)), + PreFilterSource::None, + ) + .unwrap() + .with_shared_scorer(shared.clone()), + ); + let flat = Arc::new( + FlatMatchQueryExec::new(dataset, query, FtsSearchParams::default(), residual.clone()) + .unwrap() + .with_shared_scorer(shared.clone()) + .with_candidate_filter(candidate_input(vec![], false)), + ); + let union = UnionExec::try_new(vec![indexed, flat]).unwrap(); + let input: Arc = Arc::new(CoalescePartitionsExec::new(union)); + let owner = Arc::new(SharedFtsScorerExec::new(input.clone(), shared)); + let owner = owner.with_new_children(vec![input]).unwrap(); + assert_eq!(owner.required_input_distribution().len(), 1); + assert!(matches!( + owner.required_input_distribution()[0], + datafusion_physical_expr::Distribution::SinglePartition + )); + assert_eq!(owner.schema(), FTS_SCHEMA.clone()); + let context = Arc::new(TaskContext::default()); + let stream = owner.execute(0, context.clone()).unwrap(); + if matches!(failure, ResidualFailure::Cancel) { + tokio::time::timeout(Duration::from_secs(1), residual.started.notified()) + .await + .unwrap(); + drop(stream); + tokio::time::timeout(Duration::from_secs(1), async { + while !residual.dropped.load(Ordering::SeqCst) { + tokio::task::yield_now().await; + } + }) + .await + .expect( + "dropping the output must cancel the residual stream while the owner stays alive", + ); + } else { + let error = + tokio::time::timeout(Duration::from_secs(1), stream.try_collect::>()) + .await + .unwrap() + .unwrap_err(); + assert!(error.to_string().contains("injected residual"), "{error}"); + } + let mut counts = ExecutionSummaryCounts::default(); + collect_execution_metrics(owner.as_ref(), &mut counts); + assert_eq!(counts.all_counts["residual_executions"], 1); + assert!( + owner + .metrics() + .unwrap() + .iter() + .any(|metric| metric.value().name() == "scorer_build_ms") + ); + + // Poll the next execution's real indexed consumer through index opening + // before starting its producer. It must wait, never see the old error. + let runtime = owner + .downcast_ref::() + .unwrap() + .prepare_execution() + .unwrap(); + let union = runtime.children()[0].clone(); + let indexed = union.children()[0].clone(); + let flat = union.children()[1].clone(); + let mut indexed_stream = indexed.execute(0, context.clone()).unwrap(); + tokio::time::timeout(Duration::from_secs(1), async { + loop { + tokio::select! { + result = indexed_stream.try_next() => panic!("consumer completed before its producer: {result:?}"), + _ = tokio::task::yield_now() => {} + } + if metric_value(indexed.as_ref(), PARTITIONS_SEARCHED_METRIC) > 0 { break; } + } + }).await.unwrap(); + assert!(indexed_stream.try_next().now_or_never().is_none()); + let flat_batches: Vec<_> = flat + .execute(0, context.clone()) + .unwrap() + .try_collect() + .await + .unwrap(); + assert_eq!( + flat_batches + .iter() + .map(RecordBatch::num_rows) + .sum::(), + 0 + ); + let result = tokio::time::timeout(Duration::from_secs(1), indexed_stream.try_next()) + .await + .unwrap() + .unwrap() + .unwrap(); + assert_eq!( + result[ROW_ID] + .as_any() + .downcast_ref::() + .unwrap() + .value(0), + 1 + ); + + let mut stream = owner.execute(0, context).unwrap(); + let result = tokio::time::timeout(Duration::from_secs(1), async { + loop { + let batch = stream.try_next().await.unwrap().unwrap(); + if batch.num_rows() > 0 { + break batch; + } + } + }) + .await + .unwrap(); + assert_eq!(result.num_rows(), 1); + assert_eq!( + result[ROW_ID] + .as_any() + .downcast_ref::() + .unwrap() + .value(0), + 1 + ); + assert!( + (result[SCORE_COL] + .as_any() + .downcast_ref::() + .unwrap() + .value(0) + - 2.1594842) + .abs() + < 1e-5 + ); + let mut counts = ExecutionSummaryCounts::default(); + // Inspect metrics before EOF or stream drop. Lazy registrations must + // already be visible while the owner is still executing. + collect_execution_metrics(owner.as_ref(), &mut counts); + assert_eq!( + owner.metrics().unwrap().output_rows(), + Some(result.num_rows()) + ); + assert_eq!(counts.all_counts[PARTITIONS_SEARCHED_METRIC], 2); + assert_eq!( + counts.all_counts["residual_executions"], 3, + "shared source metrics must not be counted twice" + ); + assert!(owner.metrics().unwrap().elapsed_compute().unwrap() > 0); + drop(stream); + let mut final_counts = ExecutionSummaryCounts::default(); + collect_execution_metrics(owner.as_ref(), &mut final_counts); + assert_eq!(final_counts.all_counts, counts.all_counts); + } + #[test] fn execute_without_context() { // These tests ensure we can create nodes and call execute without a tokio Runtime diff --git a/rust/lance/src/io/exec/row_addr_mask.rs b/rust/lance/src/io/exec/row_addr_mask.rs index eb7059098bc..4b36398485b 100644 --- a/rust/lance/src/io/exec/row_addr_mask.rs +++ b/rust/lance/src/io/exec/row_addr_mask.rs @@ -136,7 +136,7 @@ impl ExecutionPlan for RowAddrMaskFilterExec { /// Keep rows whose `_rowid` is selected by the mask (the mask is keyed in the /// same `_rowid` space). Null ids are dropped; they cannot be in any allow set. -fn apply_mask(mask: &RowAddrMask, batch: RecordBatch) -> DataFusionResult { +pub(super) fn apply_mask(mask: &RowAddrMask, batch: RecordBatch) -> DataFusionResult { let row_id_column = batch.column_by_name(ROW_ID).ok_or_else(|| { DataFusionError::Internal(format!( "RowAddrMaskFilterExec input missing {ROW_ID} column"