diff --git a/python/python/tests/test_scalar_index.py b/python/python/tests/test_scalar_index.py index cdc4b677132..7976a153f0c 100644 --- a/python/python/tests/test_scalar_index.py +++ b/python/python/tests/test_scalar_index.py @@ -2903,7 +2903,7 @@ def test_zonemap_index_remapping(tmp_path: Path): # Run compaction to merge fragments compaction = dataset.optimize.compact_files(target_rows_per_fragment=2000) assert compaction.fragments_removed == 5 - assert len(dataset.get_fragments()) == 3 + assert len(dataset.get_fragments()) == 2 # Check if the zone map index is no longer being used scanner = dataset.scanner(filter="values > 2500", prefilter=True) diff --git a/rust/lance-datafusion/src/chunker.rs b/rust/lance-datafusion/src/chunker.rs index f30e215e712..63523460899 100644 --- a/rust/lance-datafusion/src/chunker.rs +++ b/rust/lance-datafusion/src/chunker.rs @@ -42,8 +42,8 @@ impl BatchReaderChunker { buffer_total - self.i } - async fn fill_buffer(&mut self) -> Result<()> { - while self.buffered_len() < self.output_size { + async fn fill_buffer(&mut self, output_size: usize) -> Result<()> { + while self.buffered_len() < output_size { match self.inner.next().await { Some(Ok(batch)) => self.buffered.push_back(batch), Some(Err(e)) => return Err(e.into()), @@ -54,7 +54,11 @@ impl BatchReaderChunker { } async fn next(&mut self) -> Option>> { - match self.fill_buffer().await { + self.next_sized(self.output_size).await + } + + async fn next_sized(&mut self, output_size: usize) -> Option>> { + match self.fill_buffer(output_size).await { Ok(_) => {} Err(e) => return Some(Err(e)), }; @@ -63,7 +67,7 @@ impl BatchReaderChunker { let mut rows_collected = 0; - while rows_collected < self.output_size { + while rows_collected < output_size { if let Some(batch) = self.buffered.pop_front() { // Skip empty batch if batch.num_rows() == 0 { @@ -72,7 +76,7 @@ impl BatchReaderChunker { let rows_remaining_in_batch = batch.num_rows() - self.i; let rows_to_take = - std::cmp::min(rows_remaining_in_batch, self.output_size - rows_collected); + std::cmp::min(rows_remaining_in_batch, output_size - rows_collected); if rows_to_take == rows_remaining_in_batch { // We're taking the whole batch, so we can just move it @@ -104,6 +108,53 @@ impl BatchReaderChunker { Some(Ok(batches)) } } + + async fn next_at_most(&mut self, output_size: usize) -> Option>> { + loop { + let batch = match self.buffered.pop_front() { + Some(batch) => batch, + None => match self.inner.next().await { + Some(Ok(batch)) => batch, + Some(Err(error)) => return Some(Err(error.into())), + None => return None, + }, + }; + + if batch.num_rows() == 0 { + continue; + } + + let rows_remaining_in_batch = batch.num_rows() - self.i; + let rows_to_take = rows_remaining_in_batch.min(output_size); + if rows_to_take == rows_remaining_in_batch { + let batch = if self.i == 0 { + batch + } else { + batch.slice(self.i, rows_to_take) + }; + self.i = 0; + return Some(Ok(vec![batch])); + } + + let output = batch.slice(self.i, rows_to_take); + self.i += rows_to_take; + self.buffered.push_front(batch); + return Some(Ok(vec![output])); + } + } +} + +struct VariableBatchReaderChunker { + chunker: BatchReaderChunker, + output_sizes: I, + is_done: bool, +} + +struct VariableBreakStreamState { + chunker: BatchReaderChunker, + output_sizes: I, + rows_remaining: Option, + is_done: bool, } struct BreakStreamState { @@ -186,6 +237,210 @@ pub fn chunk_stream( .boxed() } +/// Preserve input batch boundaries while inserting the requested row boundaries. +/// +/// The requested sizes must describe the complete input. Unlike +/// [`chunk_stream_with_sizes`], this does not combine adjacent input batches. It +/// only slices a batch when it crosses a requested boundary. +/// +/// # Example +/// +/// ``` +/// # use datafusion::physical_plan::SendableRecordBatchStream; +/// # use lance_datafusion::chunker::break_stream_with_sizes; +/// # fn split_stream(stream: SendableRecordBatchStream) { +/// let batches = break_stream_with_sizes(stream, vec![512, 512, 256]); +/// # drop(batches); +/// # } +/// ``` +pub fn break_stream_with_sizes( + stream: SendableRecordBatchStream, + output_sizes: I, +) -> Pin>> + Send>> +where + I: IntoIterator, + I::IntoIter: Send + 'static, +{ + let state = VariableBreakStreamState { + chunker: BatchReaderChunker::new(stream, 1), + output_sizes: output_sizes.into_iter(), + rows_remaining: None, + is_done: false, + }; + futures::stream::unfold(state, |mut state| async move { + if state.is_done { + return None; + } + + if state.rows_remaining.is_none() { + let Some(output_size) = state.output_sizes.next() else { + return match state.chunker.next_at_most(1).await { + None => None, + Some(Ok(_)) => { + state.is_done = true; + Some(( + Err(lance_core::Error::invalid_input( + "Input contained more rows than the requested chunk sizes", + )), + state, + )) + } + Some(Err(error)) => { + state.is_done = true; + Some((Err(error), state)) + } + }; + }; + if output_size == 0 { + state.is_done = true; + return Some(( + Err(lance_core::Error::invalid_input( + "Requested chunk sizes must be greater than zero", + )), + state, + )); + } + state.rows_remaining = Some(output_size); + } + + let Some(rows_remaining) = state.rows_remaining else { + state.is_done = true; + return Some(( + Err(lance_core::Error::internal( + "Requested chunk boundary was not initialized", + )), + state, + )); + }; + match state.chunker.next_at_most(rows_remaining).await { + Some(Ok(batches)) => { + let actual_size = batches.iter().map(RecordBatch::num_rows).sum::(); + let Some(rows_remaining) = rows_remaining.checked_sub(actual_size) else { + state.is_done = true; + return Some(( + Err(lance_core::Error::internal( + "A boundary-preserving chunk exceeded its requested row count", + )), + state, + )); + }; + state.rows_remaining = (rows_remaining > 0).then_some(rows_remaining); + Some((Ok(batches), state)) + } + Some(Err(error)) => { + state.is_done = true; + Some((Err(error), state)) + } + None => { + state.is_done = true; + Some(( + Err(lance_core::Error::invalid_input(format!( + "Input ended with {rows_remaining} rows remaining in a requested chunk" + ))), + state, + )) + } + } + }) + .boxed() +} + +/// Given a stream of record batches, yield chunks with the requested row counts. +/// +/// The requested sizes must describe the complete input. An error is returned if +/// the input ends early, contains additional rows, or a requested size is zero. +/// Sizes are consumed lazily as chunks are requested. +/// +/// # Example +/// +/// ``` +/// # use datafusion::physical_plan::SendableRecordBatchStream; +/// # use lance_datafusion::chunker::chunk_stream_with_sizes; +/// # fn split_stream(stream: SendableRecordBatchStream) { +/// let chunks = chunk_stream_with_sizes(stream, vec![512, 512, 256]); +/// # drop(chunks); +/// # } +/// ``` +pub fn chunk_stream_with_sizes( + stream: SendableRecordBatchStream, + output_sizes: I, +) -> Pin>> + Send>> +where + I: IntoIterator, + I::IntoIter: Send + 'static, +{ + let state = VariableBatchReaderChunker { + chunker: BatchReaderChunker::new(stream, 1), + output_sizes: output_sizes.into_iter(), + is_done: false, + }; + futures::stream::unfold(state, |mut state| async move { + if state.is_done { + return None; + } + + let Some(output_size) = state.output_sizes.next() else { + return match state.chunker.next_sized(1).await { + None => None, + Some(Ok(_)) => { + state.is_done = true; + Some(( + Err(lance_core::Error::invalid_input( + "Input contained more rows than the requested chunk sizes", + )), + state, + )) + } + Some(Err(error)) => { + state.is_done = true; + Some((Err(error), state)) + } + }; + }; + + if output_size == 0 { + state.is_done = true; + return Some(( + Err(lance_core::Error::invalid_input( + "Requested chunk sizes must be greater than zero", + )), + state, + )); + } + + match state.chunker.next_sized(output_size).await { + Some(Ok(batches)) => { + let actual_size = batches.iter().map(RecordBatch::num_rows).sum::(); + if actual_size == output_size { + Some((Ok(batches), state)) + } else { + state.is_done = true; + Some(( + Err(lance_core::Error::invalid_input(format!( + "Input ended after {actual_size} rows while filling a requested {output_size}-row chunk" + ))), + state, + )) + } + } + Some(Err(error)) => { + state.is_done = true; + Some((Err(error), state)) + } + None => { + state.is_done = true; + Some(( + Err(lance_core::Error::invalid_input(format!( + "Input ended before a requested {output_size}-row chunk could be filled" + ))), + state, + )) + } + } + }) + .boxed() +} + /// Given a stream of record batches, this will yield batches of a fixed size. /// /// This stream _will_ combine record batches and so it can be fairly expensive as it will @@ -311,7 +566,10 @@ where #[cfg(test)] mod tests { - use std::sync::Arc; + use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }; use arrow::datatypes::{Int32Type, Int64Type}; use datafusion::physical_plan::stream::RecordBatchStreamAdapter; @@ -360,6 +618,82 @@ mod tests { assert_eq!(chunked[2].len(), 1); assert_eq!(chunked[2][0].num_rows(), 8); + let sizes_consumed = Arc::new(AtomicUsize::new(0)); + let requested_sizes = [9, 10, 9].into_iter().inspect({ + let sizes_consumed = sizes_consumed.clone(); + move |_| { + sizes_consumed.fetch_add(1, Ordering::SeqCst); + } + }); + let mut chunked = super::chunk_stream_with_sizes(make_stream(), requested_sizes); + assert_eq!(sizes_consumed.load(Ordering::SeqCst), 0); + let first_chunk = chunked.next().await.unwrap().unwrap(); + assert_eq!(sizes_consumed.load(Ordering::SeqCst), 1); + let mut chunked = chunked.try_collect::>().await.unwrap(); + chunked.insert(0, first_chunk); + assert_eq!(sizes_consumed.load(Ordering::SeqCst), 3); + assert_eq!( + chunked + .iter() + .map(|batches| batches.iter().map(|batch| batch.num_rows()).sum::()) + .collect::>(), + vec![9, 10, 9] + ); + + let error = super::chunk_stream_with_sizes(make_stream(), vec![10, 17]) + .try_collect::>() + .await + .unwrap_err(); + assert!( + error + .to_string() + .contains("more rows than the requested chunk sizes") + ); + + let error = super::chunk_stream_with_sizes(make_stream(), vec![10, 19]) + .try_collect::>() + .await + .unwrap_err(); + assert!(error.to_string().contains("ended after 18 rows")); + + let sizes_consumed = Arc::new(AtomicUsize::new(0)); + let requested_sizes = [9, 10, 9].into_iter().inspect({ + let sizes_consumed = sizes_consumed.clone(); + move |_| { + sizes_consumed.fetch_add(1, Ordering::SeqCst); + } + }); + let mut broken = super::break_stream_with_sizes(make_stream(), requested_sizes); + assert_eq!(sizes_consumed.load(Ordering::SeqCst), 0); + let first_batch = broken.next().await.unwrap().unwrap(); + assert_eq!(sizes_consumed.load(Ordering::SeqCst), 1); + let mut broken = broken.try_collect::>().await.unwrap(); + broken.insert(0, first_batch); + assert_eq!(sizes_consumed.load(Ordering::SeqCst), 3); + assert_eq!( + broken + .iter() + .map(|batches| batches.iter().map(|batch| batch.num_rows()).sum::()) + .collect::>(), + vec![9, 1, 5, 4, 9] + ); + + let error = super::break_stream_with_sizes(make_stream(), vec![27]) + .try_collect::>() + .await + .unwrap_err(); + assert!( + error + .to_string() + .contains("more rows than the requested chunk sizes") + ); + + let error = super::break_stream_with_sizes(make_stream(), vec![29]) + .try_collect::>() + .await + .unwrap_err(); + assert!(error.to_string().contains("1 rows remaining")); + let chunked = super::chunk_concat_stream(make_stream(), 10) .try_collect::>() .await diff --git a/rust/lance/src/dataset.rs b/rust/lance/src/dataset.rs index 4b6beb9e782..895fc3f556e 100644 --- a/rust/lance/src/dataset.rs +++ b/rust/lance/src/dataset.rs @@ -123,7 +123,7 @@ use self::refs::Refs; use self::scanner::{DatasetRecordBatchStream, Scanner}; use self::statistics::DatasetStatistics; use self::transaction::{Operation, Transaction, TransactionBuilder, UpdateMapEntry}; -use self::write::{cleanup_data_fragments, write_fragments_internal}; +use self::write::cleanup_data_fragments; use crate::dataset::branch_location::BranchLocation; use crate::dataset::cleanup::{CleanupOperation, CleanupPolicy, CleanupPolicyBuilder}; use crate::dataset::refs::{BranchContents, BranchIdentifier, Branches, Tags}; diff --git a/rust/lance/src/dataset/fragment/write.rs b/rust/lance/src/dataset/fragment/write.rs index eabade7afdf..b641ee16cc5 100644 --- a/rust/lance/src/dataset/fragment/write.rs +++ b/rust/lance/src/dataset/fragment/write.rs @@ -259,6 +259,7 @@ impl<'a> FragmentCreateBuilder<'a> { params, target_bases_info, Vec::new(), + None, ) .await } diff --git a/rust/lance/src/dataset/optimize.rs b/rust/lance/src/dataset/optimize.rs index 4c8b4675871..f4e0ea4ca16 100644 --- a/rust/lance/src/dataset/optimize.rs +++ b/rust/lance/src/dataset/optimize.rs @@ -95,7 +95,10 @@ use super::transaction::{ }; use super::utils::make_rowid_capture_stream; use super::versions; -use super::{WriteMode, WriteParams, cleanup_data_fragments, write_fragments_internal}; +use super::{ + WriteMode, WriteParams, cleanup_data_fragments, + write::write_fragments_internal_with_file_row_counts, +}; use crate::Dataset; use crate::Result; use crate::dataset::utils::CapturedRowIds; @@ -2390,8 +2393,54 @@ async fn rewrite_files( } } + let surviving_rows = fragments.iter().try_fold(0_u64, |total, fragment| { + let fragment_rows = fragment.num_rows().ok_or_else(|| { + Error::internal(format!( + "Fragment {} is missing row count metadata after migration", + fragment.id + )) + })?; + total.checked_add(fragment_rows as u64).ok_or_else(|| { + Error::internal("Compaction task surviving row count overflowed u64".to_string()) + }) + })?; + + // Planner-sized tasks may exceed the target, but should remain one output + // instead of producing a target-sized fragment plus a stranded tail. For + // genuinely oversized tasks, choose a target-scale output count and spread + // the tail across those outputs. + let target_rows_per_fragment = options.target_rows_per_fragment as u64; + let output_fragment_count = surviving_rows + .checked_div(target_rows_per_fragment) + .unwrap_or(1) + .max(1); + let output_fragment_count_usize = usize::try_from(output_fragment_count).map_err(|_| { + Error::internal(format!( + "Compaction output fragment count {output_fragment_count} does not fit in usize" + )) + })?; + let base_rows_per_file = surviving_rows / output_fragment_count; + let larger_file_count = + usize::try_from(surviving_rows % output_fragment_count).map_err(|_| { + Error::internal("Compaction larger output fragment count does not fit in usize") + })?; + let file_row_counts = if surviving_rows == 0 { + Vec::new() + } else { + (0..output_fragment_count_usize) + .map(|file_index| { + let file_rows = base_rows_per_file + u64::from(file_index < larger_file_count); + usize::try_from(file_rows).map_err(|_| { + Error::internal(format!( + "Compaction output row count {file_rows} does not fit in usize" + )) + }) + }) + .collect::>>()? + }; + let max_rows_per_file = file_row_counts.first().copied().unwrap_or(1); let mut params = WriteParams { - max_rows_per_file: options.target_rows_per_fragment, + max_rows_per_file, max_rows_per_group: options.max_rows_per_group, mode: WriteMode::Append, // External blobs may reference URIs outside the dataset's base_paths @@ -2444,7 +2493,7 @@ async fn rewrite_files( row_ids_rx = Some(rx); } } else { - let (frags, _) = write_fragments_internal( + let (frags, _) = write_fragments_internal_with_file_row_counts( dataset.manifest.data_storage_format.lance_file_format(), Some(dataset.as_ref()), dataset.object_store.clone(), @@ -2453,6 +2502,7 @@ async fn rewrite_files( reader.expect("reader must be prepared for non-binary-copy path"), params, None, + Some(file_row_counts), ) .await?; new_fragments = frags; @@ -3599,49 +3649,44 @@ mod tests { .unwrap(); let first_new_frag_idx = 7; - // Predicting the remap is difficult. One task will remap to fragments 7/8 and the other - // will remap to fragments 9/10 but we don't know which is which and so we just allow ourselves - // to expect both possibilities. + // The tasks execute concurrently, so either one may reserve the first + // output fragment id. let remap_a = expect_remap( &[ vec![ - // 3 small fragments are rewritten to frags 7 & 8 + // 3 small fragments are rewritten to frag 7 (row_addrs(0, 0..400), true), (row_addrs(1, 0..400), true), - (row_addrs(2, 0..200), true), + (row_addrs(2, 0..400), true), ], - vec![(row_addrs(2, 200..400), true)], // frag 3 is skipped since it does not have enough missing data - // Frags 4, 5, and 6 are rewritten to frags 9 & 10 + // Frags 4, 5, and 6 are rewritten to frag 8 vec![ - // Only 800 of the 1000 rows taken from frag 4 (row_addrs(4, 0..200), true), (row_addrs(4, 200..400), false), (row_addrs(4, 400..1000), true), - // frags 5 compacted with frag 4 - (row_addrs(5, 0..200), true), + (row_addrs(5, 0..300), true), + (row_addrs(6, 0..300), true), ], - vec![(row_addrs(5, 200..300), true), (row_addrs(6, 0..300), true)], ], first_new_frag_idx, ); let remap_b = expect_remap( &[ - // Frags 4, 5, and 6 are rewritten to frags 7 & 8 + // Frags 4, 5, and 6 are rewritten to frag 7 vec![ (row_addrs(4, 0..200), true), (row_addrs(4, 200..400), false), (row_addrs(4, 400..1000), true), - (row_addrs(5, 0..200), true), + (row_addrs(5, 0..300), true), + (row_addrs(6, 0..300), true), ], - vec![(row_addrs(5, 200..300), true), (row_addrs(6, 0..300), true)], - // 3 small fragments rewritten to frags 9 & 10 + // 3 small fragments rewritten to frag 8 vec![ (row_addrs(0, 0..400), true), (row_addrs(1, 0..400), true), - (row_addrs(2, 0..200), true), + (row_addrs(2, 0..400), true), ], - vec![(row_addrs(2, 200..400), true)], ], first_new_frag_idx, ); @@ -3682,16 +3727,155 @@ mod tests { // Assert on metrics assert_eq!(metrics.fragments_removed, 6); - assert_eq!(metrics.fragments_added, 4); + assert_eq!(metrics.fragments_added, 2); assert_eq!(metrics.files_removed, 7); // 6 data files + 1 deletion file - assert_eq!(metrics.files_added, 4); + assert_eq!(metrics.files_added, 2); let fragment_ids = dataset .get_fragments() .iter() .map(|f| f.id()) .collect::>(); - assert_eq!(fragment_ids, vec![3, 7, 8, 9, 10]); + assert_eq!(fragment_ids, vec![3, 7, 8]); + } + + #[rstest] + #[tokio::test] + async fn test_compaction_does_not_strand_small_remainders( + #[values(LanceFileVersion::Legacy, LanceFileVersion::Stable)] + data_storage_version: LanceFileVersion, + ) { + let test_dir = TempStrDir::default(); + let data = sample_data().slice(0, 2_000); + let reader = RecordBatchIterator::new(vec![Ok(data.clone())], data.schema()); + let mut dataset = Dataset::write( + reader, + &test_dir, + Some(WriteParams { + max_rows_per_file: 200, + data_storage_version: Some(data_storage_version), + ..Default::default() + }), + ) + .await + .unwrap(); + + let options = CompactionOptions { + target_rows_per_fragment: 500, + ..Default::default() + }; + let metrics = compact_files(&mut dataset, options.clone(), None) + .await + .unwrap(); + + assert_eq!(metrics.fragments_removed, 10); + assert_eq!(metrics.fragments_added, 3); + let mut fragment_sizes = dataset + .get_fragments() + .iter() + .map(|fragment| fragment.metadata.physical_rows.unwrap()) + .collect::>(); + fragment_sizes.sort_unstable(); + assert_eq!(fragment_sizes, vec![600, 600, 800]); + + let second_metrics = compact_files(&mut dataset, options, None).await.unwrap(); + assert_eq!(second_metrics, CompactionMetrics::default()); + } + + #[rstest] + #[case::legacy(LanceFileVersion::Legacy)] + #[case::stable(LanceFileVersion::Stable)] + #[tokio::test] + async fn test_compaction_rebalances_oversized_task( + #[case] data_storage_version: LanceFileVersion, + ) { + let test_dir = TempStrDir::default(); + let data = sample_data().slice(0, 5_100); + let reader = RecordBatchIterator::new(vec![Ok(data.slice(0, 5_000))], data.schema()); + let mut dataset = Dataset::write( + reader, + &test_dir, + Some(WriteParams { + max_rows_per_file: 5_000, + data_storage_version: Some(data_storage_version), + ..Default::default() + }), + ) + .await + .unwrap(); + let reader = RecordBatchIterator::new(vec![Ok(data.slice(5_000, 100))], data.schema()); + dataset.append(reader, None).await.unwrap(); + + dataset.delete("a < 1000").await.unwrap(); + + let options = CompactionOptions { + target_rows_per_fragment: 1_000, + ..Default::default() + }; + let plan = plan_compaction(&dataset, &options).await.unwrap(); + assert_eq!(plan.tasks.len(), 1); + assert_eq!(plan.tasks[0].fragments.len(), 2); + + let metrics = compact_files(&mut dataset, options.clone(), None) + .await + .unwrap(); + assert_eq!(metrics.fragments_removed, 2); + assert_eq!(metrics.fragments_added, 4); + assert_eq!( + dataset + .get_fragments() + .iter() + .map(|fragment| fragment.metadata.physical_rows.unwrap()) + .collect::>(), + vec![1_025; 4] + ); + + let second_metrics = compact_files(&mut dataset, options, None).await.unwrap(); + assert_eq!(second_metrics, CompactionMetrics::default()); + } + + #[tokio::test] + async fn test_compaction_balances_non_divisible_stable_task() { + let test_dir = TempStrDir::default(); + let data = sample_data().slice(0, 121); + let reader = RecordBatchIterator::new(vec![Ok(data.clone())], data.schema()); + let mut dataset = Dataset::write( + reader, + &test_dir, + Some(WriteParams { + max_rows_per_file: 121, + data_storage_version: Some(LanceFileVersion::Stable), + ..Default::default() + }), + ) + .await + .unwrap(); + dataset.delete("a < 20").await.unwrap(); + + let options = CompactionOptions { + target_rows_per_fragment: 10, + ..Default::default() + }; + let plan = plan_compaction(&dataset, &options).await.unwrap(); + assert_eq!(plan.tasks.len(), 1); + assert_eq!(plan.tasks[0].fragments.len(), 1); + + let metrics = compact_files(&mut dataset, options.clone(), None) + .await + .unwrap(); + assert_eq!(metrics.fragments_removed, 1); + assert_eq!(metrics.fragments_added, 10); + assert_eq!( + dataset + .get_fragments() + .iter() + .map(|fragment| fragment.metadata.physical_rows.unwrap()) + .collect::>(), + [vec![11], vec![10; 9]].concat() + ); + + let second_metrics = compact_files(&mut dataset, options, None).await.unwrap(); + assert_eq!(second_metrics, CompactionMetrics::default()); } #[rstest] diff --git a/rust/lance/src/dataset/versions/mod.rs b/rust/lance/src/dataset/versions/mod.rs index 1fbac103898..67015b5b15a 100644 --- a/rust/lance/src/dataset/versions/mod.rs +++ b/rust/lance/src/dataset/versions/mod.rs @@ -19,7 +19,9 @@ use lance_core::{ Error, Result, datatypes::{Field, Projection, Schema, SchemaCompareOptions}, }; -use lance_datafusion::chunker::{break_stream, chunk_stream}; +use lance_datafusion::chunker::{ + break_stream, break_stream_with_sizes, chunk_stream, chunk_stream_with_sizes, +}; use lance_file::{ version::ConcreteFileVersion, versions as file_versions, @@ -124,6 +126,7 @@ pub async fn write_fragments( data: SendableRecordBatchStream, params: WriteParams, target_bases_info: Option>, + file_row_counts: Option>, ) -> Result<(Vec, Schema)> { let version_name = format!("{version:?}"); let schema = write::prepare_write_schema( @@ -151,6 +154,7 @@ pub async fn write_fragments( params, target_bases_info, seed_writers, + file_row_counts, ) .await?; Ok((fragments, schema)) @@ -167,17 +171,50 @@ pub async fn write_fragments_direct( params: WriteParams, target_bases_info: Option>, seed_writers: Vec>, + file_row_counts: Option>, ) -> Result> { let adapter = SchemaAdapter::new(data.schema()); let data = adapter.to_physical_stream(data); - let buffered_reader = match version { - ConcreteFileVersion::V1 => chunk_stream(data, params.max_rows_per_group), - ConcreteFileVersion::V2_0 - | ConcreteFileVersion::V2_1 - | ConcreteFileVersion::V2_2 - | ConcreteFileVersion::V2_3 => break_stream(data, params.max_rows_per_file) - .map_ok(|batch| vec![batch]) - .boxed(), + let buffered_reader = if let Some(file_row_counts) = file_row_counts.as_ref() { + if file_row_counts.contains(&0) { + return Err(Error::invalid_input( + "File row counts must be greater than zero", + )); + } + match version { + ConcreteFileVersion::V1 => { + if params.max_rows_per_group == 0 { + return Err(Error::invalid_input( + "max_rows_per_group must be greater than zero when file row counts are specified", + )); + } + let max_rows_per_group = params.max_rows_per_group; + let batch_row_counts = + file_row_counts + .clone() + .into_iter() + .flat_map(move |file_rows| { + (0..file_rows) + .step_by(max_rows_per_group) + .map(move |offset| (file_rows - offset).min(max_rows_per_group)) + }); + chunk_stream_with_sizes(data, batch_row_counts) + } + ConcreteFileVersion::V2_0 + | ConcreteFileVersion::V2_1 + | ConcreteFileVersion::V2_2 + | ConcreteFileVersion::V2_3 => break_stream_with_sizes(data, file_row_counts.clone()), + } + } else { + match version { + ConcreteFileVersion::V1 => chunk_stream(data, params.max_rows_per_group), + ConcreteFileVersion::V2_0 + | ConcreteFileVersion::V2_1 + | ConcreteFileVersion::V2_2 + | ConcreteFileVersion::V2_3 => break_stream(data, params.max_rows_per_file) + .map_ok(|batch| vec![batch]) + .boxed(), + } }; let external_base_resolver = match version { ConcreteFileVersion::V2_2 | ConcreteFileVersion::V2_3 => { @@ -198,6 +235,7 @@ pub async fn write_fragments_direct( external_base_resolver, target_bases_info, seed_writers, + file_row_counts, ) .await } diff --git a/rust/lance/src/dataset/write.rs b/rust/lance/src/dataset/write.rs index fffe71f930a..af99fc63b7a 100644 --- a/rust/lance/src/dataset/write.rs +++ b/rust/lance/src/dataset/write.rs @@ -31,7 +31,7 @@ use lance_table::io::commit::{CommitHandler, commit_handler_from_url}; use lance_table::io::manifest::ManifestDescribing; use object_store::path::Path; use std::borrow::Cow; -use std::collections::{BTreeSet, HashMap, HashSet}; +use std::collections::{BTreeSet, HashMap, HashSet, VecDeque}; use std::future::Future; use std::num::NonZero; use std::sync::Arc; @@ -595,6 +595,44 @@ pub async fn write_fragments( .await } +fn take_batch_rows(batches: &mut VecDeque, max_rows: usize) -> Vec { + let mut output = Vec::with_capacity(batches.len()); + let mut rows_remaining = max_rows; + + while rows_remaining > 0 { + let Some(batch) = batches.pop_front() else { + break; + }; + let batch_rows = batch.num_rows(); + if batch_rows == 0 { + continue; + } + if batch_rows <= rows_remaining { + rows_remaining -= batch_rows; + output.push(batch); + } else { + output.push(batch.slice(0, rows_remaining)); + batches.push_front(batch.slice(rows_remaining, batch_rows - rows_remaining)); + rows_remaining = 0; + } + } + + output +} + +fn balanced_row_counts(total_rows: usize, max_rows_per_file: usize) -> VecDeque { + if total_rows == 0 { + return VecDeque::new(); + } + + let file_count = total_rows.div_ceil(max_rows_per_file); + let base_rows_per_file = total_rows / file_count; + let larger_file_count = total_rows % file_count; + (0..file_count) + .map(|file_index| base_rows_per_file + usize::from(file_index < larger_file_count)) + .collect() +} + #[allow(clippy::too_many_arguments)] pub(super) async fn do_write_fragments_impl( dataset: Option<&Dataset>, @@ -607,6 +645,7 @@ pub(super) async fn do_write_fragments_impl( external_base_resolver: Option>, target_bases_info: Option>, mut seed_writers: Vec>, + file_row_counts: Option>, ) -> Result> where OpenWriter: Fn(Arc, Schema, Path, WriterOptions) -> OpenWriterFuture + Send + Sync, @@ -638,69 +677,177 @@ where let mut bytes_completed: u64 = 0; let mut rows_completed: u64 = 0; let mut files_written: u32 = 0; + let has_file_row_counts = file_row_counts.is_some(); + let max_planned_file_rows = file_row_counts + .as_ref() + .and_then(|row_counts| row_counts.iter().copied().max()); + let mut planned_rows_remaining = file_row_counts + .as_ref() + .map(|row_counts| { + row_counts.iter().try_fold(0_usize, |total, &row_count| { + total + .checked_add(row_count) + .ok_or_else(|| Error::internal("Planned file row count total overflowed usize")) + }) + }) + .transpose()?; + let mut file_row_counts = file_row_counts.map(VecDeque::from); + let mut rows_remaining_in_planned_file = file_row_counts.as_mut().and_then(VecDeque::pop_front); // Wrap the loop in an async block so `?` returns into `loop_result` and we // can run cleanup before propagating the error. let loop_result: Result<()> = async { while let Some(batch_chunk) = buffered_reader.next().await { - let batch_chunk = batch_chunk?; + let mut pending_batches = VecDeque::from(batch_chunk?); + + while !pending_batches.is_empty() { + let rows_to_take = if has_file_row_counts { + rows_remaining_in_planned_file.ok_or_else(|| { + Error::internal( + "Writer received rows after all planned file boundaries were consumed", + ) + })? + } else { + usize::MAX + }; + let batch_chunk = take_batch_rows(&mut pending_batches, rows_to_take); + if batch_chunk.is_empty() { + continue; + } - if writer.is_none() { - let (new_writer, new_fragment) = writer_generator.new_writer().await?; - params.progress.begin(&new_fragment).await?; - writer = Some(new_writer); - fragments.push(new_fragment); - } + if writer.is_none() { + let (new_writer, new_fragment) = writer_generator.new_writer().await?; + params.progress.begin(&new_fragment).await?; + writer = Some(new_writer); + fragments.push(new_fragment); + } - writer.as_mut().unwrap().write(&batch_chunk).await?; - for seed_writer in seed_writers.iter_mut() { - let col_name = seed_writer.column_name().to_owned(); - for batch in &batch_chunk { - if let Some(col) = batch.column_by_name(&col_name) { - seed_writer.observe_batch(col)?; + let active_writer = writer.as_mut().ok_or_else(|| { + Error::internal("Writer was not initialized before writing a batch") + })?; + active_writer.write(&batch_chunk).await?; + for seed_writer in seed_writers.iter_mut() { + let col_name = seed_writer.column_name().to_owned(); + for batch in &batch_chunk { + if let Some(col) = batch.column_by_name(&col_name) { + seed_writer.observe_batch(col)?; + } } } - } - for batch in &batch_chunk { - num_rows_in_current_file += batch.num_rows() as u32; - } - - if let Some(cb) = ¶ms.write_progress { - let current_bytes = writer.as_mut().unwrap().tell().await?; - cb.call(WriteStats { - bytes_written: bytes_completed + current_bytes, - rows_written: rows_completed + num_rows_in_current_file as u64, - files_written, - }); - } + let batch_chunk_rows = + batch_chunk.iter().map(RecordBatch::num_rows).sum::(); + num_rows_in_current_file += batch_chunk_rows as u32; + + let reached_planned_file_boundary = if has_file_row_counts { + let rows_remaining = + rows_remaining_in_planned_file.as_mut().ok_or_else(|| { + Error::internal( + "Writer received rows without an active planned file boundary", + ) + })?; + *rows_remaining = rows_remaining.checked_sub(batch_chunk_rows).ok_or_else(|| { + Error::internal(format!( + "Writer chunk of {batch_chunk_rows} rows crossed a planned file boundary with {rows_remaining} rows remaining" + )) + })?; + let total_remaining = planned_rows_remaining.as_mut().ok_or_else(|| { + Error::internal("Writer lost the planned row count total") + })?; + *total_remaining = + total_remaining.checked_sub(batch_chunk_rows).ok_or_else(|| { + Error::internal(format!( + "Writer consumed {batch_chunk_rows} rows after the planned row count total was exhausted" + )) + })?; + *rows_remaining == 0 + } else { + false + }; - if num_rows_in_current_file >= params.max_rows_per_file as u32 - || writer.as_mut().unwrap().tell().await? >= params.max_bytes_per_file as u64 - { - let mut w = writer.take().unwrap(); - flush_seed_writers(w.as_mut(), &mut seed_writers).await?; - let (num_rows, data_file) = w.finish().await?; - info!(target: TRACE_FILE_AUDIT, mode=AUDIT_MODE_CREATE, r#type=AUDIT_TYPE_DATA, path = &data_file.path); - debug_assert_eq!(num_rows, num_rows_in_current_file); - bytes_completed += data_file.file_size_bytes.get().map_or(0, |s| s.get()); - rows_completed += num_rows as u64; - files_written += 1; - let last_fragment = fragments.last_mut().unwrap(); - last_fragment.physical_rows = Some(num_rows as usize); - last_fragment.files.push(data_file); - // Notify after pushing the data file so it's tracked for cleanup - // if the callback fails. - params.progress.complete(fragments.last().unwrap()).await?; + let current_file_bytes = writer + .as_mut() + .ok_or_else(|| Error::internal("Writer disappeared after writing a batch"))? + .tell() + .await?; if let Some(cb) = ¶ms.write_progress { cb.call(WriteStats { - bytes_written: bytes_completed, - rows_written: rows_completed, + bytes_written: bytes_completed + current_file_bytes, + rows_written: rows_completed + num_rows_in_current_file as u64, files_written, }); } - num_rows_in_current_file = 0; + + let reached_row_limit = if has_file_row_counts { + reached_planned_file_boundary + } else { + num_rows_in_current_file >= params.max_rows_per_file as u32 + }; + let reached_byte_limit = current_file_bytes >= params.max_bytes_per_file as u64; + + if reached_row_limit || reached_byte_limit { + if has_file_row_counts { + if reached_planned_file_boundary { + rows_remaining_in_planned_file = file_row_counts + .as_mut() + .and_then(VecDeque::pop_front); + } else { + // A byte-driven close is an extra physical boundary. Rebalance all + // unwritten rows under the original maximum instead of preserving a + // tiny abandoned remainder or rolling it into an oversized tail. + let total_remaining = planned_rows_remaining.ok_or_else(|| { + Error::internal("Writer lost the planned row count total") + })?; + let max_rows_per_file = max_planned_file_rows.ok_or_else(|| { + Error::internal( + "Writer cannot replan byte-limited files without a maximum planned row count", + ) + })?; + let mut replanned_counts = + balanced_row_counts(total_remaining, max_rows_per_file); + rows_remaining_in_planned_file = replanned_counts.pop_front(); + file_row_counts = Some(replanned_counts); + } + } + + let mut w = writer.take().ok_or_else(|| { + Error::internal("Writer disappeared before completing a file") + })?; + flush_seed_writers(w.as_mut(), &mut seed_writers).await?; + let (num_rows, data_file) = w.finish().await?; + info!(target: TRACE_FILE_AUDIT, mode=AUDIT_MODE_CREATE, r#type=AUDIT_TYPE_DATA, path = &data_file.path); + debug_assert_eq!(num_rows, num_rows_in_current_file); + bytes_completed += data_file.file_size_bytes.get().map_or(0, |s| s.get()); + rows_completed += num_rows as u64; + files_written += 1; + let last_fragment = fragments.last_mut().ok_or_else(|| { + Error::internal("Writer completed a file without a pending fragment") + })?; + last_fragment.physical_rows = Some(num_rows as usize); + last_fragment.files.push(data_file); + // Notify after pushing the data file so it's tracked for cleanup + // if the callback fails. + let completed_fragment = fragments.last().ok_or_else(|| { + Error::internal("Writer completed a file without a fragment") + })?; + params.progress.complete(completed_fragment).await?; + if let Some(cb) = ¶ms.write_progress { + cb.call(WriteStats { + bytes_written: bytes_completed, + rows_written: rows_completed, + files_written, + }); + } + num_rows_in_current_file = 0; + } } } + + if has_file_row_counts && planned_rows_remaining != Some(0) { + return Err(Error::internal(format!( + "Writer input ended with {} planned rows remaining", + planned_rows_remaining.unwrap_or_default() + ))); + } Ok(()) } .await; @@ -1289,6 +1436,33 @@ pub async fn write_fragments_internal( data: SendableRecordBatchStream, params: WriteParams, target_bases_info: Option>, +) -> Result<(Vec, Schema)> { + write_fragments_internal_with_file_row_counts( + storage_version, + dataset, + object_store, + base_dir, + schema, + data, + params, + target_bases_info, + None, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +#[instrument(level = "debug", skip_all)] +pub(crate) async fn write_fragments_internal_with_file_row_counts( + storage_version: ConcreteFileVersion, + dataset: Option<&Dataset>, + object_store: Arc, + base_dir: &Path, + schema: Schema, + data: SendableRecordBatchStream, + params: WriteParams, + target_bases_info: Option>, + file_row_counts: Option>, ) -> Result<(Vec, Schema)> { let mut params = params; let adapter = SchemaAdapter::new(data.schema()); @@ -1318,6 +1492,7 @@ pub async fn write_fragments_internal( data, params, target_bases_info, + file_row_counts, ) .await } @@ -1874,7 +2049,9 @@ mod tests { #[cfg(windows)] use std::path::{Component, Prefix}; - use arrow_array::{Int32Array, RecordBatchIterator, RecordBatchReader, StructArray}; + use arrow_array::{ + Int32Array, LargeBinaryArray, RecordBatchIterator, RecordBatchReader, StructArray, + }; use arrow_schema::{DataType, Field as ArrowField, Fields, Schema as ArrowSchema}; use datafusion::{error::DataFusionError, physical_plan::stream::RecordBatchStreamAdapter}; use datafusion_physical_plan::RecordBatchStream; @@ -1886,6 +2063,7 @@ mod tests { use lance_io::object_store::StorageOptionsAccessor; use lance_io::traits::Reader; use lance_table::format::BasePath; + use rstest::rstest; async fn open_v2_1_test_writer( object_store: Arc, @@ -2103,6 +2281,138 @@ mod tests { assert_eq!(fragments.len(), 2); } + #[rstest] + #[case::rebalance_pending_remainder( + &[9_999, 10_001], + &[10_000, 10_000], + 2 * 1024, + &[9_999, 5_001, 5_000] + )] + #[case::replan_pending_boundary( + &[9_999, 1, 10_000, 10_000], + &[20_000, 10_000], + 100 * 1024, + &[9_999, 10_001, 10_000] + )] + #[tokio::test] + async fn test_planned_file_boundary_with_byte_limit( + #[case] input_batch_sizes: &[usize], + #[case] file_row_counts: &[usize], + #[case] max_bytes_per_file: usize, + #[case] expected_file_rows: &[usize], + ) { + let value = vec![0_u8; 1024]; + let arrow_schema = Arc::new(ArrowSchema::new(vec![ArrowField::new( + "a", + DataType::LargeBinary, + false, + )])); + let total_rows = input_batch_sizes.iter().sum::(); + let data = RecordBatch::try_new( + arrow_schema.clone(), + vec![Arc::new(LargeBinaryArray::from_iter_values( + (0..total_rows).map(|_| value.as_slice()), + ))], + ) + .unwrap(); + let mut offset = 0; + let batches = input_batch_sizes + .iter() + .map(|&batch_rows| { + let batch = data.slice(offset, batch_rows); + offset += batch_rows; + Ok::<_, DataFusionError>(batch) + }) + .collect::>(); + let stream = + RecordBatchStreamAdapter::new(arrow_schema.clone(), futures::stream::iter(batches)); + let schema = Schema::try_from(arrow_schema.as_ref()).unwrap(); + let object_store = Arc::new(ObjectStore::memory()); + + let (fragments, _) = write_fragments_internal_with_file_row_counts( + ConcreteFileVersion::V2_0, + None, + object_store, + &Path::from("planned_byte_boundary"), + schema, + Box::pin(stream), + WriteParams { + max_rows_per_file: file_row_counts[0], + max_bytes_per_file, + mode: WriteMode::Create, + ..Default::default() + }, + None, + Some(file_row_counts.to_vec()), + ) + .await + .unwrap(); + + assert_eq!( + fragments + .iter() + .map(|fragment| fragment.physical_rows.unwrap()) + .collect::>(), + expected_file_rows + ); + } + + #[tokio::test] + async fn test_repeated_byte_closes_rebalance_planned_rows() { + let large_value = vec![0_u8; 16 * 1024 * 1024]; + let mut values = Vec::with_capacity(15); + values.extend(std::iter::repeat_n(large_value.as_slice(), 2)); + values.extend(std::iter::repeat_n(&[][..], 13)); + let arrow_schema = Arc::new(ArrowSchema::new(vec![ArrowField::new( + "a", + DataType::LargeBinary, + false, + )])); + let data = RecordBatch::try_new( + arrow_schema.clone(), + vec![Arc::new(LargeBinaryArray::from_iter_values(values))], + ) + .unwrap(); + let input_batch_sizes = [1, 1, 3, 5, 5]; + let mut offset = 0; + let batches = input_batch_sizes.map(|batch_rows| { + let batch = data.slice(offset, batch_rows); + offset += batch_rows; + Ok::<_, DataFusionError>(batch) + }); + let stream = + RecordBatchStreamAdapter::new(arrow_schema.clone(), futures::stream::iter(batches)); + let schema = Schema::try_from(arrow_schema.as_ref()).unwrap(); + let object_store = Arc::new(ObjectStore::memory()); + + let (fragments, _) = write_fragments_internal_with_file_row_counts( + ConcreteFileVersion::V2_0, + None, + object_store, + &Path::from("repeated_planned_byte_boundaries"), + schema, + Box::pin(stream), + WriteParams { + max_rows_per_file: 5, + max_bytes_per_file: 100 * 1024, + mode: WriteMode::Create, + ..Default::default() + }, + None, + Some(vec![5, 5, 5]), + ) + .await + .unwrap(); + + assert_eq!( + fragments + .iter() + .map(|fragment| fragment.physical_rows.unwrap()) + .collect::>(), + [1, 1, 5, 4, 4] + ); + } + #[tokio::test] async fn test_max_rows_per_file() { let reader_to_frags = |data_reader: Box| { @@ -3959,6 +4269,7 @@ mod tests { WriteParams::default(), None, Vec::new(), + None, ) .await; @@ -4020,6 +4331,7 @@ mod tests { }, None, Vec::new(), + None, ) .await; @@ -4248,6 +4560,7 @@ mod tests { }, Some(target_bases), vec![], + None, ) .await;