diff --git a/native/Cargo.lock b/native/Cargo.lock index fabb56e2ec8..712238849da 100644 --- a/native/Cargo.lock +++ b/native/Cargo.lock @@ -2065,6 +2065,7 @@ name = "datafusion-comet-shuffle" version = "1.1.0" dependencies = [ "arrow", + "arrow-data", "arrow-select", "async-trait", "bytes", diff --git a/native/Cargo.toml b/native/Cargo.toml index 6cf35f8676d..a21658884b2 100644 --- a/native/Cargo.toml +++ b/native/Cargo.toml @@ -38,6 +38,7 @@ rust-version = "1.94.0" [workspace.dependencies] arrow = { version = "59.2.0", features = ["prettyprint", "ffi", "chrono-tz"] } +arrow-data = { version = "59.2.0" } arrow-select = { version = "59.2.0" } async-trait = { version = "0.1" } bytes = { version = "1.11.1" } diff --git a/native/shuffle/Cargo.toml b/native/shuffle/Cargo.toml index 9504834ef4a..f0ed22ad730 100644 --- a/native/shuffle/Cargo.toml +++ b/native/shuffle/Cargo.toml @@ -30,6 +30,7 @@ publish = false [dependencies] arrow = { workspace = true } +arrow-data = { workspace = true } arrow-select = { workspace = true } async-trait = { workspace = true } bytes = { workspace = true } diff --git a/native/shuffle/benches/shuffle_reader.rs b/native/shuffle/benches/shuffle_reader.rs index 47903d002f8..43d4d44b999 100644 --- a/native/shuffle/benches/shuffle_reader.rs +++ b/native/shuffle/benches/shuffle_reader.rs @@ -16,15 +16,17 @@ // under the License. //! Shuffle read benchmarks: the per-block schema parse measured against a full block decode, -//! across column counts and rows per block. +//! across column counts, rows per block, the default codec and no codec, and a dictionary-encoded +//! string column. -use arrow::array::{Int64Array, RecordBatch, StringArray}; -use arrow::datatypes::{DataType, Field, Schema, SchemaRef}; +use arrow::array::{ArrayRef, DictionaryArray, Int64Array, RecordBatch, StringArray}; +use arrow::datatypes::{DataType, Field, Int32Type, Schema, SchemaRef}; use arrow::ipc::reader::StreamReader; use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; use datafusion::physical_plan::metrics::Time; use datafusion_comet_shuffle::{ - read_ipc_compressed, CompressionCodec, ShuffleBlockWriter, ShuffleCodecContext, + read_ipc_compressed, read_ipc_compressed_validated, reset_schema_cache, CompressionCodec, + ShuffleBlockWriter, ShuffleCodecContext, }; use std::hint::black_box; use std::io::Cursor; @@ -33,15 +35,30 @@ use std::sync::Arc; /// 8-byte compressed length plus 8-byte field count; `read_ipc_compressed` expects what follows. const BLOCK_HEADER_LEN: usize = 16; -/// Alternating `Int64` and `Utf8`. -fn schema_of(num_columns: usize) -> SchemaRef { +/// How the odd columns hold their strings. +#[derive(Clone, Copy)] +enum Strings { + Plain, + /// `Dictionary(Int32, Utf8)`: the block carries a dictionary batch before its record batch, + /// as the JVM columnar shuffle writes for strings. + Dictionary, +} + +/// Alternating `Int64` and string columns. +fn schema_of(num_columns: usize, strings: Strings) -> SchemaRef { Arc::new(Schema::new( (0..num_columns) .map(|i| { let data_type = if i % 2 == 0 { DataType::Int64 } else { - DataType::Utf8 + match strings { + Strings::Plain => DataType::Utf8, + Strings::Dictionary => DataType::Dictionary( + Box::new(DataType::Int32), + Box::new(DataType::Utf8), + ), + } }; Field::new(format!("column_{i}"), data_type, false) }) @@ -49,8 +66,8 @@ fn schema_of(num_columns: usize) -> SchemaRef { )) } -fn batch_of(num_columns: usize, num_rows: usize) -> RecordBatch { - let schema = schema_of(num_columns); +fn batch_of(num_columns: usize, num_rows: usize, strings: Strings) -> RecordBatch { + let schema = schema_of(num_columns, strings); let columns = (0..num_columns) .map(|i| { if i % 2 == 0 { @@ -58,13 +75,26 @@ fn batch_of(num_columns: usize, num_rows: usize) -> RecordBatch { (0..num_rows) .map(|r| Some(r as i64)) .collect::(), - ) as arrow::array::ArrayRef + ) as ArrayRef } else { - Arc::new( - (0..num_rows) - .map(|r| Some(format!("value_{r}"))) - .collect::(), - ) as arrow::array::ArrayRef + match strings { + Strings::Plain => Arc::new( + (0..num_rows) + .map(|r| Some(format!("value_{r}"))) + .collect::(), + ) as ArrayRef, + // a small dictionary that every row's key points into + Strings::Dictionary => { + let values: Vec = + (0..num_rows).map(|r| format!("value_{}", r % 16)).collect(); + Arc::new( + values + .iter() + .map(String::as_str) + .collect::>(), + ) as ArrayRef + } + } } }) .collect::>(); @@ -86,25 +116,43 @@ fn encode_block(batch: &RecordBatch, codec: CompressionCodec) -> Vec { fn criterion_benchmark(c: &mut Criterion) { let mut group = c.benchmark_group("shuffle_reader"); - // rows per block shrink as partition count rises, so the small cases stand in for wide shuffles - for num_columns in [5usize, 50] { - for num_rows in [64usize, 512, 8192] { - let batch = batch_of(num_columns, num_rows); - let uncompressed = encode_block(&batch, CompressionCodec::None); - - let id = format!("{num_columns}col_{num_rows}row"); + // Lz4Frame is the default codec; None isolates the decode from decompression. + for (codec_name, codec) in [ + ("none", CompressionCodec::None), + ("lz4", CompressionCodec::Lz4Frame), + ] { + // rows per block shrink as partition count rises, so the small cases stand in for wide + // shuffles + for num_columns in [5usize, 50] { + for num_rows in [64usize, 512, 8192] { + let batch = batch_of(num_columns, num_rows, Strings::Plain); + let block = encode_block(&batch, codec.clone()); + let id = format!("{codec_name}/{num_columns}col_{num_rows}row"); + bench_block(&mut group, &id, &block); + } + } - // full decode: schema parse plus record batch - group.bench_with_input( - BenchmarkId::new("decode_block", &id), - &uncompressed, - |b, block| b.iter(|| black_box(read_ipc_compressed(black_box(block)).unwrap())), - ); + // the dictionary batch before every record batch, at a narrow and a wide block + for num_rows in [64usize, 8192] { + let batch = batch_of(5, num_rows, Strings::Dictionary); + let block = encode_block(&batch, codec.clone()); + let id = format!("{codec_name}/5col_{num_rows}row_dict"); + bench_block(&mut group, &id, &block); + } + } - // schema parse alone: `try_new` stops before the record batch. Skips the codec tag. + // schema parse alone: `try_new` stops before the record batch. Skips the codec tag, so it + // only applies to uncompressed blocks. A control arm: this change does not touch it. + for num_columns in [5usize, 50] { + for num_rows in [64usize, 512, 8192] { + let batch = batch_of(num_columns, num_rows, Strings::Plain); + let block = encode_block(&batch, CompressionCodec::None); group.bench_with_input( - BenchmarkId::new("parse_schema_only", &id), - &uncompressed, + BenchmarkId::new( + "parse_schema_only", + format!("none/{num_columns}col_{num_rows}row"), + ), + &block, |b, block| { b.iter(|| { let mut ipc = &black_box(block)[4..]; @@ -118,5 +166,35 @@ fn criterion_benchmark(c: &mut Criterion) { group.finish(); } +fn bench_block( + group: &mut criterion::BenchmarkGroup<'_, criterion::measurement::WallTime>, + id: &str, + block: &[u8], +) { + // full decode with the schema served from the cache after the first iteration + group.bench_with_input(BenchmarkId::new("decode_block", id), block, |b, block| { + b.iter(|| black_box(read_ipc_compressed(black_box(block)).unwrap())) + }); + + // the remote entry point: the same decode with array validation on + group.bench_with_input( + BenchmarkId::new("decode_block_validated", id), + block, + |b, block| b.iter(|| black_box(read_ipc_compressed_validated(black_box(block)).unwrap())), + ); + + // same decode with the cache cleared each iteration, so drift moves both arms together + group.bench_with_input( + BenchmarkId::new("decode_block_uncached", id), + block, + |b, block| { + b.iter(|| { + reset_schema_cache(); + black_box(read_ipc_compressed(black_box(block)).unwrap()) + }) + }, + ); +} + criterion_group!(benches, criterion_benchmark); criterion_main!(benches); diff --git a/native/shuffle/src/ipc.rs b/native/shuffle/src/ipc.rs index 97890f50148..7e54367a330 100644 --- a/native/shuffle/src/ipc.rs +++ b/native/shuffle/src/ipc.rs @@ -15,11 +15,19 @@ // specific language governing permissions and limitations // under the License. -use arrow::array::RecordBatch; -use arrow::ipc::reader::StreamReader; +use arrow::array::{ArrayRef, RecordBatch}; +use arrow::buffer::{Buffer, MutableBuffer}; +use arrow::datatypes::SchemaRef; +use arrow::ipc::convert::fb_to_schema; +use arrow::ipc::reader::{read_dictionary_impl, RecordBatchDecoder}; +use arrow::ipc::{root_as_message, Message, MessageHeader}; +use arrow_data::UnsafeFlag; use datafusion::common::DataFusionError; use datafusion::error::Result; +use std::cell::RefCell; +use std::collections::HashMap; use std::io::{Error, ErrorKind, Read}; +use std::sync::Arc; /// Decode trusted local Comet output without revalidating every Arrow array value or offset. pub fn read_ipc_compressed(bytes: &[u8]) -> Result { @@ -31,24 +39,142 @@ pub fn read_ipc_compressed_validated(bytes: &[u8]) -> Result { read_ipc_compressed_impl(bytes, true) } +/// Arrow IPC continuation marker introducing a message length. +const CONTINUATION_MARKER: [u8; 4] = [0xff; 4]; + +/// Distinct schemas cached per thread. More than one because a reduce task can interleave blocks +/// from several shuffles, and a single entry would thrash. +const SCHEMA_CACHE_CAPACITY: usize = 4; + +/// Metadata scratch larger than this is released after the block rather than kept for the thread. +/// Real metadata is a few KiB even for wide schemas; only a corrupt length gets anywhere near. +const SCRATCH_RETAIN_LIMIT: usize = 1 << 20; + +/// Per-thread decoder state. +/// +/// Every block is a complete IPC stream that opens with a schema message. `ShuffleBlockWriter` +/// encodes that message once and writes it verbatim into every block, so consecutive blocks carry +/// byte-identical schema messages. The cache is keyed on those bytes: a hit is one memcmp, and +/// the schema message is neither verified nor parsed. +#[derive(Default)] +struct DecoderState { + /// Parsed schemas keyed on the raw schema message, most recently used first. + schemas: Vec<(Box<[u8]>, SchemaRef)>, + /// Message metadata read from a decompressor lands here, so it is not reallocated per block. + scratch: Vec, + #[cfg(test)] + stats: SchemaCacheStats, +} + +thread_local! { + static STATE: RefCell = RefCell::new(DecoderState::default()); +} + +/// Empties this thread's schema cache, so the next decode re-parses its schema. For benchmarks +/// and tests comparing the cold and warm paths; not part of the decode contract. +#[doc(hidden)] +pub fn reset_schema_cache() { + STATE.with_borrow_mut(|state| { + state.schemas.clear(); + #[cfg(test)] + { + state.stats = SchemaCacheStats::default(); + } + }); +} + +/// Schema cache hits and misses on this thread since the last [`reset_schema_cache`]. +#[cfg(test)] +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +struct SchemaCacheStats { + hits: usize, + misses: usize, +} + +#[cfg(test)] +fn schema_cache_stats() -> SchemaCacheStats { + STATE.with_borrow(|state| state.stats) +} + +#[cfg(test)] +fn scratch_capacity() -> usize { + STATE.with_borrow(|state| state.scratch.capacity()) +} + +fn cached_schema( + schemas: &mut [(Box<[u8]>, SchemaRef)], + schema_message: &[u8], +) -> Option { + let hit = schemas + .iter() + .position(|(message, _)| message.as_ref() == schema_message)?; + // most recently used first, so an alternating pair stays resident + if hit != 0 { + schemas.swap(0, hit); + } + Some(Arc::clone(&schemas[0].1)) +} + +fn cache_schema( + schemas: &mut Vec<(Box<[u8]>, SchemaRef)>, + schema_message: &[u8], + schema: SchemaRef, +) { + if schemas.len() == SCHEMA_CACHE_CAPACITY { + schemas.pop(); + } + schemas.insert(0, (schema_message.into(), schema)); +} + +fn decode_error(what: &str) -> DataFusionError { + DataFusionError::Execution(format!("Failed to decode batch: {what}")) +} + +fn parse_message(metadata: &[u8]) -> Result> { + root_as_message(metadata) + .map_err(|error| decode_error(&format!("unable to get root as message: {error:?}"))) +} + +fn body_length(message: &Message<'_>) -> Result { + usize::try_from(message.bodyLength()).map_err(|_| { + decode_error(&format!( + "invalid message body length: {}", + message.bodyLength() + )) + }) +} + fn read_ipc_compressed_impl(bytes: &[u8], validate: bool) -> Result { - let codec = bytes.get(..4).ok_or_else(|| { - DataFusionError::Execution("Failed to decode batch: truncated compression codec".to_owned()) - })?; + let codec = bytes + .get(..4) + .ok_or_else(|| decode_error("truncated compression codec"))?; let mut encoded = &bytes[4..]; let batch = match codec { - b"SNAP" => read_single_batch(snap::read::FrameDecoder::new(&mut encoded), validate)?, - b"LZ4_" => read_single_batch( - lz4_flex::frame::FrameDecoder::new(RequireLz4EndMark(&mut encoded)), + b"SNAP" => decode( + Streamed(snap::read::FrameDecoder::new(&mut encoded)), + validate, + )?, + b"LZ4_" => decode( + Streamed(lz4_flex::frame::FrameDecoder::new(RequireLz4EndMark( + &mut encoded, + ))), validate, )?, // The slice already implements BufRead. Adding another BufReader would let read-ahead // conceal compressed bytes left over after the decoder reaches its end marker. - b"ZSTD" => read_single_batch(zstd::Decoder::with_buffer(&mut encoded)?, validate)?, - b"NONE" => read_single_batch(&mut encoded, validate)?, + b"ZSTD" => decode( + Streamed(zstd::Decoder::with_buffer(&mut encoded)?), + validate, + )?, + // Uncompressed messages are located in place, so only bodies are copied. + b"NONE" => { + let batch = decode(Sliced::new(encoded), validate)?; + encoded = &[]; + batch + } other => { - return Err(DataFusionError::Execution(format!( - "Failed to decode batch: invalid compression codec: {other:?}" + return Err(decode_error(&format!( + "invalid compression codec: {other:?}" ))) } }; @@ -56,13 +182,291 @@ fn read_ipc_compressed_impl(bytes: &[u8], validate: bool) -> Result // the encoded source as well as the decoded IPC tail so an oversized outer frame cannot // silently swallow another native frame's bytes. if !encoded.is_empty() { - return Err(DataFusionError::Execution( - "Failed to decode batch: trailing data after compressed stream".to_owned(), - )); + return Err(decode_error("trailing data after compressed stream")); + } + Ok(batch) +} + +fn decode<'b, S: BlockSource<'b>>(source: S, validate: bool) -> Result { + STATE.with_borrow_mut(|state| { + let batch = read_single_batch(state, source, validate); + // a corrupt length can grow the scratch arbitrarily; do not pin that for the thread's life + if state.scratch.capacity() > SCRATCH_RETAIN_LIMIT { + state.scratch = Vec::new(); + } + batch + }) +} + +/// Reads one complete IPC stream holding exactly one record batch. Mirrors what +/// `arrow::ipc::reader::StreamReader` does message by message, except that the schema message is +/// served from the cache when its bytes match one already parsed. +fn read_single_batch<'b, S: BlockSource<'b>>( + state: &mut DecoderState, + mut source: S, + validate: bool, +) -> Result { + let DecoderState { + schemas, scratch, .. + } = state; + + let mut skip_validation = UnsafeFlag::new(); + if !validate { + // SAFETY: local blocks were written by this Comet version's ShuffleBlockWriter from arrays + // that were valid when encoded, the same trust the StreamReader path placed in them. + // Remote blocks keep full validation. + unsafe { skip_validation.set(true) }; + } + + let Some(metadata) = source.next_metadata(scratch)? else { + return Err(decode_error("empty IPC stream")); + }; + let schema = match cached_schema(schemas, metadata) { + Some(schema) => { + #[cfg(test)] + { + state.stats.hits += 1; + } + schema + } + None => { + #[cfg(test)] + { + state.stats.misses += 1; + } + let message = parse_message(metadata)?; + if message.header_type() != MessageHeader::Schema { + return Err(decode_error(&format!( + "expected a schema as the first message in the stream, got: {:?}", + message.header_type() + ))); + } + let schema = message + .header_as_schema() + .ok_or_else(|| decode_error("failed to parse schema from message header"))?; + let schema = Arc::new(fb_to_schema(schema)); + // A schema message has no body. Only bodiless ones are cached, so a hit never has a + // body to skip; anything else is read past as StreamReader does, without caching. + match body_length(&message)? { + 0 => cache_schema(schemas, metadata, Arc::clone(&schema)), + len => { + source.body(len)?; + } + } + schema + } + }; + + // dictionaries belong to the block that carries them, never to the cached schema + let mut dictionaries: HashMap = HashMap::new(); + let mut batch = None; + while let Some(metadata) = source.next_metadata(scratch)? { + let message = parse_message(metadata)?; + let version = message.version(); + let body_len = body_length(&message)?; + match message.header_type() { + MessageHeader::DictionaryBatch => { + let dictionary = message + .header_as_dictionary_batch() + .ok_or_else(|| decode_error("unable to read dictionary batch"))?; + let body = source.body(body_len)?; + read_dictionary_impl( + &body, + dictionary, + &schema, + &mut dictionaries, + &version, + false, + skip_validation.clone(), + )?; + } + MessageHeader::RecordBatch => { + // Each Comet frame contains one complete IPC stream with exactly one record + // batch. Stopping after that batch would skip codec footer/checksum validation + // and could silently discard further frames swallowed by a corrupt outer length + // prefix, so keep reading to the end-of-stream marker and reject a second batch. + if batch.is_some() { + return Err(decode_error("multiple record batches in one shuffle frame")); + } + let record_batch = message + .header_as_record_batch() + .ok_or_else(|| decode_error("unable to read record batch"))?; + let body = source.body(body_len)?; + batch = Some( + RecordBatchDecoder::try_new( + &body, + record_batch, + Arc::clone(&schema), + &dictionaries, + &version, + )? + .with_require_alignment(false) + .with_skip_validation(skip_validation.clone()) + .read_record_batch()?, + ); + } + MessageHeader::Schema => { + return Err(decode_error("expected a record batch, but found a schema")); + } + other => { + return Err(decode_error(&format!( + "unsupported message header type in IPC stream: '{other:?}'" + ))); + } + } } + + let batch = batch.ok_or_else(|| decode_error("empty IPC stream"))?; + source.expect_exhausted()?; Ok(batch) } +/// Where a block's IPC messages come from. Metadata is borrowed one message at a time; bodies +/// become exactly sized buffers that the decoded arrays keep. +/// +/// `'b` is the lifetime of an in-memory block, so [`Sliced`] can hand out metadata without +/// copying it; a streamed source uses `'static` and copies metadata into the caller's scratch. +trait BlockSource<'b> { + /// The next message's metadata, or `None` at the end of the stream: an explicit + /// end-of-stream marker, or a clean EOF on a message boundary, which is the legacy ending. + fn next_metadata<'a>(&mut self, scratch: &'a mut Vec) -> Result> + where + 'b: 'a; + + /// The next message's body, `len` bytes long. + fn body(&mut self, len: usize) -> Result; + + /// Errors unless every byte of the block has been consumed. + fn expect_exhausted(&mut self) -> Result<()>; +} + +/// Decodes the metadata length a message starts with, from its first four bytes and a reader for +/// four more should those be the continuation marker. `None` is the end-of-stream marker. +fn metadata_length( + first: [u8; 4], + next: impl FnOnce() -> Result<[u8; 4]>, +) -> Result> { + let length_bytes = if first == CONTINUATION_MARKER { + next()? + } else { + first + }; + match i32::from_le_bytes(length_bytes) { + 0 => Ok(None), + len => usize::try_from(len) + .map(Some) + .map_err(|_| decode_error(&format!("invalid metadata length: {len}"))), + } +} + +/// A block read through a decompressor. +struct Streamed(R); + +impl Streamed { + fn read_exact(&mut self, buffer: &mut [u8], what: &str) -> Result<()> { + self.0.read_exact(buffer).map_err(|error| { + if error.kind() == ErrorKind::UnexpectedEof { + decode_error(what) + } else { + error.into() + } + }) + } +} + +impl BlockSource<'static> for Streamed { + fn next_metadata<'a>(&mut self, scratch: &'a mut Vec) -> Result> + where + 'static: 'a, + { + let mut prefix = [0u8; 4]; + // EOF on a message boundary ends the stream; a partial length prefix does not + if self.0.read(&mut prefix[..1])? == 0 { + return Ok(None); + } + self.read_exact(&mut prefix[1..], "truncated IPC message length")?; + let Some(len) = metadata_length(prefix, || { + let mut bytes = [0u8; 4]; + self.read_exact(&mut bytes, "truncated IPC message length")?; + Ok(bytes) + })? + else { + return Ok(None); + }; + scratch.resize(len, 0); + self.read_exact(scratch, "truncated IPC metadata")?; + Ok(Some(scratch.as_slice())) + } + + fn body(&mut self, len: usize) -> Result { + let mut body = MutableBuffer::from_len_zeroed(len); + self.read_exact(&mut body, "truncated IPC body")?; + Ok(body.into()) + } + + fn expect_exhausted(&mut self) -> Result<()> { + if self.0.read(&mut [0])? != 0 { + return Err(decode_error("trailing data after IPC stream")); + } + Ok(()) + } +} + +/// An uncompressed block, walked in place. +struct Sliced<'b> { + block: &'b [u8], + offset: usize, +} + +impl<'b> Sliced<'b> { + fn new(block: &'b [u8]) -> Self { + Self { block, offset: 0 } + } + + fn take(&mut self, len: usize, what: &str) -> Result<&'b [u8]> { + let end = self + .offset + .checked_add(len) + .filter(|end| *end <= self.block.len()) + .ok_or_else(|| decode_error(what))?; + let bytes = &self.block[self.offset..end]; + self.offset = end; + Ok(bytes) + } +} + +impl<'b> BlockSource<'b> for Sliced<'b> { + fn next_metadata<'a>(&mut self, _scratch: &'a mut Vec) -> Result> + where + 'b: 'a, + { + if self.offset == self.block.len() { + return Ok(None); + } + let first = self.take(4, "truncated IPC message length")?; + let Some(len) = metadata_length(first.try_into().expect("four bytes"), || { + let bytes = self.take(4, "truncated IPC message length")?; + Ok(bytes.try_into().expect("four bytes")) + })? + else { + return Ok(None); + }; + Ok(Some(self.take(len, "truncated IPC metadata")?)) + } + + fn body(&mut self, len: usize) -> Result { + // an exactly sized copy, with no zero fill before it + Ok(Buffer::from(self.take(len, "truncated IPC body")?)) + } + + fn expect_exhausted(&mut self) -> Result<()> { + if self.offset != self.block.len() { + return Err(decode_error("trailing data after IPC stream")); + } + Ok(()) + } +} + // lz4_flex treats physical EOF (including a partial block header) as a clean end of frame. // Comet always writes an explicit LZ4 EndMark, so a decoder trying to read past the supplied // bytes has encountered a truncated frame. InvalidData is deliberate: UnexpectedEof is swallowed @@ -83,44 +487,27 @@ impl Read for RequireLz4EndMark { } } -fn read_single_batch(input: R, validate: bool) -> Result { - let reader = StreamReader::try_new(input, None)?; - let mut reader = if validate { - // Remote data must not escape as unchecked arrays and fail later in a native operator. - reader - } else { - // Preserve the existing local-shuffle fast path for trusted Comet-written arrays. - unsafe { reader.with_skip_validation(true) } - }; - let batch = reader.next().transpose()?.ok_or_else(|| { - DataFusionError::Execution("Failed to decode batch: empty IPC stream".to_owned()) - })?; - - // Each Comet frame contains one complete IPC stream with exactly one record batch. - // Stopping after that batch would skip codec footer/checksum validation and could silently - // discard further frames swallowed by a corrupt outer length prefix. - if reader.next().transpose()?.is_some() { - return Err(DataFusionError::Execution( - "Failed to decode batch: multiple record batches in one shuffle frame".to_owned(), - )); - } - if reader.get_mut().read(&mut [0])? != 0 { - return Err(DataFusionError::Execution( - "Failed to decode batch: trailing data after IPC stream".to_owned(), - )); - } - Ok(batch) -} - #[cfg(test)] mod tests { - use super::{read_ipc_compressed, read_ipc_compressed_validated}; - use arrow::array::{Int32Array, RecordBatch, StringArray}; - use arrow::datatypes::{DataType, Field, Schema}; + use super::{ + read_ipc_compressed, read_ipc_compressed_validated, reset_schema_cache, schema_cache_stats, + scratch_capacity, RequireLz4EndMark, SchemaCacheStats, SCHEMA_CACHE_CAPACITY, + SCRATCH_RETAIN_LIMIT, + }; + use crate::writers::rss::tests::allocations; + use arrow::array::{Array, DictionaryArray, Int32Array, RecordBatch, StringArray}; + use arrow::datatypes::{DataType, Field, Int32Type, Schema}; + use arrow::ipc::reader::StreamReader; use arrow::ipc::writer::StreamWriter; - use std::io::Write; + use std::io::{Cursor, Read, Write}; use std::sync::Arc; + const CODECS: [&[u8; 4]; 4] = [b"NONE", b"LZ4_", b"ZSTD", b"SNAP"]; + + fn stats(hits: usize, misses: usize) -> SchemaCacheStats { + SchemaCacheStats { hits, misses } + } + fn ipc_stream(batch_count: usize) -> Vec { let schema = Arc::new(Schema::new(vec![Field::new("n", DataType::Int32, false)])); let batch = RecordBatch::try_new( @@ -161,6 +548,366 @@ mod tests { bytes } + /// One batch as a complete IPC stream. + fn ipc_bytes(batch: &RecordBatch) -> Vec { + let mut payload = Vec::new(); + let mut writer = StreamWriter::try_new(&mut payload, batch.schema_ref()).unwrap(); + writer.write(batch).unwrap(); + writer.finish().unwrap(); + payload + } + + /// One encoded block, without the 16-byte Comet header. + fn block_for(batch: &RecordBatch, codec: &[u8; 4]) -> Vec { + encode(codec, &ipc_bytes(batch)) + } + + fn mixed_batch() -> RecordBatch { + let schema = Arc::new(Schema::new(vec![ + Field::new("i", DataType::Int32, true), + Field::new("s", DataType::Utf8, true), + Field::new("f", DataType::Float64, false), + ])); + RecordBatch::try_new( + schema, + vec![ + Arc::new(Int32Array::from(vec![Some(1), None, Some(3)])), + Arc::new(StringArray::from(vec![Some("a"), Some(""), None])), + Arc::new(arrow::array::Float64Array::from(vec![1.5, -0.0, 2.25])), + ], + ) + .unwrap() + } + + /// One dictionary-encoded string column; every call shares the same schema, so blocks built + /// from different values share a schema message but carry their own dictionary batch. + fn dictionary_batch(values: &[&str]) -> RecordBatch { + let dictionary: DictionaryArray = values.iter().copied().collect(); + let schema = Arc::new(Schema::new(vec![Field::new( + "d", + dictionary.data_type().clone(), + true, + )])); + RecordBatch::try_new(schema, vec![Arc::new(dictionary)]).unwrap() + } + + fn strings(batch: &RecordBatch) -> Vec { + let values = arrow::compute::cast(batch.column(0), &DataType::Utf8).unwrap(); + let values = values.as_any().downcast_ref::().unwrap(); + values.iter().map(|v| v.unwrap().to_owned()).collect() + } + + fn n_column_batch(num_columns: usize) -> RecordBatch { + let fields = (0..num_columns) + .map(|i| Field::new(format!("c{i}"), DataType::Int32, false)) + .collect::>(); + let columns = (0..num_columns) + .map(|_| Arc::new(Int32Array::from(vec![1, 2])) as arrow::array::ArrayRef) + .collect(); + RecordBatch::try_new(Arc::new(Schema::new(fields)), columns).unwrap() + } + + /// After a cold decode, the same schema is served from the cache by both entry points, and + /// the warm decodes equal the cold one on every codec. + #[test] + #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. + fn warm_decodes_hit_the_cache_and_match_the_cold_one() { + for batch in [mixed_batch(), dictionary_batch(&["x", "y", "x"])] { + for codec in CODECS { + let block = block_for(&batch, codec); + reset_schema_cache(); + + let cold = read_ipc_compressed(&block).unwrap(); + assert_eq!(schema_cache_stats(), stats(0, 1), "codec {codec:?}"); + let warm = read_ipc_compressed(&block).unwrap(); + assert_eq!(schema_cache_stats(), stats(1, 1), "codec {codec:?}"); + let validated = read_ipc_compressed_validated(&block).unwrap(); + assert_eq!(schema_cache_stats(), stats(2, 1), "codec {codec:?}"); + + for decoded in [&cold, &warm, &validated] { + assert_eq!(decoded, &batch, "codec {codec:?}"); + assert_eq!(decoded.schema(), batch.schema(), "codec {codec:?}"); + } + } + } + } + + /// Blocks that share a schema each carry their own dictionary batch. With the schema served + /// from the cache, a record batch must still be decoded against the dictionary in its own + /// block, never against a previous block's. + #[test] + #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. + fn dictionaries_are_scoped_to_their_block_under_a_cached_schema() { + let first = dictionary_batch(&["a", "b", "a"]); + let second = dictionary_batch(&["x", "y", "z"]); + assert_eq!(first.schema(), second.schema()); + + for codec in CODECS { + for validate in [false, true] { + let decode = |block: &[u8]| { + if validate { + read_ipc_compressed_validated(block).unwrap() + } else { + read_ipc_compressed(block).unwrap() + } + }; + reset_schema_cache(); + assert_eq!(strings(&decode(&block_for(&first, codec))), ["a", "b", "a"]); + assert_eq!( + strings(&decode(&block_for(&second, codec))), + ["x", "y", "z"] + ); + assert_eq!(strings(&decode(&block_for(&first, codec))), ["a", "b", "a"]); + assert_eq!( + schema_cache_stats(), + stats(2, 1), + "codec {codec:?}, validate {validate}" + ); + } + } + } + + /// Each distinct schema misses once. The cache keeps several, so blocks from two shuffles + /// can alternate without evicting each other, and only the least recently used one goes + /// when the capacity is exceeded. + #[test] + fn distinct_schemas_miss_once_and_recent_ones_stay_cached() { + let blocks: Vec> = (1..=SCHEMA_CACHE_CAPACITY + 1) + .map(|num_columns| block_for(&n_column_batch(num_columns), b"NONE")) + .collect(); + let decode = |block: &[u8]| read_ipc_compressed(block).unwrap(); + + reset_schema_cache(); + decode(&blocks[0]); + decode(&blocks[1]); + decode(&blocks[0]); + decode(&blocks[1]); + assert_eq!(schema_cache_stats(), stats(2, 2)); + + // one more schema than the capacity evicts the least recently used one + for block in &blocks { + decode(block); + } + assert_eq!(schema_cache_stats(), stats(4, 5)); + decode(&blocks[0]); + assert_eq!(schema_cache_stats(), stats(4, 6), "evicted"); + decode(&blocks[SCHEMA_CACHE_CAPACITY]); + assert_eq!(schema_cache_stats(), stats(5, 6), "most recent stays"); + } + + /// An `Int32` and a `Utf8` column, `num_rows` long. + fn wide_batch(num_rows: i32) -> RecordBatch { + let schema = Arc::new(Schema::new(vec![ + Field::new("i", DataType::Int32, false), + Field::new("s", DataType::Utf8, false), + ])); + RecordBatch::try_new( + schema, + vec![ + Arc::new((0..num_rows).collect::()), + Arc::new( + (0..num_rows) + .map(|i| Some(format!("value_{i}"))) + .collect::(), + ), + ], + ) + .unwrap() + } + + /// The reader this change replaced: a `StreamReader` per block over the decompressor, + /// exactly one batch, then the end of the stream. + fn stream_reader_decode(block: &[u8]) -> RecordBatch { + fn read(input: R) -> RecordBatch { + let mut reader = unsafe { + StreamReader::try_new(input, None) + .unwrap() + .with_skip_validation(true) + }; + let batch = reader.next().unwrap().unwrap(); + assert!(reader.next().is_none()); + batch + } + let mut encoded = &block[4..]; + match &block[..4] { + b"NONE" => read(&mut encoded), + b"LZ4_" => read(lz4_flex::frame::FrameDecoder::new(RequireLz4EndMark( + &mut encoded, + ))), + b"ZSTD" => read(zstd::Decoder::with_buffer(&mut encoded).unwrap()), + b"SNAP" => read(snap::read::FrameDecoder::new(&mut encoded)), + _ => unreachable!(), + } + } + + /// With the schema cached, a decode allocates no more than the `StreamReader` path did: + /// no more allocations, no more bytes, and no higher peak, on every codec, for a tiny block + /// and a typical one. + #[test] + #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. + fn warm_decode_allocates_no_more_than_stream_reader() { + /// (allocations, bytes requested, peak live bytes) of one decode + fn probe( + decode: impl FnOnce() -> RecordBatch, + expected: &RecordBatch, + ) -> (usize, usize, usize) { + let ((batch, (allocations, bytes)), peak) = allocations::measure(|| { + let batch = decode(); + (batch, allocations::totals()) + }); + assert_eq!(&batch, expected); + (allocations, bytes, peak) + } + + for (shape, batch) in [("3 rows", mixed_batch()), ("8192 rows", wide_batch(8192))] { + for codec in CODECS { + let block = block_for(&batch, codec); + reset_schema_cache(); + assert_eq!(read_ipc_compressed(&block).unwrap(), batch); + + let old = probe(|| stream_reader_decode(&block), &batch); + let new = probe(|| read_ipc_compressed(&block).unwrap(), &batch); + assert_eq!(schema_cache_stats(), stats(1, 1)); + + let codec = std::str::from_utf8(codec).unwrap(); + println!( + "{shape} {codec}: stream reader (allocations, bytes, peak) {old:?}, \ + cached {new:?}" + ); + assert!( + new.0 <= old.0 && new.1 <= old.1 && new.2 <= old.2, + "{shape} {codec}: {old:?} -> {new:?}" + ); + } + } + } + + /// Bodies read from a decompressor are allocated at exactly their length, as `StreamReader` + /// allocates them, so the arrays carry no growth slack and report the same memory size. + #[test] + #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. + fn decoded_arrays_report_the_same_memory_size_as_stream_reader() { + let batch = wide_batch(100_000); + let ipc = ipc_bytes(&batch); + let via_stream_reader = StreamReader::try_new(Cursor::new(&ipc), None) + .unwrap() + .next() + .unwrap() + .unwrap(); + + for codec in CODECS { + reset_schema_cache(); + // cold, then warm + for _ in 0..2 { + let decoded = read_ipc_compressed(&encode(codec, &ipc)).unwrap(); + assert_eq!(decoded, batch); + assert_eq!( + decoded.get_array_memory_size(), + via_stream_reader.get_array_memory_size(), + "codec {codec:?}" + ); + } + } + } + + /// Trailing bytes after the end-of-stream marker must stay an error with a warm cache. + #[test] + fn trailing_data_still_fails_with_a_warm_cache() { + let batch = mixed_batch(); + let payload = ipc_bytes(&batch); + + reset_schema_cache(); + assert_eq!( + read_ipc_compressed(&encode(b"NONE", &payload)).unwrap(), + batch + ); + + let mut corrupted = payload.clone(); + corrupted.extend_from_slice(&[0u8; 8]); + let error = read_ipc_compressed(&encode(b"NONE", &corrupted)).unwrap_err(); + assert!( + error.to_string().contains("trailing data"), + "unexpected error: {error}" + ); + assert_eq!(schema_cache_stats(), stats(1, 1), "failed on the warm path"); + } + + /// A block truncated inside its body must fail cold and warm. Dropping only the + /// end-of-stream marker is not truncation: a stream ending on a message boundary is valid. + #[test] + fn truncated_block_fails_with_a_warm_cache() { + let batch = mixed_batch(); + let block = block_for(&batch, b"NONE"); + reset_schema_cache(); + + // cold: the schema parses and is cached before the truncation is reached + let cut_into_body = &block[..block.len() - 24]; + assert!(read_ipc_compressed(cut_into_body).is_err()); + assert_eq!(schema_cache_stats(), stats(0, 1)); + + // warm, and the same truncation must still fail + assert_eq!(read_ipc_compressed(&block).unwrap(), batch); + assert!(read_ipc_compressed(cut_into_body).is_err()); + assert_eq!(schema_cache_stats(), stats(2, 1)); + + // dropping just the end-of-stream marker stays valid + assert_eq!( + read_ipc_compressed(&block[..block.len() - 8]).unwrap(), + batch + ); + } + + /// A partial message length after the record batch is an error on every codec, whether it + /// follows the end-of-stream marker or stands in for it. + #[test] + #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. + fn partial_length_prefix_is_an_error() { + let payload = ipc_stream(1); + for codec in CODECS { + let mut after_marker = payload.clone(); + after_marker.extend_from_slice(&[0, 0]); + let error = read_ipc_compressed(&encode(codec, &after_marker)) + .unwrap_err() + .to_string(); + assert!(error.contains("trailing data"), "{codec:?}: {error}"); + + let mut instead_of_marker = payload[..payload.len() - 8].to_vec(); + instead_of_marker.extend_from_slice(&[0, 0]); + let error = read_ipc_compressed(&encode(codec, &instead_of_marker)) + .unwrap_err() + .to_string(); + assert!( + error.contains("truncated IPC message length"), + "{codec:?}: {error}" + ); + } + } + + /// A corrupt metadata length makes the streamed reader grow its scratch before the read + /// fails. That growth must not stay pinned in the thread-local state afterwards. + #[test] + fn oversized_metadata_length_is_an_error_and_releases_the_scratch() { + let mut payload = ipc_stream(1); + // the record batch message follows the schema message: continuation marker, length, body + let schema_len = i32::from_le_bytes(payload[4..8].try_into().unwrap()) as usize; + let batch_message = 8 + schema_len; + assert_eq!(payload[batch_message..batch_message + 4], [0xff; 4]); + let forged = (2 * SCRATCH_RETAIN_LIMIT) as i32; + payload[batch_message + 4..batch_message + 8].copy_from_slice(&forged.to_le_bytes()); + + let error = read_ipc_compressed(&encode(b"LZ4_", &payload)) + .unwrap_err() + .to_string(); + assert!(error.contains("truncated IPC metadata"), "{error}"); + assert!(scratch_capacity() <= SCRATCH_RETAIN_LIMIT); + + // the in-place reader rejects the same length without allocating anything + let error = read_ipc_compressed(&encode(b"NONE", &payload)) + .unwrap_err() + .to_string(); + assert!(error.contains("truncated IPC metadata"), "{error}"); + } + #[test] fn malformed_codec_prefix_returns_error() { for prefix in [&b""[..], b"N", b"NO", b"NON", b"BAD!"] { @@ -172,7 +919,7 @@ mod tests { #[test] #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. fn empty_or_multiple_batch_stream_returns_error() { - for codec in [b"NONE", b"SNAP", b"LZ4_", b"ZSTD"] { + for codec in CODECS { for batch_count in [0, 2] { let error = read_ipc_compressed(&encode(codec, &ipc_stream(batch_count))) .unwrap_err() @@ -194,7 +941,7 @@ mod tests { fn trailing_data_after_ipc_stream_returns_error() { let mut payload = ipc_stream(1); payload.extend_from_slice(b"another shuffle frame"); - for codec in [b"NONE", b"SNAP", b"LZ4_", b"ZSTD"] { + for codec in CODECS { let error = read_ipc_compressed(&encode(codec, &payload)) .unwrap_err() .to_string(); @@ -205,7 +952,7 @@ mod tests { #[test] #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. fn trailing_data_after_compressed_stream_returns_error() { - for codec in [b"NONE", b"SNAP", b"LZ4_", b"ZSTD"] { + for codec in CODECS { let mut frame = encode(codec, &ipc_stream(1)); frame.extend_from_slice(&20_u64.to_le_bytes()); frame.extend_from_slice(b"another native frame"); @@ -224,19 +971,18 @@ mod tests { } } + /// Validation must reject a corrupt array whether the schema is parsed for this block or + /// served from the cache by an earlier valid block of the same schema. #[test] #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. - fn invalid_array_offsets_return_error() { + fn invalid_array_offsets_fail_validation_cold_and_warm() { let schema = Arc::new(Schema::new(vec![Field::new("s", DataType::Utf8, false)])); let batch = RecordBatch::try_new( Arc::clone(&schema), vec![Arc::new(StringArray::from(vec!["abc", "def"]))], ) .unwrap(); - let mut payload = Vec::new(); - let mut writer = StreamWriter::try_new(&mut payload, &schema).unwrap(); - writer.write(&batch).unwrap(); - writer.finish().unwrap(); + let mut payload = ipc_bytes(&batch); let offsets: Vec = [0_i32, 3, 6] .into_iter() @@ -250,15 +996,26 @@ mod tests { assert_eq!(positions.len(), 1); // Change [0, 3, 6] to [0, 3, 2]: the second string now has decreasing offsets. payload[positions[0] + 8..positions[0] + 12].copy_from_slice(&2_i32.to_le_bytes()); - for codec in [b"NONE", b"SNAP", b"LZ4_", b"ZSTD"] { + let valid = ipc_bytes(&batch); + for codec in CODECS { + reset_schema_cache(); assert!(read_ipc_compressed_validated(&encode(codec, &payload)).is_err()); + assert_eq!( + read_ipc_compressed_validated(&encode(codec, &valid)).unwrap(), + batch + ); + assert!( + read_ipc_compressed_validated(&encode(codec, &payload)).is_err(), + "{codec:?}: warm" + ); + assert_eq!(schema_cache_stats(), stats(2, 1), "{codec:?}"); } } #[test] #[cfg_attr(miri, ignore)] // Miri cannot call Zstd's C FFI. fn valid_single_batch_frames_decode_with_all_codecs() { - for codec in [b"NONE", b"SNAP", b"LZ4_", b"ZSTD"] { + for codec in CODECS { let frame = encode(codec, &ipc_stream(1)); let batch = read_ipc_compressed(&frame).unwrap(); let validated = read_ipc_compressed_validated(&frame).unwrap(); diff --git a/native/shuffle/src/lib.rs b/native/shuffle/src/lib.rs index 938c13e0fda..a510b7691f8 100644 --- a/native/shuffle/src/lib.rs +++ b/native/shuffle/src/lib.rs @@ -33,7 +33,7 @@ pub(crate) mod writers; pub use codec_context::ShuffleCodecContext; pub use comet_partitioning::CometPartitioning; -pub use ipc::{read_ipc_compressed, read_ipc_compressed_validated}; +pub use ipc::{read_ipc_compressed, read_ipc_compressed_validated, reset_schema_cache}; pub use remote_schema::{decode_remote_shuffle_batch, validate_remote_schema}; pub use schema_align::SchemaAlignExec; pub use shuffle_writer::{ShuffleWriterDestination, ShuffleWriterExec}; diff --git a/native/shuffle/src/writers/mod.rs b/native/shuffle/src/writers/mod.rs index fb3af2c991a..4586e46c25b 100644 --- a/native/shuffle/src/writers/mod.rs +++ b/native/shuffle/src/writers/mod.rs @@ -19,7 +19,7 @@ mod buf_batch_writer; mod checksum; mod local; mod partition_writer; -mod rss; +pub(crate) mod rss; mod shuffle_block_writer; pub(crate) use buf_batch_writer::BufBatchWriter; diff --git a/native/shuffle/src/writers/rss/mod.rs b/native/shuffle/src/writers/rss/mod.rs index 164680d4a62..061c6b53a17 100644 --- a/native/shuffle/src/writers/rss/mod.rs +++ b/native/shuffle/src/writers/rss/mod.rs @@ -18,7 +18,7 @@ pub(crate) mod rss_partition_writer; #[cfg(test)] -mod tests { +pub(crate) mod tests { use super::rss_partition_writer::RssPartitionWriter; use crate::metrics::ShufflePartitionerMetrics; use crate::writers::PartitionWriter; @@ -44,7 +44,8 @@ mod tests { /// Test-only allocation observation on a synchronous encoder thread. Production execution /// does not use thread-local state. Zstd's C allocations are covered separately by its public /// streaming-workspace estimate; this observes Rust buffers and their realloc overlap. - mod allocations { + /// Shared with the reader tests in `ipc.rs`, since a crate has one global allocator. + pub(crate) mod allocations { use std::alloc::{GlobalAlloc, Layout, System}; use std::cell::Cell; @@ -125,7 +126,7 @@ mod tests { } // Allocation/reallocation requests and requested bytes, not retained memory. - pub(super) fn totals() -> (usize, usize) { + pub(crate) fn totals() -> (usize, usize) { COUNTERS.with(|counter| { let value = counter.get().unwrap(); (value.allocations, value.allocated_bytes) @@ -155,7 +156,7 @@ mod tests { }); } - pub(super) fn measure(run: impl FnOnce() -> T) -> (T, usize) { + pub(crate) fn measure(run: impl FnOnce() -> T) -> (T, usize) { struct Reset; impl Drop for Reset { fn drop(&mut self) {