This is an automated email from the ASF dual-hosted git repository.
github-merge-queue[bot] pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datafusion-comet.git
The following commit(s) were added to refs/heads/main by this push:
new 587e387681 perf: decode shuffle blocks against a cached schema instead
of re-parsing per block (#5809)
587e387681 is described below
commit 587e387681859121f12eb9bc6c8bdce9ffe5074e
Author: Peter Lee <[email protected]>
AuthorDate: Sun Sep 20 20:11:32 2026 +0000
perf: decode shuffle blocks against a cached schema instead of re-parsing
per block (#5809)
* bench: add a shuffle read benchmark covering the per-block schema parse
Every shuffle block is a self-contained Arrow IPC stream, so
read_single_batch
builds a fresh StreamReader per block and parses the schema flatbuffer once
per
block, even though every block in a shuffle carries the same schema. The
write
side already avoids the mirror image of this, encoding the schema once in
ShuffleBlockWriter::try_new and writing the pre-encoded bytes verbatim, but
there
was no read-side benchmark to say whether the reader's half is worth
removing.
This adds one, parameterized by column count and rows per block, measuring
the
schema parse separately from the full block decode. On an M-series laptop:
shape decode schema parse share
5 col x 64 row 1.93 us 1.14 us 59%
5 col x 512 row 2.38 us 0.91 us 38%
5 col x 8192 row 10.99 us 0.86 us 8%
50 col x 64 row 12.77 us 6.03 us 47%
50 col x 512 row 17.89 us 6.05 us 34%
50 col x 8192 row 218 us 6.05 us 3%
The parse cost is constant per block and independent of row count, so its
share
is set by how many rows land in a block. That is largest exactly where the
issue
predicted: wide shuffles, where rows per partition are few, and repeated
spilling, where each spill round emits its own block per partition.
Co-Authored-By: Claude Opus 5 <[email protected]>
* perf: decode shuffle blocks against a cached schema instead of re-parsing
per block
Every shuffle block is a self-contained Arrow IPC stream, so
read_single_batch
built a fresh StreamReader per block and parsed the schema flatbuffer once
per
block, even though every block in a shuffle carries the same schema. The
write
side already avoids the mirror image of this, encoding the schema once in
ShuffleBlockWriter::try_new and writing the pre-encoded bytes verbatim.
Blocks are now decoded against a per-thread cache keyed on the raw schema
message, so a hit costs one memcmp. On a hit the block is decoded in place
with
RecordBatchDecoder; on a miss the original StreamReader path runs unchanged
and
its parsed schema is cached for later blocks. The cache holds four schemas,
since
a reduce task can interleave blocks from more than one shuffle and a single
entry
would thrash.
The fast path never reports an error of its own. A cache miss, a dictionary
message, more than one record batch, trailing bytes after the end-of-stream
marker, or a block that simply fails to decode all fall back to the general
decoder, so validation behaviour and every error message are unchanged and
the
fast path is always safe to skip.
The measured win is not where #5792 predicted. Comparing this commit
against its
parent back to back, with the parse_schema_only arm as a control that this
change
does not touch (it drifted within 5% between the runs):
shape before after change
5 col x 64 row 1.663 us 1.775 us +6.7%
5 col x 512 row 2.120 us 1.913 us -9.8%
5 col x 8192 row 11.098 us 7.841 us -29.3%
50 col x 64 row 13.479 us 12.849 us -4.7%
50 col x 512 row 18.606 us 16.090 us -13.5%
50 col x 8192 row 159.49 us 77.03 us -51.7%
The issue expected the gain at small blocks, where the constant per-block
parse is
the largest share of decode. It is the other way round: the parse is worth
under a
microsecond, while decoding in place avoids the per-body MutableBuffer that
StreamReader allocates and zero-fills before copying into it, and that cost
scales
with body size. Small blocks are marginally slower, since materializing the
block
and walking its messages is not repaid when the body is tiny.
Co-Authored-By: Claude Opus 5 <[email protected]>
* review: trim comments to what they need to say
Co-Authored-By: Claude Opus 5 <[email protected]>
* review: trim comments to what they need to say
Co-Authored-By: Claude Opus 5 <[email protected]>
* review: stream each block message by message and serve the cached schema
without parsing it
Replaces the materialize-then-probe fast path with one message loop that
mirrors StreamReader. A cached schema is matched on its raw bytes and never
verified or parsed again; the record batch and any dictionary batches are
parsed once each. Bodies are read into exactly sized buffers, so a decoded
batch reports the same memory as before, and dictionary blocks decode from
the cache with dictionaries scoped to their own block.
A #[cfg(test)] hit/miss counter proves which path each decode took; the
tests reset the cache so cold and warm phases are explicit. The benchmark
adds Lz4Frame, the default codec, and a dictionary-encoded string column.
Co-Authored-By: Claude Fable 5.1 <[email protected]>
* review: probe allocations against StreamReader, validate corrupt arrays
on a warm cache, bench the validated entry point
The RSS tests' allocation observer is shared with the reader tests, which
compare one warm decode against the StreamReader path this change replaced:
no more allocations, bytes or peak live memory on any codec. The corrupt
offsets test now also fails validation with the schema served from the
cache, and the benchmark times read_ipc_compressed_validated as well.
Co-Authored-By: Claude Fable 5.1 <[email protected]>
---------
Co-authored-by: Claude Opus 5 <[email protected]>
---
native/Cargo.lock | 1 +
native/Cargo.toml | 1 +
native/shuffle/Cargo.toml | 1 +
native/shuffle/benches/shuffle_reader.rs | 140 +++--
native/shuffle/src/ipc.rs | 873 +++++++++++++++++++++++++++++--
native/shuffle/src/lib.rs | 2 +-
native/shuffle/src/writers/mod.rs | 2 +-
native/shuffle/src/writers/rss/mod.rs | 9 +-
8 files changed, 934 insertions(+), 95 deletions(-)
diff --git a/native/Cargo.lock b/native/Cargo.lock
index d4c592e1f8..fe01bd97ec 100644
--- a/native/Cargo.lock
+++ b/native/Cargo.lock
@@ -2066,6 +2066,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 6cf35f8676..a21658884b 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 9504834ef4..f0ed22ad73 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 47903d002f..43d4d44b99 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::<Int64Array>(),
- ) as arrow::array::ArrayRef
+ ) as ArrayRef
} else {
- Arc::new(
- (0..num_rows)
- .map(|r| Some(format!("value_{r}")))
- .collect::<StringArray>(),
- ) as arrow::array::ArrayRef
+ match strings {
+ Strings::Plain => Arc::new(
+ (0..num_rows)
+ .map(|r| Some(format!("value_{r}")))
+ .collect::<StringArray>(),
+ ) as ArrayRef,
+ // a small dictionary that every row's key points into
+ Strings::Dictionary => {
+ let values: Vec<String> =
+ (0..num_rows).map(|r| format!("value_{}", r %
16)).collect();
+ Arc::new(
+ values
+ .iter()
+ .map(String::as_str)
+ .collect::<DictionaryArray<Int32Type>>(),
+ ) as ArrayRef
+ }
+ }
}
})
.collect::<Vec<_>>();
@@ -86,25 +116,43 @@ fn encode_block(batch: &RecordBatch, codec:
CompressionCodec) -> Vec<u8> {
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 97890f5014..7e54367a33 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<RecordBatch> {
@@ -31,24 +39,142 @@ pub fn read_ipc_compressed_validated(bytes: &[u8]) ->
Result<RecordBatch> {
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<u8>,
+ #[cfg(test)]
+ stats: SchemaCacheStats,
+}
+
+thread_local! {
+ static STATE: RefCell<DecoderState> =
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<SchemaRef> {
+ 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<Message<'_>> {
+ root_as_message(metadata)
+ .map_err(|error| decode_error(&format!("unable to get root as message:
{error:?}")))
+}
+
+fn body_length(message: &Message<'_>) -> Result<usize> {
+ 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<RecordBatch> {
- 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<RecordBatch>
// 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<RecordBatch> {
+ 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<RecordBatch> {
+ 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<i64, ArrayRef> = 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<u8>) ->
Result<Option<&'a [u8]>>
+ where
+ 'b: 'a;
+
+ /// The next message's body, `len` bytes long.
+ fn body(&mut self, len: usize) -> Result<Buffer>;
+
+ /// 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<Option<usize>> {
+ 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>(R);
+
+impl<R: Read> Streamed<R> {
+ 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<R: Read> BlockSource<'static> for Streamed<R> {
+ fn next_metadata<'a>(&mut self, scratch: &'a mut Vec<u8>) ->
Result<Option<&'a [u8]>>
+ 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<Buffer> {
+ 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<u8>) ->
Result<Option<&'a [u8]>>
+ 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<Buffer> {
+ // 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<R: Read> Read for RequireLz4EndMark<R> {
}
}
-fn read_single_batch<R: Read>(input: R, validate: bool) -> Result<RecordBatch>
{
- 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<u8> {
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<u8> {
+ 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<u8> {
+ 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<Int32Type> =
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<String> {
+ let values = arrow::compute::cast(batch.column(0),
&DataType::Utf8).unwrap();
+ let values = values.as_any().downcast_ref::<StringArray>().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::<Vec<_>>();
+ 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<Vec<u8>> = (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::<Int32Array>()),
+ Arc::new(
+ (0..num_rows)
+ .map(|i| Some(format!("value_{i}")))
+ .collect::<StringArray>(),
+ ),
+ ],
+ )
+ .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<R: 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<u8> = [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 0eb18b517f..1158a2b1e2 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::{PartitionOffsets, ShuffleWriterDestination,
ShuffleWriterExec};
diff --git a/native/shuffle/src/writers/mod.rs
b/native/shuffle/src/writers/mod.rs
index fb3af2c991..4586e46c25 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 164680d4a6..061c6b53a1 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<T>(run: impl FnOnce() -> T) -> (T, usize) {
+ pub(crate) fn measure<T>(run: impl FnOnce() -> T) -> (T, usize) {
struct Reset;
impl Drop for Reset {
fn drop(&mut self) {
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]