diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index 281f34b23da..d42e3b3e713 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -2200,12 +2200,12 @@ impl Scanner { /// Default value is false. /// /// This is essentially a weak consistency search, only on the indexed data. + /// Row IDs are used internally but are only returned when explicitly requested. pub fn fast_search(&mut self) -> &mut Self { if let Some(q) = self.nearest.as_mut() { q.use_index = true; } self.fast_search = true; - self.projection_plan.include_row_id(); // fast search requires _rowid self } @@ -8521,6 +8521,7 @@ mod test { use lance_index::scalar::inverted::query::{ BooleanQuery, BoostQuery, FtsQuery, MatchQuery, MultiMatchQuery, Occur, PhraseQuery, }; + use lance_index::scalar::inverted::tokenizer::InvertedIndexParams; use lance_index::vector::hnsw::builder::HnswBuildParams; use lance_index::vector::ivf::IvfBuildParams; use lance_index::vector::pq::PQBuildParams; @@ -17330,6 +17331,166 @@ full_filter=name LIKE Utf8(\"test%2\"), refine_filter=name LIKE Utf8(\"test%2\") .unwrap(); } + #[rstest] + #[case::vector(IndexType::Vector, DIST_COL, vec![2, 1])] + #[case::fts(IndexType::Inverted, SCORE_COL, vec![0, 2])] + #[tokio::test] + async fn test_fast_search_projection( + #[case] index_type: IndexType, + #[case] scoring_column: &str, + #[case] expected_ids: Vec, + #[values(false, true)] has_projection: bool, + #[values(false, true)] is_row_id_requested: bool, + #[values(false, true)] has_stable_row_ids: bool, + ) { + let mut dataset = gen_batch() + .col("id", array::step::()) + .col("text", array::cycle_utf8_literals(&["quick", "slow"])) + .col( + "vec", + array::cycle_vec(array::step::(), Dimension::from(2)), + ) + .into_ram_dataset_with_params( + FragmentCount::from(2), + FragmentRowCount::from(2), + Some(WriteParams { + max_rows_per_file: 2, + enable_stable_row_ids: has_stable_row_ids, + ..Default::default() + }), + ) + .await + .unwrap(); + + if index_type == IndexType::Vector { + let centroids = Arc::new( + FixedSizeListArray::try_new_from_values(Float32Array::from(vec![0.0, 1.0]), 2) + .unwrap(), + ); + let params = VectorIndexParams::with_ivf_flat_params( + MetricType::L2, + IvfBuildParams::try_with_centroids(1, centroids).unwrap(), + ); + dataset + .create_index(&["vec"], index_type, None, ¶ms, true) + .await + .unwrap(); + } else { + dataset + .create_index( + &["text"], + index_type, + None, + &InvertedIndexParams::default(), + true, + ) + .await + .unwrap(); + } + + let mut scanner = dataset.scan(); + if index_type == IndexType::Vector { + scanner + .nearest("vec", &Float32Array::from(vec![4.0, 5.0]), 2) + .unwrap(); + } else { + scanner + .full_text_search(FullTextSearchQuery::new("quick".to_owned())) + .unwrap(); + } + let mut expected_columns = if has_projection { + scanner.project(&["id"]).unwrap(); + vec!["id"] + } else { + vec!["id", "text", "vec"] + }; + expected_columns.push(scoring_column); + if is_row_id_requested { + scanner.with_row_id(); + expected_columns.push(ROW_ID); + } + scanner.limit(Some(2), None).unwrap(); + let normal_batch = scanner.try_into_batch().await.unwrap(); + + scanner.fast_search(); + let fast_batch = scanner.try_into_batch().await.unwrap(); + assert_eq!( + fast_batch + .schema() + .fields() + .iter() + .map(|f| f.name().as_str()) + .collect::>(), + expected_columns + ); + assert_eq!(fast_batch, normal_batch); + assert_eq!(fast_batch.num_rows(), 2); + let recall = fast_batch["id"] + .as_primitive::() + .values() + .iter() + .filter(|id| expected_ids.contains(id)) + .count() as f64 + / expected_ids.len() as f64; + assert_eq!(recall, 1.0); + + scanner.project(&["id", ROW_ID]).unwrap(); + let explicit_row_id_batch = scanner.try_into_batch().await.unwrap(); + assert_eq!( + explicit_row_id_batch + .schema() + .fields() + .iter() + .map(|f| f.name().as_str()) + .collect::>(), + if is_row_id_requested { + vec!["id", scoring_column, ROW_ID] + } else { + vec!["id", ROW_ID, scoring_column] + } + ); + assert_eq!(&explicit_row_id_batch["id"], &fast_batch["id"]); + assert_eq!(explicit_row_id_batch[ROW_ID].null_count(), 0); + } + + #[rstest] + #[tokio::test] + async fn test_fast_search_requires_row_id_for_deleted_rows( + #[values(false, true)] is_row_id_requested: bool, + ) { + let mut dataset = gen_batch() + .col("id", array::step::()) + .into_ram_dataset(FragmentCount::from(2), FragmentRowCount::from(2)) + .await + .unwrap(); + dataset.delete("id = 1").await.unwrap(); + + let mut scanner = dataset.scan(); + scanner + .project(&["id"]) + .unwrap() + .fast_search() + .include_deleted_rows(); + if is_row_id_requested { + scanner.with_row_id(); + } + + let result = scanner.try_into_batch().await; + if is_row_id_requested { + let batch = result.unwrap(); + assert_eq!(batch.num_rows(), 4); + assert_eq!(batch[ROW_ID].null_count(), 1); + } else { + let error = result.unwrap_err(); + assert!(matches!(error, Error::InvalidInput { .. })); + assert!( + error + .to_string() + .contains("include_deleted_rows is set but with_row_id is false") + ); + } + } + #[tokio::test] async fn test_fast_search_without_vector_index_returns_empty() { let dataset = TestVectorDataset::new(LanceFileVersion::Stable, true) @@ -17427,6 +17588,22 @@ full_filter=name LIKE Utf8(\"test%2\"), refine_filter=name LIKE Utf8(\"test%2\") assert_eq!(normal_batch.num_rows(), 15); assert_eq!(fast_batch.num_rows(), 5); + + let mut project_first_scanner = dataset.dataset.scan(); + project_first_scanner + .filter("i >= 395") + .unwrap() + .project(&["i"]) + .unwrap() + .fast_search(); + assert_eq!( + project_first_scanner.explain_plan(false).await.unwrap(), + scanner.explain_plan(false).await.unwrap() + ); + assert_eq!( + project_first_scanner.try_into_batch().await.unwrap(), + fast_batch + ); } fn make_scalar_filter_test_batch(schema: SchemaRef, start: i32, end: i32) -> RecordBatch {