use crate::egress::column::{
BinaryColumn, ColumnView, Decimal64Column, Decimal128Column, Decimal256Column,
DoubleArrayColumn, FixedColumn, GeohashColumn, Long256Column, LongArrayColumn, SymbolColumn,
UuidColumn, Validity, VarcharColumn,
};
use crate::egress::column_kind::ColumnKind;
use crate::egress::schema::Schema;
use crate::egress::symbol_dict::SymbolDict;
use crate::egress::wire::ByteReader;
use crate::egress::wire::header::flags;
use crate::egress::wire::msg_kind::MsgKind;
use crate::error::{Error, Result, fmt};
use bytes::Bytes;
pub(crate) const MAX_ROWS_PER_BATCH: usize = 1_048_576;
pub(crate) const MAX_COLUMNS_PER_TABLE: usize = 2048;
pub(crate) const MAX_COLUMN_NAME_LENGTH: usize = 127;
pub(crate) const MAX_TABLE_NAME_LENGTH: usize = 127;
pub(crate) const MAX_DECODED_BATCH_BYTES: usize = 1024 * 1024 * 1024;
fn read_owned(r: &mut ByteReader<'_>, parent: &Bytes, n: usize) -> Result<Bytes> {
let start = r.pos();
r.advance(n)?;
Ok(parent.slice(start..start + n))
}
#[derive(Debug, Clone)]
pub struct ColumnBuffer {
pub values: Bytes,
pub validity: Option<Bytes>,
}
#[derive(Debug, Clone)]
pub enum DecodedColumn {
Boolean(ColumnBuffer),
Byte(ColumnBuffer),
Short(ColumnBuffer),
Int(ColumnBuffer),
Long(ColumnBuffer),
Float(ColumnBuffer),
Double(ColumnBuffer),
Symbol {
codes: Vec<u32>,
validity: Option<Bytes>,
local_dict: Option<SymbolDict>,
},
Timestamp(ColumnBuffer),
Date(ColumnBuffer),
Uuid(ColumnBuffer),
Long256(ColumnBuffer),
TimestampNanos(ColumnBuffer),
Decimal64 {
buffer: ColumnBuffer,
scale: i8,
},
Char(ColumnBuffer),
Ipv4(ColumnBuffer),
Varchar {
offsets: Vec<u32>,
data: Bytes,
validity: Option<Bytes>,
},
Binary {
offsets: Vec<u32>,
data: Bytes,
validity: Option<Bytes>,
},
Geohash {
buffer: ColumnBuffer,
byte_width: u8,
precision_bits: u8,
},
Decimal128 {
buffer: ColumnBuffer,
scale: i8,
},
Decimal256 {
buffer: ColumnBuffer,
scale: i8,
},
DoubleArray(ArrayBuffers),
LongArray(ArrayBuffers),
}
#[derive(Debug, Clone)]
pub struct ArrayBuffers {
pub data_offsets: Vec<u32>,
pub data: Bytes,
pub shapes: Vec<u32>,
pub shape_offsets: Vec<u32>,
pub validity: Option<Bytes>,
}
#[derive(Debug, Clone)]
pub struct DecodedBatch {
pub request_id: i64,
pub batch_seq: u64,
pub row_count: usize,
pub columns: Vec<DecodedColumn>,
pub flags: u8,
}
impl DecodedBatch {
#[inline]
pub fn column_view<'a>(&'a self, idx: usize, dict: &'a SymbolDict) -> Result<ColumnView<'a>> {
let col = self
.columns
.get(idx)
.ok_or_else(|| fmt!(InvalidApiCall, "column index {} out of range", idx))?;
Ok(match col {
DecodedColumn::Boolean(b) => {
ColumnView::Boolean(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Byte(b) => {
ColumnView::Byte(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Short(b) => {
ColumnView::Short(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Int(b) => {
ColumnView::Int(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Long(b) => {
ColumnView::Long(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Float(b) => {
ColumnView::Float(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Double(b) => {
ColumnView::Double(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Timestamp(b) => {
ColumnView::Timestamp(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Date(b) => {
ColumnView::Date(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::TimestampNanos(b) => ColumnView::TimestampNanos(FixedColumn::new(
&b.values,
validity_of(b, self.row_count)?,
)),
DecodedColumn::Char(b) => {
ColumnView::Char(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Ipv4(b) => {
ColumnView::Ipv4(FixedColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Uuid(b) => {
ColumnView::Uuid(UuidColumn::new(&b.values, validity_of(b, self.row_count)?))
}
DecodedColumn::Long256(b) => ColumnView::Long256(Long256Column::new(
&b.values,
validity_of(b, self.row_count)?,
)),
DecodedColumn::Decimal64 { buffer, scale } => ColumnView::Decimal64(
Decimal64Column::new(&buffer.values, validity_of(buffer, self.row_count)?, *scale),
),
DecodedColumn::Symbol {
codes,
validity,
local_dict,
} => {
let active_dict = local_dict.as_ref().unwrap_or(dict);
ColumnView::Symbol(SymbolColumn::new(
codes,
validity_from_opt(validity, self.row_count)?,
active_dict,
))
}
DecodedColumn::Varchar {
offsets,
data,
validity,
} => {
let view = unsafe {
VarcharColumn::new(offsets, data, validity_from_opt(validity, self.row_count)?)
};
ColumnView::Varchar(view)
}
DecodedColumn::Binary {
offsets,
data,
validity,
} => ColumnView::Binary(BinaryColumn::new(
offsets,
data,
validity_from_opt(validity, self.row_count)?,
)),
DecodedColumn::Geohash {
buffer,
byte_width,
precision_bits,
} => ColumnView::Geohash(GeohashColumn::new(
&buffer.values,
*byte_width,
*precision_bits,
validity_of(buffer, self.row_count)?,
)),
DecodedColumn::Decimal128 { buffer, scale } => ColumnView::Decimal128(
Decimal128Column::new(&buffer.values, validity_of(buffer, self.row_count)?, *scale),
),
DecodedColumn::Decimal256 { buffer, scale } => ColumnView::Decimal256(
Decimal256Column::new(&buffer.values, validity_of(buffer, self.row_count)?, *scale),
),
DecodedColumn::DoubleArray(b) => ColumnView::DoubleArray(DoubleArrayColumn::new(
&b.data_offsets,
&b.data,
&b.shapes,
&b.shape_offsets,
validity_from_opt(&b.validity, self.row_count)?,
)),
DecodedColumn::LongArray(b) => ColumnView::LongArray(LongArrayColumn::new(
&b.data_offsets,
&b.data,
&b.shapes,
&b.shape_offsets,
validity_from_opt(&b.validity, self.row_count)?,
)),
})
}
}
#[inline]
fn validity_of<'a>(buf: &'a ColumnBuffer, row_count: usize) -> Result<Validity<'a>> {
validity_from_opt(&buf.validity, row_count)
}
#[inline]
fn validity_from_opt<'a>(validity: &'a Option<Bytes>, row_count: usize) -> Result<Validity<'a>> {
match validity {
None => Ok(Validity::None),
Some(bytes) => Validity::from_bitmap(bytes, row_count),
}
}
pub fn decode_result_batch(
payload: &Bytes,
flags_byte: u8,
dict: &mut SymbolDict,
query_schema: &mut Option<Schema>,
zstd_scratch: &mut ZstdScratch,
) -> Result<DecodedBatch> {
let mut r = ByteReader::new(payload);
let kind = r.read_u8()?;
if kind != MsgKind::ResultBatch.as_u8() {
return Err(fmt!(
ProtocolError,
"expected RESULT_BATCH (0x11), got 0x{:02X}",
kind
));
}
let request_id = r.read_i64_le()?;
let batch_seq = r.read_varint_u64()?;
let _ = &zstd_scratch;
let body: Bytes = if flags_byte & flags::ZSTD != 0 {
#[cfg(feature = "sync-reader-zstd")]
{
zstd_decompress_body(r.remaining(), zstd_scratch)?
}
#[cfg(not(feature = "sync-reader-zstd"))]
{
return Err(fmt!(
UnsupportedServer,
"server sent FLAG_ZSTD batch but client was built without the \
`sync-reader-zstd` feature"
));
}
} else {
payload.slice(r.pos()..)
};
let mut r = ByteReader::new(&body);
if flags_byte & flags::DELTA_SYMBOL_DICT != 0 {
let consumed = dict.apply_delta_from_bytes(r.remaining())?;
r.advance(consumed)?;
}
let name_len = r.read_varint_usize()?;
if name_len > MAX_TABLE_NAME_LENGTH {
return Err(fmt!(
ProtocolError,
"table name length {} exceeds max {}",
name_len,
MAX_TABLE_NAME_LENGTH
));
}
r.read_bytes(name_len)?; let row_count = r.read_varint_usize()?;
if row_count > MAX_ROWS_PER_BATCH {
return Err(fmt!(
ProtocolError,
"table block declares {} rows; max supported is {}",
row_count,
MAX_ROWS_PER_BATCH
));
}
if batch_seq == 0 {
let col_count = r.read_varint_usize()?;
if col_count > MAX_COLUMNS_PER_TABLE {
return Err(fmt!(
ProtocolError,
"table block declares {} columns; max supported is {}",
col_count,
MAX_COLUMNS_PER_TABLE
));
}
let (schema, consumed) = Schema::decode_inline(r.remaining(), col_count)?;
r.advance(consumed)?;
*query_schema = Some(schema);
}
let schema = query_schema.as_ref().ok_or_else(|| {
fmt!(
ProtocolError,
"RESULT_BATCH batch_seq={} arrived before the schema-bearing batch_seq=0",
batch_seq
)
})?;
let col_count = schema.len();
let mut columns = Vec::with_capacity(col_count);
let connection_dict_size = dict.len();
let mut dense_budget = MAX_DECODED_BATCH_BYTES;
for (i, col_meta) in schema.columns().iter().enumerate() {
let kind = col_meta.kind;
let col = decode_column(
&mut r,
&body,
kind,
row_count,
flags_byte,
connection_dict_size,
&mut dense_budget,
)
.map_err(|e| {
Error::new(
e.code(),
format!("column {}/{} ({}): {}", i, col_count, kind.name(), e.msg()),
)
})?;
columns.push(col);
}
if !r.is_empty() {
return Err(fmt!(
ProtocolError,
"RESULT_BATCH has {} trailing bytes",
r.remaining().len()
));
}
Ok(DecodedBatch {
request_id,
batch_seq,
row_count,
columns,
flags: flags_byte,
})
}
fn dense_row_bytes(kind: ColumnKind) -> usize {
match kind {
ColumnKind::Boolean | ColumnKind::Byte => 1,
ColumnKind::Short | ColumnKind::Char => 2,
ColumnKind::Int
| ColumnKind::Float
| ColumnKind::Symbol
| ColumnKind::Varchar
| ColumnKind::Binary
| ColumnKind::Ipv4 => 4,
ColumnKind::Long
| ColumnKind::Double
| ColumnKind::Timestamp
| ColumnKind::TimestampNanos
| ColumnKind::Date
| ColumnKind::Geohash
| ColumnKind::Decimal64
| ColumnKind::DoubleArray
| ColumnKind::LongArray => 8,
ColumnKind::Uuid | ColumnKind::Decimal128 => 16,
ColumnKind::Long256 | ColumnKind::Decimal256 => 32,
}
}
fn decode_column(
r: &mut ByteReader<'_>,
parent: &Bytes,
kind: ColumnKind,
row_count: usize,
flags_byte: u8,
connection_dict_size: usize,
dense_budget: &mut usize,
) -> Result<DecodedColumn> {
let cost = row_count.saturating_mul(dense_row_bytes(kind));
*dense_budget = dense_budget.checked_sub(cost).ok_or_else(|| {
fmt!(
ProtocolError,
"RESULT_BATCH densified data exceeds the {}-byte per-batch cap; \
a null-sparse batch cannot inflate client memory without bound",
MAX_DECODED_BATCH_BYTES
)
})?;
Ok(match kind {
ColumnKind::Boolean => DecodedColumn::Boolean(decode_boolean(r, parent, row_count)?),
ColumnKind::Byte => {
DecodedColumn::Byte(decode_fixed_non_nullable(r, parent, row_count, 1, "BYTE")?)
}
ColumnKind::Short => {
DecodedColumn::Short(decode_fixed_non_nullable(r, parent, row_count, 2, "SHORT")?)
}
ColumnKind::Int => DecodedColumn::Int(decode_fixed(
r,
parent,
row_count,
4,
Some(&null_sentinel::I32_LE),
)?),
ColumnKind::Long => DecodedColumn::Long(decode_fixed(
r,
parent,
row_count,
8,
Some(&null_sentinel::I64_LE),
)?),
ColumnKind::Float => DecodedColumn::Float(decode_fixed(
r,
parent,
row_count,
4,
Some(&null_sentinel::F32_NAN_LE),
)?),
ColumnKind::Double => DecodedColumn::Double(decode_fixed(
r,
parent,
row_count,
8,
Some(&null_sentinel::F64_NAN_LE),
)?),
ColumnKind::Char => {
DecodedColumn::Char(decode_fixed_non_nullable(r, parent, row_count, 2, "CHAR")?)
}
ColumnKind::Ipv4 => DecodedColumn::Ipv4(decode_fixed(r, parent, row_count, 4, None)?),
ColumnKind::Uuid => DecodedColumn::Uuid(decode_fixed(
r,
parent,
row_count,
16,
Some(&null_sentinel::UUID_LE),
)?),
ColumnKind::Long256 => DecodedColumn::Long256(decode_fixed(
r,
parent,
row_count,
32,
Some(&null_sentinel::LONG256_LE),
)?),
ColumnKind::Timestamp => {
DecodedColumn::Timestamp(decode_temporal(r, parent, row_count, flags_byte)?)
}
ColumnKind::Date => DecodedColumn::Date(decode_temporal(r, parent, row_count, flags_byte)?),
ColumnKind::TimestampNanos => {
DecodedColumn::TimestampNanos(decode_temporal(r, parent, row_count, flags_byte)?)
}
ColumnKind::Symbol => {
let (codes, validity, local_dict) =
decode_symbol(r, parent, row_count, flags_byte, connection_dict_size)?;
DecodedColumn::Symbol {
codes,
validity,
local_dict,
}
}
ColumnKind::Decimal64 => {
let (scale, buffer) = decode_decimal64(r, parent, row_count)?;
DecodedColumn::Decimal64 { buffer, scale }
}
ColumnKind::Varchar => {
let (offsets, data, validity) =
decode_varlen(r, parent, row_count, true)?;
DecodedColumn::Varchar {
offsets,
data,
validity,
}
}
ColumnKind::Binary => {
let (offsets, data, validity) =
decode_varlen(r, parent, row_count, false)?;
DecodedColumn::Binary {
offsets,
data,
validity,
}
}
ColumnKind::Geohash => {
let (buffer, byte_width, precision_bits) = decode_geohash(r, parent, row_count)?;
DecodedColumn::Geohash {
buffer,
byte_width,
precision_bits,
}
}
ColumnKind::Decimal128 => {
let (scale, buffer) = decode_decimal_wide(r, parent, row_count, 16)?;
DecodedColumn::Decimal128 { buffer, scale }
}
ColumnKind::Decimal256 => {
let (scale, buffer) = decode_decimal_wide(r, parent, row_count, 32)?;
DecodedColumn::Decimal256 { buffer, scale }
}
ColumnKind::DoubleArray => DecodedColumn::DoubleArray(decode_array(r, parent, row_count)?),
ColumnKind::LongArray => DecodedColumn::LongArray(decode_array(r, parent, row_count)?),
})
}
const MAX_ARRAY_ELEMENTS_PER_ROW: u64 = 16 * 1024 * 1024;
const MAX_ARRAY_DIMS: usize = 32;
fn decode_array(r: &mut ByteReader<'_>, parent: &Bytes, row_count: usize) -> Result<ArrayBuffers> {
let validity = decode_validity(r, parent, row_count)?;
let mut data_offsets = Vec::with_capacity(row_count + 1);
let mut data: Vec<u8> = Vec::new();
let mut shapes: Vec<u32> = Vec::new();
let mut shape_offsets = Vec::with_capacity(row_count + 1);
data_offsets.push(0u32);
shape_offsets.push(0u32);
for row in 0..row_count {
if is_null_at_opt(&validity, row) {
data_offsets.push(*data_offsets.last().unwrap());
shape_offsets.push(*shape_offsets.last().unwrap());
continue;
}
let n_dims = r.read_u8()? as usize;
if n_dims == 0 {
return Err(fmt!(
ProtocolError,
"array row {} has nDims=0 (must be >= 1)",
row
));
}
if n_dims > MAX_ARRAY_DIMS {
return Err(fmt!(
ProtocolError,
"array row {row} has nDims={n_dims}; max {MAX_ARRAY_DIMS}"
));
}
let mut total: u64 = 1;
let dims_start = shapes.len();
for d in 0..n_dims {
let dim_bytes = r.read_bytes(4)?;
let dim = u32::from_le_bytes(dim_bytes.try_into().unwrap());
shapes.push(dim);
total = total.checked_mul(dim as u64).ok_or_else(|| {
fmt!(
ProtocolError,
"array row {} shape product overflow at dim {}",
row,
d
)
})?;
if total > MAX_ARRAY_ELEMENTS_PER_ROW {
return Err(fmt!(
LimitExceeded,
"array row {} has {} elements (max {})",
row,
total,
MAX_ARRAY_ELEMENTS_PER_ROW
));
}
}
let byte_count = (total as usize)
.checked_mul(8)
.ok_or_else(|| fmt!(ProtocolError, "array row {} byte count overflow", row))?;
let elements = r.read_bytes(byte_count)?;
data.extend_from_slice(elements);
let new_data_off = u32::try_from(data.len())
.map_err(|_| fmt!(ProtocolError, "array column data exceeds u32 byte offset"))?;
data_offsets.push(new_data_off);
let new_shape_off = u32::try_from(dims_start + n_dims)
.map_err(|_| fmt!(ProtocolError, "array column shape table exceeds u32"))?;
shape_offsets.push(new_shape_off);
}
Ok(ArrayBuffers {
data_offsets,
data: Bytes::from(data),
shapes,
shape_offsets,
validity,
})
}
fn decode_geohash(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
) -> Result<(ColumnBuffer, u8, u8)> {
let validity = decode_validity(r, parent, row_count)?;
let precision_bits = r.read_varint_u64()?;
if precision_bits == 0 || precision_bits > 60 {
return Err(fmt!(
ProtocolError,
"geohash precision_bits {} outside 1..=60",
precision_bits
));
}
let byte_width = precision_bits.div_ceil(8) as u8;
let sentinel = &null_sentinel::GEOHASH_FF[..byte_width as usize];
let buffer = densify_fixed(
r,
parent,
row_count,
byte_width as usize,
validity,
Some(sentinel),
)?;
Ok((buffer, byte_width, precision_bits as u8))
}
fn decode_decimal_wide(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
width: usize,
) -> Result<(i8, ColumnBuffer)> {
let validity = decode_validity(r, parent, row_count)?;
let scale = r.read_u8()? as i8;
let per_width_max: i8 = match width {
8 => 18,
16 => 38,
_ => 76,
};
if !(0..=per_width_max).contains(&scale) {
return Err(fmt!(
ProtocolError,
"DECIMAL{} scale {} outside 0..={}",
width * 8,
scale,
per_width_max
));
}
let sentinel: &[u8] = match width {
8 => &null_sentinel::I64_LE,
16 => &null_sentinel::UUID_LE,
32 => &null_sentinel::LONG256_LE,
other => {
return Err(fmt!(
ProtocolError,
"DECIMAL width must be 8/16/32, got {other}"
));
}
};
let buffer = densify_fixed(r, parent, row_count, width, validity, Some(sentinel))?;
Ok((scale, buffer))
}
mod null_sentinel {
pub const I32_LE: [u8; 4] = i32::MIN.to_le_bytes();
pub const I64_LE: [u8; 8] = i64::MIN.to_le_bytes();
pub const F32_NAN_LE: [u8; 4] = 0x7FC0_0000u32.to_le_bytes();
pub const F64_NAN_LE: [u8; 8] = 0x7FF8_0000_0000_0000u64.to_le_bytes();
pub const UUID_LE: [u8; 16] = [
0, 0, 0, 0, 0, 0, 0, 0x80, 0, 0, 0, 0, 0, 0, 0, 0x80, ];
pub const LONG256_LE: [u8; 32] = [
0, 0, 0, 0, 0, 0, 0, 0x80, 0, 0, 0, 0, 0, 0, 0, 0x80, 0, 0, 0, 0, 0, 0, 0, 0x80, 0, 0, 0,
0, 0, 0, 0, 0x80,
];
pub const GEOHASH_FF: [u8; 8] = [0xFF; 8];
}
fn densify_fixed(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
elem_size: usize,
validity: Option<Bytes>,
null_sentinel: Option<&[u8]>,
) -> Result<ColumnBuffer> {
debug_assert!(
null_sentinel.is_none_or(|s| s.len() == elem_size),
"null_sentinel length must equal elem_size"
);
let dense_len = row_count
.checked_mul(elem_size)
.ok_or_else(|| fmt!(ProtocolError, "fixed column size overflow"))?;
match &validity {
None => {
let values = read_owned(r, parent, dense_len)?;
Ok(ColumnBuffer { values, validity })
}
Some(bitmap) => {
let non_null = row_count - count_nulls(bitmap, row_count);
let compact = r.read_bytes(non_null * elem_size)?;
let mut dense = allocate_dense_with_sentinel(dense_len, elem_size, null_sentinel);
let mut src = 0usize;
for row in 0..row_count {
if !is_null_at(bitmap, row) {
let dst = row * elem_size;
dense[dst..dst + elem_size].copy_from_slice(&compact[src..src + elem_size]);
src += elem_size;
}
}
Ok(ColumnBuffer {
values: Bytes::from(dense),
validity,
})
}
}
}
fn allocate_dense_with_sentinel(
dense_len: usize,
elem_size: usize,
sentinel: Option<&[u8]>,
) -> Vec<u8> {
debug_assert_eq!(
dense_len % elem_size,
0,
"dense_len {dense_len} not a multiple of elem_size {elem_size}"
);
match sentinel {
Some(s) if s.iter().any(|&b| b != 0) => {
let mut dense = vec![0u8; dense_len];
for chunk in dense.chunks_exact_mut(elem_size) {
chunk.copy_from_slice(s);
}
dense
}
_ => vec![0u8; dense_len],
}
}
type VarlenBuffers = (Vec<u32>, Bytes, Option<Bytes>);
fn decode_varlen(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
utf8: bool,
) -> Result<VarlenBuffers> {
let validity = decode_validity(r, parent, row_count)?;
let non_null = match &validity {
None => row_count,
Some(bitmap) => row_count - count_nulls(bitmap, row_count),
};
let offsets_byte_len = (non_null + 1)
.checked_mul(4)
.ok_or_else(|| fmt!(ProtocolError, "varlen offsets size overflow"))?;
let offsets_bytes = r.read_bytes(offsets_byte_len)?;
let count = non_null + 1;
let mut compact: Vec<u32> = Vec::with_capacity(count);
unsafe {
std::ptr::copy_nonoverlapping(
offsets_bytes.as_ptr(),
compact.as_mut_ptr().cast::<u8>(),
offsets_byte_len,
);
compact.set_len(count);
}
#[cfg(target_endian = "big")]
for v in &mut compact {
*v = v.swap_bytes();
}
if compact[0] != 0 {
return Err(fmt!(
ProtocolError,
"varlen offsets must start at 0, got {}",
compact[0]
));
}
for i in 1..compact.len() {
if compact[i] < compact[i - 1] {
return Err(fmt!(
ProtocolError,
"varlen offsets not monotonic at index {}: {} < {}",
i,
compact[i],
compact[i - 1]
));
}
}
let data_len = compact[non_null] as usize;
let data = read_owned(r, parent, data_len)?;
if utf8 {
let s = std::str::from_utf8(&data)
.map_err(|e| fmt!(InvalidUtf8, "varchar data buffer not valid UTF-8: {}", e))?;
for &off in &compact {
if !s.is_char_boundary(off as usize) {
return Err(fmt!(
InvalidUtf8,
"varchar offset {} does not lie on a UTF-8 codepoint boundary",
off
));
}
}
}
if validity.is_none() {
debug_assert_eq!(compact.len(), row_count + 1);
return Ok((compact, data, validity));
}
let mut dense = vec![0u32; row_count + 1];
let mut k = 0usize; for row in 0..row_count {
if is_null_at_opt(&validity, row) {
dense[row + 1] = dense[row];
} else {
let len = compact[k + 1] - compact[k];
dense[row + 1] = dense[row] + len;
k += 1;
}
}
Ok((dense, data, validity))
}
fn decode_validity(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
) -> Result<Option<Bytes>> {
let null_flag = r.read_u8()?;
match null_flag {
0 => Ok(None),
1 => {
let bitmap_len = row_count.div_ceil(8);
Ok(Some(read_owned(r, parent, bitmap_len)?))
}
other => Err(fmt!(
ProtocolError,
"unknown null_flag 0x{:02X}; expected 0x00 or 0x01",
other
)),
}
}
fn decode_fixed(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
elem_size: usize,
null_sentinel: Option<&[u8]>,
) -> Result<ColumnBuffer> {
let validity = decode_validity(r, parent, row_count)?;
densify_fixed(r, parent, row_count, elem_size, validity, null_sentinel)
}
fn expect_no_validity_flag(r: &mut ByteReader<'_>, kind: &str) -> Result<()> {
let null_flag = r.read_u8()?;
if null_flag != 0 {
return Err(fmt!(
ProtocolError,
"{} column has null_flag 0x{:02X}; spec requires 0x00 \
({} is not nullable on the wire)",
kind,
null_flag,
kind
));
}
Ok(())
}
fn decode_fixed_non_nullable(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
elem_size: usize,
kind: &str,
) -> Result<ColumnBuffer> {
expect_no_validity_flag(r, kind)?;
let byte_count = row_count.checked_mul(elem_size).ok_or_else(|| {
fmt!(
ProtocolError,
"{} column byte count overflow (row_count={}, elem_size={})",
kind,
row_count,
elem_size
)
})?;
let values = read_owned(r, parent, byte_count)?;
Ok(ColumnBuffer {
values,
validity: None,
})
}
fn decode_boolean(
r: &mut ByteReader<'_>,
_parent: &Bytes,
row_count: usize,
) -> Result<ColumnBuffer> {
expect_no_validity_flag(r, "BOOLEAN")?;
let bit_bytes = row_count.div_ceil(8);
let bits = r.read_bytes(bit_bytes)?;
let mut dense = vec![0u8; row_count];
for (row, slot) in dense.iter_mut().enumerate() {
let b = bits[row >> 3];
*slot = (b >> (row & 7)) & 1;
}
Ok(ColumnBuffer {
values: Bytes::from(dense),
validity: None,
})
}
fn decode_temporal(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
flags_byte: u8,
) -> Result<ColumnBuffer> {
let sentinel = Some(&null_sentinel::I64_LE[..]);
if flags_byte & flags::GORILLA == 0 {
return decode_fixed(r, parent, row_count, 8, sentinel);
}
let validity = decode_validity(r, parent, row_count)?;
let non_null = match &validity {
None => row_count,
Some(bitmap) => row_count - count_nulls(bitmap, row_count),
};
let disc = r.read_u8()?;
match disc {
0x00 => densify_fixed(r, parent, row_count, 8, validity, sentinel),
0x01 => decode_gorilla_temporal(r, row_count, non_null, validity),
other => Err(fmt!(
ProtocolError,
"unknown temporal encoding discriminator 0x{:02X}",
other
)),
}
}
fn decode_gorilla_temporal(
r: &mut ByteReader<'_>,
row_count: usize,
non_null: usize,
validity: Option<Bytes>,
) -> Result<ColumnBuffer> {
let dense_len = row_count
.checked_mul(8)
.ok_or_else(|| fmt!(ProtocolError, "gorilla temporal column size overflow"))?;
let mut dense = allocate_dense_with_sentinel(dense_len, 8, Some(&null_sentinel::I64_LE));
let mut seeds = [0i64; 2];
let seed_count = non_null.min(2);
for seed in seeds.iter_mut().take(seed_count) {
*seed = i64::from_le_bytes(r.read_bytes(8)?.try_into().unwrap());
}
let mut decoder = if non_null >= 3 {
Some(crate::egress::gorilla::GorillaDecoder::new(
seeds[0],
seeds[1],
r.remaining(),
))
} else {
None
};
let mut filled = 0usize;
for row in 0..row_count {
if is_null_at_opt(&validity, row) {
continue;
}
let v = if filled < seed_count {
seeds[filled]
} else {
let dec = decoder.as_mut().ok_or_else(|| {
fmt!(
ProtocolError,
"Gorilla decoder state: non_null={non_null}, seed_count={seed_count}, filled={filled}"
)
})?;
dec.decode_next()?
};
dense[row * 8..row * 8 + 8].copy_from_slice(&v.to_le_bytes());
filled += 1;
if filled == non_null {
break;
}
}
if let Some(d) = decoder {
r.advance(d.bytes_consumed())?;
}
Ok(ColumnBuffer {
values: Bytes::from(dense),
validity,
})
}
type SymbolBuffers = (Vec<u32>, Option<Bytes>, Option<SymbolDict>);
fn decode_symbol(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
flags_byte: u8,
connection_dict_size: usize,
) -> Result<SymbolBuffers> {
let validity = decode_validity(r, parent, row_count)?;
let (active_dict_size, local_dict) = if flags_byte & flags::DELTA_SYMBOL_DICT != 0 {
(connection_dict_size, None)
} else {
let dict_size = r.read_varint_usize()?;
if dict_size > row_count {
return Err(fmt!(
ProtocolError,
"SYMBOL column-local dict_size {} > row_count {}",
dict_size,
row_count
));
}
let mut entries: Vec<&[u8]> = Vec::with_capacity(dict_size);
for i in 0..dict_size {
let entry_len = r.read_varint_usize().map_err(|e| {
Error::new(
e.code(),
format!("SYMBOL local dict entry {} length: {}", i, e.msg()),
)
})?;
entries.push(r.read_bytes(entry_len)?);
}
let mut local = SymbolDict::new();
local.apply_delta(0, entries)?;
(dict_size, Some(local))
};
let codes = if validity.is_none() {
decode_codes_no_nulls(r, row_count, active_dict_size)?
} else {
let mut codes = vec![0u32; row_count];
for (row, slot) in codes.iter_mut().enumerate() {
if is_null_at_opt(&validity, row) {
continue;
}
let code = r.read_varint_u64().map_err(|e| {
Error::new(e.code(), format!("symbol code at row {}: {}", row, e.msg()))
})?;
let code32 = u32::try_from(code).map_err(|_| {
fmt!(
ProtocolError,
"symbol code {} at row {} exceeds u32",
code,
row
)
})?;
if (code32 as usize) >= active_dict_size {
return Err(fmt!(
ProtocolError,
"symbol id {} at row {} out of range (dict size {})",
code32,
row,
active_dict_size
));
}
*slot = code32;
}
codes
};
Ok((codes, validity, local_dict))
}
fn decode_codes_no_nulls(
r: &mut ByteReader<'_>,
row_count: usize,
active_dict_size: usize,
) -> Result<Vec<u32>> {
let mut codes = vec![0u32; row_count];
let bytes = r.remaining();
let mut pos = 0usize;
let limit = bytes.len();
for slot in codes.iter_mut() {
if pos + 3 <= limit {
let b0 = bytes[pos];
if b0 < 0x80 {
*slot = b0 as u32;
pos += 1;
continue;
}
let b1 = bytes[pos + 1];
if b1 < 0x80 {
*slot = (b0 & 0x7F) as u32 | ((b1 as u32) << 7);
pos += 2;
continue;
}
let b2 = bytes[pos + 2];
if b2 < 0x80 {
*slot = (b0 & 0x7F) as u32 | (((b1 & 0x7F) as u32) << 7) | ((b2 as u32) << 14);
pos += 3;
continue;
}
}
let (v, n) = crate::egress::wire::varint::decode_u64(&bytes[pos..])
.map_err(|e| Error::new(e.code(), format!("symbol code: {}", e.msg())))?;
*slot =
u32::try_from(v).map_err(|_| fmt!(ProtocolError, "symbol code {} exceeds u32", v))?;
pos += n;
}
r.advance(pos)?;
let dict_size_u32 = u32::try_from(active_dict_size).map_err(|_| {
fmt!(
ProtocolError,
"active dict size {} exceeds u32",
active_dict_size
)
})?;
if codes.iter().any(|&c| c >= dict_size_u32) {
let (row, &bad) = codes
.iter()
.enumerate()
.find(|&(_, &c)| c >= dict_size_u32)
.expect("any() reported a match");
return Err(fmt!(
ProtocolError,
"symbol id {} at row {} out of range (dict size {})",
bad,
row,
active_dict_size
));
}
Ok(codes)
}
fn decode_decimal64(
r: &mut ByteReader<'_>,
parent: &Bytes,
row_count: usize,
) -> Result<(i8, ColumnBuffer)> {
let (scale, buffer) = decode_decimal_wide(r, parent, row_count, 8)?;
Ok((scale, buffer))
}
#[cfg(feature = "sync-reader-zstd")]
const MAX_ZSTD_DECOMPRESSED: u64 = 64 * 1024 * 1024;
#[cfg(feature = "sync-reader-zstd")]
const ZSTD_POOL_CAPACITY: usize = 2;
#[cfg(feature = "sync-reader-zstd")]
#[derive(Default)]
struct ZstdBufferPool {
buffers: std::sync::Mutex<Vec<Vec<u8>>>,
}
#[cfg(feature = "sync-reader-zstd")]
struct PooledZstdBuffer {
buf: Vec<u8>,
pool: std::sync::Arc<ZstdBufferPool>,
}
#[cfg(feature = "sync-reader-zstd")]
impl AsRef<[u8]> for PooledZstdBuffer {
#[inline]
fn as_ref(&self) -> &[u8] {
&self.buf
}
}
#[cfg(feature = "sync-reader-zstd")]
impl Drop for PooledZstdBuffer {
fn drop(&mut self) {
let Ok(mut guard) = self.pool.buffers.lock() else {
return;
};
if guard.len() >= ZSTD_POOL_CAPACITY {
return;
}
let buf = std::mem::take(&mut self.buf);
if buf.capacity() > 0 {
guard.push(buf);
}
}
}
#[derive(Default)]
pub struct ZstdScratch {
#[cfg(feature = "sync-reader-zstd")]
decompressor: Option<zstd::bulk::Decompressor<'static>>,
#[cfg(feature = "sync-reader-zstd")]
pool: std::sync::Arc<ZstdBufferPool>,
}
impl ZstdScratch {
pub fn new() -> Self {
Self::default()
}
}
#[cfg(feature = "sync-reader-zstd")]
fn zstd_decompress_body(compressed: &[u8], scratch: &mut ZstdScratch) -> Result<Bytes> {
let size = match zstd::zstd_safe::get_frame_content_size(compressed) {
Ok(Some(n)) => n,
Ok(None) => {
return Err(fmt!(
ProtocolError,
"zstd frame missing content size (protocol violation)"
));
}
Err(_) => {
return Err(fmt!(
ProtocolError,
"invalid zstd frame header (truncated, bad magic, or content size > u64::MAX)"
));
}
};
if size > MAX_ZSTD_DECOMPRESSED {
return Err(fmt!(
LimitExceeded,
"zstd frame content size {} exceeds client cap {}",
size,
MAX_ZSTD_DECOMPRESSED
));
}
let usize_size = usize::try_from(size).map_err(|_| {
fmt!(
LimitExceeded,
"zstd frame content size {} does not fit in usize",
size
)
})?;
let decompressor = match scratch.decompressor.as_mut() {
Some(d) => d,
None => {
scratch.decompressor = Some(
zstd::bulk::Decompressor::new()
.map_err(|e| fmt!(ProtocolError, "zstd decompressor init failed: {}", e))?,
);
scratch.decompressor.as_mut().unwrap()
}
};
let mut buf = scratch
.pool
.buffers
.lock()
.ok()
.and_then(|mut g| g.pop())
.unwrap_or_default();
buf.clear();
buf.reserve(usize_size);
let written = decompressor
.decompress_to_buffer(compressed, &mut buf)
.map_err(|e| fmt!(ProtocolError, "zstd decompress failed: {}", e))?;
if written != usize_size {
return Err(fmt!(
ProtocolError,
"zstd decompressed size {} != frame content size {}",
written,
size
));
}
buf.truncate(usize_size);
let owner = PooledZstdBuffer {
buf,
pool: std::sync::Arc::clone(&scratch.pool),
};
Ok(Bytes::from_owner(owner))
}
fn count_nulls(bitmap: &[u8], row_count: usize) -> usize {
let full_bytes = row_count >> 3;
let tail_bits = row_count & 7;
let body = &bitmap[..full_bytes];
let mut chunks = body.chunks_exact(8);
let mut nulls: usize = 0;
for c in chunks.by_ref() {
let w = u64::from_ne_bytes(c.try_into().unwrap());
nulls += w.count_ones() as usize;
}
for b in chunks.remainder() {
nulls += b.count_ones() as usize;
}
if tail_bits != 0 {
let mask = (1u8 << tail_bits) - 1;
nulls += (bitmap[full_bytes] & mask).count_ones() as usize;
}
nulls
}
fn is_null_at(bitmap: &[u8], row: usize) -> bool {
(bitmap[row >> 3] >> (row & 7)) & 1 != 0
}
fn is_null_at_opt(validity: &Option<Bytes>, row: usize) -> bool {
match validity {
None => false,
Some(bitmap) => is_null_at(bitmap, row),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::egress::schema::{Schema, SchemaColumn};
use crate::egress::wire::varint::encode_u64;
use crate::error::ErrorCode;
fn count_nulls_naive(bitmap: &[u8], row_count: usize) -> usize {
let full_bytes = row_count >> 3;
let tail_bits = row_count & 7;
let mut nulls = 0usize;
for b in &bitmap[..full_bytes] {
nulls += b.count_ones() as usize;
}
if tail_bits != 0 {
let mask = (1u8 << tail_bits) - 1;
nulls += (bitmap[full_bytes] & mask).count_ones() as usize;
}
nulls
}
#[test]
fn count_nulls_chunked_matches_naive_across_boundaries() {
let mut bitmap = vec![0u8; 32];
for (i, b) in bitmap.iter_mut().enumerate() {
*b = match i % 4 {
0 => 0x00,
1 => 0xFF,
2 => 0xA5,
_ => 0x5A,
};
}
let row_counts: &[usize] = &[
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 15, 16, 23, 24, 56, 57, 62, 63, 64, 65, 66, 67, 71, 72, 73, 79, 80, 81, 87, 88, 127, 128, 129, 191, 192, 193, 255, 256, ];
for &rc in row_counts {
let got = count_nulls(&bitmap, rc);
let want = count_nulls_naive(&bitmap, rc);
assert_eq!(got, want, "count_nulls mismatch at row_count={}", rc);
}
}
#[cfg(not(feature = "sync-reader-zstd"))]
#[test]
fn zstd_flag_rejected_without_feature() {
let mut payload = vec![MsgKind::ResultBatch.as_u8()];
payload.extend_from_slice(&0i64.to_le_bytes());
payload.push(0u8); let payload = Bytes::from(payload);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags::ZSTD,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.expect_err("decoder must reject FLAG_ZSTD when built without sync-reader-zstd");
assert_eq!(err.code(), ErrorCode::UnsupportedServer);
assert!(
err.msg().contains("sync-reader-zstd"),
"rejection message should name the missing feature: {}",
err.msg()
);
}
#[test]
fn count_nulls_tight_buffer_matches_naive() {
for row_count in [0usize, 1, 7, 8, 9, 63, 64, 65, 100, 1000] {
let bytes_needed = row_count.div_ceil(8);
let mut bitmap = vec![0u8; bytes_needed];
for (i, b) in bitmap.iter_mut().enumerate() {
*b = ((i.wrapping_mul(31) ^ 0xA5) & 0xFF) as u8;
}
let got = count_nulls(&bitmap, row_count);
let want = count_nulls_naive(&bitmap, row_count);
assert_eq!(
got, want,
"tight-buffer mismatch at row_count={} ({} bytes)",
row_count, bytes_needed
);
}
}
struct BatchBuilder {
flags: u8,
request_id: i64,
batch_seq: u64,
delta: Option<Vec<&'static str>>, delta_start: u64,
row_count: usize,
cols: Vec<(String, ColumnKind)>,
column_data: Vec<Vec<u8>>,
}
impl BatchBuilder {
fn new(row_count: usize) -> Self {
Self {
flags: 0,
request_id: 1,
batch_seq: 0,
delta: None,
delta_start: 0,
row_count,
cols: Vec::new(),
column_data: Vec::new(),
}
}
fn with_flags(mut self, f: u8) -> Self {
self.flags = f;
self
}
fn with_dict_delta(mut self, start: u64, entries: Vec<&'static str>) -> Self {
self.flags |= flags::DELTA_SYMBOL_DICT;
self.delta_start = start;
self.delta = Some(entries);
self
}
fn with_batch_seq(mut self, seq: u64) -> Self {
self.batch_seq = seq;
self
}
fn add_column(mut self, name: &str, kind: ColumnKind, data: Vec<u8>) -> Self {
self.cols.push((name.to_string(), kind));
self.column_data.push(data);
self
}
fn build(self) -> (u8, Bytes) {
let mut out = Vec::new();
out.push(MsgKind::ResultBatch.as_u8());
out.extend_from_slice(&self.request_id.to_le_bytes());
encode_u64(self.batch_seq, &mut out);
if let Some(entries) = self.delta {
encode_u64(self.delta_start, &mut out);
encode_u64(entries.len() as u64, &mut out);
for e in entries {
encode_u64(e.len() as u64, &mut out);
out.extend_from_slice(e.as_bytes());
}
}
encode_u64(0, &mut out); encode_u64(self.row_count as u64, &mut out);
if self.batch_seq == 0 {
encode_u64(self.cols.len() as u64, &mut out);
for (name, kind) in &self.cols {
encode_u64(name.len() as u64, &mut out);
out.extend_from_slice(name.as_bytes());
out.push(kind.as_u8());
}
}
for data in self.column_data {
out.extend_from_slice(&data);
}
(self.flags, Bytes::from(out))
}
}
fn col_no_nulls(values: &[u8]) -> Vec<u8> {
let mut out = vec![0x00]; out.extend_from_slice(values);
out
}
fn col_with_bitmap(bitmap: &[u8], values: &[u8]) -> Vec<u8> {
let mut out = vec![0x01]; out.extend_from_slice(bitmap);
out.extend_from_slice(values);
out
}
fn le_i64s(vs: &[i64]) -> Vec<u8> {
let mut o = Vec::new();
for v in vs {
o.extend_from_slice(&v.to_le_bytes());
}
o
}
#[test]
fn decode_simple_long_no_nulls() {
let (flags_byte, payload) = BatchBuilder::new(3)
.add_column("v", ColumnKind::Long, col_no_nulls(&le_i64s(&[1, 2, 3])))
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert_eq!(batch.row_count, 3);
assert_eq!(batch.columns.len(), 1);
let view = batch.column_view(0, &dict).unwrap();
match view {
ColumnView::Long(c) => {
assert_eq!(c.len(), 3);
assert_eq!(c.value(0), 1);
assert_eq!(c.value(1), 2);
assert_eq!(c.value(2), 3);
}
other => panic!("unexpected view: {:?}", other.kind()),
}
}
#[test]
fn decode_long_with_nulls_densifies() {
let (flags_byte, payload) = BatchBuilder::new(4)
.add_column(
"v",
ColumnKind::Long,
col_with_bitmap(&[0x02], &le_i64s(&[10, 30, 40])),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Long(c) = view else { panic!() };
assert!(!c.is_null(0));
assert!(c.is_null(1));
assert!(!c.is_null(2));
assert!(!c.is_null(3));
assert_eq!(c.value(0), 10);
assert_eq!(c.value(1), i64::MIN);
assert_eq!(c.value(2), 30);
assert_eq!(c.value(3), 40);
}
#[test]
fn decode_long_densifies_multiple_nulls() {
let (flags_byte, payload) = BatchBuilder::new(8)
.add_column(
"v",
ColumnKind::Long,
col_with_bitmap(&[0x92], &le_i64s(&[100, 102, 103, 105, 106])),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Long(c) = view else { panic!() };
let expected: Vec<Option<i64>> = vec![
Some(100),
None,
Some(102),
Some(103),
None,
Some(105),
Some(106),
None,
];
let got: Vec<Option<i64>> = (0..8)
.map(|r| if c.is_null(r) { None } else { Some(c.value(r)) })
.collect();
assert_eq!(got, expected);
}
fn le_f64s(vs: &[f64]) -> Vec<u8> {
let mut o = Vec::new();
for v in vs {
o.extend_from_slice(&v.to_le_bytes());
}
o
}
fn le_f32s(vs: &[f32]) -> Vec<u8> {
let mut o = Vec::new();
for v in vs {
o.extend_from_slice(&v.to_le_bytes());
}
o
}
fn le_i32s(vs: &[i32]) -> Vec<u8> {
let mut o = Vec::new();
for v in vs {
o.extend_from_slice(&v.to_le_bytes());
}
o
}
fn le_i16s(vs: &[i16]) -> Vec<u8> {
let mut o = Vec::new();
for v in vs {
o.extend_from_slice(&v.to_le_bytes());
}
o
}
#[test]
fn decode_double_with_nulls_densifies() {
let (flags_byte, payload) = BatchBuilder::new(4)
.add_column(
"v",
ColumnKind::Double,
col_with_bitmap(&[0x02], &le_f64s(&[1.5, 3.5, 4.5])),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Double(c) = view else {
panic!()
};
assert!(!c.is_null(0));
assert!(c.is_null(1));
assert!(!c.is_null(2));
assert!(!c.is_null(3));
assert_eq!(c.value(0), 1.5);
assert_eq!(c.value(1).to_bits(), 0x7FF8_0000_0000_0000u64);
assert_eq!(c.value(2), 3.5);
assert_eq!(c.value(3), 4.5);
}
#[test]
fn decode_float_with_nulls_densifies() {
let (flags_byte, payload) = BatchBuilder::new(4)
.add_column(
"v",
ColumnKind::Float,
col_with_bitmap(&[0x02], &le_f32s(&[1.5_f32, 3.5_f32, 4.5_f32])),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Float(c) = view else { panic!() };
assert_eq!(c.value(0), 1.5_f32);
assert_eq!(c.value(1).to_bits(), 0x7FC0_0000u32);
assert_eq!(c.value(2), 3.5_f32);
assert_eq!(c.value(3), 4.5_f32);
}
#[test]
fn decode_int_with_nulls_densifies() {
let (flags_byte, payload) = BatchBuilder::new(8)
.add_column(
"v",
ColumnKind::Int,
col_with_bitmap(&[0x92], &le_i32s(&[10, 12, 13, 15, 16])),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Int(c) = view else { panic!() };
let expected: Vec<Option<i32>> = vec![
Some(10),
None,
Some(12),
Some(13),
None,
Some(15),
Some(16),
None,
];
let got: Vec<Option<i32>> = (0..8)
.map(|r| if c.is_null(r) { None } else { Some(c.value(r)) })
.collect();
assert_eq!(got, expected);
assert_eq!(c.value(1), i32::MIN);
assert_eq!(c.value(4), i32::MIN);
assert_eq!(c.value(7), i32::MIN);
}
#[test]
fn null_sentinels_per_spec_11_5() {
let bitmap = vec![0x02];
let mut uuid_vals = vec![0u8; 16];
uuid_vals[..8].copy_from_slice(&1i64.to_le_bytes());
uuid_vals[8..16].copy_from_slice(&2i64.to_le_bytes());
let mut long256_vals = vec![0u8; 32];
for chunk in 0..4 {
long256_vals[chunk * 8..chunk * 8 + 8]
.copy_from_slice(&((chunk + 1) as i64).to_le_bytes());
}
let cases: &[(ColumnKind, Vec<u8>, &[u8])] = &[
(ColumnKind::Int, le_i32s(&[7]), &i32::MIN.to_le_bytes()),
(ColumnKind::Long, le_i64s(&[7]), &i64::MIN.to_le_bytes()),
(
ColumnKind::Float,
le_f32s(&[1.5]),
&0x7FC0_0000u32.to_le_bytes(),
),
(
ColumnKind::Double,
le_f64s(&[1.5]),
&0x7FF8_0000_0000_0000u64.to_le_bytes(),
),
(
ColumnKind::Uuid,
uuid_vals,
&[0, 0, 0, 0, 0, 0, 0, 0x80, 0, 0, 0, 0, 0, 0, 0, 0x80],
),
(
ColumnKind::Long256,
long256_vals,
&[
0, 0, 0, 0, 0, 0, 0, 0x80, 0, 0, 0, 0, 0, 0, 0, 0x80, 0, 0, 0, 0, 0, 0, 0,
0x80, 0, 0, 0, 0, 0, 0, 0, 0x80,
],
),
];
for (kind, value_bytes, expected_null) in cases {
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("v", *kind, col_with_bitmap(&bitmap, value_bytes))
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_or_else(|e| panic!("{:?}: {}", kind, e.msg()));
let view = batch.column_view(0, &dict).unwrap();
let null_bytes: Vec<u8> = match view {
ColumnView::Int(c) => c.value(1).to_le_bytes().to_vec(),
ColumnView::Long(c) => c.value(1).to_le_bytes().to_vec(),
ColumnView::Float(c) => c.value(1).to_bits().to_le_bytes().to_vec(),
ColumnView::Double(c) => c.value(1).to_bits().to_le_bytes().to_vec(),
ColumnView::Uuid(c) => c.value(1).to_vec(),
ColumnView::Long256(c) => c.value(1).to_vec(),
_ => panic!("unexpected view for {:?}", kind),
};
assert_eq!(
null_bytes, *expected_null,
"spec §11.5 NULL sentinel mismatch for {:?}: got {:02X?}, expected {:02X?}",
kind, null_bytes, expected_null
);
}
let mut geo_payload = vec![0x01]; geo_payload.extend_from_slice(&[0x02]); geo_payload.push(8); geo_payload.push(0x12); let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("g", ColumnKind::Geohash, geo_payload)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ColumnView::Geohash(c) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
assert_eq!(c.value(0), 0x12);
assert_eq!(
c.value(1),
0xFF,
"spec §11.5: GEOHASH NULL = 0xFF * byte_width"
);
}
#[test]
fn decode_short_no_nulls() {
let (flags_byte, payload) = BatchBuilder::new(4)
.add_column(
"v",
ColumnKind::Short,
col_no_nulls(&le_i16s(&[-1, -2, -3, 32767])),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Short(c) = view else { panic!() };
assert_eq!(c.value(0), -1);
assert_eq!(c.value(1), -2);
assert_eq!(c.value(2), -3);
assert_eq!(c.value(3), 32767);
for r in 0..4 {
assert!(!c.is_null(r));
}
}
#[test]
fn decode_byte_no_nulls() {
let (flags_byte, payload) = BatchBuilder::new(5)
.add_column(
"v",
ColumnKind::Byte,
col_no_nulls(&[0x00, 0x7F, 0x80, 0xFF, 0x01]),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Byte(c) = view else { panic!() };
assert_eq!(c.value(0), 0);
assert_eq!(c.value(1), 0x7F);
assert_eq!(c.value(2), -128); assert_eq!(c.value(3), -1); assert_eq!(c.value(4), 1);
for r in 0..5 {
assert!(!c.is_null(r));
}
}
#[test]
fn decode_boolean_bit_packed() {
let (flags_byte, payload) = BatchBuilder::new(5)
.add_column("b", ColumnKind::Boolean, col_no_nulls(&[0x0D]))
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Boolean(c) = view else {
panic!()
};
assert_eq!(c.len(), 5);
assert_eq!(c.value(0), 1);
assert_eq!(c.value(1), 0);
assert_eq!(c.value(2), 1);
assert_eq!(c.value(3), 1);
assert_eq!(c.value(4), 0);
}
#[test]
fn decode_boolean_rejects_validity_bitmap() {
assert_non_nullable_rejects_bitmap(
ColumnKind::Boolean,
"BOOLEAN",
col_with_bitmap(&[0b0001_1111], &[0x0D]),
);
}
#[test]
fn decode_byte_rejects_validity_bitmap() {
assert_non_nullable_rejects_bitmap(
ColumnKind::Byte,
"BYTE",
col_with_bitmap(&[0b0001_1111], &[1, 2, 3, 4, 5]),
);
}
#[test]
fn decode_short_rejects_validity_bitmap() {
assert_non_nullable_rejects_bitmap(
ColumnKind::Short,
"SHORT",
col_with_bitmap(&[0b0001_1111], &le_i16s(&[1, 2, 3, 4, 5])),
);
}
#[test]
fn decode_char_rejects_validity_bitmap() {
assert_non_nullable_rejects_bitmap(
ColumnKind::Char,
"CHAR",
col_with_bitmap(&[0b0001_1111], &le_u16s(&[b'a' as u16; 5])),
);
}
fn assert_non_nullable_rejects_bitmap(kind: ColumnKind, kind_name: &str, body: Vec<u8>) {
let (flags_byte, payload) = BatchBuilder::new(5).add_column("c", kind, body).build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(
err.msg().contains(kind_name) && err.msg().contains("null_flag"),
"unexpected error message for {}: {}",
kind_name,
err.msg()
);
}
fn le_u16s(vs: &[u16]) -> Vec<u8> {
let mut o = Vec::new();
for v in vs {
o.extend_from_slice(&v.to_le_bytes());
}
o
}
fn symbol_column_local(
bitmap: Option<&[u8]>,
dict: &[&str],
codes_per_non_null: &[u64],
) -> Vec<u8> {
let mut col = Vec::new();
if let Some(bm) = bitmap {
col.push(0x01);
col.extend_from_slice(bm);
} else {
col.push(0x00);
}
encode_u64(dict.len() as u64, &mut col); for entry in dict {
encode_u64(entry.len() as u64, &mut col);
col.extend_from_slice(entry.as_bytes());
}
for code in codes_per_non_null {
encode_u64(*code, &mut col);
}
col
}
#[test]
fn decode_symbol_column_local_no_nulls() {
let col = symbol_column_local(None, &["AAPL", "MSFT", "GOOG"], &[0, 1, 2]);
let (flags_byte, payload) = BatchBuilder::new(3)
.add_column("s", ColumnKind::Symbol, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert_eq!(dict.len(), 0);
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Symbol(s) = view else {
panic!()
};
assert_eq!(s.resolve(0), Some("AAPL"));
assert_eq!(s.resolve(1), Some("MSFT"));
assert_eq!(s.resolve(2), Some("GOOG"));
}
#[test]
fn decode_symbol_column_local_with_nulls() {
let col = symbol_column_local(Some(&[0x02]), &["X", "Y"], &[1, 0, 0]);
let (flags_byte, payload) = BatchBuilder::new(4)
.add_column("s", ColumnKind::Symbol, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Symbol(s) = view else {
panic!()
};
assert_eq!(s.resolve(0), Some("Y"));
assert!(s.is_null(1));
assert_eq!(s.resolve(1), None);
assert_eq!(s.resolve(2), Some("X"));
assert_eq!(s.resolve(3), Some("X"));
}
#[test]
fn decode_symbol_column_local_independent_per_column() {
let col_a = symbol_column_local(None, &["alpha", "beta"], &[0, 1]);
let col_b = symbol_column_local(None, &["one", "two"], &[1, 0]);
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("a", ColumnKind::Symbol, col_a)
.add_column("b", ColumnKind::Symbol, col_b)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ColumnView::Symbol(a) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
let ColumnView::Symbol(b) = batch.column_view(1, &dict).unwrap() else {
panic!()
};
assert_eq!(a.resolve(0), Some("alpha"));
assert_eq!(a.resolve(1), Some("beta"));
assert_eq!(b.resolve(0), Some("two"));
assert_eq!(b.resolve(1), Some("one"));
}
#[test]
fn decode_symbol_column_local_id_out_of_range_rejected() {
let col = symbol_column_local(None, &["a", "b"], &[0, 5]);
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("s", ColumnKind::Symbol, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(err.msg().contains("out of range"));
}
#[test]
fn decode_symbol_column_local_dict_size_exceeds_rows_rejected() {
let mut col = vec![0x00u8]; encode_u64(5, &mut col); let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("s", ColumnKind::Symbol, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn decode_symbol_delta_id_out_of_range_rejected() {
let mut col_data = vec![0x00u8]; encode_u64(9, &mut col_data); let (flags_byte, payload) = BatchBuilder::new(1)
.with_dict_delta(0, vec!["AAPL", "MSFT"])
.add_column("s", ColumnKind::Symbol, col_data)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(err.msg().contains("out of range"));
}
#[test]
fn decode_symbol_with_dict_delta() {
let mut col_data = vec![0x01u8, 0x02]; encode_u64(0, &mut col_data);
encode_u64(1, &mut col_data);
let (flags_byte, payload) = BatchBuilder::new(3)
.with_dict_delta(0, vec!["AAPL", "MSFT"])
.add_column("sym", ColumnKind::Symbol, col_data)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert_eq!(dict.len(), 2);
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Symbol(s) = view else {
panic!()
};
assert_eq!(s.len(), 3);
assert_eq!(s.resolve(0), Some("AAPL"));
assert_eq!(s.resolve(1), None);
assert_eq!(s.resolve(2), Some("MSFT"));
}
#[test]
fn decode_decimal64_with_scale() {
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("p", ColumnKind::Decimal64, {
let mut d = vec![0x00u8, 0x02]; d.extend_from_slice(&le_i64s(&[12345, 6789]));
d
})
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Decimal64(d) = view else {
panic!()
};
assert_eq!(d.scale(), 2);
assert_eq!(d.value(0), 12345);
assert_eq!(d.value(1), 6789);
}
#[test]
fn decode_decimal_rejects_negative_scale() {
for kind in [
ColumnKind::Decimal64,
ColumnKind::Decimal128,
ColumnKind::Decimal256,
] {
let width = match kind {
ColumnKind::Decimal64 => 8,
ColumnKind::Decimal128 => 16,
ColumnKind::Decimal256 => 32,
_ => unreachable!(),
};
let mut data = vec![0x00u8, 0xFF]; data.extend(std::iter::repeat_n(0u8, width)); let (flags_byte, payload) = BatchBuilder::new(1).add_column("p", kind, data).build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::ProtocolError);
assert!(
err.msg().contains("scale"),
"expected scale error msg, got: {}",
err.msg()
);
}
}
#[test]
fn decode_decimal_rejects_scale_above_max() {
let mut data = vec![0x00u8, 39u8];
data.extend(std::iter::repeat_n(0u8, 8));
let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("p", ColumnKind::Decimal64, data)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::ProtocolError);
}
fn decimal_col_data(scale: u8, width: usize) -> Vec<u8> {
let mut data = vec![0x00u8, scale]; data.extend(std::iter::repeat_n(0u8, width)); data
}
fn decode_decimal_one_row(kind: ColumnKind, scale: u8, width: usize) -> Result<()> {
let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("p", kind, decimal_col_data(scale, width))
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.map(|_| ())
}
#[test]
fn decode_decimal64_per_width_scale_boundary() {
assert!(decode_decimal_one_row(ColumnKind::Decimal64, 18, 8).is_ok());
let err = decode_decimal_one_row(ColumnKind::Decimal64, 19, 8).unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::ProtocolError);
assert!(err.msg().contains("DECIMAL"), "{}", err.msg());
}
#[test]
fn decode_decimal128_per_width_scale_boundary() {
assert!(decode_decimal_one_row(ColumnKind::Decimal128, 38, 16).is_ok());
let err = decode_decimal_one_row(ColumnKind::Decimal128, 39, 16).unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::ProtocolError);
assert!(err.msg().contains("DECIMAL"), "{}", err.msg());
}
#[test]
fn decode_decimal256_per_width_scale_boundary() {
assert!(decode_decimal_one_row(ColumnKind::Decimal256, 39, 32).is_ok());
assert!(decode_decimal_one_row(ColumnKind::Decimal256, 76, 32).is_ok());
let err = decode_decimal_one_row(ColumnKind::Decimal256, 77, 32).unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::ProtocolError);
assert!(err.msg().contains("DECIMAL"), "{}", err.msg());
}
#[test]
fn schema_reused_across_continuation_batches() {
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let (f1, p1) = BatchBuilder::new(2)
.add_column("v", ColumnKind::Long, col_no_nulls(&le_i64s(&[1, 2])))
.build();
decode_result_batch(&p1, f1, &mut dict, &mut schema, &mut ZstdScratch::new()).unwrap();
assert!(schema.is_some());
let (f2, p2) = BatchBuilder::new(1)
.with_batch_seq(1)
.add_column("v", ColumnKind::Long, col_no_nulls(&le_i64s(&[42])))
.build();
let b2 =
decode_result_batch(&p2, f2, &mut dict, &mut schema, &mut ZstdScratch::new()).unwrap();
assert_eq!(b2.batch_seq, 1);
let view = b2.column_view(0, &dict).unwrap();
let ColumnView::Long(c) = view else { panic!() };
assert_eq!(c.value(0), 42);
}
#[test]
fn continuation_before_schema_rejected() {
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let (f, p) = BatchBuilder::new(1)
.with_batch_seq(1)
.add_column("v", ColumnKind::Long, col_no_nulls(&le_i64s(&[42])))
.build();
let err = decode_result_batch(&p, f, &mut dict, &mut schema, &mut ZstdScratch::new())
.unwrap_err();
assert_eq!(err.code(), crate::ErrorCode::ProtocolError);
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn zstd_round_trips_simple_long_batch() {
let (_, raw_payload) = BatchBuilder::new(3)
.add_column("v", ColumnKind::Long, col_no_nulls(&le_i64s(&[10, 20, 30])))
.build();
let prefix_len = {
let mut r = ByteReader::new(&raw_payload);
r.read_u8().unwrap();
r.read_i64_le().unwrap();
r.read_varint_u64().unwrap();
raw_payload.len() - r.remaining().len()
};
let prefix = &raw_payload[..prefix_len];
let body = &raw_payload[prefix_len..];
let compressed_body = zstd::bulk::compress(body, 0).expect("zstd compress");
let mut zstd_payload = Vec::new();
zstd_payload.extend_from_slice(prefix);
zstd_payload.extend_from_slice(&compressed_body);
let zstd_payload = Bytes::from(zstd_payload);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&zstd_payload,
flags::ZSTD,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert_eq!(batch.row_count, 3);
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Long(c) = view else { panic!() };
assert_eq!(c.value(0), 10);
assert_eq!(c.value(1), 20);
assert_eq!(c.value(2), 30);
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn zstd_scratch_pool_recycles_buffer_across_batches() {
fn build_zstd_payload(seed: i64) -> Bytes {
let (_, raw_payload) = BatchBuilder::new(3)
.add_column(
"v",
ColumnKind::Long,
col_no_nulls(&le_i64s(&[seed, seed + 1, seed + 2])),
)
.build();
let prefix_len = {
let mut r = ByteReader::new(&raw_payload);
r.read_u8().unwrap();
r.read_i64_le().unwrap();
r.read_varint_u64().unwrap();
raw_payload.len() - r.remaining().len()
};
let prefix = &raw_payload[..prefix_len];
let body = &raw_payload[prefix_len..];
let compressed = zstd::bulk::compress(body, 0).expect("compress");
let mut out = Vec::with_capacity(prefix.len() + compressed.len());
out.extend_from_slice(prefix);
out.extend_from_slice(&compressed);
Bytes::from(out)
}
let mut scratch = ZstdScratch::new();
assert_eq!(
scratch.pool.buffers.lock().unwrap().len(),
0,
"pool starts empty"
);
let body1 = zstd_decompress_body(
{
let p = build_zstd_payload(100);
p.slice(10..)
}
.as_ref(),
&mut scratch,
)
.expect("decompress 1");
assert_eq!(
scratch.pool.buffers.lock().unwrap().len(),
0,
"pool empty while body1 holds the buffer"
);
let body1_len = body1.len();
drop(body1);
let pool_len = scratch.pool.buffers.lock().unwrap().len();
assert_eq!(
pool_len, 1,
"pool should hold the recycled buffer after the first Bytes drops"
);
let recycled_capacity = scratch.pool.buffers.lock().unwrap()[0].capacity();
assert!(
recycled_capacity >= body1_len,
"recycled buffer retained capacity >= body length ({} >= {})",
recycled_capacity,
body1_len
);
let body2 = zstd_decompress_body(
{
let p = build_zstd_payload(200);
p.slice(10..)
}
.as_ref(),
&mut scratch,
)
.expect("decompress 2");
assert_eq!(
scratch.pool.buffers.lock().unwrap().len(),
0,
"pool emptied by the second decompress drawing from it"
);
assert_eq!(body2.len(), body1_len, "second body decoded successfully");
let body3 = zstd_decompress_body(
{
let p = build_zstd_payload(300);
p.slice(10..)
}
.as_ref(),
&mut scratch,
)
.expect("decompress 3");
drop(body2);
drop(body3);
let final_pool_len = scratch.pool.buffers.lock().unwrap().len();
assert!(
final_pool_len <= ZSTD_POOL_CAPACITY,
"pool stays bounded by ZSTD_POOL_CAPACITY (got {})",
final_pool_len
);
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn zstd_invalid_frame_is_protocol_error() {
let (_, raw_payload) = BatchBuilder::new(0).build();
let prefix_len = {
let mut r = ByteReader::new(&raw_payload);
r.read_u8().unwrap();
r.read_i64_le().unwrap();
r.read_varint_u64().unwrap();
raw_payload.len() - r.remaining().len()
};
let mut payload = raw_payload[..prefix_len].to_vec();
payload.extend_from_slice(&[0u8, 0, 0, 0]); let payload = Bytes::from(payload);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags::ZSTD,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[cfg(feature = "sync-reader-zstd")]
fn zstd_payload_with_body(body: &[u8]) -> Bytes {
let (_, raw) = BatchBuilder::new(0).build();
let prefix_len = {
let mut r = ByteReader::new(&raw);
r.read_u8().unwrap();
r.read_i64_le().unwrap();
r.read_varint_u64().unwrap();
raw.len() - r.remaining().len()
};
let mut out = raw[..prefix_len].to_vec();
out.extend_from_slice(body);
Bytes::from(out)
}
#[cfg(feature = "sync-reader-zstd")]
fn forged_fcs_zstd_frame(forged: u64) -> Vec<u8> {
let mut frame = vec![0x28, 0xB5, 0x2F, 0xFD]; frame.push(0xE0);
frame.extend_from_slice(&forged.to_le_bytes());
frame.extend_from_slice(&[0x01, 0x00, 0x00]);
frame
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn zstd_frame_without_content_size_is_protocol_error() {
use std::io::Write;
let mut encoder = zstd::stream::write::Encoder::new(Vec::new(), 0).unwrap();
encoder
.write_all(b"some bytes that will never be read")
.unwrap();
let body = encoder.finish().expect("zstd encode");
assert!(
matches!(zstd::zstd_safe::get_frame_content_size(&body), Ok(None)),
"zstd::Encoder default must produce a frame without FCS; \
header bytes: {:02x?}",
&body[..body.len().min(16)]
);
let payload = zstd_payload_with_body(&body);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags::ZSTD,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(
err.msg().contains("missing content size"),
"expected missing-content-size message, got: {}",
err.msg()
);
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn zstd_frame_exceeding_cap_is_limit_exceeded() {
let oversized = MAX_ZSTD_DECOMPRESSED + 1;
let frame = forged_fcs_zstd_frame(oversized);
assert_eq!(
zstd::zstd_safe::get_frame_content_size(&frame).ok(),
Some(Some(oversized)),
"forged FCS bytes must round-trip through zstd"
);
let payload = zstd_payload_with_body(&frame);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags::ZSTD,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::LimitExceeded);
assert!(
err.msg().contains("exceeds client cap"),
"expected cap-exceeded message, got: {}",
err.msg()
);
}
#[cfg(feature = "sync-reader-zstd")]
#[test]
fn zstd_frame_with_size_mismatch_is_protocol_error() {
use std::io::Write;
let mut encoder = zstd::stream::write::Encoder::new(Vec::new(), 0).unwrap();
encoder.set_pledged_src_size(Some(100)).ok();
encoder.write_all(b"only ten!!").unwrap(); let body = match encoder.finish() {
Ok(b) => b,
Err(_) => {
return;
}
};
assert_eq!(
zstd::zstd_safe::get_frame_content_size(&body).ok(),
Some(Some(100))
);
let payload = zstd_payload_with_body(&body);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags::ZSTD,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn rejects_unknown_temporal_discriminator() {
let mut col_data = vec![0x00u8]; col_data.push(0x02); let (_, payload) = BatchBuilder::new(1)
.with_flags(flags::GORILLA)
.add_column("ts", ColumnKind::TimestampNanos, col_data)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags::GORILLA,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(err.msg().to_lowercase().contains("discriminator"));
}
#[test]
fn decodes_gorilla_with_few_non_null() {
let mut col_data = vec![0x00u8]; col_data.push(0x01); col_data.extend_from_slice(&0i64.to_le_bytes());
col_data.extend_from_slice(&100i64.to_le_bytes());
let (_, payload) = BatchBuilder::new(2)
.with_flags(flags::GORILLA)
.add_column("ts", ColumnKind::TimestampNanos, col_data)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags::GORILLA,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::TimestampNanos(c) = view else {
panic!("expected TimestampNanos column")
};
assert_eq!(c.value(0), 0);
assert_eq!(c.value(1), 100);
}
#[test]
fn decodes_gorilla_with_one_non_null() {
let mut col_data = vec![0x01u8, 0b0000_0010];
col_data.push(0x01); col_data.extend_from_slice(&42i64.to_le_bytes());
let (_, payload) = BatchBuilder::new(2)
.with_flags(flags::GORILLA)
.add_column("ts", ColumnKind::TimestampNanos, col_data)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags::GORILLA,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::TimestampNanos(c) = view else {
panic!("expected TimestampNanos column")
};
assert!(!c.is_null(0));
assert_eq!(c.value(0), 42);
assert!(c.is_null(1));
}
#[test]
fn decodes_gorilla_with_zero_non_null() {
let mut col_data = vec![0x01u8, 0b0000_0011];
col_data.push(0x01); let (_, payload) = BatchBuilder::new(2)
.with_flags(flags::GORILLA)
.add_column("ts", ColumnKind::TimestampNanos, col_data)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags::GORILLA,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::TimestampNanos(c) = view else {
panic!("expected TimestampNanos column")
};
assert!(c.is_null(0));
assert!(c.is_null(1));
}
#[test]
fn raw_temporal_under_gorilla_flag_decodes() {
let mut col_data = vec![0x00u8]; col_data.push(0x00); col_data.extend_from_slice(&le_i64s(&[10, 20, 30]));
let (_, payload) = BatchBuilder::new(3)
.with_flags(flags::GORILLA)
.add_column("ts", ColumnKind::TimestampNanos, col_data)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags::GORILLA,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::TimestampNanos(c) = view else {
panic!()
};
assert_eq!(c.value(0), 10);
assert_eq!(c.value(1), 20);
assert_eq!(c.value(2), 30);
}
fn encode_gorilla_temporal_bitstream(timestamps: &[i64]) -> Vec<u8> {
assert!(timestamps.len() >= 2, "need at least two seeds");
let mut prev_delta = timestamps[1] - timestamps[0];
let mut prev_ts = timestamps[1];
let mut bytes = Vec::new();
let mut cur: u8 = 0;
let mut bits: u32 = 0;
let write_bit = |b: u8, bytes: &mut Vec<u8>, cur: &mut u8, bits: &mut u32| {
*cur |= (b & 1) << *bits;
*bits += 1;
if *bits == 8 {
bytes.push(*cur);
*cur = 0;
*bits = 0;
}
};
let write_bits = |val: u64, n: u32, bytes: &mut Vec<u8>, cur: &mut u8, bits: &mut u32| {
for i in 0..n {
write_bit(((val >> i) & 1) as u8, bytes, cur, bits);
}
};
for &ts in ×tamps[2..] {
let delta = ts - prev_ts;
let dod = delta - prev_delta;
if dod == 0 {
write_bit(0, &mut bytes, &mut cur, &mut bits);
} else if (-64..=63).contains(&dod) {
write_bits(0b01, 2, &mut bytes, &mut cur, &mut bits);
write_bits((dod as u64) & 0x7F, 7, &mut bytes, &mut cur, &mut bits);
} else if (-256..=255).contains(&dod) {
write_bits(0b011, 3, &mut bytes, &mut cur, &mut bits);
write_bits((dod as u64) & 0x1FF, 9, &mut bytes, &mut cur, &mut bits);
} else if (-2048..=2047).contains(&dod) {
write_bits(0b0111, 4, &mut bytes, &mut cur, &mut bits);
write_bits((dod as u64) & 0xFFF, 12, &mut bytes, &mut cur, &mut bits);
} else {
write_bits(0b1111, 4, &mut bytes, &mut cur, &mut bits);
write_bits(
(dod as u64) & 0xFFFF_FFFF,
32,
&mut bytes,
&mut cur,
&mut bits,
);
}
prev_delta = delta;
prev_ts = ts;
}
if bits > 0 {
bytes.push(cur);
}
bytes
}
fn build_gorilla_temporal_column_body(timestamps: &[i64]) -> Vec<u8> {
let bitstream = encode_gorilla_temporal_bitstream(timestamps);
let mut col = Vec::with_capacity(2 + 16 + bitstream.len());
col.push(0x00); col.push(0x01); col.extend_from_slice(×tamps[0].to_le_bytes());
col.extend_from_slice(×tamps[1].to_le_bytes());
col.extend_from_slice(&bitstream);
col
}
fn assert_gorilla_temporal_round_trip(
kind: ColumnKind,
timestamps: &[i64],
view_to_values: fn(ColumnView<'_>) -> Vec<i64>,
) {
let body = build_gorilla_temporal_column_body(timestamps);
let (_, payload) = BatchBuilder::new(timestamps.len())
.with_flags(flags::GORILLA)
.add_column("ts", kind, body)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags::GORILLA,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_or_else(|e| panic!("decode failed for {:?}: {}", kind, e));
let view = batch.column_view(0, &dict).unwrap();
let got = view_to_values(view);
assert_eq!(
got.len(),
timestamps.len(),
"row count mismatch for {:?}",
kind
);
for (i, (g, e)) in got.iter().zip(timestamps.iter()).enumerate() {
assert_eq!(
g, e,
"{:?} row {} mismatch (got {}, expected {})",
kind, i, g, e
);
}
}
#[test]
fn decode_gorilla_temporal_round_trip() {
let timestamps: [i64; 6] = [1_000, 1_100, 1_200, 1_310, 1_405, 1_488];
type Extract = fn(ColumnView<'_>) -> Vec<i64>;
let cases: &[(ColumnKind, Extract)] = &[
(ColumnKind::Timestamp, |v| {
let ColumnView::Timestamp(c) = v else {
panic!("expected ColumnView::Timestamp")
};
(0..c.len()).map(|i| c.value(i)).collect()
}),
(ColumnKind::TimestampNanos, |v| {
let ColumnView::TimestampNanos(c) = v else {
panic!("expected ColumnView::TimestampNanos")
};
(0..c.len()).map(|i| c.value(i)).collect()
}),
(ColumnKind::Date, |v| {
let ColumnView::Date(c) = v else {
panic!("expected ColumnView::Date")
};
(0..c.len()).map(|i| c.value(i)).collect()
}),
];
for &(kind, view_to_values) in cases {
assert_gorilla_temporal_round_trip(kind, ×tamps, view_to_values);
}
}
fn build_double_array_row(shape: &[u32], elements: &[f64]) -> Vec<u8> {
let mut out = Vec::new();
out.push(shape.len() as u8);
for d in shape {
out.extend_from_slice(&d.to_le_bytes());
}
for e in elements {
out.extend_from_slice(&e.to_le_bytes());
}
out
}
fn build_long_array_row(shape: &[u32], elements: &[i64]) -> Vec<u8> {
let mut out = Vec::new();
out.push(shape.len() as u8);
for d in shape {
out.extend_from_slice(&d.to_le_bytes());
}
for e in elements {
out.extend_from_slice(&e.to_le_bytes());
}
out
}
#[test]
fn decode_double_array_1d_no_nulls() {
let mut col = vec![0x00u8]; col.extend_from_slice(&build_double_array_row(&[3], &[1.0, 2.0, 3.0]));
col.extend_from_slice(&build_double_array_row(&[2], &[10.0, 20.0]));
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("a", ColumnKind::DoubleArray, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::DoubleArray(c) = view else {
panic!()
};
assert_eq!(c.len(), 2);
assert_eq!(c.shape(0), Some(&[3u32][..]));
assert_eq!(c.element_count(0), 3);
assert_eq!(c.element(0, 0), Some(1.0));
assert_eq!(c.element(0, 2), Some(3.0));
assert_eq!(c.shape(1), Some(&[2u32][..]));
assert_eq!(c.element(1, 1), Some(20.0));
}
#[test]
fn decode_long_array_2d_with_nulls() {
let mut col = vec![0x01u8, 0x02];
col.extend_from_slice(&build_long_array_row(&[2, 2], &[1, 2, 3, 4]));
col.extend_from_slice(&build_long_array_row(&[1, 3], &[7, 8, 9]));
let (flags_byte, payload) = BatchBuilder::new(3)
.add_column("a", ColumnKind::LongArray, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::LongArray(c) = view else {
panic!()
};
assert_eq!(c.len(), 3);
assert_eq!(c.shape(0), Some(&[2u32, 2][..]));
assert_eq!(c.element_count(0), 4);
assert_eq!(c.element(0, 3), Some(4));
assert!(c.is_null(1));
assert_eq!(c.shape(1), None);
assert_eq!(c.shape(2), Some(&[1u32, 3][..]));
assert_eq!(c.element(2, 0), Some(7));
assert_eq!(c.element(2, 2), Some(9));
}
#[test]
fn decode_array_empty_vs_null_distinct() {
let mut col = vec![0x01u8, 0x04]; col.extend_from_slice(&build_double_array_row(&[2], &[1.0, 2.0]));
col.extend_from_slice(&build_double_array_row(&[0], &[])); let (flags_byte, payload) = BatchBuilder::new(3)
.add_column("a", ColumnKind::DoubleArray, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::DoubleArray(c) = view else {
panic!()
};
assert_eq!(c.len(), 3);
assert!(!c.is_null(0));
assert_eq!(c.shape(0), Some(&[2u32][..]));
assert_eq!(c.element_count(0), 2);
assert_eq!(c.element(0, 0), Some(1.0));
assert_eq!(c.element(0, 1), Some(2.0));
assert!(!c.is_null(1));
assert_eq!(c.shape(1), Some(&[0u32][..]));
assert_eq!(c.element_count(1), 0);
assert_eq!(c.raw(1), Some(&[][..]));
assert!(c.is_null(2));
assert_eq!(c.shape(2), None);
assert_eq!(c.element_count(2), 0);
assert_eq!(c.raw(2), None);
}
#[test]
fn decode_array_zero_dims_rejected() {
let mut col = vec![0x00u8];
col.push(0u8); let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("a", ColumnKind::DoubleArray, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn decode_array_huge_row_rejected() {
let mut col = vec![0x00u8, 1]; let big = (MAX_ARRAY_ELEMENTS_PER_ROW + 1) as u32;
col.extend_from_slice(&big.to_le_bytes());
let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("a", ColumnKind::LongArray, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::LimitExceeded);
}
#[test]
fn decode_array_degenerate_prefix_dim_rejected() {
let mut col = vec![0x00u8, 2]; col.extend_from_slice(&(1u32 << 31).to_le_bytes()); col.extend_from_slice(&0u32.to_le_bytes()); let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("a", ColumnKind::DoubleArray, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::LimitExceeded);
}
fn varchar_col_no_nulls(values: &[&str]) -> Vec<u8> {
let mut out = vec![0x00u8]; let mut total = 0u32;
out.extend_from_slice(&total.to_le_bytes());
for v in values {
total += v.len() as u32;
out.extend_from_slice(&total.to_le_bytes());
}
for v in values {
out.extend_from_slice(v.as_bytes());
}
out
}
fn varchar_col_with_bitmap(bitmap: &[u8], non_null_values: &[&str]) -> Vec<u8> {
let mut out = vec![0x01u8];
out.extend_from_slice(bitmap);
let mut total = 0u32;
out.extend_from_slice(&total.to_le_bytes());
for v in non_null_values {
total += v.len() as u32;
out.extend_from_slice(&total.to_le_bytes());
}
for v in non_null_values {
out.extend_from_slice(v.as_bytes());
}
out
}
#[test]
fn decode_varchar_no_nulls() {
let (flags_byte, payload) = BatchBuilder::new(3)
.add_column(
"s",
ColumnKind::Varchar,
varchar_col_no_nulls(&["foo", "", "café"]),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Varchar(c) = view else {
panic!()
};
assert_eq!(c.len(), 3);
assert_eq!(c.value(0), Some("foo"));
assert_eq!(c.value(1), Some(""));
assert_eq!(c.value(2), Some("café"));
}
#[test]
fn decode_varchar_with_nulls_densifies_offsets() {
let (flags_byte, payload) = BatchBuilder::new(4)
.add_column(
"s",
ColumnKind::Varchar,
varchar_col_with_bitmap(&[0x0A], &["hello", "world"]),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Varchar(c) = view else {
panic!()
};
assert_eq!(c.len(), 4);
assert_eq!(c.value(0), Some("hello"));
assert_eq!(c.value(1), None);
assert_eq!(c.value(2), Some("world"));
assert_eq!(c.value(3), None);
assert_eq!(c.offsets(), &[0u32, 5, 5, 10, 10]);
}
#[test]
fn decode_varchar_invalid_utf8_rejected() {
let mut col = vec![0x00u8]; col.extend_from_slice(&0u32.to_le_bytes());
col.extend_from_slice(&2u32.to_le_bytes());
col.extend_from_slice(&[0xFF, 0xFE]); let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("s", ColumnKind::Varchar, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::InvalidUtf8);
}
#[test]
fn decode_varchar_offset_splitting_codepoint_rejected() {
let mut col = vec![0x00u8]; for o in [0u32, 1, 2] {
col.extend_from_slice(&o.to_le_bytes());
}
col.extend_from_slice(&[0xC3, 0xB1]);
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("s", ColumnKind::Varchar, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::InvalidUtf8);
}
#[test]
fn decode_binary_no_nulls() {
let mut col = vec![0x00u8];
for o in [0u32, 3, 5] {
col.extend_from_slice(&o.to_le_bytes());
}
col.extend_from_slice(&[0xDE, 0xAD, 0xBE, 0xEF, 0x42]);
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("b", ColumnKind::Binary, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Binary(c) = view else {
panic!()
};
assert_eq!(c.len(), 2);
assert_eq!(c.value(0), Some([0xDEu8, 0xAD, 0xBE].as_slice()));
assert_eq!(c.value(1), Some([0xEFu8, 0x42].as_slice()));
}
#[test]
fn decode_binary_invalid_utf8_accepted() {
let mut col = vec![0x00u8];
col.extend_from_slice(&0u32.to_le_bytes());
col.extend_from_slice(&2u32.to_le_bytes());
col.extend_from_slice(&[0xFF, 0xFE]);
let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("b", ColumnKind::Binary, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Binary(c) = view else {
panic!()
};
assert_eq!(c.value(0), Some([0xFFu8, 0xFE].as_slice()));
}
#[test]
fn decode_varlen_non_monotonic_rejected() {
let mut col = vec![0x00u8];
for o in [0u32, 5, 3] {
col.extend_from_slice(&o.to_le_bytes());
}
col.extend_from_slice(&[0u8; 5]);
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("s", ColumnKind::Varchar, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
fn le_i128s(vs: &[i128]) -> Vec<u8> {
let mut o = Vec::new();
for v in vs {
o.extend_from_slice(&v.to_le_bytes());
}
o
}
#[test]
fn decode_geohash_8bit() {
let mut col = vec![0x00u8]; encode_u64(8, &mut col); col.extend_from_slice(&[0xAA, 0xBB, 0xCC]);
let (flags_byte, payload) = BatchBuilder::new(3)
.add_column("g", ColumnKind::Geohash, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Geohash(c) = view else {
panic!()
};
assert_eq!(c.precision_bits(), 8);
assert_eq!(c.byte_width(), 1);
assert_eq!(c.len(), 3);
assert_eq!(c.value(0), 0xAA);
assert_eq!(c.value(1), 0xBB);
assert_eq!(c.value(2), 0xCC);
}
#[test]
fn decode_column_rejects_densified_batch_over_budget() {
let mut r = ByteReader::new(&[]);
let parent = Bytes::new();
let mut budget = 10usize;
let err =
decode_column(&mut r, &parent, ColumnKind::Long256, 16, 0, 0, &mut budget).unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(
err.msg().contains("per-batch cap"),
"unexpected error: {}",
err.msg()
);
}
#[test]
fn decode_geohash_60bit_with_nulls() {
let mut col = vec![0x01u8, 0x02]; encode_u64(60, &mut col);
col.extend_from_slice(&0x0102_0304_0506_0708u64.to_le_bytes());
col.extend_from_slice(&0xAAAA_BBBB_CCCC_DDDDu64.to_le_bytes());
col.extend_from_slice(&0x1111_2222_3333_4444u64.to_le_bytes());
let (flags_byte, payload) = BatchBuilder::new(4)
.add_column("g", ColumnKind::Geohash, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Geohash(c) = view else {
panic!()
};
assert_eq!(c.precision_bits(), 60);
assert_eq!(c.byte_width(), 8);
assert!(!c.is_null(0));
assert!(c.is_null(1));
assert_eq!(c.value(0), 0x0102_0304_0506_0708);
assert_eq!(c.value(2), 0xAAAA_BBBB_CCCC_DDDD);
assert_eq!(c.value(3), 0x1111_2222_3333_4444);
}
#[test]
fn decode_geohash_invalid_precision_rejected() {
let mut col = vec![0x00u8];
encode_u64(0, &mut col); let (flags_byte, payload) = BatchBuilder::new(0)
.add_column("g", ColumnKind::Geohash, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn decode_decimal128_with_scale() {
let mut col = vec![0x00u8, 0x04]; col.extend_from_slice(&le_i128s(&[100_000i128, -42i128]));
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("p", ColumnKind::Decimal128, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Decimal128(c) = view else {
panic!()
};
assert_eq!(c.scale(), 4);
assert_eq!(c.value(0), 100_000i128);
assert_eq!(c.value(1), -42i128);
}
#[test]
fn decode_decimal256_passes_raw_bytes() {
let mut col = vec![0x00u8, 0x06]; let row0: [u8; 32] = std::array::from_fn(|i| i as u8);
let row1: [u8; 32] = std::array::from_fn(|i| (255 - i) as u8);
col.extend_from_slice(&row0);
col.extend_from_slice(&row1);
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("p", ColumnKind::Decimal256, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Decimal256(c) = view else {
panic!()
};
assert_eq!(c.scale(), 6);
assert_eq!(c.value(0), &row0);
assert_eq!(c.value(1), &row1);
}
#[test]
fn decode_varchar_all_null_column() {
let mut col = vec![0x01u8, 0x07];
col.extend_from_slice(&0u32.to_le_bytes());
let (flags_byte, payload) = BatchBuilder::new(3)
.add_column("s", ColumnKind::Varchar, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let view = batch.column_view(0, &dict).unwrap();
let ColumnView::Varchar(c) = view else {
panic!()
};
assert_eq!(c.len(), 3);
assert_eq!(c.value(0), None);
assert_eq!(c.value(1), None);
assert_eq!(c.value(2), None);
assert_eq!(c.offsets(), &[0u32, 0, 0, 0]);
}
#[test]
fn trailing_bytes_rejected() {
let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("v", ColumnKind::Long, col_no_nulls(&le_i64s(&[7])))
.build();
let mut bytes_vec: Vec<u8> = payload.to_vec();
bytes_vec.push(0xAA); let payload = Bytes::from(bytes_vec);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
assert!(err.msg().contains("trailing"));
}
#[test]
fn truncated_column_rejected() {
let (flags_byte, mut payload) = BatchBuilder::new(1)
.add_column("v", ColumnKind::Long, col_no_nulls(&le_i64s(&[7])))
.build();
payload.truncate(payload.len() - 4); let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn multi_column_batch() {
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("a", ColumnKind::Long, col_no_nulls(&le_i64s(&[10, 20])))
.add_column(
"b",
ColumnKind::Double,
col_no_nulls(&{
let mut o = Vec::new();
o.extend_from_slice(&1.5f64.to_le_bytes());
o.extend_from_slice(&2.5f64.to_le_bytes());
o
}),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert_eq!(batch.columns.len(), 2);
let ColumnView::Long(a) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
let ColumnView::Double(b) = batch.column_view(1, &dict).unwrap() else {
panic!()
};
assert_eq!(a.value(0), 10);
assert_eq!(a.value(1), 20);
assert_eq!(b.value(0), 1.5);
assert_eq!(b.value(1), 2.5);
}
#[test]
fn decode_zero_row_batch_long() {
let (flags_byte, payload) = BatchBuilder::new(0)
.add_column("v", ColumnKind::Long, col_no_nulls(&le_i64s(&[])))
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert_eq!(batch.row_count, 0);
assert_eq!(batch.columns.len(), 1);
let ColumnView::Long(c) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
assert_eq!(c.len(), 0);
}
#[test]
fn decode_zero_row_batch_varchar() {
let (flags_byte, payload) = BatchBuilder::new(0)
.add_column("s", ColumnKind::Varchar, varchar_col_no_nulls(&[]))
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert_eq!(batch.row_count, 0);
let ColumnView::Varchar(c) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
assert_eq!(c.len(), 0);
assert_eq!(c.offsets(), &[0u32]);
}
#[test]
fn decode_zero_row_batch_multi_kind() {
let (flags_byte, payload) = BatchBuilder::new(0)
.add_column("i", ColumnKind::Int, col_no_nulls(&le_i32s(&[])))
.add_column("l", ColumnKind::Long, col_no_nulls(&le_i64s(&[])))
.add_column("d", ColumnKind::Double, col_no_nulls(&le_f64s(&[])))
.add_column("s", ColumnKind::Varchar, varchar_col_no_nulls(&[]))
.add_column("b", ColumnKind::Binary, {
let mut out = vec![0x00u8];
out.extend_from_slice(&0u32.to_le_bytes());
out
})
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
assert_eq!(batch.row_count, 0);
assert_eq!(batch.columns.len(), 5);
for col_idx in 0..5 {
let v = batch.column_view(col_idx, &dict).unwrap();
let len = match v {
ColumnView::Int(c) => c.len(),
ColumnView::Long(c) => c.len(),
ColumnView::Double(c) => c.len(),
ColumnView::Varchar(c) => c.len(),
ColumnView::Binary(c) => c.len(),
_ => unreachable!(),
};
assert_eq!(len, 0, "column {} reported non-zero rows", col_idx);
}
}
#[test]
fn decode_varchar_multi_mb_value() {
let big = "x".repeat(2 * 1024 * 1024); let (flags_byte, payload) = BatchBuilder::new(1)
.add_column(
"s",
ColumnKind::Varchar,
varchar_col_no_nulls(&[big.as_str()]),
)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ColumnView::Varchar(c) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
assert_eq!(c.len(), 1);
let v = c.value(0).expect("non-null");
assert_eq!(v.len(), big.len());
assert_eq!(v.as_bytes()[0], b'x');
assert_eq!(v.as_bytes()[v.len() - 1], b'x');
}
#[test]
fn decode_binary_multi_mb_value() {
let big_a = vec![0xABu8; 2 * 1024 * 1024];
let big_b = vec![0xCDu8; 1024 * 1024 + 7];
let mut col = vec![0x00u8]; let off_a: u32 = big_a.len() as u32;
let off_b: u32 = (big_a.len() + big_b.len()) as u32;
col.extend_from_slice(&0u32.to_le_bytes());
col.extend_from_slice(&off_a.to_le_bytes());
col.extend_from_slice(&off_b.to_le_bytes());
col.extend_from_slice(&big_a);
col.extend_from_slice(&big_b);
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("b", ColumnKind::Binary, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ColumnView::Binary(c) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
assert_eq!(c.len(), 2);
let v0 = c.value(0).expect("non-null");
let v1 = c.value(1).expect("non-null");
assert_eq!(v0.len(), big_a.len());
assert_eq!(v0[0], 0xAB);
assert_eq!(v0[v0.len() - 1], 0xAB);
assert_eq!(v1.len(), big_b.len());
assert_eq!(v1[0], 0xCD);
assert_eq!(v1[v1.len() - 1], 0xCD);
}
#[test]
fn decode_binary_empty_value_distinct_from_null() {
let mut col = vec![0x01u8]; col.push(0x02); col.extend_from_slice(&0u32.to_le_bytes());
col.extend_from_slice(&0u32.to_le_bytes());
col.extend_from_slice(&2u32.to_le_bytes());
col.extend_from_slice(&[0xAA, 0xBB]);
let (flags_byte, payload) = BatchBuilder::new(3)
.add_column("b", ColumnKind::Binary, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ColumnView::Binary(c) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
assert_eq!(c.len(), 3);
let v0 = c.value(0);
assert!(
matches!(v0, Some(s) if s.is_empty()),
"zero-length non-null binary must be Some(empty slice), not None: got {:?}",
v0
);
assert_eq!(c.value(1), None);
assert!(c.is_null(1));
assert_eq!(c.value(2), Some(&[0xAA, 0xBB][..]));
}
#[test]
fn decode_symbol_column_large_dict_multibyte_codes() {
const N: usize = 17_000;
let dict_entries: Vec<String> = (0..N).map(|i| format!("s{}", i)).collect();
let dict_refs: Vec<&str> = dict_entries.iter().map(String::as_str).collect();
let codes_per_row: Vec<u64> = (0..N as u64).collect();
let col = symbol_column_local(None, &dict_refs, &codes_per_row);
let (flags_byte, payload) = BatchBuilder::new(N)
.add_column("s", ColumnKind::Symbol, col)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let batch = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.unwrap();
let ColumnView::Symbol(s) = batch.column_view(0, &dict).unwrap() else {
panic!()
};
for &code in &[0u64, 127, 128, 16_383, 16_384, 16_999] {
let row = code as usize;
let expected = format!("s{}", code);
assert_eq!(
s.resolve(row),
Some(expected.as_str()),
"row {} (code {}) misresolved — multi-byte LEB128 boundary regression",
row,
code
);
}
}
#[allow(dead_code)]
fn _unused(_: &Schema, _: &SchemaColumn) {}
mod hardening {
use super::*;
fn array_col_body(row_dims: &[&[u32]]) -> Vec<u8> {
let mut out = vec![0x00u8]; for dims in row_dims {
out.push(dims.len() as u8);
for &d in *dims {
out.extend_from_slice(&d.to_le_bytes());
}
let total: usize = dims.iter().map(|&d| d as usize).product();
out.resize(out.len() + total * 8, 0);
}
out
}
#[test]
fn array_dim_zero_is_valid_empty_array() {
let body = array_col_body(&[&[0u32, 5u32]]);
let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("a", ColumnKind::DoubleArray, body)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.expect("ARRAY row with dim==0 must decode as a valid empty array");
}
#[test]
fn array_valid_dims_accepted() {
let body = array_col_body(&[&[2u32, 3u32]]);
let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("a", ColumnKind::DoubleArray, body)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.expect("2D array with all non-zero dims must decode cleanly");
}
#[test]
fn geohash_precision_below_min_rejected() {
let mut body = vec![0x00u8];
encode_u64(0, &mut body); let (flags_byte, payload) = BatchBuilder::new(0)
.add_column("g", ColumnKind::Geohash, body)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.expect_err("decoder must reject GEOHASH precision_bits=0");
assert!(
err.msg().contains("precision"),
"error must mention precision, got: {}",
err.msg()
);
}
#[test]
fn geohash_precision_above_max_rejected() {
let mut body = vec![0x00u8];
encode_u64(61, &mut body); let (flags_byte, payload) = BatchBuilder::new(0)
.add_column("g", ColumnKind::Geohash, body)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.expect_err("decoder must reject GEOHASH precision_bits > 60");
assert!(
err.msg().contains("precision"),
"error must mention precision, got: {}",
err.msg()
);
}
#[test]
fn table_name_len_overflow_rejected() {
let mut out = Vec::new();
out.push(MsgKind::ResultBatch.as_u8());
out.extend_from_slice(&1i64.to_le_bytes()); encode_u64(0, &mut out);
encode_u64(u32::MAX as u64, &mut out);
let payload = Bytes::from(out);
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err =
decode_result_batch(&payload, 0, &mut dict, &mut schema, &mut ZstdScratch::new())
.expect_err("decoder must reject huge table name length");
assert!(
err.msg().contains("table name length"),
"error must mention table name length, got: {}",
err.msg()
);
}
#[test]
fn symbol_non_delta_huge_dict_rejected() {
let mut body = vec![0x00u8];
encode_u64(1000, &mut body); let (flags_byte, payload) = BatchBuilder::new(3)
.add_column("s", ColumnKind::Symbol, body)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.expect_err("decoder must reject SYMBOL dict_size > row_count");
assert!(
err.msg().contains("dict_size"),
"error must mention dict_size, got: {}",
err.msg()
);
}
#[test]
fn string_non_monotonic_offsets_rejected() {
let mut body = vec![0x00u8]; for &o in &[0u32, 10, 5] {
body.extend_from_slice(&o.to_le_bytes());
}
body.extend_from_slice(&[b'a'; 10]);
let (flags_byte, payload) = BatchBuilder::new(2)
.add_column("v", ColumnKind::Varchar, body)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.expect_err("decoder must reject non-monotonic varlen offsets");
assert!(
err.msg().contains("not monotonic"),
"error must say 'not monotonic', got: {}",
err.msg()
);
}
#[test]
fn string_first_offset_nonzero_rejected() {
let mut body = vec![0x00u8];
for &o in &[5u32, 12] {
body.extend_from_slice(&o.to_le_bytes());
}
body.extend_from_slice(&[b'a'; 12]);
let (flags_byte, payload) = BatchBuilder::new(1)
.add_column("v", ColumnKind::Varchar, body)
.build();
let mut dict = SymbolDict::new();
let mut schema: Option<Schema> = None;
let err = decode_result_batch(
&payload,
flags_byte,
&mut dict,
&mut schema,
&mut ZstdScratch::new(),
)
.expect_err("decoder must reject non-zero first offset");
assert!(
err.msg().contains("must start at 0"),
"error must say first offset must start at 0, got: {}",
err.msg()
);
}
}
}