Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 89 additions & 0 deletions ballista/client/tests/sort_shuffle.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ mod sort_shuffle_tests {
BALLISTA_SHUFFLE_READER_MAX_BYTES_IN_FLIGHT,
BALLISTA_SHUFFLE_READER_REMOTE_PREFER_FLIGHT,
};
use datafusion::arrow::datatypes::DataType;
use datafusion::arrow::util::pretty::pretty_format_batches;
use datafusion::common::Result;
use datafusion::execution::SessionStateBuilder;
Expand Down Expand Up @@ -201,6 +202,94 @@ mod sort_shuffle_tests {
Ok(())
}

/// Shuffles a variable-width column. Every other query here groups on a
/// fixed-width primitive, so no offsets buffer otherwise crosses a shuffle.
/// Offset buffers are part of what the skipped validation covers.
#[rstest]
#[case::local(ReadMode::Local)]
#[case::remote_flight(ReadMode::RemoteFlight)]
#[case::remote_block_io(ReadMode::RemoteBlockIo)]
#[tokio::test]
async fn test_sort_shuffle_group_by_binary_column(
#[case] read_mode: ReadMode,
) -> Result<()> {
let ctx = create_sort_shuffle_context(read_mode).await;
register_test_data(&ctx).await;

let df = ctx
.sql(
"SELECT date_string_col, COUNT(*) as cnt
FROM test
GROUP BY date_string_col
ORDER BY date_string_col",
)
.await?;
let results = df.collect().await?;

// `date_string_col` is `Binary` (no UTF8 annotation in the fixture), so
// it renders as hex: these are "01/01/09" .. "04/01/09".
let expected = vec![
"+------------------+-----+",
"| date_string_col | cnt |",
"+------------------+-----+",
"| 30312f30312f3039 | 2 |",
"| 30322f30312f3039 | 2 |",
"| 30332f30312f3039 | 2 |",
"| 30342f30312f3039 | 2 |",
"+------------------+-----+",
];
assert_result_eq(expected, &results);
Ok(())
}

/// Shuffles view-typed keys. Validating those bounds-checks each view's
/// buffer index and offset, a separate path from walking offsets. Keys of
/// 12 bytes or fewer live inside the view; longer ones reference a data
/// buffer, so one of each is grouped on. No fixture column is a view type,
/// hence `arrow_cast`.
#[rstest]
#[case::local(ReadMode::Local)]
#[case::remote_flight(ReadMode::RemoteFlight)]
#[case::remote_block_io(ReadMode::RemoteBlockIo)]
#[tokio::test]
async fn test_sort_shuffle_group_by_view_columns(
#[case] read_mode: ReadMode,
) -> Result<()> {
let ctx = create_sort_shuffle_context(read_mode).await;
register_test_data(&ctx).await;

let df = ctx
.sql(
"SELECT arrow_cast(CAST(date_string_col AS VARCHAR), 'Utf8View') AS inline_key,
arrow_cast(
CAST(date_string_col AS VARCHAR) || ' well past the inline limit',
'Utf8View'
) AS spilled_key,
COUNT(*) as cnt
FROM test
GROUP BY inline_key, spilled_key
ORDER BY inline_key",
)
.await?;
let results = df.collect().await?;
let schema = results[0].schema();
assert_eq!(schema.field(0).data_type(), &DataType::Utf8View);
assert_eq!(schema.field(1).data_type(), &DataType::Utf8View);

let expected = vec![
"+------------+-------------------------------------+-----+",
"| inline_key | spilled_key | cnt |",
"+------------+-------------------------------------+-----+",
"| 01/01/09 | 01/01/09 well past the inline limit | 2 |",
"| 02/01/09 | 02/01/09 well past the inline limit | 2 |",
"| 03/01/09 | 03/01/09 well past the inline limit | 2 |",
"| 04/01/09 | 04/01/09 well past the inline limit | 2 |",
"+------------+-------------------------------------+-----+",
];
assert_result_eq(expected, &results);
Ok(())
}

#[rstest]
#[case::local(ReadMode::Local)]
#[case::remote_flight(ReadMode::RemoteFlight)]
Expand Down
17 changes: 15 additions & 2 deletions ballista/core/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -441,6 +441,19 @@ impl RecordBatchStream for FlightDataStream {
self.schema.clone()
}
}

/// Decoder for shuffle bytes streamed by [`BlockDataStream`].
///
/// The producing executor wrote these from arrays Arrow had already validated,
/// so re-validating on the consumer only costs a scan.
fn new_decoder() -> StreamDecoder {
// Safety: setting `skip_validation` requires `unsafe`, user assures data is valid
unsafe {
StreamDecoder::new()
.with_skip_validation(cfg!(feature = "arrow-ipc-optimizations"))
}
}

#[allow(rustdoc::private_intra_doc_links)]
/// [BlockDataStream] facilitates the transfer of original shuffle files in a block-by-block manner.
/// This implementation utilizes a custom `do_action` method on the Arrow Flight server.
Expand Down Expand Up @@ -487,7 +500,7 @@ impl<S: Stream<Item = Result<prost::bytes::Bytes>> + Unpin> BlockDataStream<S> {
match try_schema_from_ipc_buffer(state_buffer.as_slice()) {
Ok(schema) => {
return Ok(Self {
decoder: StreamDecoder::new(),
decoder: new_decoder(),
transmitted: state_buffer.len(),
state_buffer,
ipc_stream,
Expand Down Expand Up @@ -567,7 +580,7 @@ impl<S: Stream<Item = Result<prost::bytes::Bytes>> + Unpin> Stream
// stream followed by the requested partition's streams).
// Reset the decoder; the schema captured at construction
// time stays authoritative for downstream consumers.
self.decoder = StreamDecoder::new();
self.decoder = new_decoder();
continue;
}
Err(e) => {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -99,7 +99,13 @@ impl MultiStreamPartitionStream {
}
let mut file = File::open(&self.data_path)?;
file.seek(SeekFrom::Start(next))?;
self.state = State::Reading(StreamReader::try_new(file, None)?);
// Safety: setting `skip_validation` requires `unsafe`, user assures data is valid
let reader = unsafe {
StreamReader::try_new(file, None)?.with_skip_validation(cfg!(
feature = "arrow-ipc-optimizations"
))
};
self.state = State::Reading(reader);
}
}
}
Expand Down