diff --git a/native/core/src/execution/operators/shuffle_scan.rs b/native/core/src/execution/operators/shuffle_scan.rs index f3209b0c1b0..bbdbce397b3 100644 --- a/native/core/src/execution/operators/shuffle_scan.rs +++ b/native/core/src/execution/operators/shuffle_scan.rs @@ -349,6 +349,7 @@ mod tests { use crate::execution::shuffle::{CompressionCodec, ShuffleBlockWriter}; use arrow::array::{Int32Array, StringArray}; use arrow::datatypes::{DataType, Field, Schema}; + use arrow::ipc::writer::CompressionContext; use arrow::record_batch::RecordBatch; use datafusion::physical_plan::metrics::Time; use std::io::Cursor; @@ -377,7 +378,14 @@ mod tests { ShuffleBlockWriter::try_new(&batch.schema(), CompressionCodec::Zstd(1)).unwrap(); let mut buf = Cursor::new(Vec::new()); let ipc_time = Time::new(); - writer.write_batch(&batch, &mut buf, &ipc_time).unwrap(); + writer + .write_batch( + &batch, + &mut buf, + &mut CompressionContext::default(), + &ipc_time, + ) + .unwrap(); // Read back (skip 16-byte header: 8 compressed_length + 8 field_count) let bytes = buf.into_inner(); @@ -443,7 +451,12 @@ mod tests { let mut buf = Cursor::new(Vec::new()); let ipc_time = Time::new(); writer - .write_batch(&dict_batch, &mut buf, &ipc_time) + .write_batch( + &dict_batch, + &mut buf, + &mut CompressionContext::default(), + &ipc_time, + ) .unwrap(); let bytes = buf.into_inner(); let body = &bytes[16..]; diff --git a/native/shuffle/benches/shuffle_writer.rs b/native/shuffle/benches/shuffle_writer.rs index 2f1f193cdae..ec4fcd5373a 100644 --- a/native/shuffle/benches/shuffle_writer.rs +++ b/native/shuffle/benches/shuffle_writer.rs @@ -18,6 +18,7 @@ use arrow::array::builder::{Date32Builder, Decimal128Builder, Int32Builder}; use arrow::array::{builder::StringBuilder, Array, Int32Array, RecordBatch}; use arrow::datatypes::{DataType, Field, Schema}; +use arrow::ipc::writer::CompressionContext; use arrow::row::{RowConverter, SortField}; use criterion::{criterion_group, criterion_main, Criterion}; use datafusion::datasource::memory::MemorySourceConfig; @@ -53,10 +54,12 @@ fn criterion_benchmark(c: &mut Criterion) { let ipc_time = Time::default(); let w = ShuffleBlockWriter::try_new(&batch.schema(), compression_codec.clone()).unwrap(); + let mut compression_context = CompressionContext::default(); b.iter(|| { buffer.clear(); let mut cursor = Cursor::new(&mut buffer); - w.write_batch(&batch, &mut cursor, &ipc_time).unwrap(); + w.write_batch(&batch, &mut cursor, &mut compression_context, &ipc_time) + .unwrap(); }); }); } @@ -214,12 +217,15 @@ fn schema_encoding_benchmark(c: &mut Criterion) { let writer = ShuffleBlockWriter::try_new(batch.schema().as_ref(), CompressionCodec::None).unwrap(); let ipc_time = Time::default(); + let mut compression_context = CompressionContext::default(); group.bench_function(format!("write_batch ({name} schema)"), |b| { let mut buffer = vec![]; b.iter(|| { buffer.clear(); let mut cursor = Cursor::new(&mut buffer); - writer.write_batch(&batch, &mut cursor, &ipc_time).unwrap(); + writer + .write_batch(&batch, &mut cursor, &mut compression_context, &ipc_time) + .unwrap(); }); }); } diff --git a/native/shuffle/src/shuffle_writer.rs b/native/shuffle/src/shuffle_writer.rs index 12025fe80d3..66894188680 100644 --- a/native/shuffle/src/shuffle_writer.rs +++ b/native/shuffle/src/shuffle_writer.rs @@ -273,6 +273,7 @@ mod test { use crate::{read_ipc_compressed, ShuffleBlockWriter}; use arrow::array::{Array, Int64Array, StringArray, StringBuilder}; use arrow::datatypes::{DataType, Field, Schema}; + use arrow::ipc::writer::CompressionContext; use arrow::record_batch::RecordBatch; use arrow::row::{RowConverter, SortField}; use datafusion::datasource::memory::MemorySourceConfig; @@ -302,8 +303,14 @@ mod test { let mut cursor = Cursor::new(&mut output); let writer = ShuffleBlockWriter::try_new(batch.schema().as_ref(), codec.clone()).unwrap(); + let mut compression_context = CompressionContext::default(); let length = writer - .write_batch(&batch, &mut cursor, &Time::default()) + .write_batch( + &batch, + &mut cursor, + &mut compression_context, + &Time::default(), + ) .unwrap(); assert_eq!(length, output.len()); @@ -342,8 +349,14 @@ mod test { let mut output = vec![]; let mut cursor = Cursor::new(&mut output); let writer = ShuffleBlockWriter::try_new(schema.as_ref(), codec.clone()).unwrap(); + let mut compression_context = CompressionContext::default(); writer - .write_batch(&batch, &mut cursor, &Time::default()) + .write_batch( + &batch, + &mut cursor, + &mut compression_context, + &Time::default(), + ) .unwrap(); let batch2 = read_ipc_compressed(&output[16..]).unwrap(); diff --git a/native/shuffle/src/spark_unsafe/row.rs b/native/shuffle/src/spark_unsafe/row.rs index e0ebf8e7d01..77e44024fbf 100644 --- a/native/shuffle/src/spark_unsafe/row.rs +++ b/native/shuffle/src/spark_unsafe/row.rs @@ -38,6 +38,7 @@ use arrow::array::{ use arrow::compute::cast; use arrow::datatypes::{DataType, Field, Schema, TimeUnit}; use arrow::error::ArrowError; +use arrow::ipc::writer::CompressionContext; use datafusion::physical_plan::metrics::Time; use datafusion_comet_jni_bridge::errors::CometError; use jni::sys::{jint, jlong}; @@ -1387,6 +1388,7 @@ pub fn process_sorted_row_partition( // Single ipc_time accumulates encode + compression time across all batches. let ipc_time = Time::default(); + let mut compression_context = CompressionContext::default(); while current_row < row_num { let n = std::cmp::min(batch_size, row_num - current_row); @@ -1420,7 +1422,8 @@ pub fn process_sorted_row_partition( let mut cursor = Cursor::new(&mut frozen); let block_writer = ShuffleBlockWriter::try_new(batch.schema().as_ref(), codec.clone())?; - written += block_writer.write_batch(&batch, &mut cursor, &ipc_time)?; + written += + block_writer.write_batch(&batch, &mut cursor, &mut compression_context, &ipc_time)?; if let Some(checksum) = &mut current_checksum { checksum.update(&mut cursor)?; diff --git a/native/shuffle/src/writers/buf_batch_writer.rs b/native/shuffle/src/writers/buf_batch_writer.rs index 72db86529d8..55d88a4ba48 100644 --- a/native/shuffle/src/writers/buf_batch_writer.rs +++ b/native/shuffle/src/writers/buf_batch_writer.rs @@ -18,6 +18,7 @@ use super::ShuffleBlockWriter; use arrow::array::RecordBatch; use arrow::compute::kernels::coalesce::BatchCoalescer; +use arrow::ipc::writer::CompressionContext; use datafusion::physical_plan::metrics::Time; use std::borrow::Borrow; use std::io::{Cursor, Seek, SeekFrom, Write}; @@ -37,6 +38,7 @@ pub(crate) struct BufBatchWriter, W: Write> { writer: W, buffer: Vec, buffer_max_size: usize, + compression_context: CompressionContext, /// Coalesces small batches into target_batch_size before serialization. /// Lazily initialized on first write to capture the schema. coalescer: Option, @@ -56,6 +58,7 @@ impl, W: Write> BufBatchWriter { writer, buffer: vec![], buffer_max_size, + compression_context: CompressionContext::default(), coalescer: None, batch_size, } @@ -107,10 +110,12 @@ impl, W: Write> BufBatchWriter { ) -> datafusion::common::Result { let mut cursor = Cursor::new(&mut self.buffer); cursor.seek(SeekFrom::End(0))?; - let bytes_written = - self.shuffle_block_writer - .borrow() - .write_batch(batch, &mut cursor, encode_time)?; + let bytes_written = self.shuffle_block_writer.borrow().write_batch( + batch, + &mut cursor, + &mut self.compression_context, + encode_time, + )?; let pos = cursor.position(); if pos >= self.buffer_max_size as u64 { let mut write_timer = write_time.timer(); diff --git a/native/shuffle/src/writers/shuffle_block_writer.rs b/native/shuffle/src/writers/shuffle_block_writer.rs index 7b6846b3ba7..65e7bf4570f 100644 --- a/native/shuffle/src/writers/shuffle_block_writer.rs +++ b/native/shuffle/src/writers/shuffle_block_writer.rs @@ -142,7 +142,12 @@ impl ShuffleBlockWriter { } /// Serialize `batch` as a standalone Arrow IPC stream into `out`. - fn encode_ipc_stream(&self, batch: &RecordBatch, out: &mut W) -> Result<()> { + fn encode_ipc_stream( + &self, + batch: &RecordBatch, + out: &mut W, + compression_context: &mut CompressionContext, + ) -> Result<()> { let schema_message = match &self.schema_encoding { SchemaEncoding::Fallback(schema) => { // Dictionary encoding requires the schema and record batch to share a dictionary @@ -159,12 +164,11 @@ impl ShuffleBlockWriter { // Fast path: reuse the pre-encoded schema message and write the record batch manually. let data_gen = IpcDataGenerator::default(); let mut dictionary_tracker = DictionaryTracker::new(true); - let mut compression_context = CompressionContext::default(); let (encoded_dictionaries, encoded_batch) = data_gen.encode( batch, &mut dictionary_tracker, &self.write_options, - &mut compression_context, + compression_context, )?; debug_assert!(encoded_dictionaries.is_empty()); @@ -180,6 +184,7 @@ impl ShuffleBlockWriter { &self, batch: &RecordBatch, output: &mut W, + compression_context: &mut CompressionContext, ipc_time: &Time, ) -> Result { if batch.num_rows() == 0 { @@ -194,25 +199,25 @@ impl ShuffleBlockWriter { match &self.codec { CompressionCodec::None => { - self.encode_ipc_stream(batch, output)?; + self.encode_ipc_stream(batch, output, compression_context)?; } CompressionCodec::Lz4Frame => { let mut wtr = lz4_flex::frame::FrameEncoder::new(&mut *output); - self.encode_ipc_stream(batch, &mut wtr)?; + self.encode_ipc_stream(batch, &mut wtr, compression_context)?; wtr.finish().map_err(|e| { DataFusionError::Execution(format!("lz4 compression error: {e}")) })?; } CompressionCodec::Snappy => { let mut wtr = snap::write::FrameEncoder::new(&mut *output); - self.encode_ipc_stream(batch, &mut wtr)?; + self.encode_ipc_stream(batch, &mut wtr, compression_context)?; wtr.into_inner().map_err(|e| { DataFusionError::Execution(format!("snappy compression error: {e}")) })?; } CompressionCodec::Zstd(level) => { let mut encoder = zstd::Encoder::new(&mut *output, *level)?; - self.encode_ipc_stream(batch, &mut encoder)?; + self.encode_ipc_stream(batch, &mut encoder, compression_context)?; encoder.finish()?; } }