use bytes::{Buf, BufMut, Bytes, BytesMut};
pub const ALIGNMENT: usize = 8;
pub const NULL_ROW_MARKER: u64 = u64::MAX;
pub const MAX_ROWS_PER_ROWSET: u64 = 5 * 1024 * 1024;
pub const MAX_VALUES_PER_ROW: u64 = 1024;
pub const MAX_VALUE_LENGTH: u32 = 16 * 1024 * 1024;
pub const MAX_ROWSET_SIZE: usize = crate::bus::packet::MAX_PART_SIZE as usize;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[repr(u8)]
pub enum ValueType {
Null = 0x02,
Int64 = 0x03,
Uint64 = 0x04,
Double = 0x05,
Boolean = 0x06,
String = 0x10,
Any = 0x11,
Composite = 0x12,
}
impl ValueType {
fn from_wire(value: u8) -> Option<Self> {
match value {
0x02 => Some(Self::Null),
0x03 => Some(Self::Int64),
0x04 => Some(Self::Uint64),
0x05 => Some(Self::Double),
0x06 => Some(Self::Boolean),
0x10 => Some(Self::String),
0x11 => Some(Self::Any),
0x12 => Some(Self::Composite),
_ => None,
}
}
pub fn is_string_like(self) -> bool {
matches!(self, Self::String | Self::Any | Self::Composite)
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Value {
Null,
Int64(i64),
Uint64(u64),
Double(f64),
Boolean(bool),
String(Bytes),
Any(Bytes),
Composite(Bytes),
}
impl Value {
pub fn value_type(&self) -> ValueType {
match self {
Self::Null => ValueType::Null,
Self::Int64(_) => ValueType::Int64,
Self::Uint64(_) => ValueType::Uint64,
Self::Double(_) => ValueType::Double,
Self::Boolean(_) => ValueType::Boolean,
Self::String(_) => ValueType::String,
Self::Any(_) => ValueType::Any,
Self::Composite(_) => ValueType::Composite,
}
}
fn blob(&self) -> Option<&Bytes> {
match self {
Self::String(bytes) | Self::Any(bytes) | Self::Composite(bytes) => Some(bytes),
_ => None,
}
}
fn scalar(&self) -> Option<u64> {
match self {
Self::Int64(value) => Some(*value as u64),
Self::Uint64(value) => Some(*value),
Self::Double(value) => Some(value.to_bits()),
Self::Boolean(value) => Some(u64::from(*value)),
_ => None,
}
}
pub fn wire_size(&self) -> usize {
let payload = match self {
Self::Null => 0,
Self::String(bytes) | Self::Any(bytes) | Self::Composite(bytes) => {
bytes.len() + padding_for(bytes.len())
}
_ => ALIGNMENT,
};
ALIGNMENT + payload
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct UnversionedValue {
pub id: u16,
pub aggregate: bool,
pub value: Value,
}
impl UnversionedValue {
pub fn new(id: u16, value: Value) -> Self {
Self {
id,
aggregate: false,
value,
}
}
pub fn wire_size(&self) -> usize {
self.value.wire_size()
}
}
pub type Row = Vec<UnversionedValue>;
pub type MaybeRow = Option<Row>;
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum WireError {
#[error("rowset is truncated: need {needed} more bytes at offset {offset}")]
Truncated { offset: usize, needed: usize },
#[error("rowset declares {count} rows, more than the {MAX_ROWS_PER_ROWSET} allowed")]
TooManyRows { count: u64 },
#[error("row {row} declares {count} values, more than the {MAX_VALUES_PER_ROW} allowed")]
TooManyValues { row: usize, count: u64 },
#[error("value is {length} bytes, more than the {MAX_VALUE_LENGTH} allowed")]
ValueTooLong { length: u32 },
#[error("rowset is {size} bytes, more than the {MAX_ROWSET_SIZE}-byte RPC attachment limit")]
RowsetTooLarge { size: usize },
#[error("could not reserve {size} bytes for a rowset")]
AllocationFailed { size: usize },
#[error("unknown value type {0:#04x}")]
UnknownValueType(u8),
#[error("{0} bytes are left over after the last row")]
TrailingBytes(usize),
}
const fn padding_for(length: usize) -> usize {
(ALIGNMENT - (length % ALIGNMENT)) % ALIGNMENT
}
pub fn encoded_size(rows: &[MaybeRow]) -> usize {
let mut size = ALIGNMENT;
for row in rows {
size += ALIGNMENT;
if let Some(row) = row {
size += row.iter().map(UnversionedValue::wire_size).sum::<usize>();
}
}
size
}
pub fn encode_rowset(rows: &[MaybeRow]) -> Result<Bytes, WireError> {
let size = validate_and_measure(rows)?;
let mut storage = Vec::new();
storage
.try_reserve_exact(size)
.map_err(|_| WireError::AllocationFailed { size })?;
let mut buffer = BytesMut::from(Bytes::from(storage));
encode_rowset_unchecked(rows, &mut buffer);
Ok(buffer.freeze())
}
pub fn encode_rowset_into(rows: &[MaybeRow], out: &mut BytesMut) -> Result<(), WireError> {
validate_and_measure(rows)?;
encode_rowset_unchecked(rows, out);
Ok(())
}
fn validate_and_measure(rows: &[MaybeRow]) -> Result<usize, WireError> {
if rows.len() as u64 > MAX_ROWS_PER_ROWSET {
return Err(WireError::TooManyRows {
count: rows.len() as u64,
});
}
let mut size = ALIGNMENT;
for (index, row) in rows.iter().enumerate() {
add_to_rowset_size(&mut size, ALIGNMENT)?;
let Some(row) = row else {
continue;
};
if row.len() as u64 > MAX_VALUES_PER_ROW {
return Err(WireError::TooManyValues {
row: index,
count: row.len() as u64,
});
}
for value in row {
validate_value(value)?;
add_to_rowset_size(&mut size, value.wire_size())?;
}
}
Ok(size)
}
fn add_to_rowset_size(size: &mut usize, additional: usize) -> Result<(), WireError> {
*size = size
.checked_add(additional)
.ok_or(WireError::RowsetTooLarge { size: usize::MAX })?;
if *size > MAX_ROWSET_SIZE {
return Err(WireError::RowsetTooLarge { size: *size });
}
Ok(())
}
fn encode_rowset_unchecked(rows: &[MaybeRow], out: &mut BytesMut) {
out.put_u64_le(rows.len() as u64);
for row in rows {
let Some(row) = row else {
out.put_u64_le(NULL_ROW_MARKER);
continue;
};
out.put_u64_le(row.len() as u64);
for value in row {
encode_value_unchecked(value, out);
}
}
}
fn validate_value(value: &UnversionedValue) -> Result<(), WireError> {
let blob = value.value.blob();
if let Some(blob) = blob
&& blob.len() as u64 > u64::from(MAX_VALUE_LENGTH)
{
return Err(WireError::ValueTooLong {
length: blob.len().min(u32::MAX as usize) as u32,
});
}
Ok(())
}
fn encode_value_unchecked(value: &UnversionedValue, out: &mut BytesMut) {
let blob = value.value.blob();
out.put_u16_le(value.id);
out.put_u8(value.value.value_type() as u8);
out.put_u8(u8::from(value.aggregate));
out.put_u32_le(blob.map_or(0, |bytes| bytes.len() as u32));
if let Some(scalar) = value.value.scalar() {
out.put_u64_le(scalar);
} else if let Some(blob) = blob {
out.put_slice(blob);
out.put_bytes(0, padding_for(blob.len()));
}
}
pub fn decode_rowset(input: &Bytes) -> Result<Vec<MaybeRow>, WireError> {
let mut reader = Reader { input, offset: 0 };
let row_count = reader.read_u64()?;
if row_count > MAX_ROWS_PER_ROWSET {
return Err(WireError::TooManyRows { count: row_count });
}
let mut rows = Vec::new();
for index in 0..row_count as usize {
let value_count = reader.read_u64()?;
if value_count == NULL_ROW_MARKER {
rows.push(None);
continue;
}
if value_count > MAX_VALUES_PER_ROW {
return Err(WireError::TooManyValues {
row: index,
count: value_count,
});
}
let mut row = Row::with_capacity(value_count as usize);
for _ in 0..value_count {
row.push(reader.read_value()?);
}
rows.push(Some(row));
}
if reader.offset != input.len() {
return Err(WireError::TrailingBytes(input.len() - reader.offset));
}
Ok(rows)
}
struct Reader<'a> {
input: &'a Bytes,
offset: usize,
}
impl Reader<'_> {
fn need(&self, count: usize) -> Result<(), WireError> {
if self.input.len() - self.offset < count {
return Err(WireError::Truncated {
offset: self.offset,
needed: count - (self.input.len() - self.offset),
});
}
Ok(())
}
fn read_u64(&mut self) -> Result<u64, WireError> {
self.need(8)?;
let mut slice = &self.input[self.offset..self.offset + 8];
self.offset += 8;
Ok(slice.get_u64_le())
}
fn read_value(&mut self) -> Result<UnversionedValue, WireError> {
self.need(8)?;
let header = &self.input[self.offset..self.offset + 8];
let id = u16::from_le_bytes(header[0..2].try_into().unwrap());
let raw_type = header[2];
let aggregate = header[3] != 0;
let length = u32::from_le_bytes(header[4..8].try_into().unwrap());
self.offset += 8;
let value_type =
ValueType::from_wire(raw_type).ok_or(WireError::UnknownValueType(raw_type))?;
let value = match value_type {
ValueType::Null => Value::Null,
ValueType::Int64 | ValueType::Uint64 | ValueType::Double | ValueType::Boolean => {
let word = self.read_u64()?;
match value_type {
ValueType::Int64 => Value::Int64(word as i64),
ValueType::Uint64 => Value::Uint64(word),
ValueType::Double => Value::Double(f64::from_bits(word)),
_ => Value::Boolean(word != 0),
}
}
ValueType::String | ValueType::Any | ValueType::Composite => {
if length > MAX_VALUE_LENGTH {
return Err(WireError::ValueTooLong { length });
}
let length = length as usize;
let padded = length + padding_for(length);
self.need(padded)?;
let blob = self.input.slice(self.offset..self.offset + length);
self.offset += padded;
match value_type {
ValueType::String => Value::String(blob),
ValueType::Any => Value::Any(blob),
_ => Value::Composite(blob),
}
}
};
Ok(UnversionedValue {
id,
aggregate,
value,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn unwrap_encode(rows: &[MaybeRow]) -> Bytes {
encode_rowset(rows).expect("this rowset is within every limit")
}
fn round_trip(rows: &[MaybeRow]) -> Vec<MaybeRow> {
let encoded = unwrap_encode(rows);
assert_eq!(
encoded.len(),
encoded_size(rows),
"encoded_size disagrees with what encode_rowset wrote"
);
assert_eq!(
encoded.len() % ALIGNMENT,
0,
"a rowset is 8-byte aligned throughout"
);
decode_rowset(&encoded).expect("what this encoder wrote must decode")
}
fn sample_row() -> Row {
vec![
UnversionedValue::new(1, Value::Null),
UnversionedValue::new(2, Value::Boolean(true)),
UnversionedValue::new(3, Value::Boolean(false)),
UnversionedValue::new(4, Value::Int64(-42)),
UnversionedValue::new(5, Value::Uint64(42)),
UnversionedValue::new(6, Value::Double(1.25)),
UnversionedValue::new(7, Value::String(Bytes::from_static(b"foobar"))),
UnversionedValue::new(8, Value::Any(Bytes::from_static(b"[1;2;3]"))),
UnversionedValue::new(9, Value::String(Bytes::new())),
]
}
#[test]
fn value_types_carry_the_documented_numbers() {
assert_eq!(ValueType::Null as u8, 0x02);
assert_eq!(ValueType::Int64 as u8, 0x03);
assert_eq!(ValueType::Uint64 as u8, 0x04);
assert_eq!(ValueType::Double as u8, 0x05);
assert_eq!(ValueType::Boolean as u8, 0x06);
assert_eq!(ValueType::String as u8, 0x10);
assert_eq!(ValueType::Any as u8, 0x11);
assert_eq!(ValueType::Composite as u8, 0x12);
}
#[test]
fn the_rowset_and_row_headers_are_single_words() {
let encoded = unwrap_encode(&[Some(vec![UnversionedValue::new(0, Value::Int64(7))])]);
assert_eq!(&encoded[0..8], &1u64.to_le_bytes(), "row count");
assert_eq!(&encoded[8..16], &1u64.to_le_bytes(), "value count");
assert_eq!(encoded.len(), 8 + 8 + 8 + 8);
}
#[test]
fn a_value_header_is_id_type_aggregate_length() {
let value = UnversionedValue {
id: 0x1234,
aggregate: true,
value: Value::String(Bytes::from_static(b"abc")),
};
let encoded = unwrap_encode(&[Some(vec![value])]);
let header = &encoded[16..24];
assert_eq!(&header[0..2], &0x1234u16.to_le_bytes(), "id");
assert_eq!(header[2], ValueType::String as u8, "type");
assert_eq!(header[3], 1, "aggregate");
assert_eq!(&header[4..8], &3u32.to_le_bytes(), "length");
assert_eq!(&encoded[24..27], b"abc");
assert_eq!(
&encoded[27..32],
&[0, 0, 0, 0, 0],
"padded to eight with zeroes"
);
}
#[test]
fn a_null_value_has_no_payload_at_all() {
let encoded = unwrap_encode(&[Some(vec![UnversionedValue::new(1, Value::Null)])]);
assert_eq!(encoded.len(), 24);
assert_eq!(
&encoded[20..24],
&0u32.to_le_bytes(),
"length word stays zero"
);
}
#[test]
fn scalars_occupy_exactly_one_word() {
for value in [
Value::Int64(-1),
Value::Uint64(u64::MAX),
Value::Double(-0.0),
Value::Boolean(true),
] {
let encoded = unwrap_encode(&[Some(vec![UnversionedValue::new(0, value.clone())])]);
assert_eq!(
encoded.len(),
32,
"{value:?} should be header plus one word"
);
assert_eq!(
&encoded[20..24],
&0u32.to_le_bytes(),
"{value:?} must leave the length word zero"
);
}
}
#[test]
fn everything_round_trips() {
let rows = vec![None, Some(Vec::new()), Some(sample_row())];
assert_eq!(round_trip(&rows), rows);
}
#[test]
fn a_null_row_is_not_an_empty_row() {
let encoded_null = unwrap_encode(&[None]);
let encoded_empty = unwrap_encode(&[Some(Vec::new())]);
assert_ne!(encoded_null, encoded_empty);
assert_eq!(&encoded_null[8..16], &NULL_ROW_MARKER.to_le_bytes());
assert_eq!(&encoded_empty[8..16], &0u64.to_le_bytes());
assert_eq!(decode_rowset(&encoded_null).unwrap(), vec![None]);
assert_eq!(
decode_rowset(&encoded_empty).unwrap(),
vec![Some(Vec::new())]
);
}
#[test]
fn strings_of_every_length_modulo_eight_round_trip() {
for length in 0..24usize {
let blob = Bytes::from(vec![b'x'; length]);
let rows = vec![Some(vec![UnversionedValue::new(
0,
Value::String(blob.clone()),
)])];
let encoded = unwrap_encode(&rows);
assert_eq!(
encoded.len() % ALIGNMENT,
0,
"length {length} left the stream unaligned"
);
assert_eq!(
round_trip(&rows),
rows,
"length {length} did not round-trip"
);
}
}
#[test]
fn composite_values_keep_their_payload() {
let rows = vec![Some(vec![UnversionedValue::new(
3,
Value::Composite(Bytes::from_static(b"[1;2;3]")),
)])];
let decoded = round_trip(&rows);
assert_eq!(decoded, rows);
match &decoded[0].as_ref().unwrap()[0].value {
Value::Composite(blob) => assert_eq!(blob, &Bytes::from_static(b"[1;2;3]")),
other => panic!("expected a composite value, got {other:?}"),
}
}
#[test]
fn doubles_survive_bit_for_bit() {
for value in [
0.0,
-0.0,
1.25,
f64::MIN,
f64::MAX,
f64::INFINITY,
f64::NEG_INFINITY,
] {
let rows = vec![Some(vec![UnversionedValue::new(0, Value::Double(value))])];
let decoded = round_trip(&rows);
match decoded[0].as_ref().unwrap()[0].value {
Value::Double(read) => assert_eq!(read.to_bits(), value.to_bits()),
ref other => panic!("expected a double, got {other:?}"),
}
}
let rows = vec![Some(vec![UnversionedValue::new(
0,
Value::Double(f64::NAN),
)])];
let encoded = unwrap_encode(&rows);
match decode_rowset(&encoded).unwrap()[0].as_ref().unwrap()[0].value {
Value::Double(read) => assert!(read.is_nan()),
ref other => panic!("expected a double, got {other:?}"),
}
}
#[test]
fn negative_integers_use_two_s_complement_in_the_word() {
let rows = vec![Some(vec![UnversionedValue::new(0, Value::Int64(-42))])];
let encoded = unwrap_encode(&rows);
assert_eq!(&encoded[24..32], &(-42i64 as u64).to_le_bytes());
assert_eq!(round_trip(&rows), rows);
}
#[test]
fn the_aggregate_flag_survives() {
let rows = vec![Some(vec![UnversionedValue {
id: 5,
aggregate: true,
value: Value::Int64(1),
}])];
assert_eq!(round_trip(&rows), rows);
}
#[test]
fn truncated_input_is_an_error_not_a_panic() {
let rows = vec![Some(sample_row())];
let whole = unwrap_encode(&rows);
for length in 0..whole.len() {
let truncated = whole.slice(0..length);
match decode_rowset(&truncated) {
Err(WireError::Truncated { offset, needed }) => {
assert!(
needed > 0,
"a truncation that needs no more bytes is not one"
);
assert!(
offset <= length,
"reported offset {offset} is past the {length} bytes given"
);
}
Err(other) => panic!("cut to {length} bytes gave {other:?}, not a truncation"),
Ok(rows) => panic!("a rowset cut to {length} bytes decoded to {rows:?}"),
}
}
}
#[test]
fn trailing_bytes_are_rejected() {
let mut encoded = BytesMut::from(&unwrap_encode(&[Some(sample_row())])[..]);
encoded.put_u64_le(0);
assert_eq!(
decode_rowset(&encoded.freeze()),
Err(WireError::TrailingBytes(8))
);
}
#[test]
fn an_unknown_value_type_is_rejected() {
let mut encoded = BytesMut::from(
&unwrap_encode(&[Some(vec![UnversionedValue::new(0, Value::Int64(1))])])[..],
);
encoded[18] = 0x7f;
assert_eq!(
decode_rowset(&encoded.freeze()),
Err(WireError::UnknownValueType(0x7f))
);
}
#[test]
fn an_absurd_row_count_is_rejected_before_anything_is_reserved() {
let mut buffer = BytesMut::new();
buffer.put_u64_le(MAX_ROWS_PER_ROWSET + 1);
assert_eq!(
decode_rowset(&buffer.freeze()),
Err(WireError::TooManyRows {
count: MAX_ROWS_PER_ROWSET + 1
})
);
}
#[test]
fn an_absurd_value_count_is_rejected() {
let mut buffer = BytesMut::new();
buffer.put_u64_le(1);
buffer.put_u64_le(MAX_VALUES_PER_ROW + 1);
assert_eq!(
decode_rowset(&buffer.freeze()),
Err(WireError::TooManyValues {
row: 0,
count: MAX_VALUES_PER_ROW + 1
})
);
}
#[test]
fn an_absurd_value_length_is_rejected_before_the_bytes_are_read() {
let mut buffer = BytesMut::new();
buffer.put_u64_le(1);
buffer.put_u64_le(1);
buffer.put_u16_le(0);
buffer.put_u8(ValueType::String as u8);
buffer.put_u8(0);
buffer.put_u32_le(MAX_VALUE_LENGTH + 1);
assert_eq!(
decode_rowset(&buffer.freeze()),
Err(WireError::ValueTooLong {
length: MAX_VALUE_LENGTH + 1
})
);
}
#[test]
fn a_large_but_legal_row_count_with_no_rows_behind_it_fails_cheaply() {
let mut buffer = BytesMut::new();
buffer.put_u64_le(MAX_ROWS_PER_ROWSET);
assert!(matches!(
decode_rowset(&buffer.freeze()),
Err(WireError::Truncated { .. })
));
}
#[test]
fn the_encoder_refuses_what_the_decoder_would() {
let too_many_values = vec![Some(
(0..MAX_VALUES_PER_ROW as u16 + 1)
.map(|id| UnversionedValue::new(id, Value::Int64(0)))
.collect::<Row>(),
)];
assert_eq!(
encode_rowset(&too_many_values),
Err(WireError::TooManyValues {
row: 0,
count: MAX_VALUES_PER_ROW + 1
})
);
let too_long = vec![Some(vec![UnversionedValue::new(
0,
Value::String(Bytes::from(vec![0u8; MAX_VALUE_LENGTH as usize + 1])),
)])];
assert_eq!(
encode_rowset(&too_long),
Err(WireError::ValueTooLong {
length: MAX_VALUE_LENGTH + 1
})
);
}
#[test]
fn too_many_shared_large_values_are_refused_before_allocation() {
let shared = Bytes::from(vec![0; MAX_VALUE_LENGTH as usize]);
let row = (0..=MAX_VALUES_PER_ROW)
.map(|_| UnversionedValue::new(0, Value::String(shared.clone())))
.collect::<Row>();
assert_eq!(
encode_rowset(&[Some(row)]),
Err(WireError::TooManyValues {
row: 0,
count: MAX_VALUES_PER_ROW + 1,
})
);
}
#[test]
fn a_rowset_larger_than_one_rpc_attachment_is_refused_before_allocation() {
let shared = Bytes::from(vec![0; MAX_VALUE_LENGTH as usize]);
let row = (0..65)
.map(|_| UnversionedValue::new(0, Value::String(shared.clone())))
.collect::<Row>();
assert!(matches!(
encode_rowset(&[Some(row)]),
Err(WireError::RowsetTooLarge { .. })
));
}
#[test]
fn a_rejected_encode_does_not_append_a_partial_rowset() {
let mut output = BytesMut::from(&b"prefix"[..]);
let rows = vec![Some(
(0..=MAX_VALUES_PER_ROW as u16)
.map(|id| UnversionedValue::new(id, Value::Int64(0)))
.collect::<Row>(),
)];
assert!(matches!(
encode_rowset_into(&rows, &mut output),
Err(WireError::TooManyValues { .. })
));
assert_eq!(&output[..], b"prefix");
}
#[test]
fn a_big_rowset_round_trips() {
let rows: Vec<MaybeRow> = (0..1000)
.map(|index| {
Some(vec![
UnversionedValue::new(0, Value::Int64(index)),
UnversionedValue::new(1, Value::String(Bytes::from(format!("row {index}")))),
])
})
.collect();
assert_eq!(round_trip(&rows), rows);
}
}