diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8537cd89..04469e81 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -244,10 +244,11 @@ jobs: # publishes the length afterwards, including on the unwind paths its tests # drive. `PreSized` (encode_sink.rs) writes a message into a sink's # uninitialised spare capacity and advances the sink by the bytes written. + # The table interpreters (table/) access message fields through raw pointers. # Miri tracks per-byte init state, so this job mechanically verifies these # invariants on every path those test modules exercise. Scoped to them to - # keep the (interpreted, slow) run fast — under a minute steady state, so it - # runs on every PR as a blocking gate rather than on a nightly schedule. + # keep the (interpreted, slow) run fast, so it runs on every PR as a blocking + # gate rather than on a nightly schedule. # # The nightly is PINNED (not floating): a required check must not be broken by # an unrelated bad nightly. Bump MIRI_TOOLCHAIN occasionally; the sysroot @@ -284,6 +285,15 @@ jobs: cargo +${{ env.MIRI_TOOLCHAIN }} miri test -p buffa -- size_cache copy_into_spare encode_sink::tests contiguous_sink_ + # The table interpreters (table/) read and write message fields through + # raw pointers at offsets from a `Table`. Miri checks the pointer + # provenance, alignment and initialisation on the kinds and field + # shapes the module's tests build, including message fields stored + # inline, boxed and in a `Vec`, whose child pointers point into the + # parent or the heap. + - name: Miri (table interpreter soundness) + run: cargo +${{ env.MIRI_TOOLCHAIN }} miri test -p buffa --lib -- 'table::' + # OwnedView: the view borrows from the Bytes it is stored beside, with # forged 'static lifetimes hidden behind a MaybeDangling wrapper. Miri # turns a wrong field order (a dangling borrow while the view's diff --git a/buffa/src/encode_sink.rs b/buffa/src/encode_sink.rs index f31ffc71..e7fad911 100644 --- a/buffa/src/encode_sink.rs +++ b/buffa/src/encode_sink.rs @@ -155,6 +155,30 @@ pub trait EncodeSink { #[doc(hidden)] const __PRE_SIZED: bool = false; + /// Run `f` on this sink if it is a [`PreSized`] cursor, and return + /// whether it did. + /// + /// Lets the table interpreters switch, once per message, from an instance + /// generic over the sink to one compiled in this crate. `f` gets the + /// cursor itself, and a `PreSized` cannot be built or altered outside this + /// crate, so `f` cannot claim more bytes than were written: + /// + /// ```compile_fail,E0616 + /// use buffa::EncodeSink; + /// + /// let mut vec: Vec = Vec::new(); + /// vec.__with_pre_sized(&mut |cursor| cursor.pos = 4); + /// ``` + /// + /// Implementations outside this crate keep the default, which does not + /// call `f`. + #[doc(hidden)] + #[inline] + fn __with_pre_sized(&mut self, f: &mut dyn FnMut(&mut PreSized<'_>)) -> bool { + let _ = f; + false + } + /// Run `fill` over `len` bytes of contiguous space at the end of the /// sink and append the bytes it wrote, or give `fill` back unrun if the /// sink's current chunk is shorter than `len`. @@ -271,6 +295,12 @@ impl<'a> PreSized<'a> { } impl EncodeSink for PreSized<'_> { + #[inline] + fn __with_pre_sized(&mut self, f: &mut dyn FnMut(&mut PreSized<'_>)) -> bool { + f(self); + true + } + #[inline] fn put_u8(&mut self, value: u8) { let Some(slot) = self.dst.get_mut(self.pos) else { diff --git a/buffa/src/lib.rs b/buffa/src/lib.rs index a5c5ceec..1afc87c0 100644 --- a/buffa/src/lib.rs +++ b/buffa/src/lib.rs @@ -236,6 +236,10 @@ pub mod message_field; pub mod message_set; pub mod oneof; mod size_cache; +// Runtime for table-driven message codecs, called by generated code; see the +// module docs. +#[doc(hidden)] +pub mod table; #[cfg(test)] pub(crate) mod test_doubles; #[cfg(feature = "text")] diff --git a/buffa/src/table/decode.rs b/buffa/src/table/decode.rs new file mode 100644 index 00000000..cce1bc19 --- /dev/null +++ b/buffa/src/table/decode.rs @@ -0,0 +1,442 @@ +//! The decode pass: [`merge_to_limit`] and its per-kind arms. +//! +//! Decoding runs over one contiguous `&[u8]`, so every arm reads through a +//! non-generic function compiled in this crate. + +use super::scalar::Sc; +use super::{ + Bool, Double, Entry, EnumVt, Fixed32, Fixed64, Float, Int32, Int64, Kind, MessageTable, + Sfixed32, Sfixed64, Sint32, Sint64, Uint32, Uint64, IMPLICIT, NO_UNKNOWN, OPTIONAL, PACKED, + REPEATED, REQUIRED, +}; +use crate::alloc::{string::String, vec::Vec}; +use crate::bytes::Buf; +use crate::encoding::{ + check_wire_type, decode_unknown_field, decode_varint, skip_field_depth, wire_type_mismatch, + Tag, WireType, +}; +use crate::message::MAX_MESSAGE_BYTES; +use crate::{types, DecodeContext, DecodeError, UnknownField, UnknownFieldData, UnknownFields}; + +/// Decode into the message at `base` until `buf` has `limit` bytes remaining. +/// +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +pub(super) unsafe fn merge_to_limit( + table: &MessageTable, + base: *mut u8, + buf: &mut B, + ctx: DecodeContext<'_>, + limit: usize, +) -> Result<(), DecodeError> { + let n = buf.remaining().saturating_sub(limit); + let chunk = buf.chunk(); + if chunk.len() >= n { + let mut payload = &chunk[..n]; + // SAFETY: forwarded from the caller. + unsafe { merge_slice(table, base, &mut payload, ctx)? }; + buf.advance(n); + Ok(()) + } else { + let gathered = buf.copy_to_bytes(n); + let mut payload = &gathered[..]; + // SAFETY: forwarded from the caller. + unsafe { merge_slice(table, base, &mut payload, ctx) } + } +} + +/// Decode a length-prefixed message into the message at `base`. +/// +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +pub(super) unsafe fn merge_length_delimited( + table: &MessageTable, + base: *mut u8, + buf: &mut B, + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + let ctx = ctx.descend()?; + let len = decode_varint(buf)?; + if len > u64::from(MAX_MESSAGE_BYTES) { + return Err(DecodeError::MessageTooLarge); + } + let len = usize::try_from(len).map_err(|_| DecodeError::MessageTooLarge)?; + if buf.remaining() < len { + return Err(DecodeError::UnexpectedEof); + } + let limit = buf.remaining() - len; + // SAFETY: forwarded from the caller. + unsafe { merge_to_limit(table, base, buf, ctx, limit) } +} + +/// Decode one field, whose `tag` has been read, into the message at `base`. +/// The rest of `buf` must be one chunk. +/// +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +pub(super) unsafe fn merge_field( + table: &MessageTable, + base: *mut u8, + tag: Tag, + buf: &mut B, + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + let remaining = buf.remaining(); + let chunk = buf.chunk(); + if chunk.len() < remaining { + return Err(DecodeError::UnexpectedEof); + } + let mut payload = chunk; + // SAFETY: forwarded from the caller. + unsafe { merge_one(table, base, tag, &mut payload, ctx)? }; + let consumed = remaining - payload.len(); + buf.advance(consumed); + Ok(()) +} + +/// Decode every field in `buf` into the message at `base`. +/// +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +unsafe fn merge_slice( + table: &MessageTable, + base: *mut u8, + buf: &mut &[u8], + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + while !buf.is_empty() { + let tag = Tag::decode(buf)?; + // SAFETY: forwarded from the caller. + unsafe { merge_one(table, base, tag, buf, ctx)? }; + } + Ok(()) +} + +/// Decode a length-prefixed sub-message into the message at `base`. +/// +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +unsafe fn merge_sub( + table: &MessageTable, + base: *mut u8, + buf: &mut &[u8], + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + let ctx = ctx.descend()?; + let len = decode_varint(buf)?; + if len > u64::from(MAX_MESSAGE_BYTES) { + return Err(DecodeError::MessageTooLarge); + } + let len = usize::try_from(len).map_err(|_| DecodeError::MessageTooLarge)?; + if buf.len() < len { + return Err(DecodeError::UnexpectedEof); + } + let (mut payload, rest) = buf.split_at(len); + *buf = rest; + // SAFETY: forwarded from the caller. + unsafe { merge_slice(table, base, &mut payload, ctx) } +} + +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +#[inline] +unsafe fn merge_one( + table: &MessageTable, + base: *mut u8, + tag: Tag, + buf: &mut &[u8], + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + let Some(e) = table.find(tag.field_number()) else { + // SAFETY: forwarded from the caller. + return unsafe { merge_unknown(table, base, tag, buf, ctx) }; + }; + // SAFETY: the offset is within the message, per the table's contract. + unsafe { merge_kind(table, e, base, base.add(e.offset as usize), tag, buf, ctx) } +} + +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +#[cold] +unsafe fn merge_unknown( + table: &MessageTable, + base: *mut u8, + tag: Tag, + buf: &mut &[u8], + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + if table.unknown == NO_UNKNOWN { + return skip_field_depth(tag, buf, ctx.depth()); + } + let field = decode_unknown_field(tag, buf, ctx)?; + // SAFETY: `unknown` is the offset of the message's `UnknownFields`. + unsafe { (*base.add(table.unknown as usize).cast::()).push(field) }; + Ok(()) +} + +macro_rules! merge_dispatch { + ($($name:ident: $fam:ident $ty:ident $card:ident;)*) => { + /// # Safety + /// + /// `slot` points to the field `e` describes, inside the live message + /// at `base` of the type `table` describes. + #[inline] + unsafe fn merge_kind( + table: &MessageTable, + e: &Entry, + base: *mut u8, + slot: *mut u8, + tag: Tag, + buf: &mut &[u8], + ctx: DecodeContext<'_>, + ) -> Result<(), DecodeError> { + // SAFETY: each arm writes the slot as the type its kind names. + unsafe { + match e.kind { + $(Kind::$name => merge_dispatch!(@arm $fam $ty $card table e base slot tag buf ctx),)* + } + } + } + }; + (@arm Scalar $ty:ident $card:ident $table:ident $e:ident $base:ident $slot:ident $tag:ident $buf:ident $ctx:ident) => { + merge_scalar::<$ty, $card>($slot, $tag, $buf) + }; + (@arm Str $ty:ident $card:ident $table:ident $e:ident $base:ident $slot:ident $tag:ident $buf:ident $ctx:ident) => { + merge_str::<$card>($slot, $tag, $buf, $ctx) + }; + (@arm Bytes $ty:ident $card:ident $table:ident $e:ident $base:ident $slot:ident $tag:ident $buf:ident $ctx:ident) => { + merge_bytes::<$card>($slot, $tag, $buf, $ctx) + }; + (@arm Enum $ty:ident $card:ident $table:ident $e:ident $base:ident $slot:ident $tag:ident $buf:ident $ctx:ident) => { + merge_enum::<$card>($table, $e, $base, $slot, $tag, $buf, $ctx) + }; + (@arm Msg $ty:ident $card:ident $table:ident $e:ident $base:ident $slot:ident $tag:ident $buf:ident $ctx:ident) => { + merge_msg::<$card>($table, $e, $slot, $tag, $buf, $ctx) + }; +} + +kind_table!(merge_dispatch); + +/// # Safety +/// +/// `slot` points to a field of scalar type `S` in the shape `C` names. +#[inline] +unsafe fn merge_scalar( + slot: *mut u8, + tag: Tag, + buf: &mut &[u8], +) -> Result<(), DecodeError> { + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT | REQUIRED => { + check_wire_type(tag, S::WIRE)?; + *slot.cast::() = S::read(buf)?; + } + OPTIONAL => { + check_wire_type(tag, S::WIRE)?; + *slot.cast::>() = Some(S::read(buf)?); + } + _ => { + let out = &mut *slot.cast::>(); + let wire = tag.wire_type(); + if wire == WireType::LengthDelimited { + let payload = take_len_delimited(buf)?; + S::extend(payload, out)?; + } else if wire == S::WIRE { + out.push(S::read(buf)?); + } else { + return Err(wire_type_mismatch(tag, WireType::LengthDelimited)); + } + } + } + } + Ok(()) +} + +/// Split a length-prefixed payload off the front of `buf`. +#[inline] +fn take_len_delimited<'a>(buf: &mut &'a [u8]) -> Result<&'a [u8], DecodeError> { + let len = decode_varint(buf)?; + let len = usize::try_from(len).map_err(|_| DecodeError::MessageTooLarge)?; + if buf.len() < len { + return Err(DecodeError::UnexpectedEof); + } + let (payload, rest) = buf.split_at(len); + *buf = rest; + Ok(payload) +} + +/// # Safety +/// +/// `slot` points to a `String` field in the shape `C` names. +#[inline] +unsafe fn merge_str( + slot: *mut u8, + tag: Tag, + buf: &mut &[u8], + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + check_wire_type(tag, WireType::LengthDelimited)?; + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT | REQUIRED => types::merge_string(&mut *slot.cast::(), buf), + OPTIONAL => types::merge_string( + (*slot.cast::>()).get_or_insert_with(String::new), + buf, + ), + _ => { + let elem = types::decode_string(buf)?; + ctx.register_element_memory(core::mem::size_of::())?; + (*slot.cast::>()).push(elem); + Ok(()) + } + } + } +} + +/// # Safety +/// +/// `slot` points to a `Vec` field in the shape `C` names. +#[inline] +unsafe fn merge_bytes( + slot: *mut u8, + tag: Tag, + buf: &mut &[u8], + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + check_wire_type(tag, WireType::LengthDelimited)?; + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT | REQUIRED => types::merge_bytes(&mut *slot.cast::>(), buf), + OPTIONAL => types::merge_bytes( + (*slot.cast::>>()).get_or_insert_with(Vec::new), + buf, + ), + _ => { + let elem = types::decode_bytes(buf)?; + ctx.register_element_memory(core::mem::size_of::>())?; + (*slot.cast::>>()).push(elem); + Ok(()) + } + } + } +} + +/// # Safety +/// +/// `slot` points to the message field `e` describes, in the shape `C` names. +#[inline] +unsafe fn merge_msg( + table: &MessageTable, + e: &Entry, + slot: *mut u8, + tag: Tag, + buf: &mut &[u8], + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + check_wire_type(tag, WireType::LengthDelimited)?; + // SAFETY: the descriptor was built for the slot's shape and child type. + unsafe { + if C == REPEATED { + let vt = table.rep_vt(e); + ctx.register_element_memory(vt.size)?; + let elem = (vt.push)(slot); + let decoded = merge_sub(vt.table, elem, buf, ctx); + if decoded.is_err() { + // Like unrolled code, which decodes into a local and pushes + // it only on success, leave no partial element behind. + (vt.pop)(slot); + } + decoded + } else { + let vt = table.msg_vt(e); + let child = (vt.place)(slot); + merge_sub(vt.table, child, buf, ctx) + } + } +} + +/// # Safety +/// +/// `slot` points to the enum field `e` describes, in the shape `C` names, +/// inside the live message at `base`. +#[inline] +unsafe fn merge_enum( + table: &MessageTable, + e: &Entry, + base: *mut u8, + slot: *mut u8, + tag: Tag, + buf: &mut &[u8], + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + let vt = table.enum_vt(e); + // SAFETY: `vt` was built for the slot's shape. + unsafe { + match C { + REPEATED | PACKED => { + let wire = tag.wire_type(); + if wire == WireType::LengthDelimited { + let mut payload = take_len_delimited(buf)?; + while !payload.is_empty() { + let raw = types::decode_int32_packed(&mut payload)?; + enum_store(table, e, base, vt, slot, raw, ctx)?; + } + Ok(()) + } else if wire == WireType::Varint { + let raw = types::decode_int32(buf)?; + enum_store(table, e, base, vt, slot, raw, ctx) + } else { + Err(wire_type_mismatch(tag, WireType::LengthDelimited)) + } + } + _ => { + check_wire_type(tag, WireType::Varint)?; + let raw = types::decode_int32(buf)?; + enum_store(table, e, base, vt, slot, raw, ctx) + } + } + } +} + +/// Store `raw` in the enum field at `slot`; a value a closed enum rejects goes +/// to the message's unknown fields, or is dropped if it keeps none. +/// +/// # Safety +/// +/// `slot` is a live slot of the shape `vt` was built for, inside the live +/// message at `base` of the type `table` describes. +#[inline] +unsafe fn enum_store( + table: &MessageTable, + e: &Entry, + base: *mut u8, + vt: &EnumVt, + slot: *mut u8, + raw: i32, + ctx: DecodeContext<'_>, +) -> Result<(), DecodeError> { + // SAFETY: `slot` matches the shape `vt` was built for. + if unsafe { (vt.set)(slot, raw) } || table.unknown == NO_UNKNOWN { + return Ok(()); + } + ctx.register_unknown_field()?; + // SAFETY: `unknown` is the offset of the message's `UnknownFields`. + unsafe { + (*base.add(table.unknown as usize).cast::()).push(UnknownField { + number: e.tag >> 3, + data: UnknownFieldData::Varint(raw as u64), + }); + } + Ok(()) +} diff --git a/buffa/src/table/encode.rs b/buffa/src/table/encode.rs new file mode 100644 index 00000000..aeb8ef42 --- /dev/null +++ b/buffa/src/table/encode.rs @@ -0,0 +1,330 @@ +//! The write pass: [`write_to`] and its per-kind arms. + +use super::scalar::Sc; +use super::{ + Bool, Double, Entry, Fixed32, Fixed64, Float, Int32, Int64, Kind, MessageTable, Sfixed32, + Sfixed64, Sint32, Sint64, Uint32, Uint64, IMPLICIT, NO_UNKNOWN, OPTIONAL, PACKED, REPEATED, + REQUIRED, +}; +use crate::alloc::{string::String, vec::Vec}; +use crate::encode_sink::PreSized; +use crate::encoding::encode_varint; +use crate::{types, EncodeSink, SizeCache, UnknownFields}; + +/// Write the message at `base` to `buf`, consuming nested sizes from `cache`, +/// which [`compute_size`](super::size::compute_size) must have filled for the +/// message; a cache that does not match makes the write panic. +/// +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +pub(super) unsafe fn write_to( + table: &MessageTable, + base: *const u8, + cache: &mut SizeCache, + buf: &mut K, +) { + let mut through_cursor = |cursor: &mut PreSized<'_>| { + // SAFETY: forwarded from the caller. + unsafe { write_pre_sized(table, base, &mut *cache, cursor) } + }; + if !buf.__with_pre_sized(&mut through_cursor) { + // SAFETY: forwarded from the caller. + unsafe { write_message(table, base, cache, buf) }; + } +} + +/// [`write_message`] for a [`PreSized`] sink, which every `BufMut` is written +/// through. Not generic, so it is compiled once, in this crate, however many +/// sink types the caller's crate encodes into. +/// +/// # Safety +/// +/// As for [`write_to`]. +#[inline(never)] +unsafe fn write_pre_sized( + table: &MessageTable, + base: *const u8, + cache: &mut SizeCache, + buf: &mut PreSized<'_>, +) { + // SAFETY: forwarded from the caller. + unsafe { write_message(table, base, cache, buf) } +} + +/// # Safety +/// +/// As for [`write_to`]. +unsafe fn write_message( + table: &MessageTable, + base: *const u8, + cache: &mut SizeCache, + buf: &mut K, +) { + for e in table.entries { + // SAFETY: the offset is within the message, per the table's contract. + unsafe { write_kind(table, e, base.add(e.offset as usize), cache, buf) }; + } + if table.unknown != NO_UNKNOWN { + // SAFETY: `unknown` is the offset of the message's `UnknownFields`. + unsafe { (*base.add(table.unknown as usize).cast::()).write_to(buf) }; + } +} + +#[inline(always)] +fn put_tag(e: &Entry, buf: &mut K) { + if e.tag_len == 1 { + buf.put_u8(e.tag as u8); + } else { + encode_varint(u64::from(e.tag), buf); + } +} + +macro_rules! write_dispatch { + ($($name:ident: $fam:ident $ty:ident $card:ident;)*) => { + /// # Safety + /// + /// `slot` points to the field `e` describes, in a live message of the + /// type `table` describes. + #[inline] + unsafe fn write_kind( + table: &MessageTable, + e: &Entry, + slot: *const u8, + cache: &mut SizeCache, + buf: &mut K, + ) { + // SAFETY: each arm reads the slot as the type its kind names. + unsafe { + match e.kind { + $(Kind::$name => write_dispatch!(@arm $fam $ty $card table e slot cache buf),)* + } + } + } + }; + (@arm Scalar $ty:ident $card:ident $table:ident $e:ident $slot:ident $cache:ident $buf:ident) => { + write_scalar::<$ty, $card, K>($e, $slot, $buf) + }; + (@arm Str $ty:ident $card:ident $table:ident $e:ident $slot:ident $cache:ident $buf:ident) => { + write_str::<$card, K>($e, $slot, $buf) + }; + (@arm Bytes $ty:ident $card:ident $table:ident $e:ident $slot:ident $cache:ident $buf:ident) => { + write_bytes::<$card, K>($e, $slot, $buf) + }; + (@arm Enum $ty:ident $card:ident $table:ident $e:ident $slot:ident $cache:ident $buf:ident) => { + write_enum::<$card, K>($table, $e, $slot, $buf) + }; + (@arm Msg $ty:ident $card:ident $table:ident $e:ident $slot:ident $cache:ident $buf:ident) => { + write_msg::<$card, K>($table, $e, $slot, $cache, $buf) + }; +} + +kind_table!(write_dispatch); + +/// # Safety +/// +/// `slot` points to a field of scalar type `S` in the shape `C` names. +#[inline] +unsafe fn write_scalar(e: &Entry, slot: *const u8, buf: &mut K) { + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT => { + let v = *slot.cast::(); + if !S::is_default(v) { + put_tag(e, buf); + S::encode(v, buf); + } + } + REQUIRED => { + put_tag(e, buf); + S::encode(*slot.cast::(), buf); + } + OPTIONAL => { + if let Some(v) = *slot.cast::>() { + put_tag(e, buf); + S::encode(v, buf); + } + } + REPEATED => { + for &v in &*slot.cast::>() { + put_tag(e, buf); + S::encode(v, buf); + } + } + _ => { + let v = &*slot.cast::>(); + if !v.is_empty() { + let payload: u64 = v.iter().map(|&x| S::len(x)).sum(); + put_tag(e, buf); + encode_varint(payload, buf); + for &x in v { + S::encode(x, buf); + } + } + } + } + } +} + +/// # Safety +/// +/// `slot` points to a `String` field in the shape `C` names. +#[inline] +unsafe fn write_str(e: &Entry, slot: *const u8, buf: &mut K) { + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT => { + let s = &*slot.cast::(); + if !s.is_empty() { + put_tag(e, buf); + types::encode_string(s, buf); + } + } + REQUIRED => { + put_tag(e, buf); + types::encode_string(&*slot.cast::(), buf); + } + OPTIONAL => { + if let Some(s) = &*slot.cast::>() { + put_tag(e, buf); + types::encode_string(s, buf); + } + } + _ => { + for s in &*slot.cast::>() { + put_tag(e, buf); + types::encode_string(s, buf); + } + } + } + } +} + +/// # Safety +/// +/// `slot` points to a `Vec` field in the shape `C` names. +#[inline] +unsafe fn write_bytes(e: &Entry, slot: *const u8, buf: &mut K) { + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT => { + let s = &*slot.cast::>(); + if !s.is_empty() { + put_tag(e, buf); + types::encode_shared_bytes(s, buf); + } + } + REQUIRED => { + put_tag(e, buf); + types::encode_shared_bytes(&*slot.cast::>(), buf); + } + OPTIONAL => { + if let Some(s) = &*slot.cast::>>() { + put_tag(e, buf); + types::encode_shared_bytes(s, buf); + } + } + _ => { + for s in &*slot.cast::>>() { + put_tag(e, buf); + types::encode_shared_bytes(s, buf); + } + } + } + } +} + +/// # Safety +/// +/// `slot` points to the enum field `e` describes, in the shape `C` names. +#[inline] +unsafe fn write_enum( + table: &MessageTable, + e: &Entry, + slot: *const u8, + buf: &mut K, +) { + let vt = table.enum_vt(e); + // SAFETY: `vt` was built for the slot's shape. + unsafe { + match C { + IMPLICIT => { + if let Some(v) = (vt.get)(slot, 0) { + if v != 0 { + put_tag(e, buf); + types::encode_int32(v, buf); + } + } + } + REQUIRED | OPTIONAL => { + if let Some(v) = (vt.get)(slot, 0) { + put_tag(e, buf); + types::encode_int32(v, buf); + } + } + REPEATED => { + for i in 0..(vt.len)(slot) { + if let Some(v) = (vt.get)(slot, i) { + put_tag(e, buf); + types::encode_int32(v, buf); + } + } + } + _ => { + let n = (vt.len)(slot); + if n == 0 { + return; + } + let mut payload = 0; + for i in 0..n { + if let Some(v) = (vt.get)(slot, i) { + payload += types::int32_encoded_len(v) as u64; + } + } + put_tag(e, buf); + encode_varint(payload, buf); + for i in 0..n { + if let Some(v) = (vt.get)(slot, i) { + types::encode_int32(v, buf); + } + } + } + } + } +} + +/// # Safety +/// +/// `slot` points to the message field `e` describes, in the shape `C` names. +#[inline] +unsafe fn write_msg( + table: &MessageTable, + e: &Entry, + slot: *const u8, + cache: &mut SizeCache, + buf: &mut K, +) { + // SAFETY: the descriptor was built for the slot's shape and child type. + unsafe { + if C == REPEATED { + let vt = table.rep_vt(e); + let (ptr, len) = (vt.parts)(slot); + for i in 0..len { + put_tag(e, buf); + encode_varint(u64::from(cache.consume_next()), buf); + write_message(vt.table, ptr.add(i * vt.size), cache, buf); + } + } else { + let vt = table.msg_vt(e); + let child = (vt.get)(slot); + if !child.is_null() { + put_tag(e, buf); + encode_varint(u64::from(cache.consume_next()), buf); + write_message(vt.table, child, cache, buf); + } + } + } +} diff --git a/buffa/src/table/mod.rs b/buffa/src/table/mod.rs new file mode 100644 index 00000000..1fdd784c --- /dev/null +++ b/buffa/src/table/mod.rs @@ -0,0 +1,846 @@ +//! Table-driven message codec: the runtime behind the table codec strategy of +//! code generation. +//! +//! Generated code normally emits `compute_size`, `write_to` and `merge_field` +//! bodies specialised to each message. With the table strategy it emits one +//! static [`Table`] per message instead: a sorted array of twelve-byte +//! [`Entry`] values, one per field, giving the field number, its byte offset +//! in the message struct, and a [`Kind`] (field type crossed with +//! cardinality). The `Message` methods forward to the interpreters here, +//! which every message shares. +//! +//! This module is support code for generated code, is not meant to be called +//! directly, and may change in any release, so generated code must be +//! regenerated with the `buffa-codegen` that matches the `buffa` it builds +//! against. [`ABI`] enforces that when the table is built. Building a +//! [`Table`] is `unsafe`, because the table's offsets and kinds must describe +//! the message struct's actual layout; the `__table!` and `__table_entry!` macros +//! that generated code uses check what the compiler can, and every other entry +//! point is safe and sound given a correct table. +//! +//! Generated code should name everything by absolute path (`::buffa::table::…`): +//! `Entry`, `Kind` and `Table` are also plausible message names. +//! +//! # Compared with unrolled code +//! +//! The wire format, the two-pass [`SizeCache`] protocol, the recursion, +//! unknown-field and element-memory limits of [`DecodeContext`], and whether +//! a given input decodes are the same, so a message can change strategy +//! without changing behaviour. The differences are: +//! +//! - Decoding runs over one contiguous slice, so a non-contiguous [`Buf`] is +//! gathered into a buffer first, and a field that declares a length past +//! the end of its enclosing message fails at once with +//! [`DecodeError::UnexpectedEof`], where unrolled code reads on into the +//! enclosing message and can report a different error. +//! - [`Table::merge_field`] decodes one field and cannot gather, so it +//! returns [`DecodeError::UnexpectedEof`] for a buffer that is not one +//! chunk. Only a caller that drives `merge_field` itself is affected, such +//! as the default `Message::merge_group`, so a message that is the type of a +//! group field must not use the table strategy. +//! - [`clear`](crate::Message::clear) resets to `Default`, which releases +//! allocations that unrolled code keeps. +//! +//! # Where the code is compiled +//! +//! `compute_size`, decoding from a contiguous buffer and encoding into a +//! [`BufMut`](crate::bytes::BufMut) through `Message::encode` and its +//! siblings are non-generic functions compiled in this crate, at this crate's +//! optimisation level, once for all messages. A build can therefore optimise +//! this crate for speed and its own generated code for size. Encoding into a +//! sink that is not a `BufMut`, such as [`Rope`](crate::Rope), and the +//! generic wrappers around decoding are instantiated in the crate that calls +//! them. + +use core::marker::PhantomData; + +use crate::alloc::{string::String, vec::Vec}; +use crate::bytes::Buf; +use crate::encoding::{Tag, WireType}; +use crate::{DecodeContext, DecodeError, EncodeSink, SizeCache, UnknownFields}; + +pub use shape::{ + EnumShape, EnumVt, ImplicitClosed, ImplicitOpen, MsgSlot, MsgVt, OptionalClosed, OptionalOpen, + RepVt, RepeatedClosed, RepeatedOpen, +}; + +use scalar::{ + Bool, Double, Fixed32, Fixed64, Float, Int32, Int64, Sc, Sfixed32, Sfixed64, Sint32, Sint64, + Uint32, Uint64, +}; + +/// The version of the contract between generated tables and this module. +/// +/// Generated code passes the version it was generated for to [`Table::new`], +/// which refuses to build a table for any other, because a table built for +/// different rules could make the interpreters read fields at the wrong +/// offsets. +pub const ABI: u32 = 1; + +/// The `unknown` offset of a message that does not preserve unknown fields. +const NO_UNKNOWN: u32 = u32::MAX; + +/// Re-exported for generated code, which needs it to build a [`Table`] and +/// must work on the crate's minimum supported Rust version. +/// +/// Requires Rust 1.77; on older compilers using it is a compile error. +#[doc(hidden)] +#[rustversion::since(1.77)] +pub use core::mem::offset_of; + +/// Stand-in for `core::mem::offset_of!` on compilers that lack it. +#[doc(hidden)] +#[rustversion::before(1.77)] +#[macro_export] +macro_rules! __buffa_offset_of_unavailable { + ($($tt:tt)*) => { + compile_error!( + "the table codec strategy requires Rust 1.77 or newer (`core::mem::offset_of!`)" + ) + }; +} + +#[doc(hidden)] +#[rustversion::before(1.77)] +pub use __buffa_offset_of_unavailable as offset_of; + +// Cardinalities. The `Msg` kinds use `IMPLICIT` for a singular field. +const IMPLICIT: u8 = 0; +const REQUIRED: u8 = 1; +const OPTIONAL: u8 = 2; +const REPEATED: u8 = 3; +const PACKED: u8 = 4; + +/// Calls `$callback!` with every [`Kind`] as `Name: Family Type Cardinality;`. +/// +/// The families are `Scalar`, `Str`, `Bytes`, `Enum` and `Msg`. The one list +/// generates the enum and the three dispatch functions, so a kind cannot be +/// added to one and not the others. +macro_rules! kind_table { + ($callback:ident) => { + $callback! { + Int32Implicit: Scalar Int32 IMPLICIT; + Int32Required: Scalar Int32 REQUIRED; + Int32Optional: Scalar Int32 OPTIONAL; + Int32Repeated: Scalar Int32 REPEATED; + Int32Packed: Scalar Int32 PACKED; + Int64Implicit: Scalar Int64 IMPLICIT; + Int64Required: Scalar Int64 REQUIRED; + Int64Optional: Scalar Int64 OPTIONAL; + Int64Repeated: Scalar Int64 REPEATED; + Int64Packed: Scalar Int64 PACKED; + Uint32Implicit: Scalar Uint32 IMPLICIT; + Uint32Required: Scalar Uint32 REQUIRED; + Uint32Optional: Scalar Uint32 OPTIONAL; + Uint32Repeated: Scalar Uint32 REPEATED; + Uint32Packed: Scalar Uint32 PACKED; + Uint64Implicit: Scalar Uint64 IMPLICIT; + Uint64Required: Scalar Uint64 REQUIRED; + Uint64Optional: Scalar Uint64 OPTIONAL; + Uint64Repeated: Scalar Uint64 REPEATED; + Uint64Packed: Scalar Uint64 PACKED; + Sint32Implicit: Scalar Sint32 IMPLICIT; + Sint32Required: Scalar Sint32 REQUIRED; + Sint32Optional: Scalar Sint32 OPTIONAL; + Sint32Repeated: Scalar Sint32 REPEATED; + Sint32Packed: Scalar Sint32 PACKED; + Sint64Implicit: Scalar Sint64 IMPLICIT; + Sint64Required: Scalar Sint64 REQUIRED; + Sint64Optional: Scalar Sint64 OPTIONAL; + Sint64Repeated: Scalar Sint64 REPEATED; + Sint64Packed: Scalar Sint64 PACKED; + BoolImplicit: Scalar Bool IMPLICIT; + BoolRequired: Scalar Bool REQUIRED; + BoolOptional: Scalar Bool OPTIONAL; + BoolRepeated: Scalar Bool REPEATED; + BoolPacked: Scalar Bool PACKED; + Fixed32Implicit: Scalar Fixed32 IMPLICIT; + Fixed32Required: Scalar Fixed32 REQUIRED; + Fixed32Optional: Scalar Fixed32 OPTIONAL; + Fixed32Repeated: Scalar Fixed32 REPEATED; + Fixed32Packed: Scalar Fixed32 PACKED; + Fixed64Implicit: Scalar Fixed64 IMPLICIT; + Fixed64Required: Scalar Fixed64 REQUIRED; + Fixed64Optional: Scalar Fixed64 OPTIONAL; + Fixed64Repeated: Scalar Fixed64 REPEATED; + Fixed64Packed: Scalar Fixed64 PACKED; + Sfixed32Implicit: Scalar Sfixed32 IMPLICIT; + Sfixed32Required: Scalar Sfixed32 REQUIRED; + Sfixed32Optional: Scalar Sfixed32 OPTIONAL; + Sfixed32Repeated: Scalar Sfixed32 REPEATED; + Sfixed32Packed: Scalar Sfixed32 PACKED; + Sfixed64Implicit: Scalar Sfixed64 IMPLICIT; + Sfixed64Required: Scalar Sfixed64 REQUIRED; + Sfixed64Optional: Scalar Sfixed64 OPTIONAL; + Sfixed64Repeated: Scalar Sfixed64 REPEATED; + Sfixed64Packed: Scalar Sfixed64 PACKED; + FloatImplicit: Scalar Float IMPLICIT; + FloatRequired: Scalar Float REQUIRED; + FloatOptional: Scalar Float OPTIONAL; + FloatRepeated: Scalar Float REPEATED; + FloatPacked: Scalar Float PACKED; + DoubleImplicit: Scalar Double IMPLICIT; + DoubleRequired: Scalar Double REQUIRED; + DoubleOptional: Scalar Double OPTIONAL; + DoubleRepeated: Scalar Double REPEATED; + DoublePacked: Scalar Double PACKED; + StrImplicit: Str Str IMPLICIT; + StrRequired: Str Str REQUIRED; + StrOptional: Str Str OPTIONAL; + StrRepeated: Str Str REPEATED; + BytesImplicit: Bytes Bytes IMPLICIT; + BytesRequired: Bytes Bytes REQUIRED; + BytesOptional: Bytes Bytes OPTIONAL; + BytesRepeated: Bytes Bytes REPEATED; + EnumImplicit: Enum Enum IMPLICIT; + EnumRequired: Enum Enum REQUIRED; + EnumOptional: Enum Enum OPTIONAL; + EnumRepeated: Enum Enum REPEATED; + EnumPacked: Enum Enum PACKED; + MsgSingular: Msg Msg IMPLICIT; + MsgRepeated: Msg Msg REPEATED; + } + }; +} + +macro_rules! define_kind { + ($($name:ident: $fam:ident $ty:ident $card:ident;)*) => { + /// The type and cardinality of a field: one interpreter arm each. + #[derive(Clone, Copy, PartialEq, Eq, Debug)] + #[repr(u8)] + pub enum Kind { + $($name,)* + } + + impl Kind { + /// The wire type of the field's tag, as a number. + const fn wire_type(self) -> u32 { + match self { + $(Kind::$name => define_kind!(@wire $fam $ty $card),)* + } + } + + /// Whether the entry carries an index into the table's aux array. + const fn aux_kind(self) -> Option { + match self { + $(Kind::$name => define_kind!(@aux $fam $card),)* + } + } + + /// The cardinality a field of this kind has in an enum shape: + /// `IMPLICIT` (also for a required field), `OPTIONAL` or `REPEATED` + /// (also for a packed field). + const fn shape_card(self) -> u8 { + match self { + $(Kind::$name => define_kind!(@shape $card),)* + } + } + } + + /// One zero-sized type per [`Kind`], named the same, for the type + /// checks that `__table_entry!` makes. + pub mod kinds { + #[allow(clippy::wildcard_imports)] + use super::*; + + $( + #[doc = concat!("The type-level name of [`Kind::", stringify!($name), "`](super::Kind::", stringify!($name), ").")] + pub struct $name; + define_kind!(@slot $name $fam $ty $card); + )* + } + }; + (@slot $name:ident Scalar $ty:ident IMPLICIT) => { impl KindSlot for $name { type Slot = <$ty as Sc>::V; } }; + (@slot $name:ident Scalar $ty:ident REQUIRED) => { impl KindSlot for $name { type Slot = <$ty as Sc>::V; } }; + (@slot $name:ident Scalar $ty:ident OPTIONAL) => { impl KindSlot for $name { type Slot = Option<<$ty as Sc>::V>; } }; + (@slot $name:ident Scalar $ty:ident REPEATED) => { impl KindSlot for $name { type Slot = Vec<<$ty as Sc>::V>; } }; + (@slot $name:ident Scalar $ty:ident PACKED) => { impl KindSlot for $name { type Slot = Vec<<$ty as Sc>::V>; } }; + (@slot $name:ident Str $ty:ident IMPLICIT) => { impl KindSlot for $name { type Slot = String; } }; + (@slot $name:ident Str $ty:ident REQUIRED) => { impl KindSlot for $name { type Slot = String; } }; + (@slot $name:ident Str $ty:ident OPTIONAL) => { impl KindSlot for $name { type Slot = Option; } }; + (@slot $name:ident Str $ty:ident REPEATED) => { impl KindSlot for $name { type Slot = Vec; } }; + (@slot $name:ident Bytes $ty:ident IMPLICIT) => { impl KindSlot for $name { type Slot = Vec; } }; + (@slot $name:ident Bytes $ty:ident REQUIRED) => { impl KindSlot for $name { type Slot = Vec; } }; + (@slot $name:ident Bytes $ty:ident OPTIONAL) => { impl KindSlot for $name { type Slot = Option>; } }; + (@slot $name:ident Bytes $ty:ident REPEATED) => { impl KindSlot for $name { type Slot = Vec>; } }; + // Enum and message fields are checked against the type in their aux + // descriptor, which the generated entry names. + (@slot $name:ident Enum $ty:ident $card:ident) => {}; + (@slot $name:ident Msg $ty:ident $card:ident) => {}; + (@wire Scalar $ty:ident PACKED) => { WireType::LengthDelimited as u32 }; + (@wire Scalar $ty:ident $card:ident) => { <$ty as Sc>::WIRE as u32 }; + (@wire Str $ty:ident $card:ident) => { WireType::LengthDelimited as u32 }; + (@wire Bytes $ty:ident $card:ident) => { WireType::LengthDelimited as u32 }; + (@wire Msg $ty:ident $card:ident) => { WireType::LengthDelimited as u32 }; + (@wire Enum $ty:ident PACKED) => { WireType::LengthDelimited as u32 }; + (@wire Enum $ty:ident $card:ident) => { WireType::Varint as u32 }; + (@shape IMPLICIT) => { IMPLICIT }; + (@shape REQUIRED) => { IMPLICIT }; + (@shape OPTIONAL) => { OPTIONAL }; + (@shape REPEATED) => { REPEATED }; + (@shape PACKED) => { REPEATED }; + (@aux Scalar $card:ident) => { None }; + (@aux Str $card:ident) => { None }; + (@aux Bytes $card:ident) => { None }; + (@aux Enum $card:ident) => { Some(AuxKind::Enum) }; + (@aux Msg REPEATED) => { Some(AuxKind::Rep) }; + (@aux Msg $card:ident) => { Some(AuxKind::Msg) }; +} + +/// The type of the field that an entry of a [`kinds`] type describes, for the +/// scalar, string and bytes kinds. Enum and message kinds have none, because +/// their field type depends on the enum or message. +pub trait KindSlot { + /// The type of the field. + type Slot; +} + +kind_table!(define_kind); + +// After the macros above, whose textual scope covers only what follows them. +mod decode; +mod encode; +mod scalar; +mod shape; +mod size; + +/// Which [`Aux`] variant a kind needs. +#[derive(Clone, Copy, PartialEq, Eq)] +enum AuxKind { + Msg, + Rep, + Enum, +} + +/// Per-field data that a kind needs beyond the field's offset. +pub enum Aux { + /// The descriptor of a singular message field ([`Kind::MsgSingular`]). + Msg(&'static MsgVt), + /// The descriptor of a repeated message field ([`Kind::MsgRepeated`]). + Rep(&'static RepVt), + /// The descriptor of an enum field (the `Enum*` kinds). + Enum(&'static EnumVt), +} + +impl Aux { + const fn kind(&self) -> AuxKind { + match self { + Aux::Msg(_) => AuxKind::Msg, + Aux::Rep(_) => AuxKind::Rep, + Aux::Enum(_) => AuxKind::Enum, + } + } +} + +const _: () = assert!(core::mem::size_of::() == 12); + +/// One field of a message. +#[derive(Clone, Copy, Debug)] +#[repr(C)] +pub struct Entry { + /// The full wire tag, `(number << 3) | wire_type`, as written on encode + /// (the length-delimited tag for a packed field). + tag: u32, + /// The byte offset of the field within the message struct. + offset: u32, + kind: Kind, + /// The encoded length of `tag`. + tag_len: u8, + /// The index into [`Table`]'s aux array, for kinds that use one. + aux: u16, +} + +impl Entry { + /// An entry for the field `number` of kind `kind`, stored `offset` bytes + /// into the message struct, with the aux data at index `aux` if the kind + /// uses any. + /// + /// # Panics + /// + /// Panics, at compile time when used to initialise a `static`, if + /// `number` is not a valid field number or `offset` does not fit in a + /// `u32`. + #[must_use] + pub const fn new(kind: Kind, number: u32, offset: usize, aux: u16) -> Self { + assert!( + number >= 1 && number < (1 << 29), + "field number out of range" + ); + assert!(offset <= u32::MAX as usize, "field offset out of range"); + let tag = (number << 3) | kind.wire_type(); + let mut tag_len = 1u8; + let mut rest = tag >> 7; + while rest != 0 { + tag_len += 1; + rest >>= 7; + } + Self { + tag, + offset: offset as u32, + kind, + tag_len, + aux, + } + } + + const fn number(&self) -> u32 { + self.tag >> 3 + } +} + +/// The untyped part of a [`Table`], which the interpreters recurse through. +pub struct MessageTable { + /// The fields, sorted by number. + entries: &'static [Entry], + /// `dense[n]` is one plus the index into `entries` of field number `n`, + /// or `0` if there is none, for the numbers below `dense.len()`. + dense: &'static [u8], + aux: &'static [Aux], + /// The offset of the message's `UnknownFields`, or [`NO_UNKNOWN`]. + unknown: u32, +} + +impl MessageTable { + #[inline] + fn find(&self, number: u32) -> Option<&Entry> { + if let Some(&i) = self.dense.get(number as usize) { + return match i { + 0 => None, + i => self.entries.get(usize::from(i) - 1), + }; + } + self.entries + .binary_search_by_key(&number, Entry::number) + .ok() + .map(|i| &self.entries[i]) + } + + #[inline] + fn msg_vt(&self, e: &Entry) -> &'static MsgVt { + match &self.aux[usize::from(e.aux)] { + Aux::Msg(vt) => vt, + _ => unreachable!("`Table::new` checked that message entries index `MsgVt`s"), + } + } + + #[inline] + fn rep_vt(&self, e: &Entry) -> &'static RepVt { + match &self.aux[usize::from(e.aux)] { + Aux::Rep(vt) => vt, + _ => unreachable!("`Table::new` checked that repeated message entries index `RepVt`s"), + } + } + + #[inline] + fn enum_vt(&self, e: &Entry) -> &'static EnumVt { + match &self.aux[usize::from(e.aux)] { + Aux::Enum(vt) => vt, + _ => unreachable!("`Table::new` checked that enum entries index enum descriptors"), + } + } +} + +/// The static description of message type `M`, from which the interpreters +/// in this module encode, size and decode an `M`. +pub struct Table { + raw: MessageTable, + _marker: PhantomData, +} + +impl Table { + /// Describe `M`. Generated code builds tables with `__table!`. + /// + /// `abi` is the [`ABI`] the caller was generated for. `entries` are the + /// fields sorted by number; `dense` is the lookup array described on + /// [`MessageTable`], which must be empty if there are 255 entries or more; + /// `aux` holds the descriptors that entries index by their `aux` value; + /// `unknown` is the offset of the message's `UnknownFields`, if it keeps + /// any. + /// + /// # Panics + /// + /// Panics, at compile time when used to initialise a `static`, if `abi` + /// is not [`ABI`], if an entry lies outside `M`, if the entries are not in + /// strictly increasing field-number order, if an entry's aux index is out + /// of range or names a descriptor of the wrong variant or, for an enum, + /// the wrong cardinality, if `dense` disagrees with `entries`, or if + /// `unknown` does not leave room for an `UnknownFields`. + /// + /// # Safety + /// + /// For every entry, `offset` must be the offset within `M` of a field + /// whose Rust type is exactly the one the entry's [`Kind`] names, in the + /// default representation: + /// + /// - `Int32*` and the other scalar kinds: the scalar type itself + /// (`i32`, `u64`, `bool`, `f32`, ...) for `Implicit` and `Required`, + /// `Option` for `Optional`, and `Vec` for `Repeated` and `Packed`; + /// - `Str*`: `String`, `Option` or `Vec` likewise, and + /// `Bytes*`: `Vec`, `Option>` or `Vec>`; + /// - `Enum*`: the storage the entry's [`EnumVt`] was built for; + /// - `MsgSingular`: the storage the [`MsgVt`] was built for, and + /// `MsgRepeated`: a `Vec` of the messages the [`RepVt`]'s table + /// describes. + /// + /// `unknown`, if present, must be the offset of a field of type + /// `UnknownFields`. The `__table_entry!` macro checks the field types + /// that the compiler can; the pairing of an entry with its aux index and + /// the `dense` lookup are the generator's to get right. + #[must_use] + pub const unsafe fn new( + abi: u32, + entries: &'static [Entry], + dense: &'static [u8], + aux: &'static [Aux], + unknown: Option, + ) -> Self { + assert!( + abi == ABI, + "buffa table: the generated code is for a different table ABI; \ + regenerate it with the buffa-codegen that matches this buffa" + ); + let size = core::mem::size_of::(); + let mut i = 0; + while i < entries.len() { + let e = &entries[i]; + assert!( + i == 0 || entries[i - 1].number() < e.number(), + "buffa table: entries are not in strictly increasing field-number order" + ); + assert!( + (e.offset as usize) < size, + "buffa table: an entry's offset is outside the message struct" + ); + if let Some(want) = e.kind.aux_kind() { + assert!( + (e.aux as usize) < aux.len(), + "buffa table: an entry's aux index is out of range" + ); + let a = &aux[e.aux as usize]; + assert!( + a.kind() as u8 == want as u8, + "buffa table: an entry's aux descriptor is the wrong variant for its kind" + ); + if let Aux::Enum(vt) = a { + assert!( + vt.card == e.kind.shape_card(), + "buffa table: an enum entry's descriptor has the wrong cardinality for its kind" + ); + } + } + i += 1; + } + assert!( + dense.is_empty() || entries.len() < 255, + "buffa table: the dense lookup needs fewer than 255 entries" + ); + // Every dense slot names the entry with its number, and every entry + // below the dense range has a slot, so the two agree. + let mut named = 0; + let mut n = 0; + while n < dense.len() { + let d = dense[n] as usize; + if d != 0 { + assert!( + d <= entries.len() && entries[d - 1].number() as usize == n, + "buffa table: the dense lookup names an entry with the wrong field number" + ); + named += 1; + } + n += 1; + } + let mut covered = 0; + i = 0; + while i < entries.len() { + if (entries[i].number() as usize) < dense.len() { + covered += 1; + } + i += 1; + } + assert!( + named == covered, + "buffa table: the dense lookup omits an entry" + ); + let unknown = match unknown { + None => NO_UNKNOWN, + Some(offset) => { + assert!( + offset < NO_UNKNOWN as usize + && offset <= size + && size - offset >= core::mem::size_of::() + && offset % core::mem::align_of::() == 0, + "buffa table: the unknown-fields offset does not fit an `UnknownFields` in the message struct" + ); + offset as u32 + } + }; + Self { + raw: MessageTable { + entries, + dense, + aux, + unknown, + }, + _marker: PhantomData, + } + } + + /// The encoded size of `msg`, recording nested message sizes in `cache`. + /// + /// See [`Message::compute_size`](crate::Message::compute_size). + #[inline] + pub fn compute_size(&self, msg: &M, cache: &mut SizeCache) -> u32 { + // SAFETY: `Table::new`'s contract says the table describes `M`. + unsafe { size::compute_size(&self.raw, (msg as *const M).cast(), cache) } + } + + /// Write `msg` to `buf`, taking nested message sizes from `cache`, which + /// `compute_size` filled. + /// + /// A [`PreSized`](crate::encode_sink::PreSized) cursor, which is what + /// `Message::encode` and its siblings write a `BufMut` through, is + /// written by one function compiled in this crate, and any other sink by an + /// instance compiled in the caller's crate. + /// + /// See [`Message::write_to`](crate::Message::write_to). + #[inline] + pub fn write_to(&self, msg: &M, cache: &mut SizeCache, buf: &mut K) { + let base = (msg as *const M).cast::(); + // SAFETY: `Table::new`'s contract says the table describes `M`. + unsafe { encode::write_to(&self.raw, base, cache, buf) } + } + + /// Decode into `msg` until `buf` has `limit` bytes remaining. + /// + /// A `buf` whose current chunk is shorter than the message is gathered + /// into one contiguous buffer first. + /// + /// See [`Message::merge_to_limit`](crate::Message::merge_to_limit). + /// + /// # Errors + /// + /// Returns a [`DecodeError`] for malformed input or an exhausted limit. + #[inline] + pub fn merge_to_limit( + &self, + msg: &mut M, + buf: &mut B, + ctx: DecodeContext<'_>, + limit: usize, + ) -> Result<(), DecodeError> { + let base = (msg as *mut M).cast::(); + // SAFETY: `Table::new`'s contract says the table describes `M`. + unsafe { decode::merge_to_limit(&self.raw, base, buf, ctx, limit) } + } + + /// Decode a length-prefixed message into `msg`. + /// + /// See [`Message::merge_length_delimited`](crate::Message::merge_length_delimited). + /// + /// # Errors + /// + /// Returns a [`DecodeError`] for malformed input or an exhausted limit. + #[inline] + pub fn merge_length_delimited( + &self, + msg: &mut M, + buf: &mut B, + ctx: DecodeContext<'_>, + ) -> Result<(), DecodeError> { + let base = (msg as *mut M).cast::(); + // SAFETY: `Table::new`'s contract says the table describes `M`. + unsafe { decode::merge_length_delimited(&self.raw, base, buf, ctx) } + } + + /// Decode one field, whose `tag` has been read, into `msg`. + /// + /// Unlike the other decode entry points this does not gather a + /// non-contiguous `buf`: it requires everything remaining in `buf` to be + /// one chunk, and returns [`DecodeError::UnexpectedEof`] if it is not. + /// + /// See [`Message::merge_field`](crate::Message::merge_field). + /// + /// # Errors + /// + /// Returns a [`DecodeError`] for malformed input or an exhausted limit. + #[inline] + pub fn merge_field( + &self, + msg: &mut M, + tag: Tag, + buf: &mut B, + ctx: DecodeContext<'_>, + ) -> Result<(), DecodeError> { + let base = (msg as *mut M).cast::(); + // SAFETY: `Table::new`'s contract says the table describes `M`. + unsafe { decode::merge_field(&self.raw, base, tag, buf, ctx) } + } +} + +/// Build the `static` [`Table`] of message `$msg`, for generated code. +/// +/// `unknown` is `none` for a message that drops unknown fields, or the name of +/// its `UnknownFields` field. The `unsafe` is written here, in this crate, so +/// the generated code compiles under `#![forbid(unsafe_code)]`. +/// +/// # Example +/// +/// A message with one `int32` field, as generated code builds it: +/// +/// ``` +/// #![forbid(unsafe_code)] +/// use buffa::encoding::Tag; +/// use buffa::{DecodeContext, DecodeError, EncodeSink, Message, SizeCache, UnknownFields}; +/// +/// #[derive(Clone, Debug, Default, PartialEq)] +/// struct Point { +/// x: i32, +/// unknown: UnknownFields, +/// } +/// buffa::impl_default_instance!(Point); +/// +/// static POINT: buffa::table::Table = buffa::__table!( +/// Point, +/// abi = buffa::table::ABI, +/// entries = [buffa::__table_entry!(Point, x, Int32Implicit, 1)], +/// dense = &[0, 1], +/// aux = [], +/// unknown = unknown, +/// ); +/// +/// impl Message for Point { +/// fn compute_size(&self, cache: &mut SizeCache) -> u32 { +/// POINT.compute_size(self, cache) +/// } +/// fn write_to(&self, cache: &mut SizeCache, buf: &mut impl EncodeSink) { +/// POINT.write_to(self, cache, buf); +/// } +/// fn merge_field( +/// &mut self, +/// tag: Tag, +/// buf: &mut impl buffa::bytes::Buf, +/// ctx: DecodeContext<'_>, +/// ) -> Result<(), DecodeError> { +/// POINT.merge_field(self, tag, buf, ctx) +/// } +/// fn clear(&mut self) { +/// *self = Self::default(); +/// } +/// } +/// +/// let point = Point { x: 150, ..Point::default() }; +/// assert_eq!(point.encode_to_vec(), [0x08, 0x96, 0x01]); +/// assert_eq!(Point::decode_from_slice(&[0x08, 0x96, 0x01]).unwrap(), point); +/// ``` +#[doc(hidden)] +#[macro_export] +macro_rules! __table { + ( + $msg:ty, + abi = $abi:expr, + entries = [$($entry:expr),* $(,)?], + dense = $dense:expr, + aux = [$($aux:expr),* $(,)?], + unknown = none $(,)? + ) => { + // SAFETY: each entry's field type is checked by `__table_entry!`; the + // generator pairs entries with their aux descriptors and builds the + // dense lookup, and `Table::new` checks the rest at compile time. + unsafe { + $crate::table::Table::<$msg>::new( + $abi, + &[$($entry),*], + $dense, + &[$($aux),*], + ::core::option::Option::None, + ) + } + }; + ( + $msg:ty, + abi = $abi:expr, + entries = [$($entry:expr),* $(,)?], + dense = $dense:expr, + aux = [$($aux:expr),* $(,)?], + unknown = $unknown:ident $(,)? + ) => { + { + const _: fn(&$msg) -> *const $crate::UnknownFields = + |m| ::core::ptr::addr_of!(m.$unknown); + // SAFETY: as above, and the field just checked is an + // `UnknownFields`. + unsafe { + $crate::table::Table::<$msg>::new( + $abi, + &[$($entry),*], + $dense, + &[$($aux),*], + ::core::option::Option::Some($crate::table::offset_of!($msg, $unknown)), + ) + } + } + }; +} + +/// Build one [`Entry`] of message `$msg`'s table, for generated code. +/// +/// The field `$field` of `$msg` must have the type `$kind` names, or the +/// generated code does not compile: +/// +/// ```compile_fail,E0308 +/// struct Point { +/// x: i32, +/// } +/// // `x` is not a `String`. +/// let _ = buffa::__table_entry!(Point, x, StrImplicit, 1); +/// ``` +/// +/// The check is exact, so a field whose type only derefs to the expected one +/// does not pass: +/// +/// ```compile_fail,E0308 +/// struct Point { +/// x: Box, +/// } +/// let _ = buffa::__table_entry!(Point, x, StrImplicit, 1); +/// ``` +/// +/// The scalar, string and bytes kinds have a fixed field type. The enum and +/// message kinds take the type explicitly, as `aux = , slot = `, +/// where the type is the one their aux descriptor was built for. +#[doc(hidden)] +#[macro_export] +macro_rules! __table_entry { + ($msg:ty, $field:ident, $kind:ident, $number:expr $(,)?) => {{ + // A raw pointer, unlike a reference, cannot be deref-coerced, so this + // needs the field's type to be exactly the slot type. + const _: fn( + &$msg, + ) + -> *const <$crate::table::kinds::$kind as $crate::table::KindSlot>::Slot = + |m| ::core::ptr::addr_of!(m.$field); + $crate::table::Entry::new( + $crate::table::Kind::$kind, + $number, + $crate::table::offset_of!($msg, $field), + 0, + ) + }}; + ($msg:ty, $field:ident, $kind:ident, $number:expr, aux = $aux:expr, slot = $slot:ty $(,)?) => {{ + const _: fn(&$msg) -> *const $slot = |m| ::core::ptr::addr_of!(m.$field); + $crate::table::Entry::new( + $crate::table::Kind::$kind, + $number, + $crate::table::offset_of!($msg, $field), + $aux, + ) + }}; +} + +// The tests build tables with `offset_of!`, which needs Rust 1.77. The module +// is inline because `rustversion` cannot gate an out-of-line one. +#[cfg(test)] +#[rustversion::since(1.77)] +mod tests { + include!("tests.rs"); +} diff --git a/buffa/src/table/scalar.rs b/buffa/src/table/scalar.rs new file mode 100644 index 00000000..e66e0a9b --- /dev/null +++ b/buffa/src/table/scalar.rs @@ -0,0 +1,205 @@ +//! The thirteen numeric and bool protobuf scalar types as zero-sized markers, +//! so the interpreters are written once per cardinality and instantiated per +//! type. + +use crate::alloc::vec::Vec; +use crate::encoding::WireType; +use crate::{types, DecodeError, EncodeSink}; + +/// One protobuf scalar type: its wire type and the codec functions for a +/// single value. +pub trait Sc { + /// The Rust type of a field of this scalar type. + type V: Copy; + const WIRE: WireType; + fn read(buf: &mut &[u8]) -> Result; + fn encode(v: Self::V, buf: &mut K); + /// Encoded size of `v`, without the tag. + fn len(v: Self::V) -> u64; + /// Whether `v` is the proto3 implicit-presence default (not written). + fn is_default(v: Self::V) -> bool; + /// Append every element of a packed payload to `out`. + fn extend(payload: &[u8], out: &mut Vec) -> Result<(), DecodeError>; +} + +macro_rules! scalar { + ($name:ident, $v:ty, $wire:expr, $read:path, $enc:path, $len:expr, $default:expr, $extend:expr) => { + pub struct $name; + impl Sc for $name { + type V = $v; + const WIRE: WireType = $wire; + #[inline] + fn read(buf: &mut &[u8]) -> Result<$v, DecodeError> { + $read(buf) + } + #[inline] + fn encode(v: $v, buf: &mut K) { + $enc(v, buf); + } + #[inline] + fn len(v: $v) -> u64 { + ($len)(v) + } + #[inline] + fn is_default(v: $v) -> bool { + ($default)(v) + } + #[inline] + fn extend(payload: &[u8], out: &mut Vec<$v>) -> Result<(), DecodeError> { + ($extend)(payload, out) + } + } + }; +} + +macro_rules! varint_scalar { + ($name:ident, $v:ty, $read:path, $enc:path, $len:path, $extend:path) => { + scalar!( + $name, + $v, + WireType::Varint, + $read, + $enc, + |v| $len(v) as u64, + |v| v == <$v>::default(), + |p: &[u8], out: &mut Vec<$v>| $extend(p, out, p.len()) + ); + }; +} + +/// `$bits` maps a value to its bit pattern, so that a float is the default +/// only when it is `+0.0`: `-0.0` is not the default and is written, as in +/// unrolled code. +macro_rules! fixed_scalar { + ($name:ident, $v:ty, $wire:expr, $width:expr, $read:path, $enc:path, $extend:path, $bits:expr) => { + scalar!( + $name, + $v, + $wire, + $read, + $enc, + |_| $width as u64, + |v| ($bits)(v) == 0, + |p: &[u8], out: &mut Vec<$v>| $extend(p, out) + ); + }; +} + +varint_scalar!( + Int32, + i32, + types::decode_int32, + types::encode_int32, + types::int32_encoded_len, + types::extend_packed_int32 +); +varint_scalar!( + Int64, + i64, + types::decode_int64, + types::encode_int64, + types::int64_encoded_len, + types::extend_packed_int64 +); +varint_scalar!( + Uint32, + u32, + types::decode_uint32, + types::encode_uint32, + types::uint32_encoded_len, + types::extend_packed_uint32 +); +varint_scalar!( + Uint64, + u64, + types::decode_uint64, + types::encode_uint64, + types::uint64_encoded_len, + types::extend_packed_uint64 +); +varint_scalar!( + Sint32, + i32, + types::decode_sint32, + types::encode_sint32, + types::sint32_encoded_len, + types::extend_packed_sint32 +); +varint_scalar!( + Sint64, + i64, + types::decode_sint64, + types::encode_sint64, + types::sint64_encoded_len, + types::extend_packed_sint64 +); +scalar!( + Bool, + bool, + WireType::Varint, + types::decode_bool, + types::encode_bool, + |_| types::BOOL_ENCODED_LEN as u64, + |v: bool| !v, + |p: &[u8], out: &mut Vec| types::extend_packed_bool(p, out, p.len()) +); +fixed_scalar!( + Fixed32, + u32, + WireType::Fixed32, + 4, + types::decode_fixed32, + types::encode_fixed32, + types::extend_packed_fixed32, + |v: u32| v +); +fixed_scalar!( + Sfixed32, + i32, + WireType::Fixed32, + 4, + types::decode_sfixed32, + types::encode_sfixed32, + types::extend_packed_sfixed32, + |v: i32| v +); +fixed_scalar!( + Float, + f32, + WireType::Fixed32, + 4, + types::decode_float, + types::encode_float, + types::extend_packed_float, + |v: f32| v.to_bits() +); +fixed_scalar!( + Fixed64, + u64, + WireType::Fixed64, + 8, + types::decode_fixed64, + types::encode_fixed64, + types::extend_packed_fixed64, + |v: u64| v +); +fixed_scalar!( + Sfixed64, + i64, + WireType::Fixed64, + 8, + types::decode_sfixed64, + types::encode_sfixed64, + types::extend_packed_sfixed64, + |v: i64| v +); +fixed_scalar!( + Double, + f64, + WireType::Fixed64, + 8, + types::decode_double, + types::encode_double, + types::extend_packed_double, + |v: f64| v.to_bits() +); diff --git a/buffa/src/table/shape.rs b/buffa/src/table/shape.rs new file mode 100644 index 00000000..214d47ba --- /dev/null +++ b/buffa/src/table/shape.rs @@ -0,0 +1,321 @@ +//! Accessors for the field shapes a table entry cannot reach by offset alone: +//! message-typed fields (whose storage is a pointer or an `Option`) and +//! enum-typed fields (whose storage is an `i32` newtype or a closed enum). + +use core::marker::PhantomData; + +use super::{Table, IMPLICIT, OPTIONAL, REPEATED}; +use crate::alloc::vec::Vec; +use crate::{EnumValue, Enumeration, MessageField, ProtoBox}; + +// --------------------------------------------------------------------------- +// Singular message fields +// --------------------------------------------------------------------------- + +/// Storage of a singular message field. +pub trait MsgSlot { + /// The message type the field holds. + type Msg; + /// The message, creating the default if the field is unset. + fn place(&mut self) -> &mut Self::Msg; + /// The message, or `None` if the field is unset. + fn get(&self) -> Option<&Self::Msg>; +} + +impl> MsgSlot for MessageField { + type Msg = T; + + #[inline] + fn place(&mut self) -> &mut T { + self.get_or_insert_default() + } + + #[inline] + fn get(&self) -> Option<&T> { + self.as_option() + } +} + +/// Descriptor of a singular message field: the child's table and how to reach +/// the child through the field's storage. +pub struct MsgVt { + pub(super) table: &'static super::MessageTable, + /// The message in the field, created with its default if unset. + /// + /// # Safety + /// + /// The argument points to a live `F`. + pub(super) place: unsafe fn(*mut u8) -> *mut u8, + /// The message in the field, or null if unset. + /// + /// # Safety + /// + /// The argument points to a live `F`. + pub(super) get: unsafe fn(*const u8) -> *const u8, +} + +/// # Safety +/// +/// `slot` points to a live `F`. +unsafe fn place_impl(slot: *mut u8) -> *mut u8 { + // SAFETY: the caller passes a pointer to a live `F`. + unsafe { (*slot.cast::()).place() as *mut F::Msg as *mut u8 } +} + +/// # Safety +/// +/// `slot` points to a live `F`. +unsafe fn get_impl(slot: *const u8) -> *const u8 { + // SAFETY: the caller passes a pointer to a live `F`. + match unsafe { (*slot.cast::()).get() } { + Some(m) => (m as *const F::Msg).cast::(), + None => core::ptr::null(), + } +} + +impl MsgVt { + /// Describe a field of type `F`, whose messages `table` describes. + #[must_use] + pub const fn new(table: &'static Table) -> Self { + Self { + table: &table.raw, + place: place_impl::, + get: get_impl::, + } + } +} + +// --------------------------------------------------------------------------- +// Repeated message fields +// --------------------------------------------------------------------------- + +/// Descriptor of a repeated message field, a `Vec`. +pub struct RepVt { + pub(super) table: &'static super::MessageTable, + /// The size in bytes of one element. + pub(super) size: usize, + /// Append a default element and return a pointer to it. + /// + /// # Safety + /// + /// The argument points to a live `Vec`. + pub(super) push: unsafe fn(*mut u8) -> *mut u8, + /// Remove the last element, which `push` added. + /// + /// # Safety + /// + /// The argument points to a live `Vec`. + pub(super) pop: unsafe fn(*mut u8), + /// The element storage: a pointer to the first element and the count. + /// + /// # Safety + /// + /// The argument points to a live `Vec`. + pub(super) parts: unsafe fn(*const u8) -> (*const u8, usize), +} + +/// # Safety +/// +/// `slot` points to a live `Vec`. +unsafe fn push_impl(slot: *mut u8) -> *mut u8 { + // SAFETY: the caller passes a pointer to a live `Vec`. + let v = unsafe { &mut *slot.cast::>() }; + v.push(T::default()); + let last = v.len() - 1; + // SAFETY: `last` is in bounds. + unsafe { v.as_mut_ptr().add(last).cast::() } +} + +/// # Safety +/// +/// `slot` points to a live `Vec`. +unsafe fn pop_impl(slot: *mut u8) { + // SAFETY: the caller passes a pointer to a live `Vec`. + unsafe { (*slot.cast::>()).pop() }; +} + +/// # Safety +/// +/// `slot` points to a live `Vec`. +unsafe fn parts_impl(slot: *const u8) -> (*const u8, usize) { + // SAFETY: the caller passes a pointer to a live `Vec`. + let v = unsafe { &*slot.cast::>() }; + (v.as_ptr().cast::(), v.len()) +} + +impl RepVt { + /// Describe a `Vec` field, whose messages `table` describes. + #[must_use] + pub const fn new(table: &'static Table) -> Self { + Self { + table: &table.raw, + size: core::mem::size_of::(), + push: push_impl::, + pop: pop_impl::, + parts: parts_impl::, + } + } +} + +// --------------------------------------------------------------------------- +// Enum fields +// --------------------------------------------------------------------------- + +/// Descriptor of an enum field: how to store and read its `i32` values. +pub struct EnumVt { + /// The cardinality the shape stores: `IMPLICIT`, `OPTIONAL` or + /// `REPEATED`. + pub(super) card: u8, + /// Store `raw` (append, for a repeated field). `false` if a closed enum + /// has no variant with that number, in which case nothing is stored. + /// + /// # Safety + /// + /// The argument points to a live slot of the shape the descriptor was + /// built for. + pub(super) set: unsafe fn(*mut u8, i32) -> bool, + /// The value at `idx` (ignored for singular shapes), or `None` if unset. + /// + /// # Safety + /// + /// As for `set`. + pub(super) get: unsafe fn(*const u8, usize) -> Option, + /// The element count of a repeated shape, `0` otherwise. + /// + /// # Safety + /// + /// As for `set`. + pub(super) len: unsafe fn(*const u8) -> usize, +} + +/// A way of storing an enum field, implemented by the marker types below. +/// +/// # Safety +/// +/// [`Slot`](Self::Slot) is the type of the field, and its functions read and +/// write a field of that type. [`CARD`](Self::CARD) is its cardinality: +/// `IMPLICIT` for a singular field with no presence, `OPTIONAL` for +/// `Option<_>`, and `REPEATED` for `Vec<_>`; the table checks it against the +/// entry's kind. +pub unsafe trait EnumShape { + /// The type of the field. + type Slot; + + /// The cardinality of the field. + const CARD: u8; + + /// Store `raw`, appending it if the shape is repeated. Returns `false`, + /// and stores nothing, if a closed enum has no variant numbered `raw`. + /// + /// # Safety + /// + /// `slot` points to a live [`Slot`](Self::Slot). + unsafe fn set(slot: *mut u8, raw: i32) -> bool; + + /// The value at `idx` (ignored for singular shapes), or `None` if unset. + /// + /// # Safety + /// + /// `slot` points to a live [`Slot`](Self::Slot). + unsafe fn get(slot: *const u8, idx: usize) -> Option; + + /// The element count of a repeated shape, `0` otherwise. + /// + /// # Safety + /// + /// `slot` points to a live [`Slot`](Self::Slot). + unsafe fn len(_slot: *const u8) -> usize { + 0 + } +} + +impl EnumVt { + /// Describe a field stored as `S`. + #[must_use] + pub const fn new() -> Self { + Self { + card: S::CARD, + set: S::set, + get: S::get, + len: S::len, + } + } +} + +macro_rules! enum_shape { + ($(#[$m:meta])* $name:ident, $card:ident, $slot:ty, $set:expr, $get:expr, $len:expr) => { + $(#[$m])* + pub struct $name(PhantomData); + + // SAFETY: each function reads or writes the slot as `$slot`, which + // the trait's contract says it is. + unsafe impl EnumShape for $name { + type Slot = $slot; + const CARD: u8 = $card; + + #[inline] + unsafe fn set(slot: *mut u8, raw: i32) -> bool { + // SAFETY: the caller passes a pointer to a live slot of this shape. + let s = unsafe { &mut *slot.cast::<$slot>() }; + ($set)(s, raw) + } + + #[inline] + unsafe fn get(slot: *const u8, idx: usize) -> Option { + // SAFETY: as above. + let s = unsafe { &*slot.cast::<$slot>() }; + ($get)(s, idx) + } + + #[inline] + unsafe fn len(slot: *const u8) -> usize { + // SAFETY: as above. + let s = unsafe { &*slot.cast::<$slot>() }; + ($len)(s) + } + } + }; +} + +enum_shape!( + /// An open enum with implicit presence: `EnumValue`. + ImplicitOpen, IMPLICIT, EnumValue, + |s: &mut EnumValue, raw| { *s = EnumValue::from(raw); true }, + |s: &EnumValue, _| Some(s.to_i32()), + |_: &EnumValue| 0 +); +enum_shape!( + /// A closed enum with implicit presence: `E`. + ImplicitClosed, IMPLICIT, E, + |s: &mut E, raw| match E::from_i32(raw) { Some(v) => { *s = v; true } None => false }, + |s: &E, _| Some(s.to_i32()), + |_: &E| 0 +); +enum_shape!( + /// An open enum with explicit presence: `Option>`. + OptionalOpen, OPTIONAL, Option>, + |s: &mut Option>, raw| { *s = Some(EnumValue::from(raw)); true }, + |s: &Option>, _| s.as_ref().map(EnumValue::to_i32), + |_: &Option>| 0 +); +enum_shape!( + /// A closed enum with explicit presence: `Option`. + OptionalClosed, OPTIONAL, Option, + |s: &mut Option, raw| match E::from_i32(raw) { Some(v) => { *s = Some(v); true } None => false }, + |s: &Option, _| s.as_ref().map(Enumeration::to_i32), + |_: &Option| 0 +); +enum_shape!( + /// A repeated open enum: `Vec>`. + RepeatedOpen, REPEATED, Vec>, + |s: &mut Vec>, raw| { s.push(EnumValue::from(raw)); true }, + |s: &Vec>, i| s.get(i).map(EnumValue::to_i32), + |s: &Vec>| s.len() +); +enum_shape!( + /// A repeated closed enum: `Vec`. + RepeatedClosed, REPEATED, Vec, + |s: &mut Vec, raw| match E::from_i32(raw) { Some(v) => { s.push(v); true } None => false }, + |s: &Vec, i| s.get(i).map(Enumeration::to_i32), + |s: &Vec| s.len() +); diff --git a/buffa/src/table/size.rs b/buffa/src/table/size.rs new file mode 100644 index 00000000..e2430f32 --- /dev/null +++ b/buffa/src/table/size.rs @@ -0,0 +1,253 @@ +//! The size pass: [`compute_size`] and its per-kind arms. + +use super::scalar::Sc; +use super::{ + Bool, Double, Entry, Fixed32, Fixed64, Float, Int32, Int64, Kind, MessageTable, Sfixed32, + Sfixed64, Sint32, Sint64, Uint32, Uint64, IMPLICIT, NO_UNKNOWN, OPTIONAL, PACKED, REPEATED, + REQUIRED, +}; +use crate::alloc::{string::String, vec::Vec}; +use crate::encoding::varint_len; +use crate::{types, SizeCache, UnknownFields}; + +/// The encoded size of the message at `base`, recording nested sizes in +/// `cache`. +/// +/// # Safety +/// +/// `base` points to a live message of the type `table` describes. +pub(super) unsafe fn compute_size( + table: &MessageTable, + base: *const u8, + cache: &mut SizeCache, +) -> u32 { + let mut size = 0u64; + for e in table.entries { + // SAFETY: the offset is within the message, per the table's contract. + size += unsafe { size_kind(table, e, base.add(e.offset as usize), cache) }; + } + if table.unknown != NO_UNKNOWN { + // SAFETY: `unknown` is the offset of the message's `UnknownFields`. + size += unsafe { (*base.add(table.unknown as usize).cast::()).encoded_len() } + as u64; + } + crate::saturate_size(size) +} + +macro_rules! size_dispatch { + ($($name:ident: $fam:ident $ty:ident $card:ident;)*) => { + /// # Safety + /// + /// `slot` points to the field `e` describes, in a live message of the + /// type `table` describes. + #[inline] + unsafe fn size_kind( + table: &MessageTable, + e: &Entry, + slot: *const u8, + cache: &mut SizeCache, + ) -> u64 { + let tl = u64::from(e.tag_len); + // SAFETY: each arm reads the slot as the type its kind names. + unsafe { + match e.kind { + $(Kind::$name => size_dispatch!(@arm $fam $ty $card table e tl slot cache),)* + } + } + } + }; + (@arm Scalar $ty:ident $card:ident $table:ident $e:ident $tl:ident $slot:ident $cache:ident) => { + size_scalar::<$ty, $card>($tl, $slot) + }; + (@arm Str $ty:ident $card:ident $table:ident $e:ident $tl:ident $slot:ident $cache:ident) => { + size_str::<$card>($tl, $slot) + }; + (@arm Bytes $ty:ident $card:ident $table:ident $e:ident $tl:ident $slot:ident $cache:ident) => { + size_bytes::<$card>($tl, $slot) + }; + (@arm Enum $ty:ident $card:ident $table:ident $e:ident $tl:ident $slot:ident $cache:ident) => { + size_enum::<$card>($table, $e, $tl, $slot) + }; + (@arm Msg $ty:ident $card:ident $table:ident $e:ident $tl:ident $slot:ident $cache:ident) => { + size_msg::<$card>($table, $e, $tl, $slot, $cache) + }; +} + +kind_table!(size_dispatch); + +/// # Safety +/// +/// `slot` points to a field of scalar type `S` in the shape `C` names. +#[inline] +unsafe fn size_scalar(tl: u64, slot: *const u8) -> u64 { + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT => { + let v = *slot.cast::(); + if S::is_default(v) { + 0 + } else { + tl + S::len(v) + } + } + REQUIRED => tl + S::len(*slot.cast::()), + OPTIONAL => match *slot.cast::>() { + Some(v) => tl + S::len(v), + None => 0, + }, + REPEATED => (*slot.cast::>()) + .iter() + .map(|&v| tl + S::len(v)) + .sum(), + _ => { + let v = &*slot.cast::>(); + if v.is_empty() { + 0 + } else { + let payload: u64 = v.iter().map(|&x| S::len(x)).sum(); + tl + varint_len(payload) as u64 + payload + } + } + } + } +} + +/// # Safety +/// +/// `slot` points to a `String` field in the shape `C` names. +#[inline] +unsafe fn size_str(tl: u64, slot: *const u8) -> u64 { + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT => { + let s = &*slot.cast::(); + if s.is_empty() { + 0 + } else { + tl + types::string_encoded_len(s) as u64 + } + } + REQUIRED => tl + types::string_encoded_len(&*slot.cast::()) as u64, + OPTIONAL => match &*slot.cast::>() { + Some(s) => tl + types::string_encoded_len(s) as u64, + None => 0, + }, + _ => (*slot.cast::>()) + .iter() + .map(|s| tl + types::string_encoded_len(s) as u64) + .sum(), + } + } +} + +/// # Safety +/// +/// `slot` points to a `Vec` field in the shape `C` names. +#[inline] +unsafe fn size_bytes(tl: u64, slot: *const u8) -> u64 { + // SAFETY: the caller's contract gives the slot's type. + unsafe { + match C { + IMPLICIT => { + let s = &*slot.cast::>(); + if s.is_empty() { + 0 + } else { + tl + types::bytes_encoded_len(s) as u64 + } + } + REQUIRED => tl + types::bytes_encoded_len(&*slot.cast::>()) as u64, + OPTIONAL => match &*slot.cast::>>() { + Some(s) => tl + types::bytes_encoded_len(s) as u64, + None => 0, + }, + _ => (*slot.cast::>>()) + .iter() + .map(|s| tl + types::bytes_encoded_len(s) as u64) + .sum(), + } + } +} + +/// # Safety +/// +/// `slot` points to the enum field `e` describes, in the shape `C` names. +#[inline] +unsafe fn size_enum(table: &MessageTable, e: &Entry, tl: u64, slot: *const u8) -> u64 { + let vt = table.enum_vt(e); + // SAFETY: `vt` was built for the slot's shape. + unsafe { + match C { + IMPLICIT => match (vt.get)(slot, 0) { + Some(v) if v != 0 => tl + types::int32_encoded_len(v) as u64, + _ => 0, + }, + REQUIRED | OPTIONAL => match (vt.get)(slot, 0) { + Some(v) => tl + types::int32_encoded_len(v) as u64, + None => 0, + }, + REPEATED => { + let mut size = 0; + for i in 0..(vt.len)(slot) { + if let Some(v) = (vt.get)(slot, i) { + size += tl + types::int32_encoded_len(v) as u64; + } + } + size + } + _ => { + let n = (vt.len)(slot); + if n == 0 { + return 0; + } + let mut payload = 0; + for i in 0..n { + if let Some(v) = (vt.get)(slot, i) { + payload += types::int32_encoded_len(v) as u64; + } + } + tl + varint_len(payload) as u64 + payload + } + } + } +} + +/// # Safety +/// +/// `slot` points to the message field `e` describes, in the shape `C` names. +#[inline] +unsafe fn size_msg( + table: &MessageTable, + e: &Entry, + tl: u64, + slot: *const u8, + cache: &mut SizeCache, +) -> u64 { + // SAFETY: the descriptor was built for the slot's shape and child type. + unsafe { + if C == REPEATED { + let vt = table.rep_vt(e); + let (ptr, len) = (vt.parts)(slot); + let mut size = 0; + for i in 0..len { + let idx = cache.reserve(); + let inner = compute_size(vt.table, ptr.add(i * vt.size), cache); + cache.set(idx, inner); + size += tl + varint_len(u64::from(inner)) as u64 + u64::from(inner); + } + size + } else { + let vt = table.msg_vt(e); + let child = (vt.get)(slot); + if child.is_null() { + return 0; + } + let idx = cache.reserve(); + let inner = compute_size(vt.table, child, cache); + cache.set(idx, inner); + tl + varint_len(u64::from(inner)) as u64 + u64::from(inner) + } + } +} diff --git a/buffa/src/table/tests.rs b/buffa/src/table/tests.rs new file mode 100644 index 00000000..05346705 --- /dev/null +++ b/buffa/src/table/tests.rs @@ -0,0 +1,1003 @@ +// Tests of the interpreters over hand-written table messages. + +use super::*; +use crate::alloc::{string::String, vec, vec::Vec}; +use crate::bytes::Buf; +use crate::{ + DecodeOptions, EnumValue, Enumeration, Inline, Message, MessageField, Rope, UnknownFieldData, + UnknownFields, +}; + +// --------------------------------------------------------------------------- +// Test messages +// --------------------------------------------------------------------------- + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)] +enum Color { + #[default] + Red = 0, + Green = 1, + Blue = 2, +} + +impl Enumeration for Color { + fn from_i32(value: i32) -> Option { + match value { + 0 => Some(Self::Red), + 1 => Some(Self::Green), + 2 => Some(Self::Blue), + _ => None, + } + } + + fn to_i32(&self) -> i32 { + *self as i32 + } + + fn proto_name(&self) -> &'static str { + match self { + Self::Red => "RED", + Self::Green => "GREEN", + Self::Blue => "BLUE", + } + } +} + +/// `int32 id = 1; string label = 2; Inner next = 3;`, keeping unknown fields. +#[derive(Clone, Debug, Default, PartialEq)] +struct Inner { + id: i32, + label: String, + next: MessageField, + unknown: UnknownFields, +} + +/// Fields of most shapes the interpreters handle, plus a field number above +/// the dense lookup. `Wide` has the rest. +#[derive(Clone, Debug, Default, PartialEq)] +struct Outer { + a: i32, + b: Option, + c: String, + d: Vec, + e: Vec, + f: Vec, + g: MessageField, + h: Vec, + open: EnumValue, + closed: Color, + packed_closed: Vec, + single: f32, + opt_double: Option, + flag: bool, + unpacked_sfixed: Vec, + zigzag: Vec, + opt_open: Option>, + opt_bytes: Option>, + opt_str: Option, + rep_bytes: Vec>, + required_int: i64, + packed_fixed: Vec, + inl: MessageField>, + high: u32, + unknown: UnknownFields, +} + +/// A message with only the first field of `Inner`, which drops unknown fields. +#[derive(Clone, Debug, Default, PartialEq)] +struct Lossy { + id: i32, +} + +/// The field shapes that `Outer` does not use, in a message that drops +/// unknown fields. +#[derive(Clone, Debug, Default, PartialEq)] +struct Wide { + sint64: i64, + fixed64: Option, + sfixed32: Vec, + fixed64s: Vec, + req_str: String, + req_bytes: Vec, + req_enum: EnumValue, + rep_enum_open: Vec>, + opt_closed: Option, + req_bool: bool, + opt_u32: Option, + rep_double: Vec, + rep_bool: Vec, + rep_float: Vec, + opt_i64: Option, + opt_sint64: Option, + packed_u64: Vec, + opt_sint32: Option, + packed_sint64: Vec, + rep_enum_closed: Vec, +} + +/// The `dense` lookup for entries numbered `numbers`, which are ascending. +const fn dense(numbers: &[u32]) -> [u8; N] { + let mut d = [0u8; N]; + let mut i = 0; + while i < numbers.len() { + if (numbers[i] as usize) < N { + d[numbers[i] as usize] = (i + 1) as u8; + } + i += 1; + } + d +} + +const OUTER_NUMBERS: [u32; 24] = [ + 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 1000, +]; +const OUTER_DENSE: [u8; 24] = dense(&OUTER_NUMBERS); +const WIDE_NUMBERS: [u32; 20] = [ + 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, +]; +const WIDE_DENSE: [u8; 21] = dense(&WIDE_NUMBERS); + +static INNER: Table = crate::__table!( + Inner, + abi = ABI, + entries = [ + crate::__table_entry!(Inner, id, Int32Implicit, 1), + crate::__table_entry!(Inner, label, StrImplicit, 2), + crate::__table_entry!( + Inner, + next, + MsgSingular, + 3, + aux = 0, + slot = MessageField + ), + ], + dense = &dense::<4>(&[1, 2, 3]), + aux = [Aux::Msg(&MsgVt::new::>(&INNER))], + unknown = unknown, +); + +static LOSSY: Table = crate::__table!( + Lossy, + abi = ABI, + entries = [crate::__table_entry!(Lossy, id, Int32Implicit, 1)], + dense = &dense::<2>(&[1]), + aux = [], + unknown = none, +); + +static WIDE: Table = crate::__table!( + Wide, + abi = ABI, + entries = [ + crate::__table_entry!(Wide, sint64, Sint64Implicit, 1), + crate::__table_entry!(Wide, fixed64, Fixed64Optional, 2), + crate::__table_entry!(Wide, sfixed32, Sfixed32Packed, 3), + crate::__table_entry!(Wide, fixed64s, Fixed64Repeated, 4), + crate::__table_entry!(Wide, req_str, StrRequired, 5), + crate::__table_entry!(Wide, req_bytes, BytesRequired, 6), + crate::__table_entry!( + Wide, + req_enum, + EnumRequired, + 7, + aux = 0, + slot = as EnumShape>::Slot + ), + crate::__table_entry!( + Wide, + rep_enum_open, + EnumRepeated, + 8, + aux = 1, + slot = as EnumShape>::Slot + ), + crate::__table_entry!( + Wide, + opt_closed, + EnumOptional, + 9, + aux = 2, + slot = as EnumShape>::Slot + ), + crate::__table_entry!(Wide, req_bool, BoolRequired, 10), + crate::__table_entry!(Wide, opt_u32, Uint32Optional, 11), + crate::__table_entry!(Wide, rep_double, DoublePacked, 12), + crate::__table_entry!(Wide, rep_bool, BoolRepeated, 13), + crate::__table_entry!(Wide, rep_float, FloatPacked, 14), + crate::__table_entry!(Wide, opt_i64, Int64Optional, 15), + crate::__table_entry!(Wide, opt_sint64, Sint64Optional, 16), + crate::__table_entry!(Wide, packed_u64, Uint64Packed, 17), + crate::__table_entry!(Wide, opt_sint32, Sint32Optional, 18), + crate::__table_entry!(Wide, packed_sint64, Sint64Packed, 19), + crate::__table_entry!( + Wide, + rep_enum_closed, + EnumRepeated, + 20, + aux = 3, + slot = as EnumShape>::Slot + ), + ], + dense = &WIDE_DENSE, + aux = [ + Aux::Enum(&EnumVt::new::>()), + Aux::Enum(&EnumVt::new::>()), + Aux::Enum(&EnumVt::new::>()), + Aux::Enum(&EnumVt::new::>()), + ], + unknown = none, +); + +static OUTER: Table = crate::__table!( + Outer, + abi = ABI, + entries = [ + crate::__table_entry!(Outer, a, Int32Implicit, 1), + crate::__table_entry!(Outer, b, Uint64Optional, 2), + crate::__table_entry!(Outer, c, StrImplicit, 3), + crate::__table_entry!(Outer, d, BytesImplicit, 4), + crate::__table_entry!(Outer, e, Int32Packed, 5), + crate::__table_entry!(Outer, f, StrRepeated, 6), + crate::__table_entry!( + Outer, + g, + MsgSingular, + 7, + aux = 0, + slot = MessageField + ), + crate::__table_entry!(Outer, h, MsgRepeated, 8, aux = 1, slot = Vec), + crate::__table_entry!( + Outer, + open, + EnumImplicit, + 9, + aux = 2, + slot = as EnumShape>::Slot + ), + crate::__table_entry!( + Outer, + closed, + EnumImplicit, + 10, + aux = 3, + slot = as EnumShape>::Slot + ), + crate::__table_entry!( + Outer, + packed_closed, + EnumPacked, + 11, + aux = 4, + slot = as EnumShape>::Slot + ), + crate::__table_entry!(Outer, single, FloatImplicit, 12), + crate::__table_entry!(Outer, opt_double, DoubleOptional, 13), + crate::__table_entry!(Outer, flag, BoolImplicit, 14), + crate::__table_entry!(Outer, unpacked_sfixed, Sfixed64Repeated, 15), + crate::__table_entry!(Outer, zigzag, Sint32Packed, 16), + crate::__table_entry!( + Outer, + opt_open, + EnumOptional, + 17, + aux = 5, + slot = as EnumShape>::Slot + ), + crate::__table_entry!(Outer, opt_bytes, BytesOptional, 18), + crate::__table_entry!(Outer, opt_str, StrOptional, 19), + crate::__table_entry!(Outer, rep_bytes, BytesRepeated, 20), + crate::__table_entry!(Outer, required_int, Int64Required, 21), + crate::__table_entry!(Outer, packed_fixed, Fixed32Packed, 22), + crate::__table_entry!( + Outer, + inl, + MsgSingular, + 23, + aux = 6, + slot = MessageField> + ), + crate::__table_entry!(Outer, high, Uint32Implicit, 1000), + ], + dense = &OUTER_DENSE, + aux = [ + Aux::Msg(&MsgVt::new::>(&INNER)), + Aux::Rep(&RepVt::new::(&INNER)), + Aux::Enum(&EnumVt::new::>()), + Aux::Enum(&EnumVt::new::>()), + Aux::Enum(&EnumVt::new::>()), + Aux::Enum(&EnumVt::new::>()), + Aux::Msg(&MsgVt::new::>>(&INNER)), + ], + unknown = unknown, +); + +macro_rules! table_message { + ($ty:ty, $table:ident) => { + crate::impl_default_instance!($ty); + + impl Message for $ty { + fn compute_size(&self, cache: &mut SizeCache) -> u32 { + $table.compute_size(self, cache) + } + + fn write_to(&self, cache: &mut SizeCache, buf: &mut impl EncodeSink) { + $table.write_to(self, cache, buf); + } + + fn merge_field( + &mut self, + tag: Tag, + buf: &mut impl Buf, + ctx: DecodeContext<'_>, + ) -> Result<(), DecodeError> { + $table.merge_field(self, tag, buf, ctx) + } + + fn merge_to_limit( + &mut self, + buf: &mut impl Buf, + ctx: DecodeContext<'_>, + limit: usize, + ) -> Result<(), DecodeError> { + $table.merge_to_limit(self, buf, ctx, limit) + } + + fn merge_length_delimited( + &mut self, + buf: &mut impl Buf, + ctx: DecodeContext<'_>, + ) -> Result<(), DecodeError> { + $table.merge_length_delimited(self, buf, ctx) + } + + fn clear(&mut self) { + *self = Self::default(); + } + } + }; +} + +table_message!(Inner, INNER); +table_message!(Outer, OUTER); +table_message!(Lossy, LOSSY); +table_message!(Wide, WIDE); + +fn populated() -> Outer { + Outer { + a: -7, + b: Some(0), + c: "hello".into(), + d: vec![0, 1, 2], + e: vec![1, -1, 300], + f: vec!["x".into(), String::new(), "yz".into()], + g: MessageField::some(Inner { + id: 5, + label: "inner".into(), + next: MessageField::some(Inner { + id: 6, + ..Inner::default() + }), + unknown: UnknownFields::new(), + }), + h: vec![ + Inner { + id: 1, + ..Inner::default() + }, + Inner::default(), + ], + open: EnumValue::from(9), + closed: Color::Blue, + packed_closed: vec![Color::Green, Color::Blue, Color::Red], + single: 1.5, + opt_double: Some(-2.25), + flag: true, + unpacked_sfixed: vec![-1, 2], + zigzag: vec![-3, 4], + opt_open: Some(EnumValue::from(0)), + opt_bytes: Some(Vec::new()), + opt_str: Some("opt".into()), + rep_bytes: vec![vec![9], Vec::new()], + required_int: 0, + packed_fixed: vec![7, 8], + inl: MessageField::some(Inner { + id: 11, + next: MessageField::some(Inner { + label: "deep".into(), + ..Inner::default() + }), + ..Inner::default() + }), + high: 70_000, + unknown: UnknownFields::new(), + } +} + +fn round_trip(msg: &Outer) -> Outer { + let bytes = msg.encode_to_vec(); + assert_eq!(bytes.len() as u32, msg.encoded_len()); + Outer::decode_from_slice(&bytes).unwrap() +} + +// --------------------------------------------------------------------------- +// Encoding +// --------------------------------------------------------------------------- + +#[test] +fn encodes_known_wire_bytes() { + let msg = Outer { + a: 150, + c: "hi".into(), + e: vec![3, 270], + h: vec![Inner { + id: 1, + ..Inner::default() + }], + high: 1, + ..Outer::default() + }; + // a = 150; c = "hi"; e packed; h submessage {id = 1}; required_int (21) = 0 + // is always written; high (1000) has a two-byte tag. + let expected: &[u8] = &[ + 0x08, 0x96, 0x01, // 1: varint 150 + 0x1a, 0x02, b'h', b'i', // 3: "hi" + 0x2a, 0x03, 0x03, 0x8e, 0x02, // 5: packed [3, 270] + 0x42, 0x02, 0x08, 0x01, // 8: {1: 1} + 0xa8, 0x01, 0x00, // 21: varint 0 (required) + 0xc0, 0x3e, 0x01, // 1000: varint 1 + ]; + assert_eq!(msg.encode_to_vec(), expected); +} + +#[test] +fn an_empty_message_encodes_to_only_its_required_fields() { + assert_eq!(Outer::default().encode_to_vec(), [0xa8, 0x01, 0x00]); + assert_eq!(Inner::default().encode_to_vec(), Vec::::new()); +} + +#[test] +fn round_trips_every_shape() { + let msg = populated(); + assert_eq!(round_trip(&msg), msg); + assert_eq!(round_trip(&Outer::default()), Outer::default()); +} + +#[test] +fn every_sink_receives_the_same_bytes() { + let msg = populated(); + let expected = msg.encode_to_vec(); + + let mut rope = Rope::new(); + msg.encode(&mut rope); + assert_eq!(&rope.to_contiguous_bytes()[..], &expected[..]); + + let mut bytes_mut = crate::bytes::BytesMut::new(); + msg.encode(&mut bytes_mut); + assert_eq!(&bytes_mut[..], &expected[..]); + + let mut roomy = Vec::with_capacity(expected.len()); + msg.encode_length_delimited(&mut roomy); + let mut framed = Vec::new(); + crate::encoding::encode_varint(expected.len() as u64, &mut framed); + framed.extend_from_slice(&expected); + assert_eq!(roomy, framed); +} + +#[test] +fn a_large_bytes_field_encodes_into_a_rope() { + let payload = vec![0xab; 64 * 1024]; + let msg = Outer { + d: payload.clone(), + ..Outer::default() + }; + let mut rope = Rope::new(); + msg.encode(&mut rope); + assert_eq!(&rope.to_contiguous_bytes()[..], &msg.encode_to_vec()[..]); + assert!(rope.len() > payload.len()); +} + +#[test] +fn a_closed_enum_value_it_does_not_know_becomes_an_unknown_field() { + // closed (10) = 7; packed_closed (11) = 7 unpacked, then packed [1, 7]. + let wire = [0x50, 0x07, 0x58, 0x07, 0x5a, 0x02, 0x01, 0x07]; + let msg = Outer::decode_from_slice(&wire).unwrap(); + assert_eq!(msg.closed, Color::Red); + assert_eq!(msg.packed_closed, [Color::Green]); + let unknown: Vec<_> = msg.unknown.iter().map(|u| u.number).collect(); + assert_eq!(unknown, [10, 11, 11]); + assert!(msg + .unknown + .iter() + .all(|u| matches!(u.data, UnknownFieldData::Varint(7)))); + // Known fields are written first, in field order, then the unknown ones. + assert_eq!( + msg.encode_to_vec(), + [0x5a, 0x01, 0x01, 0xa8, 0x01, 0x00, 0x50, 0x07, 0x58, 0x07, 0x58, 0x07] + ); +} + +#[test] +fn a_closed_enum_value_is_dropped_by_a_message_that_drops_unknown_fields() { + // rep_enum_closed (20) = [1, 7, 2] unpacked; optional closed (9) = 7. + let wire = [ + 0xa0, 0x01, 0x01, 0xa0, 0x01, 0x07, 0xa0, 0x01, 0x02, 0x48, 0x07, + ]; + let msg = Wide::decode_from_slice(&wire).unwrap(); + assert_eq!(msg.rep_enum_closed, [Color::Green, Color::Blue]); + assert_eq!(msg.opt_closed, None); +} + +#[test] +fn an_open_enum_keeps_a_value_it_does_not_know() { + let msg = Outer::decode_from_slice(&[0x48, 0x09]).unwrap(); + assert_eq!(msg.open.to_i32(), 9); + assert!(msg.unknown.is_empty()); +} + +#[test] +fn unknown_fields_are_preserved_or_skipped() { + // field 500: varint 1; field 501: length-delimited "ab" + let wire = [0xa0, 0x1f, 0x01, 0xaa, 0x1f, 0x02, b'a', b'b']; + let msg = Inner::decode_from_slice(&wire).unwrap(); + assert_eq!(msg.unknown.len(), 2); + assert_eq!(msg.encode_to_vec(), wire); + assert!(matches!( + &msg.unknown.iter().nth(1).unwrap().data, + UnknownFieldData::LengthDelimited(v) if v == b"ab" + )); + + let lossy = Lossy::decode_from_slice(&wire).unwrap(); + assert_eq!(lossy, Lossy::default()); + assert!(lossy.encode_to_vec().is_empty()); +} + +#[test] +fn repeated_scalars_accept_packed_and_unpacked_forms() { + // e (5) is packed: [1] unpacked, [2, 3] packed, [4] unpacked. + let wire = [0x28, 0x01, 0x2a, 0x02, 0x02, 0x03, 0x28, 0x04]; + assert_eq!(Outer::decode_from_slice(&wire).unwrap().e, [1, 2, 3, 4]); + // unpacked_sfixed (15) is unpacked: a packed payload is also accepted. + let mut wire = vec![0x7a, 16]; + wire.extend_from_slice(&5i64.to_le_bytes()); + wire.extend_from_slice(&(-6i64).to_le_bytes()); + assert_eq!( + Outer::decode_from_slice(&wire).unwrap().unpacked_sfixed, + [5, -6] + ); +} + +#[test] +fn a_singular_field_takes_the_last_value_and_messages_merge() { + // a = 1, a = 2; g = {id = 5}, g = {label = "l"} + let wire = [ + 0x08, 0x01, 0x08, 0x02, 0x3a, 0x02, 0x08, 0x05, 0x3a, 0x03, 0x12, 0x01, b'l', + ]; + let msg = Outer::decode_from_slice(&wire).unwrap(); + assert_eq!(msg.a, 2); + assert_eq!(msg.g.id, 5); + assert_eq!(msg.g.label, "l"); +} + +#[test] +fn a_wire_type_mismatch_is_an_error() { + // a (varint) sent as length-delimited; c (string) sent as varint. + assert!(matches!( + Outer::decode_from_slice(&[0x0a, 0x00]), + Err(DecodeError::WireTypeMismatch { .. }) + )); + assert!(matches!( + Outer::decode_from_slice(&[0x18, 0x01]), + Err(DecodeError::WireTypeMismatch { .. }) + )); +} + +#[test] +fn invalid_utf8_in_a_string_is_an_error() { + assert!(matches!( + Outer::decode_from_slice(&[0x1a, 0x01, 0xff]), + Err(DecodeError::InvalidUtf8) + )); +} + +#[test] +fn a_truncated_message_is_an_error_or_ends_on_a_field_boundary() { + let bytes = populated().encode_to_vec(); + let mut ok = 0; + for end in 0..bytes.len() { + match Outer::decode_from_slice(&bytes[..end]) { + Ok(_) => ok += 1, + Err(e) => assert!( + matches!(e, DecodeError::UnexpectedEof | DecodeError::VarintTooLong), + "prefix {end}: {e}" + ), + } + } + // The empty prefix and each top-level field boundary decode; the rest fail. + assert!(ok > 1 && ok < bytes.len() / 2); + assert!(Outer::decode_from_slice(&bytes).is_ok()); +} + +#[test] +fn a_length_that_runs_past_its_message_is_an_error() { + // g = a submessage of 2 bytes whose label declares 5. + assert_eq!( + Outer::decode_from_slice(&[0x3a, 0x02, 0x12, 0x05, b'a', b'b', b'c']), + Err(DecodeError::UnexpectedEof) + ); +} + +#[test] +fn a_non_contiguous_buffer_is_gathered() { + let msg = populated(); + let bytes = msg.encode_to_vec(); + for split in [1, 2, bytes.len() / 2, bytes.len() - 1] { + let (head, tail) = bytes.split_at(split); + let mut chained = head.chain(tail); + let mut decoded = Outer::default(); + with_ctx(|ctx| decoded.merge(&mut chained, ctx)).unwrap(); + assert_eq!(decoded, msg); + } +} + +#[test] +fn length_delimited_decode_leaves_the_rest_of_the_buffer() { + let msg = populated(); + let mut framed = Vec::new(); + msg.encode_length_delimited(&mut framed); + framed.extend_from_slice(b"tail"); + let mut buf = &framed[..]; + let decoded = Outer::decode_length_delimited(&mut buf).unwrap(); + assert_eq!(decoded, msg); + assert_eq!(buf, b"tail"); +} + +fn with_ctx(f: impl FnOnce(DecodeContext<'_>) -> R) -> R { + let limit = core::cell::Cell::new(1000); + f(DecodeContext::new(crate::RECURSION_LIMIT, &limit)) +} + +#[test] +fn merge_field_decodes_one_field_from_a_contiguous_buffer() { + let wire = [0x08, 0x2a, 0x12, 0x02, b'o', b'k']; + let mut msg = Inner::default(); + let mut buf = &wire[..]; + let tag = Tag::decode(&mut buf).unwrap(); + with_ctx(|ctx| msg.merge_field(tag, &mut buf, ctx)).unwrap(); + assert_eq!(msg.id, 42); + assert_eq!(buf, &wire[2..]); +} + +#[test] +fn merge_field_rejects_a_non_contiguous_buffer() { + let wire = [0x2a, 0x01]; + let mut chained = (&wire[..1]).chain(&wire[1..]); + let mut msg = Inner::default(); + assert!(matches!( + with_ctx(|ctx| msg.merge_field(Tag::new(1, WireType::Varint), &mut chained, ctx)), + Err(DecodeError::UnexpectedEof) + )); +} + +#[test] +fn every_remaining_shape_round_trips() { + let msg = Wide { + sint64: -5, + fixed64: Some(0), + sfixed32: vec![-1, 2], + fixed64s: vec![3, 4], + req_str: String::new(), + req_bytes: vec![1], + req_enum: EnumValue::from(4), + rep_enum_open: vec![EnumValue::from(0), EnumValue::from(9)], + opt_closed: Some(Color::Green), + req_bool: false, + opt_u32: Some(0), + rep_double: vec![0.0, -1.5], + rep_bool: vec![true, false], + rep_float: vec![2.5], + opt_i64: Some(-7), + opt_sint64: Some(-8), + packed_u64: vec![u64::MAX], + opt_sint32: Some(i32::MIN), + packed_sint64: vec![i64::MIN, 1], + rep_enum_closed: vec![Color::Blue], + }; + let bytes = msg.encode_to_vec(); + assert_eq!(bytes.len() as u32, msg.encoded_len()); + assert_eq!(Wide::decode_from_slice(&bytes).unwrap(), msg); + // The required fields are written even when they hold the default. + let empty = Wide::default().encode_to_vec(); + assert_eq!(empty, [0x2a, 0x00, 0x32, 0x00, 0x38, 0x00, 0x50, 0x00]); +} + +// --------------------------------------------------------------------------- +// Limits +// --------------------------------------------------------------------------- + +fn nested(depth: usize) -> Vec { + // next (3) nested `depth` levels deep, innermost empty. + let mut wire = Vec::new(); + for _ in 0..depth { + let mut outer = vec![0x1a]; + crate::encoding::encode_varint(wire.len() as u64, &mut outer); + outer.extend_from_slice(&wire); + wire = outer; + } + wire +} + +#[test] +fn the_recursion_limit_applies() { + let deep = nested(150); + assert!(matches!( + Inner::decode_from_slice(&deep), + Err(DecodeError::RecursionLimitExceeded) + )); + assert!(Inner::decode_from_slice(&nested(50)).is_ok()); + assert!(matches!( + DecodeOptions::new() + .with_recursion_limit(10) + .decode_from_slice::(&nested(50)), + Err(DecodeError::RecursionLimitExceeded) + )); +} + +#[test] +fn the_element_memory_limit_applies_to_repeated_messages_and_strings() { + // 1000 empty elements of a repeated message field (8), then of strings (6). + let messages: Vec = (0..1000).flat_map(|_| [0x42, 0x00]).collect(); + let strings: Vec = (0..1000).flat_map(|_| [0x32, 0x00]).collect(); + for wire in [&messages, &strings] { + assert!(Outer::decode_from_slice(wire).is_ok()); + assert!(matches!( + DecodeOptions::new() + .with_element_memory_limit(100) + .decode_from_slice::(wire), + Err(DecodeError::ElementMemoryLimitExceeded) + )); + } +} + +#[test] +fn the_unknown_field_limit_applies() { + let wire: Vec = (0..100).flat_map(|_| [0xa0, 0x1f, 0x01]).collect(); + assert!(Inner::decode_from_slice(&wire).is_ok()); + assert!(matches!( + DecodeOptions::new() + .with_unknown_field_limit(10) + .decode_from_slice::(&wire), + Err(DecodeError::UnknownFieldLimitExceeded) + )); +} + +#[test] +fn a_closed_enum_value_counts_against_the_unknown_field_limit() { + // closed (10) = 7, a hundred times. + let wire: Vec = (0..100).flat_map(|_| [0x50, 0x07]).collect(); + assert!(Outer::decode_from_slice(&wire).is_ok()); + assert!(matches!( + DecodeOptions::new() + .with_unknown_field_limit(10) + .decode_from_slice::(&wire), + Err(DecodeError::UnknownFieldLimitExceeded) + )); +} + +#[test] +fn a_sub_message_longer_than_the_size_limit_is_an_error() { + // g (7) declares a length of 2^31. + let wire = [0x3a, 0x80, 0x80, 0x80, 0x80, 0x08]; + assert_eq!( + Outer::decode_from_slice(&wire), + Err(DecodeError::MessageTooLarge) + ); + // The same length as the prefix of a length-delimited message. + let framed = [0x80, 0x80, 0x80, 0x80, 0x08]; + assert_eq!( + Outer::decode_length_delimited(&mut &framed[..]), + Err(DecodeError::MessageTooLarge) + ); +} + +#[test] +fn a_repeated_message_element_that_fails_to_decode_is_not_kept() { + // h (8): one valid element, then one whose string label is not UTF-8. + let wire = [0x42, 0x02, 0x08, 0x01, 0x42, 0x03, 0x12, 0x01, 0xff]; + let mut msg = Outer::default(); + let result = with_ctx(|ctx| msg.merge(&mut &wire[..], ctx)); + assert_eq!(result, Err(DecodeError::InvalidUtf8)); + assert_eq!(msg.h.len(), 1); + assert_eq!(msg.h[0].id, 1); +} + +// --------------------------------------------------------------------------- +// Table construction +// --------------------------------------------------------------------------- + +#[test] +fn an_entry_records_its_tag_and_the_tag_length() { + let one = Entry::new(Kind::Int32Implicit, 15, 0, 0); + assert_eq!((one.tag, one.tag_len), (15 << 3, 1)); + let two = Entry::new(Kind::Int32Implicit, 16, 0, 0); + assert_eq!((two.tag, two.tag_len), (16 << 3, 2)); + let packed = Entry::new(Kind::Int32Packed, 1, 0, 0); + assert_eq!(packed.tag, (1 << 3) | 2); + let fixed = Entry::new(Kind::DoubleImplicit, 1, 0, 0); + assert_eq!(fixed.tag, (1 << 3) | 1); + let big = Entry::new(Kind::Int32Implicit, (1 << 29) - 1, 0, 0); + assert_eq!(big.tag_len, 5); +} + +#[test] +#[should_panic(expected = "field number out of range")] +fn an_entry_rejects_field_number_zero() { + let _ = Entry::new(Kind::Int32Implicit, 0, 0, 0); +} + +#[test] +fn find_uses_the_dense_array_and_falls_back_to_search() { + let raw = &OUTER.raw; + assert_eq!(raw.find(1).unwrap().kind, Kind::Int32Implicit); + assert_eq!(raw.find(22).unwrap().kind, Kind::Fixed32Packed); + assert_eq!(raw.find(23).unwrap().kind, Kind::MsgSingular); + assert!(raw.find(24).is_none()); + assert!(raw.find(63).is_none()); + assert_eq!(raw.find(1000).unwrap().kind, Kind::Uint32Implicit); + assert!(raw.find(1001).is_none()); +} + +/// Building a table that violates the checks in `Table::new` is a compile +/// error in a `static`, so these run the checks at run time. +mod invalid_tables { + use super::*; + + const ENTRY_1: Entry = Entry::new(Kind::Int32Implicit, 1, 0, 0); + const ENTRY_2: Entry = Entry::new(Kind::Int32Implicit, 2, 0, 0); + + /// `Table::new` on a `Lossy` (one `i32`), which is never used to access a + /// message. + fn lossy( + abi: u32, + entries: &'static [Entry], + dense: &'static [u8], + aux: &'static [Aux], + unknown: Option, + ) -> Table { + // SAFETY: the table is dropped without being used. + unsafe { Table::new(abi, entries, dense, aux, unknown) } + } + + #[test] + #[should_panic(expected = "different table ABI")] + fn the_abi_must_match() { + let _ = lossy(ABI + 1, &[ENTRY_1], &[], &[], None); + } + + #[test] + #[should_panic(expected = "strictly increasing")] + fn entries_must_be_sorted() { + let _ = lossy(ABI, &[ENTRY_2, ENTRY_1], &[], &[], None); + } + + #[test] + #[should_panic(expected = "strictly increasing")] + fn field_numbers_must_be_distinct() { + let _ = lossy(ABI, &[ENTRY_1, ENTRY_1], &[], &[], None); + } + + #[test] + #[should_panic(expected = "offset is outside the message struct")] + fn an_offset_must_lie_inside_the_message() { + const E: Entry = Entry::new(Kind::Int32Implicit, 1, 4, 0); + let _ = lossy(ABI, &[E], &[], &[], None); + } + + #[test] + #[should_panic(expected = "aux index is out of range")] + fn an_enum_entry_needs_an_aux() { + const E: Entry = Entry::new(Kind::EnumImplicit, 1, 0, 0); + let _ = lossy(ABI, &[E], &[], &[], None); + } + + static CLOSED_VT: EnumVt = EnumVt::new::>(); + static CLOSED_AUX: [Aux; 1] = [Aux::Enum(&CLOSED_VT)]; + + #[test] + #[should_panic(expected = "wrong variant")] + fn an_aux_must_be_of_the_kinds_variant() { + const E: Entry = Entry::new(Kind::MsgSingular, 1, 0, 0); + let _ = lossy(ABI, &[E], &[], &CLOSED_AUX, None); + } + + #[test] + #[should_panic(expected = "wrong cardinality")] + fn an_enum_shape_must_have_the_cardinality_of_its_kind() { + const E: Entry = Entry::new(Kind::EnumPacked, 1, 0, 0); + let _ = lossy(ABI, &[E], &[], &CLOSED_AUX, None); + } + + #[test] + fn an_enum_shape_of_the_right_cardinality_is_accepted() { + const E: Entry = Entry::new(Kind::EnumRequired, 1, 0, 0); + let _ = lossy(ABI, &[E], &[], &CLOSED_AUX, None); + } + + #[test] + #[should_panic(expected = "wrong field number")] + fn a_dense_slot_must_name_its_entry() { + let _ = lossy(ABI, &[ENTRY_1], &[0, 0, 1], &[], None); + } + + #[test] + #[should_panic(expected = "omits an entry")] + fn the_dense_array_must_cover_every_entry_in_its_range() { + let _ = lossy(ABI, &[ENTRY_1], &[0, 0], &[], None); + } + + #[test] + #[should_panic(expected = "unknown-fields offset")] + fn the_unknown_fields_must_fit_in_the_message() { + let _ = lossy(ABI, &[ENTRY_1], &[], &[], Some(0)); + } +} + +// --------------------------------------------------------------------------- +// The `EncodeSink` hooks +// --------------------------------------------------------------------------- + +#[test] +fn the_pre_sized_hook_runs_only_on_a_pre_sized_cursor() { + use core::mem::MaybeUninit; + + let mut ran = false; + let mut vec: Vec = Vec::new(); + assert!(!vec.__with_pre_sized(&mut |_| ran = true)); + let mut rope = Rope::new(); + assert!(!rope.__with_pre_sized(&mut |_| ran = true)); + assert!(!ran); + + let mut storage = [MaybeUninit::::uninit(); 8]; + let mut cursor = crate::encode_sink::PreSized::new(&mut storage); + assert!(cursor.__with_pre_sized(&mut |c| c.put_u8(7))); + assert_eq!(cursor.written(), 1); +} + +#[test] +fn writing_through_a_nested_cursor_continues_after_the_bytes_already_written() { + // A manual `write_to` that writes a prefix and then a table message must + // append the message after the prefix, through the same cursor. + #[derive(Clone, Default, PartialEq)] + struct Prefixed(Inner); + crate::impl_default_instance!(Prefixed); + impl Message for Prefixed { + fn compute_size(&self, cache: &mut SizeCache) -> u32 { + 2 + self.0.compute_size(cache) + } + fn write_to(&self, cache: &mut SizeCache, buf: &mut impl EncodeSink) { + buf.put_slice(&[0xaa, 0xbb]); + self.0.write_to(cache, buf); + } + fn merge_field( + &mut self, + _: Tag, + _: &mut impl Buf, + _: DecodeContext<'_>, + ) -> Result<(), DecodeError> { + unreachable!() + } + fn clear(&mut self) {} + } + let msg = Prefixed(Inner { + id: 3, + label: "x".into(), + ..Inner::default() + }); + let mut expected = vec![0xaa, 0xbb]; + expected.extend(msg.0.encode_to_vec()); + assert_eq!(msg.encode_to_vec(), expected); + let mut rope = Rope::new(); + msg.encode(&mut rope); + assert_eq!(&rope.to_contiguous_bytes()[..], &expected[..]); +}