use arrow::array::ArrayData;
use arrow::buffer::Buffer as ArrowBuffer;
use arrow::ipc::writer::{DictionaryTracker, IpcDataGenerator, IpcWriteOptions};
use arrow::ipc::{Buffer as IpcBuffer, FieldNode, MessageHeader, MetadataVersion};
use arrow_schema::{DataType, Field, Schema};
use eyre::{Context, bail, eyre};
use super::ARROW_BUFFER_ALIGNMENT as ALIGN;
const CONTINUATION_MARKER: [u8; 4] = [0xff, 0xff, 0xff, 0xff];
const PREFIX_LEN: usize = 8;
#[inline]
fn round_up(n: usize, align: usize) -> usize {
debug_assert!(align.is_power_of_two());
(n + align - 1) & !(align - 1)
}
fn is_fast_path_type(data_type: &DataType) -> bool {
use DataType::*;
matches!(
data_type,
Null | Boolean
| Int8
| Int16
| Int32
| Int64
| UInt8
| UInt16
| UInt32
| UInt64
| Float16
| Float32
| Float64
| Timestamp(_, _)
| Date32
| Date64
| Time32(_)
| Time64(_)
| Duration(_)
| Interval(_)
| Decimal128(_, _)
| Decimal256(_, _)
| FixedSizeBinary(_)
| Binary
| LargeBinary
| Utf8
| LargeUtf8
| List(_)
| LargeList(_)
| FixedSizeList(_, _)
| Struct(_)
)
}
struct Desc {
offset: usize,
len: usize,
src: BufferSrc,
}
enum BufferSrc {
Bytes(ArrowBuffer, usize),
AllOnes,
}
#[derive(Default)]
struct Layout {
nodes: Vec<FieldNode>,
ipc_buffers: Vec<IpcBuffer>,
descs: Vec<Desc>,
body_len: usize,
}
impl Layout {
fn push_buffer(&mut self, off: &mut usize, len: usize, src: BufferSrc) -> Option<()> {
let padded = len.checked_add(ALIGN - 1).map(|n| n & !(ALIGN - 1))?;
let next = off.checked_add(padded)?;
if next > super::MAX_IPC_BYTES {
return None;
}
self.ipc_buffers
.push(IpcBuffer::new(*off as i64, len as i64));
self.descs.push(Desc {
offset: *off,
len,
src,
});
*off = next;
Some(())
}
}
fn build_layout(array: &ArrayData) -> Option<Layout> {
let mut layout = Layout::default();
let mut off = 0usize;
build_layout_rec(array, &mut layout, &mut off)?;
layout.body_len = off;
Some(layout)
}
fn build_layout_rec(array: &ArrayData, layout: &mut Layout, off: &mut usize) -> Option<()> {
let data_type = array.data_type();
if !is_fast_path_type(data_type) || array.offset() != 0 {
return None;
}
let len = array.len();
let null_count = if matches!(data_type, DataType::Null) {
len
} else {
array.null_count()
};
layout
.nodes
.push(FieldNode::new(len as i64, null_count as i64));
if !matches!(data_type, DataType::Null) {
match array.nulls() {
Some(nulls) => {
if nulls.inner().offset() != 0 {
return None;
}
let sliced = nulls.inner().sliced();
let bytes = sliced.len();
layout.push_buffer(off, bytes, BufferSrc::Bytes(sliced, 0))?;
}
None => {
let bytes = len.div_ceil(8);
layout.push_buffer(off, bytes, BufferSrc::AllOnes)?;
}
}
}
for buffer in array.buffers() {
layout.push_buffer(off, buffer.len(), BufferSrc::Bytes(buffer.clone(), 0))?;
}
match data_type {
DataType::FixedSizeList(_, value_size) => {
let n = len.checked_mul(*value_size as usize)?;
let child = array.child_data().first()?;
if child.len() < n {
return None;
}
build_layout_rec(&child.slice(0, n), layout, off)?;
}
DataType::Struct(_) => {
for child in array.child_data() {
if child.len() < len {
return None;
}
build_layout_rec(&child.slice(0, len), layout, off)?;
}
}
_ => {
for child in array.child_data() {
build_layout_rec(child, layout, off)?;
}
}
}
Some(())
}
struct Prepared {
layout: Layout,
schema_message: Vec<u8>,
record_batch_message: Vec<u8>,
schema_block: usize,
record_batch_block: usize,
total: usize,
}
fn ipc_write_options() -> eyre::Result<IpcWriteOptions> {
IpcWriteOptions::try_new(ALIGN, false, MetadataVersion::V5)
.map_err(|e| eyre!("failed to build Arrow IPC write options: {e}"))
}
fn build_schema_message(data_type: &DataType) -> eyre::Result<Vec<u8>> {
let schema = Schema::new(vec![Field::new("data", data_type.clone(), true)]);
let options = ipc_write_options()?;
let mut tracker = DictionaryTracker::new(false);
let encoded = IpcDataGenerator {}.schema_to_bytes_with_dictionary_tracker(
&schema,
&mut tracker,
&options,
);
Ok(encoded.ipc_message)
}
fn build_record_batch_message(
num_rows: usize,
nodes: &[FieldNode],
buffers: &[IpcBuffer],
body_len: usize,
) -> Vec<u8> {
use flatbuffers::FlatBufferBuilder;
let mut fbb = FlatBufferBuilder::new();
let buffers_fb = fbb.create_vector(buffers);
let nodes_fb = fbb.create_vector(nodes);
let record_batch = {
let mut builder = arrow::ipc::RecordBatchBuilder::new(&mut fbb);
builder.add_length(num_rows as i64);
builder.add_nodes(nodes_fb);
builder.add_buffers(buffers_fb);
builder.finish()
};
let message = {
let mut builder = arrow::ipc::MessageBuilder::new(&mut fbb);
builder.add_version(MetadataVersion::V5);
builder.add_header_type(MessageHeader::RecordBatch);
builder.add_bodyLength(body_len as i64);
builder.add_header(record_batch.as_union_value());
builder.finish()
};
fbb.finish(message, None);
fbb.finished_data().to_vec()
}
fn prepare(array: &ArrayData) -> Option<Prepared> {
let layout = build_layout(array)?;
let schema_message = build_schema_message(array.data_type()).ok()?;
let record_batch_message = build_record_batch_message(
array.len(),
&layout.nodes,
&layout.ipc_buffers,
layout.body_len,
);
let schema_block = round_up(PREFIX_LEN + schema_message.len(), ALIGN);
let record_batch_block = round_up(PREFIX_LEN + record_batch_message.len(), ALIGN);
let total = schema_block + record_batch_block + layout.body_len + PREFIX_LEN;
Some(Prepared {
layout,
schema_message,
record_batch_message,
schema_block,
record_batch_block,
total,
})
}
pub fn ipc_fast_path_len(array: &ArrayData) -> Option<usize> {
prepare(array).map(|p| p.total)
}
fn write_framed_message(dst: &mut [u8], at: usize, flatbuffer: &[u8]) -> usize {
let block = round_up(PREFIX_LEN + flatbuffer.len(), ALIGN);
let metadata_len = (block - PREFIX_LEN) as i32;
dst[at..at + 4].copy_from_slice(&CONTINUATION_MARKER);
dst[at + 4..at + 8].copy_from_slice(&metadata_len.to_le_bytes());
dst[at + 8..at + 8 + flatbuffer.len()].copy_from_slice(flatbuffer);
dst[at + 8 + flatbuffer.len()..at + block].fill(0);
block
}
pub fn encode_ipc_into(array: &ArrayData, dst: &mut [u8]) -> eyre::Result<()> {
let prepared =
prepare(array).ok_or_else(|| eyre!("array is not Arrow IPC fast-path eligible"))?;
if dst.len() != prepared.total {
bail!(
"destination size {} does not match required IPC length {}",
dst.len(),
prepared.total
);
}
let mut at = 0;
at += write_framed_message(dst, at, &prepared.schema_message);
debug_assert_eq!(at, prepared.schema_block);
at += write_framed_message(dst, at, &prepared.record_batch_message);
debug_assert_eq!(at, prepared.schema_block + prepared.record_batch_block);
let body_start = at;
write_body(dst, body_start, &prepared.layout);
at = body_start + prepared.layout.body_len;
dst[at..at + 4].copy_from_slice(&CONTINUATION_MARKER);
dst[at + 4..at + 8].copy_from_slice(&0i32.to_le_bytes());
debug_assert_eq!(at + PREFIX_LEN, prepared.total);
Ok(())
}
fn write_body(dst: &mut [u8], body_start: usize, layout: &Layout) {
for desc in &layout.descs {
let start = body_start + desc.offset;
let end = start + desc.len;
match &desc.src {
BufferSrc::Bytes(buffer, src_off) => {
dst[start..end].copy_from_slice(&buffer.as_slice()[*src_off..*src_off + desc.len]);
}
BufferSrc::AllOnes => dst[start..end].fill(0xff),
}
let padded_end = body_start + desc.offset + round_up(desc.len, ALIGN);
dst[end..padded_end].fill(0);
}
}
pub fn encode_schema_message(data_type: &DataType) -> eyre::Result<Vec<u8>> {
let schema_message = build_schema_message(data_type)?;
let block = round_up(PREFIX_LEN + schema_message.len(), ALIGN);
let mut dst = vec![0u8; block];
write_framed_message(&mut dst, 0, &schema_message);
Ok(dst)
}
pub fn batch_fast_path_len(array: &ArrayData) -> Option<usize> {
let prepared = prepare(array)?;
Some(prepared.record_batch_block + prepared.layout.body_len + PREFIX_LEN)
}
pub fn encode_batch_into(array: &ArrayData, dst: &mut [u8]) -> eyre::Result<()> {
let prepared =
prepare(array).ok_or_else(|| eyre!("array is not Arrow IPC fast-path eligible"))?;
let expected = prepared.record_batch_block + prepared.layout.body_len + PREFIX_LEN;
if dst.len() != expected {
bail!(
"destination size {} does not match required batch length {expected}",
dst.len(),
);
}
let at = write_framed_message(dst, 0, &prepared.record_batch_message);
debug_assert_eq!(at, prepared.record_batch_block);
write_body(dst, at, &prepared.layout);
let body_end = at + prepared.layout.body_len;
dst[body_end..body_end + 4].copy_from_slice(&CONTINUATION_MARKER);
dst[body_end + 4..body_end + PREFIX_LEN].copy_from_slice(&0i32.to_le_bytes());
Ok(())
}
pub fn encode_ipc_to_vec(array: &ArrayData) -> eyre::Result<Vec<u8>> {
super::encode_arrow_ipc(array).context("Arrow IPC fallback encode")
}
struct Uint8Layout {
schema_message: Vec<u8>,
record_batch_message: Vec<u8>,
validity_len: usize,
validity_padded: usize,
body_len: usize,
total: usize,
data_offset: usize,
}
fn uint8_layout(data_len: usize) -> eyre::Result<Uint8Layout> {
if data_len > super::MAX_IPC_BYTES {
bail!(
"UInt8 payload too large: {data_len} bytes (max {})",
super::MAX_IPC_BYTES
);
}
let validity_len = data_len.div_ceil(8);
let validity_padded = round_up(validity_len, ALIGN);
let body_len = validity_padded + round_up(data_len, ALIGN);
let nodes = [FieldNode::new(data_len as i64, 0)];
let buffers = [
IpcBuffer::new(0, validity_len as i64),
IpcBuffer::new(validity_padded as i64, data_len as i64),
];
let record_batch_message = build_record_batch_message(data_len, &nodes, &buffers, body_len);
let schema_message = build_schema_message(&DataType::UInt8)?;
let schema_block = round_up(PREFIX_LEN + schema_message.len(), ALIGN);
let record_batch_block = round_up(PREFIX_LEN + record_batch_message.len(), ALIGN);
let total = schema_block + record_batch_block + body_len + PREFIX_LEN;
let data_offset = schema_block + record_batch_block + validity_padded;
Ok(Uint8Layout {
schema_message,
record_batch_message,
validity_len,
validity_padded,
body_len,
total,
data_offset,
})
}
pub fn uint8_ipc_len(data_len: usize) -> eyre::Result<usize> {
Ok(uint8_layout(data_len)?.total)
}
pub fn encode_uint8_ipc_header(dst: &mut [u8], data_len: usize) -> eyre::Result<usize> {
let layout = uint8_layout(data_len)?;
if dst.len() != layout.total {
bail!(
"destination size {} does not match required UInt8 IPC length {}",
dst.len(),
layout.total
);
}
let mut at = 0;
at += write_framed_message(dst, at, &layout.schema_message);
at += write_framed_message(dst, at, &layout.record_batch_message);
let body_start = at;
dst[body_start..body_start + layout.validity_len].fill(0xff);
dst[body_start + layout.validity_len..body_start + layout.validity_padded].fill(0);
let data_end = layout.data_offset + data_len;
let body_end = body_start + layout.body_len;
dst[data_end..body_end].fill(0);
dst[body_end..body_end + 4].copy_from_slice(&CONTINUATION_MARKER);
dst[body_end + 4..body_end + 8].copy_from_slice(&0i32.to_le_bytes());
debug_assert_eq!(body_end + PREFIX_LEN, layout.total);
Ok(layout.data_offset)
}
pub fn schema_block_len(stream: &[u8]) -> Option<usize> {
if stream.len() < PREFIX_LEN || stream[0..4] != CONTINUATION_MARKER {
return None;
}
let metadata_len = i32::from_le_bytes(stream[4..8].try_into().ok()?);
let block = PREFIX_LEN.checked_add(usize::try_from(metadata_len).ok()?)?;
(block <= stream.len()).then_some(block)
}
pub fn schema_block_and_hash(stream: &[u8]) -> Option<(u64, &[u8])> {
let block = schema_block_len(stream)?;
let schema = stream.get(..block)?;
Some((dora_message::metadata::fnv1a(schema), schema))
}
pub fn batch_slice(stream: &[u8]) -> Option<&[u8]> {
let block = schema_block_len(stream)?;
(stream.len() >= block + PREFIX_LEN).then(|| &stream[block..])
}
const MAX_RETAINED_SCHEMAS: usize = 8;
pub struct InputDecoder {
decoder: arrow::ipc::reader::StreamDecoder,
schema_hash: Option<u64>,
schemas: Vec<(u64, ArrowBuffer)>,
}
impl Default for InputDecoder {
fn default() -> Self {
Self::new()
}
}
impl InputDecoder {
pub fn new() -> Self {
Self {
decoder: arrow::ipc::reader::StreamDecoder::new(),
schema_hash: None,
schemas: Vec::new(),
}
}
pub fn reset(&mut self) {
self.decoder = arrow::ipc::reader::StreamDecoder::new();
self.schema_hash = None;
self.schemas.clear();
}
pub fn knows_schema(&self, hash: u64) -> bool {
self.schema_hash == Some(hash) || self.schemas.iter().any(|(h, _)| *h == hash)
}
pub fn set_schema(&mut self, hash: u64, schema: ArrowBuffer) -> eyre::Result<()> {
check_ipc_size(schema.len())?;
if self.schema_hash == Some(hash) {
return Ok(());
}
self.prime(hash, schema)
}
fn prime(&mut self, hash: u64, schema: ArrowBuffer) -> eyre::Result<()> {
let mut decoder = arrow::ipc::reader::StreamDecoder::new();
prime_with_schema(&mut decoder, schema.clone())?;
self.decoder = decoder;
self.schema_hash = Some(hash);
self.schemas.retain(|(h, _)| *h != hash);
self.schemas.push((hash, schema));
if self.schemas.len() > MAX_RETAINED_SCHEMAS {
self.schemas.remove(0);
}
Ok(())
}
pub fn decode_batch(
&mut self,
buffer: ArrowBuffer,
hash: u64,
) -> eyre::Result<Option<arrow::array::ArrayData>> {
check_ipc_size(buffer.len())?;
if self.schema_hash != Some(hash) {
match self.schemas.iter().find(|(h, _)| *h == hash) {
Some((_, schema)) => {
let schema = schema.clone();
self.prime(hash, schema)?;
}
None => return Ok(None),
}
}
match decode_one_batch(&mut self.decoder, buffer) {
Ok(array) => {
self.decoder = arrow::ipc::reader::StreamDecoder::new();
self.schema_hash = None;
Ok(Some(array))
}
Err(e) => {
self.decoder = arrow::ipc::reader::StreamDecoder::new();
self.schema_hash = None;
Err(e)
}
}
}
}
fn check_ipc_size(len: usize) -> eyre::Result<()> {
if len > super::MAX_IPC_BYTES {
bail!(
"Arrow IPC payload too large: {len} bytes (max {})",
super::MAX_IPC_BYTES
);
}
Ok(())
}
fn prime_with_schema(
decoder: &mut arrow::ipc::reader::StreamDecoder,
mut buffer: ArrowBuffer,
) -> eyre::Result<()> {
while !buffer.is_empty() {
let before = buffer.len();
if decoder
.decode(&mut buffer)
.map_err(|e| eyre!("failed to decode IPC schema message: {e}"))?
.is_some()
{
bail!("expected a schema message but got a record batch");
}
if buffer.len() == before {
bail!("IPC schema decoder made no progress on a partial/corrupt message");
}
}
Ok(())
}
fn decode_one_batch(
decoder: &mut arrow::ipc::reader::StreamDecoder,
mut buffer: ArrowBuffer,
) -> eyre::Result<arrow::array::ArrayData> {
while !buffer.is_empty() {
let before = buffer.len();
if let Some(batch) = decoder
.decode(&mut buffer)
.map_err(|e| eyre!("failed to decode IPC record batch: {e}"))?
{
if batch.num_columns() != 1 {
bail!(
"expected 1 column in IPC record batch, got {}",
batch.num_columns()
);
}
return Ok(batch.column(0).to_data());
}
if buffer.len() == before {
bail!("IPC batch decoder made no progress on a partial/corrupt message");
}
}
bail!("IPC batch message yielded no record batch")
}
#[cfg(test)]
mod tests {
use super::*;
use crate::arrow_utils::decode_arrow_ipc_zero_copy;
use arrow::array::{
Array, ArrayRef, BooleanArray, FixedSizeBinaryArray, Float32Array, Int32Array,
LargeStringArray, ListArray, NullArray, StringArray, StructArray, UInt8Array, UInt64Array,
};
use arrow::buffer::Buffer;
use arrow::ipc::reader::{StreamDecoder, StreamReader};
use arrow_schema::{DataType, Field};
use std::io::Cursor;
use std::sync::Arc;
fn fast_encode(array: &ArrayData) -> Vec<u8> {
let len = ipc_fast_path_len(array).expect("array should be fast-path eligible");
let mut buf = vec![0u8; len];
encode_ipc_into(array, &mut buf).expect("fast-path encode");
buf
}
fn read_official(bytes: &[u8]) -> ArrayData {
let mut reader = StreamReader::try_new(Cursor::new(bytes), None).expect("open IPC stream");
let batch = reader
.next()
.expect("one batch")
.expect("batch decodes via official reader");
assert_eq!(batch.num_columns(), 1);
batch.column(0).to_data()
}
fn aligned_buffer(bytes: &[u8]) -> (Buffer, usize, usize) {
use aligned_vec::{AVec, ConstAlign};
use std::ptr::NonNull;
let mut aligned: AVec<u8, ConstAlign<128>> = AVec::__from_elem(128, 0, bytes.len());
aligned.copy_from_slice(bytes);
let base = aligned.as_ptr() as usize;
let len = aligned.len();
let ptr = NonNull::new(aligned.as_ptr() as *mut u8).unwrap();
let buffer =
unsafe { Buffer::from_custom_allocation(ptr, len, std::sync::Arc::new(aligned)) };
(buffer, base, len)
}
fn assert_fast_roundtrip(array: &ArrayData) {
let encoded = fast_encode(array);
let decoded = read_official(&encoded);
assert_eq!(array, &decoded, "fast-path stream must decode to the input");
let (buffer, _, _) = aligned_buffer(&encoded);
let zc = decode_arrow_ipc_zero_copy(buffer).expect("zero-copy decode");
assert_eq!(array, &zc, "zero-copy decode must equal the input");
}
fn batch_bytes(array: &ArrayData) -> Vec<u8> {
let len = batch_fast_path_len(array).unwrap();
let mut buf = vec![0u8; len];
encode_batch_into(array, &mut buf).unwrap();
buf
}
fn batch_buf(array: &ArrayData) -> Buffer {
Buffer::from_vec(batch_bytes(array))
}
#[test]
fn input_decoder_schema_then_batches() {
let f32_schema = || Buffer::from_vec(encode_schema_message(&DataType::Float32).unwrap());
let mut dec = InputDecoder::new();
let early = Float32Array::from(vec![9.0]).into_data();
assert!(dec.decode_batch(batch_buf(&early), 7).unwrap().is_none());
dec.set_schema(7, f32_schema()).unwrap();
for vals in [vec![1.0f32, 2.0, 3.0], vec![4.0], vec![5.0, 6.0, 7.0]] {
let array = Float32Array::from(vals).into_data();
assert_eq!(
dec.decode_batch(batch_buf(&array), 7).unwrap().unwrap(),
array
);
}
let other = Float32Array::from(vec![8.0]).into_data();
assert!(dec.decode_batch(batch_buf(&other), 99).unwrap().is_none());
dec.set_schema(99, f32_schema()).unwrap();
let after = Float32Array::from(vec![10.0, 11.0]).into_data();
assert_eq!(
dec.decode_batch(batch_buf(&after), 99).unwrap().unwrap(),
after
);
}
#[test]
fn input_decoder_retains_multiple_schemas() {
let schema_msg = |dt: &DataType| Buffer::from_vec(encode_schema_message(dt).unwrap());
let mut dec = InputDecoder::new();
dec.set_schema(1, schema_msg(&DataType::Float32)).unwrap();
dec.set_schema(2, schema_msg(&DataType::Int32)).unwrap();
let ints = Int32Array::from(vec![1, 2, 3]).into_data();
assert_eq!(
dec.decode_batch(batch_buf(&ints), 2).unwrap().unwrap(),
ints
);
let floats = Float32Array::from(vec![4.0, 5.0]).into_data();
assert_eq!(
dec.decode_batch(batch_buf(&floats), 1).unwrap().unwrap(),
floats,
"a schema installed earlier must be retained across later primes"
);
let more_ints = Int32Array::from(vec![6]).into_data();
assert_eq!(
dec.decode_batch(batch_buf(&more_ints), 2).unwrap().unwrap(),
more_ints
);
}
#[test]
fn input_decoder_handles_sequential_batches_arrow_59_terminal_state() {
let f32_schema = || Buffer::from_vec(encode_schema_message(&DataType::Float32).unwrap());
let mut dec = InputDecoder::new();
dec.set_schema(7, f32_schema()).unwrap();
let batches = vec![
Float32Array::from(vec![1.0, 2.0]).into_data(),
Float32Array::from(vec![3.0]).into_data(),
Float32Array::from(vec![4.0, 5.0, 6.0]).into_data(),
Float32Array::from(vec![7.0, 8.0]).into_data(),
Float32Array::from(vec![9.0, 10.0, 11.0, 12.0]).into_data(),
];
for (i, batch) in batches.iter().enumerate() {
let result = dec.decode_batch(batch_buf(batch), 7);
assert!(
result.is_ok(),
"batch {} decode failed: {:?}",
i,
result.err()
);
let decoded = result.unwrap().unwrap();
assert_eq!(
&decoded, batch,
"batch {} mismatch: expected {:?}, got {:?}",
i, batch, decoded
);
}
}
#[test]
fn input_decoder_evicts_oldest_schema_beyond_cap() {
let f32_schema = || Buffer::from_vec(encode_schema_message(&DataType::Float32).unwrap());
let mut dec = InputDecoder::new();
for hash in 0..=(MAX_RETAINED_SCHEMAS as u64) {
dec.set_schema(hash, f32_schema()).unwrap();
}
let array = Float32Array::from(vec![1.0]).into_data();
assert!(
dec.decode_batch(batch_buf(&array), 0).unwrap().is_none(),
"the oldest schema must be evicted beyond the cap"
);
assert_eq!(
dec.decode_batch(batch_buf(&array), 1).unwrap().unwrap(),
array,
"schemas within the cap must be retained"
);
}
#[test]
fn decode_batch_resets_on_error_then_reprimes() {
let mut dec = InputDecoder::new();
dec.set_schema(
7,
Buffer::from_vec(encode_schema_message(&DataType::Float32).unwrap()),
)
.unwrap();
let good = Float32Array::from(vec![4.0, 5.0]).into_data();
assert_eq!(
dec.decode_batch(batch_buf(&good), 7).unwrap().unwrap(),
good
);
let dropped = Float32Array::from(vec![6.0, 7.0, 8.0, 9.0]).into_data();
let mut truncated = batch_bytes(&dropped);
truncated.truncate(truncated.len() / 2);
assert!(dec.decode_batch(Buffer::from_vec(truncated), 7).is_err());
let after = Float32Array::from(vec![10.0, 11.0]).into_data();
assert_eq!(
dec.decode_batch(batch_buf(&after), 7).unwrap().unwrap(),
after,
"after a failed batch the decoder must re-prime from the retained schema"
);
dec.reset();
let final_batch = Float32Array::from(vec![12.0]).into_data();
assert!(
dec.decode_batch(batch_buf(&final_batch), 7)
.unwrap()
.is_none(),
"after a full reset the decoder must drop until a schema is re-installed"
);
}
#[test]
fn input_decoder_dictionary_fallback_batch_sequence() {
use arrow::array::DictionaryArray;
use arrow::datatypes::Int32Type;
fn dict(words: &[&str]) -> ArrayData {
let mut values: Vec<&str> = Vec::new();
let mut keys: Vec<i32> = Vec::new();
for w in words {
let idx = values.iter().position(|v| v == w).unwrap_or_else(|| {
values.push(*w);
values.len() - 1
});
keys.push(idx as i32);
}
DictionaryArray::<Int32Type>::try_new(
Int32Array::from(keys),
Arc::new(StringArray::from(values)),
)
.unwrap()
.into_data()
}
let first = dict(&["a", "b", "a", "c", "b"]);
assert!(
ipc_fast_path_len(&first).is_none(),
"dictionary must route to the official-writer fallback"
);
let mut dec = InputDecoder::new();
let full0 = encode_ipc_to_vec(&first).unwrap();
let block = schema_block_len(&full0).unwrap();
dec.set_schema(1, Buffer::from(&full0[..block])).unwrap();
let slice0 = batch_slice(&full0).expect("fallback stream is a valid IPC stream");
assert_eq!(
dec.decode_batch(Buffer::from(slice0), 1)
.unwrap()
.expect("first batch decodes against the primed decoder"),
first
);
for words in [
["x", "y", "x", "z"].as_slice(),
["b", "b"].as_slice(),
["new", "values", "entirely"].as_slice(),
] {
let arr = dict(words);
let full = encode_ipc_to_vec(&arr).unwrap();
let slice = batch_slice(&full).expect("fallback stream is a valid IPC stream");
let got = dec
.decode_batch(Buffer::from(slice), 1)
.unwrap()
.expect("batch must decode against the primed decoder");
assert_eq!(
got, arr,
"replacement-dictionary batch must decode correctly"
);
}
}
#[test]
fn schema_primed_decoder_decodes_batch_sequence() {
let schema = encode_schema_message(&DataType::Float32).unwrap();
let mut decoder = StreamDecoder::new();
let mut sbuf = Buffer::from_vec(schema);
while !sbuf.is_empty() {
assert!(
decoder.decode(&mut sbuf).unwrap().is_none(),
"schema message must not yield a batch"
);
}
for vals in [vec![1.0f32, 2.0, 3.0], vec![4.0, 5.0], vec![6.0]] {
let array = Float32Array::from(vals).into_data();
let len = batch_fast_path_len(&array).unwrap();
let mut buf = vec![0u8; len];
encode_batch_into(&array, &mut buf).unwrap();
let mut bbuf = Buffer::from_vec(buf);
let mut got = None;
while !bbuf.is_empty() {
if let Some(b) = decoder.decode(&mut bbuf).unwrap() {
got = Some(b);
break;
}
}
assert_eq!(
got.expect("batch message must decode against the primed decoder")
.column(0)
.to_data(),
array
);
}
}
#[test]
fn uint8_ipc_header_constructs_in_place() {
for data_len in [0usize, 1, 7, 8, 9, 1000] {
let bytes: Vec<u8> = (0..data_len).map(|i| (i % 251) as u8).collect();
let total = uint8_ipc_len(data_len).unwrap();
let mut dst = vec![0u8; total];
let offset = encode_uint8_ipc_header(&mut dst, data_len).unwrap();
dst[offset..offset + data_len].copy_from_slice(&bytes);
let arr = arrow::array::make_array(read_official(&dst));
let u8 = arr.as_any().downcast_ref::<UInt8Array>().unwrap();
assert_eq!(u8.values(), bytes.as_slice(), "len {data_len}");
assert_eq!(u8.null_count(), 0);
let array = UInt8Array::from(bytes).into_data();
assert_eq!(dst, fast_encode(&array), "len {data_len}");
}
}
#[test]
fn roundtrip_primitive_no_nulls() {
let array = Float32Array::from((0..1000).map(|i| i as f32).collect::<Vec<_>>()).into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_primitive_with_nulls() {
let array = UInt64Array::from(vec![Some(1), None, Some(3), None, Some(5)]).into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_empty_primitive() {
let array = Int32Array::from(Vec::<i32>::new()).into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn uint8_header_zero_len_full_stream_roundtrip() {
let len = uint8_ipc_len(0).unwrap();
let mut buf = vec![0u8; len];
let off = encode_uint8_ipc_header(&mut buf, 0).unwrap();
assert_eq!(off, len - PREFIX_LEN);
assert_eq!(read_official(&buf).len(), 0);
let (buffer, _, _) = aligned_buffer(&buf);
let decoded = decode_arrow_ipc_zero_copy(buffer).unwrap();
assert_eq!(decoded.data_type(), &DataType::UInt8);
assert_eq!(decoded.len(), 0);
}
#[test]
fn schema_once_zero_len_batch_roundtrip() {
let schema = || Buffer::from_vec(encode_schema_message(&DataType::UInt8).unwrap());
let batch = |vals: &[u8]| {
let a = UInt8Array::from(vals.to_vec()).into_data();
let len = batch_fast_path_len(&a).unwrap();
let mut b = vec![0u8; len];
encode_batch_into(&a, &mut b).unwrap();
Buffer::from_vec(b)
};
let mut dec = InputDecoder::new();
dec.set_schema(1, schema()).unwrap();
let decoded = dec
.decode_batch(batch(&[]), 1)
.unwrap()
.expect("0-row batch must decode, not drop");
assert_eq!(decoded.data_type(), &DataType::UInt8);
assert_eq!(decoded.len(), 0);
for vals in [vec![1u8, 2, 3], vec![], vec![9u8], vec![]] {
let d = dec.decode_batch(batch(&vals), 1).unwrap().unwrap();
assert_eq!(d.len(), vals.len());
assert_eq!(d.data_type(), &DataType::UInt8);
}
}
#[test]
fn schema_once_zero_len_via_batch_slice_roundtrip() {
let total = uint8_ipc_len(0).unwrap();
let mut full = vec![0u8; total];
encode_uint8_ipc_header(&mut full, 0).unwrap();
let sblock = schema_block_len(&full).unwrap();
let schema = Buffer::from(&full[..sblock]);
let batch = batch_slice(&full).expect("batch slice of a valid stream");
let mut dec = InputDecoder::new();
dec.set_schema(7, schema).unwrap();
let decoded = dec
.decode_batch(Buffer::from(batch), 7)
.unwrap()
.expect("0-row batch via batch_slice must decode, not drop");
assert_eq!(decoded.data_type(), &DataType::UInt8);
assert_eq!(decoded.len(), 0);
}
#[test]
fn schema_block_plus_batch_slice_reconstructs_full_stream() {
let array = Int32Array::from(vec![1, 2, 3]).into_data();
let full = fast_encode(&array);
let sblock = schema_block_len(&full).unwrap();
let schema = &full[..sblock];
let batch = batch_slice(&full).expect("batch slice of a valid stream");
let rebuilt = [schema, batch].concat();
assert_eq!(
rebuilt, full,
"schema_block ++ batch_slice must equal the original stream"
);
assert_eq!(read_official(&rebuilt), array);
}
#[test]
fn roundtrip_boolean() {
let array =
BooleanArray::from(vec![true, false, true, true, false, false, true]).into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_utf8() {
let array =
StringArray::from(vec![Some("hello"), None, Some(""), Some("world!")]).into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_large_utf8_64bit_offsets() {
let array = LargeStringArray::from(vec!["a", "bb", "ccc"]).into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_fixed_size_binary() {
let values = vec![vec![1u8, 2, 3], vec![4, 5, 6], vec![7, 8, 9]];
let array = FixedSizeBinaryArray::try_from_iter(values.into_iter())
.unwrap()
.into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_fixed_size_list() {
use arrow::array::FixedSizeListArray;
let values = Int32Array::from((0..12).collect::<Vec<_>>());
let field = Arc::new(Field::new("item", DataType::Int32, true));
let array = FixedSizeListArray::try_new(field, 3, Arc::new(values), None)
.unwrap()
.into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_fixed_size_list_with_nulls() {
use arrow::array::FixedSizeListArray;
use arrow::buffer::NullBuffer;
let values = Int32Array::from((0..12).collect::<Vec<_>>());
let field = Arc::new(Field::new("item", DataType::Int32, true));
let nulls = NullBuffer::from(vec![true, false, true, true]);
let array = FixedSizeListArray::try_new(field, 3, Arc::new(values), Some(nulls))
.unwrap()
.into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_decimal128() {
use arrow::array::Decimal128Array;
let array = Decimal128Array::from(vec![Some(12_345i128), None, Some(-9_876), Some(0)])
.with_precision_and_scale(20, 4)
.unwrap()
.into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_timestamp_temporal() {
use arrow::array::TimestampMicrosecondArray;
let array =
TimestampMicrosecondArray::from(vec![Some(1_000_000i64), None, Some(2_500_000)])
.into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_struct_with_multilevel_nulls() {
let array = StructArray::from(vec![
(
Arc::new(Field::new("a", DataType::UInt64, true)),
Arc::new(UInt64Array::from(vec![Some(1), None, Some(3)])) as ArrayRef,
),
(
Arc::new(Field::new("b", DataType::Utf8, true)),
Arc::new(StringArray::from(vec![Some("x"), Some("yy"), None])) as ArrayRef,
),
])
.into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_struct_with_oversized_child() {
use arrow_schema::Fields;
let child = Int32Array::from(vec![10, 20, 30]).into_data(); let fields: Fields = vec![Field::new("v", DataType::Int32, false)].into();
let struct_data = ArrayData::builder(DataType::Struct(fields))
.len(2) .add_child_data(child)
.build()
.unwrap();
assert!(
ipc_fast_path_len(&struct_data).is_some(),
"struct with an oversized child should stay on the fast path"
);
let decoded = read_official(&fast_encode(&struct_data));
assert_eq!(decoded.len(), 2);
let arr = arrow::array::make_array(decoded);
let sa = arr.as_any().downcast_ref::<StructArray>().unwrap();
let col = sa.column(0).as_any().downcast_ref::<Int32Array>().unwrap();
assert_eq!(
col.values(),
&[10, 20],
"child must be truncated to the struct's len"
);
}
#[test]
fn roundtrip_list_of_primitive() {
let data = vec![
Some(vec![Some(0), Some(1), Some(2)]),
None,
Some(vec![Some(3), None, Some(5)]),
Some(vec![]),
];
let array =
ListArray::from_iter_primitive::<arrow::datatypes::Int32Type, _, _>(data).into_data();
assert_fast_roundtrip(&array);
}
#[test]
fn roundtrip_nullarray_zero_and_n() {
assert_fast_roundtrip(&NullArray::new(0).into_data());
assert_fast_roundtrip(&NullArray::new(7).into_data());
}
#[test]
fn fast_path_decodes_zero_copy() {
let array = UInt64Array::from((0..50_000u64).collect::<Vec<_>>()).into_data();
let encoded = fast_encode(&array);
{
let (mut buffer, _, _) = aligned_buffer(&encoded);
let mut decoder = StreamDecoder::new().with_require_alignment(true);
let mut got = None;
while !buffer.is_empty() {
if let Some(b) = decoder
.decode(&mut buffer)
.expect("aligned fast-path stream must decode without realignment")
{
got = Some(b);
break;
}
}
assert_eq!(got.unwrap().column(0).to_data(), array);
}
{
let (buffer, base, len) = aligned_buffer(&encoded);
let decoded = decode_arrow_ipc_zero_copy(buffer).unwrap();
let ptr = decoded.buffers()[0].as_ptr() as usize;
assert!(
ptr >= base && ptr < base + len,
"decoded data buffer at {ptr:#x} is outside input [{base:#x}, {:#x}) — a copy happened",
base + len
);
}
}
#[test]
fn fast_path_len_matches_official_decode() {
let array = Float32Array::from(vec![1.0, 2.0, 3.0, 4.0]).into_data();
let encoded = fast_encode(&array);
let mut reader = StreamReader::try_new(Cursor::new(&encoded[..]), None).unwrap();
let _ = reader.next().unwrap().unwrap();
assert!(reader.next().is_none(), "exactly one batch, fully consumed");
}
#[test]
fn sliced_array_routes_to_fallback_and_roundtrips() {
let array = UInt64Array::from(vec![10, 20, 30, 40, 50])
.into_data()
.slice(2, 2); assert_eq!(array.offset(), 2);
assert!(ipc_fast_path_len(&array).is_none());
let encoded = encode_ipc_to_vec(&array).unwrap();
let decoded = read_official(&encoded);
assert_eq!(array.len(), decoded.len());
let dec = arrow::array::make_array(decoded);
let dec = dec.as_any().downcast_ref::<UInt64Array>().unwrap();
assert_eq!(dec.values(), &[30, 40]);
}
#[test]
fn view_type_routes_to_fallback() {
use arrow::array::StringViewArray;
let array = StringViewArray::from(vec!["a", "bb", "ccc"]).into_data();
assert!(
ipc_fast_path_len(&array).is_none(),
"Utf8View is not fast-path eligible"
);
let encoded = encode_ipc_to_vec(&array).unwrap();
let decoded = read_official(&encoded);
assert_eq!(array, decoded);
}
#[test]
fn fuzz_roundtrip_many_shapes() {
let mut state: u64 = 0x1234_5678_9abc_def0;
let mut next = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
state
};
for _ in 0..200 {
let len = (next() % 64) as usize;
let kind = next() % 6;
let array: ArrayData = match kind {
0 => UInt8Array::from(
(0..len)
.map(|i| {
if next().is_multiple_of(4) {
None
} else {
Some((i as u8).wrapping_add(1))
}
})
.collect::<Vec<_>>(),
)
.into_data(),
1 => Float32Array::from((0..len).map(|i| i as f32 * 0.5).collect::<Vec<_>>())
.into_data(),
2 => BooleanArray::from(
(0..len)
.map(|i| (i + next() as usize).is_multiple_of(2))
.collect::<Vec<_>>(),
)
.into_data(),
3 => StringArray::from(
(0..len)
.map(|i| {
if next().is_multiple_of(5) {
None
} else {
Some("x".repeat(i % 7))
}
})
.collect::<Vec<_>>(),
)
.into_data(),
4 => Int32Array::from(
(0..len)
.map(|i| {
if next().is_multiple_of(3) {
None
} else {
Some(i as i32 - 10)
}
})
.collect::<Vec<_>>(),
)
.into_data(),
_ => StructArray::from(vec![(
Arc::new(Field::new("v", DataType::Int32, true)),
Arc::new(Int32Array::from(
(0..len).map(|i| Some(i as i32)).collect::<Vec<_>>(),
)) as ArrayRef,
)])
.into_data(),
};
assert_fast_roundtrip(&array);
}
}
}