diff --git a/rust/lance-index/src/scalar/bloomfilter.rs b/rust/lance-index/src/scalar/bloomfilter.rs index b861340b9ee..8ba6fc281b5 100644 --- a/rust/lance-index/src/scalar/bloomfilter.rs +++ b/rust/lance-index/src/scalar/bloomfilter.rs @@ -19,7 +19,8 @@ use crate::{Any, pb}; use arrow_array::{Array, UInt64Array}; mod as_bytes; pub mod sbbf; -use arrow_schema::{DataType, Field}; +use arrow_schema::{DataType, Field, Schema}; +use futures::TryStreamExt; use serde::{Deserialize, Serialize}; use std::sync::LazyLock; @@ -44,6 +45,7 @@ use super::zoned::{ZoneBound, ZoneProcessor, ZoneTrainer, rebuild_zones, search_ const BLOOMFILTER_FILENAME: &str = "bloomfilter.lance"; const BLOOMFILTER_ITEM_META_KEY: &str = "bloomfilter_item"; const BLOOMFILTER_PROBABILITY_META_KEY: &str = "bloomfilter_probability"; +const MAX_BLOOMFILTER_ARRAY_LENGTH: usize = i32::MAX as usize - 1024 * 1024; const BLOOMFILTER_INDEX_VERSION: u32 = 0; #[derive(Debug, Clone)] @@ -94,9 +96,6 @@ impl BloomFilterIndex { _index_cache: &LanceCache, ) -> Result> { let index_file = store.open_index_file(BLOOMFILTER_FILENAME).await?; - let bloom_data = index_file - .read_range(0..index_file.num_rows(), None) - .await?; let file_schema = index_file.schema(); let number_of_items: u64 = file_schema @@ -111,25 +110,52 @@ impl BloomFilterIndex { .and_then(|bs| bs.parse().ok()) .unwrap_or(*DEFAULT_PROBABILITY); - Ok(Arc::new(Self::try_from_serialized( - bloom_data, + let read_batch_size = + Self::read_batch_size(number_of_items, probability, MAX_BLOOMFILTER_ARRAY_LENGTH)?; + + let mut zones = Vec::new(); + for start in (0..index_file.num_rows()).step_by(read_batch_size) { + let end = (start + read_batch_size).min(index_file.num_rows()); + let mut bloom_data = index_file.read_range_stream(start..end, None).await?; + while let Some(batch) = bloom_data.try_next().await? { + zones.extend(Self::try_from_serialized(batch)?); + } + } + + Ok(Arc::new(Self { + zones, number_of_items, probability, - )?)) + })) } - fn try_from_serialized( - data: RecordBatch, + fn read_batch_size( number_of_items: u64, probability: f64, - ) -> Result { + max_array_length: usize, + ) -> Result { + // Bloom filters are stored in an Arrow BinaryArray, whose offsets are i32. + // The serialized filter size is fixed by the index parameters, so bound + // reads by total serialized bytes instead of row count alone. + let params = BloomFilterIndexBuilderParams { + number_of_items, + probability, + }; + let filter_size = BloomFilterProcessor::build_filter(¶ms)? + .to_bytes() + .len(); + if filter_size > max_array_length { + return Err(Error::invalid_input(format!( + "Serialized bloom filter size {} exceeds max supported batch bytes {}", + filter_size, max_array_length + ))); + } + Ok((max_array_length / filter_size).max(1)) + } + + fn try_from_serialized(data: RecordBatch) -> Result> { if data.num_rows() == 0 { - // Return empty index for empty data - return Ok(Self { - zones: Vec::new(), - number_of_items, - probability, - }); + return Ok(Vec::new()); } let fragment_id_col = data @@ -189,7 +215,6 @@ impl BloomFilterIndex { Vec::new() }; - // Convert bytes back to Sbbf let bloom_filter = Sbbf::new(&bloom_filter_bytes).map_err(|e| { Error::invalid_input(format!("Failed to deserialize bloom filter: {:?}", e)) })?; @@ -205,11 +230,7 @@ impl BloomFilterIndex { }); } - Ok(Self { - zones: blocks, - number_of_items, - probability, - }) + Ok(blocks) } fn evaluate_block_against_query( @@ -458,7 +479,9 @@ impl ScalarIndex for BloomFilterIndex { // Write the combined zones back to storage let mut builder = BloomFilterIndexBuilder::try_new(params)?; builder.blocks = updated_blocks; - builder.write_index(dest_store).await?; + builder + .write_index(dest_store, MAX_BLOOMFILTER_ARRAY_LENGTH) + .await?; Ok(CreatedIndex { index_details: prost_types::Any::from_msg(&pb::BloomFilterIndexDetails::default()) @@ -567,33 +590,30 @@ impl BloomFilterIndexBuilder { Ok(()) } - fn bloomfilter_stats_as_batch(&self) -> Result { - let fragment_ids = - UInt64Array::from_iter_values(self.blocks.iter().map(|block| block.bound.fragment_id)); - - let zone_starts = - UInt64Array::from_iter_values(self.blocks.iter().map(|block| block.bound.start)); - - let zone_lengths = UInt64Array::from_iter_values( - self.blocks.iter().map(|block| block.bound.length as u64), - ); - - let has_nulls = arrow_array::BooleanArray::from( - self.blocks - .iter() - .map(|block| block.has_null) - .collect::>(), - ); + fn bloomfilter_schema() -> Arc { + Arc::new(Schema::new(vec![ + Field::new("fragment_id", DataType::UInt64, false), + Field::new("zone_start", DataType::UInt64, false), + Field::new("zone_length", DataType::UInt64, false), + Field::new("has_null", DataType::Boolean, false), + Field::new("bloom_filter_data", DataType::Binary, false), + ])) + } - // Convert bloom filters to binary data for serialization - let bloom_filter_data = if self.blocks.is_empty() { + fn bloomfilter_stats_as_batch( + fragment_ids: Vec, + zone_starts: Vec, + zone_lengths: Vec, + has_nulls: Vec, + binary_data: Vec>, + ) -> Result { + let fragment_ids = UInt64Array::from(fragment_ids); + let zone_starts = UInt64Array::from(zone_starts); + let zone_lengths = UInt64Array::from(zone_lengths); + let has_nulls = arrow_array::BooleanArray::from(has_nulls); + let bloom_filter_data = if binary_data.is_empty() { Arc::new(arrow_array::BinaryArray::new_null(0)) as ArrayRef } else { - let binary_data: Vec> = self - .blocks - .iter() - .map(|block| block.bloom_filter.to_bytes()) - .collect(); let binary_refs: Vec> = binary_data .iter() .map(|bytes| Some(bytes.as_slice())) @@ -601,14 +621,6 @@ impl BloomFilterIndexBuilder { Arc::new(arrow_array::BinaryArray::from_opt_vec(binary_refs)) as ArrayRef }; - let schema = Arc::new(arrow_schema::Schema::new(vec![ - Field::new("fragment_id", DataType::UInt64, false), - Field::new("zone_start", DataType::UInt64, false), - Field::new("zone_length", DataType::UInt64, false), - Field::new("has_null", DataType::Boolean, false), - Field::new("bloom_filter_data", DataType::Binary, false), - ])); - let columns: Vec = vec![ Arc::new(fragment_ids) as ArrayRef, Arc::new(zone_starts) as ArrayRef, @@ -617,13 +629,15 @@ impl BloomFilterIndexBuilder { bloom_filter_data, ]; - Ok(RecordBatch::try_new(schema, columns)?) + Ok(RecordBatch::try_new(Self::bloomfilter_schema(), columns)?) } - pub async fn write_index(self, index_store: &dyn IndexStore) -> Result<()> { - let record_batch = self.bloomfilter_stats_as_batch()?; - - let mut file_schema = record_batch.schema().as_ref().clone(); + pub async fn write_index( + self, + index_store: &dyn IndexStore, + max_array_length: usize, + ) -> Result<()> { + let mut file_schema = Self::bloomfilter_schema().as_ref().clone(); file_schema.metadata.insert( BLOOMFILTER_ITEM_META_KEY.to_string(), self.params.number_of_items.to_string(), @@ -634,11 +648,113 @@ impl BloomFilterIndexBuilder { self.params.probability.to_string(), ); - let mut index_file = index_store + let index_file = index_store .new_index_file(BLOOMFILTER_FILENAME, Arc::new(file_schema)) .await?; - index_file.write_record_batch(record_batch).await?; - index_file.finish().await?; + + let mut writer = BloomFilterBatchWriter::new(index_file, max_array_length); + for block in self.blocks { + writer.emit(block).await?; + } + writer.finish().await + } +} + +/// Buffers serialized bloom filter zone statistics and flushes them as record batches +/// to the index file, respecting the `max_array_length` limit. +struct BloomFilterBatchWriter { + file: Box, + max_array_length: usize, + fragment_ids: Vec, + zone_starts: Vec, + zone_lengths: Vec, + has_nulls: Vec, + bloom_filter_data: Vec>, + current_bytes: usize, + has_written: bool, +} + +impl BloomFilterBatchWriter { + fn new(file: Box, max_array_length: usize) -> Self { + Self { + file, + max_array_length, + fragment_ids: Vec::new(), + zone_starts: Vec::new(), + zone_lengths: Vec::new(), + has_nulls: Vec::new(), + bloom_filter_data: Vec::new(), + current_bytes: 0, + has_written: false, + } + } + + async fn emit(&mut self, block: BloomFilterStatistics) -> Result<()> { + let serialized_filter = block.bloom_filter.to_bytes(); + let serialized_len = serialized_filter.len(); + + if serialized_len > self.max_array_length { + return Err(Error::invalid_input(format!( + "Serialized bloom filter size {} exceeds max supported batch bytes {}", + serialized_len, self.max_array_length + ))); + } + + let next_bytes = self + .current_bytes + .checked_add(serialized_len) + .ok_or_else(|| { + Error::invalid_input(format!( + "Bloom filter batch size overflow when adding {} bytes to {} bytes", + serialized_len, self.current_bytes + )) + })?; + + if !self.bloom_filter_data.is_empty() && next_bytes > self.max_array_length { + self.flush().await?; + } + + self.fragment_ids.push(block.bound.fragment_id); + self.zone_starts.push(block.bound.start); + self.zone_lengths.push(block.bound.length as u64); + self.has_nulls.push(block.has_null); + self.bloom_filter_data.push(serialized_filter); + self.current_bytes += serialized_len; + Ok(()) + } + + async fn flush(&mut self) -> Result<()> { + if self.bloom_filter_data.is_empty() { + return Ok(()); + } + + let batch = BloomFilterIndexBuilder::bloomfilter_stats_as_batch( + std::mem::take(&mut self.fragment_ids), + std::mem::take(&mut self.zone_starts), + std::mem::take(&mut self.zone_lengths), + std::mem::take(&mut self.has_nulls), + std::mem::take(&mut self.bloom_filter_data), + )?; + self.file.write_record_batch(batch).await?; + self.current_bytes = 0; + self.has_written = true; + Ok(()) + } + + async fn finish(mut self) -> Result<()> { + self.flush().await?; + if !self.has_written { + self.file + .write_record_batch(BloomFilterIndexBuilder::bloomfilter_stats_as_batch( + Vec::new(), + Vec::new(), + Vec::new(), + Vec::new(), + Vec::new(), + )?) + .await?; + } + self.file.finish().await?; Ok(()) } } @@ -981,7 +1097,9 @@ impl BloomFilterIndexPlugin { builder.train(batches_source).await?; - builder.write_index(index_store).await?; + builder + .write_index(index_store, MAX_BLOOMFILTER_ARRAY_LENGTH) + .await?; Ok(()) } } @@ -1163,7 +1281,9 @@ mod tests { use crate::scalar::{ BloomFilterQuery, ScalarIndex, SearchResult, - bloomfilter::{BloomFilterIndex, BloomFilterIndexBuilderParams}, + bloomfilter::{ + BloomFilterIndex, BloomFilterIndexBuilder, BloomFilterIndexBuilderParams, sbbf::Sbbf, + }, lance_format::LanceIndexStore, }; @@ -2126,4 +2246,122 @@ mod tests { _ => panic!("Expected AtMost search result from bloomfilter"), } } + + #[test] + fn test_bloomfilter_read_batch_size_is_byte_bounded() { + let number_of_items = 1; + let probability = 0.25; + let max_test_batch_bytes = 48; + let filter_bytes = Sbbf::with_ndv_fpp(number_of_items, probability) + .unwrap() + .to_bytes(); + + assert!(filter_bytes.len() <= max_test_batch_bytes); + assert!(filter_bytes.len() * 2 > max_test_batch_bytes); + + assert_eq!( + BloomFilterIndex::read_batch_size(number_of_items, probability, max_test_batch_bytes,) + .unwrap(), + 1 + ); + } + + #[test] + fn test_bloomfilter_read_batch_size_rejects_oversized_filter() { + let number_of_items = 1; + let probability = 0.25; + let filter_bytes = Sbbf::with_ndv_fpp(number_of_items, probability) + .unwrap() + .to_bytes(); + + let error = + BloomFilterIndex::read_batch_size(number_of_items, probability, filter_bytes.len() - 1) + .unwrap_err(); + assert!( + error + .to_string() + .contains("exceeds max supported batch bytes"), + "unexpected error: {error}" + ); + } + + #[tokio::test] + async fn test_bloomfilter_chunked_write_and_load() { + let tmpdir = TempObjDir::default(); + let test_store = Arc::new(LanceIndexStore::new( + Arc::new(ObjectStore::local()), + tmpdir.clone(), + Arc::new(LanceCache::no_cache()), + )); + + let row_count = 5_000; + let data = arrow_array::Int32Array::from_iter_values(0..row_count); + let schema = Arc::new(Schema::new(vec![Field::new( + VALUE_COLUMN_NAME, + DataType::Int32, + false, + )])); + let data = RecordBatch::try_new(schema.clone(), vec![Arc::new(data)]).unwrap(); + let data_stream: SendableRecordBatchStream = Box::pin(RecordBatchStreamAdapter::new( + schema, + stream::once(std::future::ready(Ok(data))), + )); + let data_stream = add_row_addr(data_stream); + + let mut builder = + BloomFilterIndexBuilder::try_new(BloomFilterIndexBuilderParams::new(1, 0.25)).unwrap(); + builder.train(data_stream).await.unwrap(); + assert_eq!(builder.blocks.len(), row_count as usize); + + builder.write_index(test_store.as_ref(), 64).await.unwrap(); + + let index = BloomFilterIndex::load(test_store.clone(), None, &LanceCache::no_cache()) + .await + .expect("Failed to load chunked BloomFilterIndex"); + + assert_eq!(index.zones.len(), row_count as usize); + assert_eq!(index.zones[0].bound.start, 0); + assert_eq!(index.zones[4096].bound.start, 4096); + assert_eq!(index.zones[4999].bound.start, 4999); + assert_eq!(index.zones[4999].bound.length, 1); + assert!(!index.zones[4096].bloom_filter.to_bytes().is_empty()); + } + + #[tokio::test] + async fn test_bloomfilter_chunked_write_rejects_oversized_filter() { + let tmpdir = TempObjDir::default(); + let test_store = Arc::new(LanceIndexStore::new( + Arc::new(ObjectStore::local()), + tmpdir.clone(), + Arc::new(LanceCache::no_cache()), + )); + + let data = arrow_array::Int32Array::from_iter_values(0..1); + let schema = Arc::new(Schema::new(vec![Field::new( + VALUE_COLUMN_NAME, + DataType::Int32, + false, + )])); + let data = RecordBatch::try_new(schema.clone(), vec![Arc::new(data)]).unwrap(); + let data_stream: SendableRecordBatchStream = Box::pin(RecordBatchStreamAdapter::new( + schema, + stream::once(std::future::ready(Ok(data))), + )); + let data_stream = add_row_addr(data_stream); + + let mut builder = + BloomFilterIndexBuilder::try_new(BloomFilterIndexBuilderParams::new(1, 0.25)).unwrap(); + builder.train(data_stream).await.unwrap(); + + let error = builder + .write_index(test_store.as_ref(), 16) + .await + .unwrap_err(); + assert!( + error + .to_string() + .contains("exceeds max supported batch bytes 16"), + "unexpected error: {error}" + ); + } }