Add RevBuilder for reverse message encoding
Generate RevBuilder structs in codegen alongside ProtoBuilder. RevBuilder writes fields into a buffer starting from the end, allowing fields to be set in reverse declaration order. This avoids reallocation and produces valid protobuf data matching the forward layout. Update generated interop files and add comprehensive tests for wire format, nesting, map entries, overflow handling, and buffer marking utilities.
This commit is contained in:
+385
-3
@@ -6,8 +6,8 @@ extern crate alloc;
|
||||
#[cfg(feature = "std")]
|
||||
extern crate std;
|
||||
|
||||
use core::fmt;
|
||||
use bytes::BufMut;
|
||||
use core::fmt;
|
||||
|
||||
pub struct MapFieldIterator<'a> {
|
||||
inner: RepeatedFieldIterator<'a>,
|
||||
@@ -65,7 +65,9 @@ impl std::error::Error for RotoError {}
|
||||
pub type Result<T> = core::result::Result<T, RotoError>;
|
||||
|
||||
pub trait RotoOwned {
|
||||
type Reader<'a> where Self: 'a;
|
||||
type Reader<'a>
|
||||
where
|
||||
Self: 'a;
|
||||
fn reader(&self) -> Self::Reader<'_>;
|
||||
}
|
||||
|
||||
@@ -442,7 +444,7 @@ impl<'a> Iterator for RawFieldIterator<'a> {
|
||||
mod tests {
|
||||
use super::*;
|
||||
#[cfg(feature = "alloc")]
|
||||
use alloc::{vec, vec::{Vec}};
|
||||
use alloc::{vec, vec::Vec};
|
||||
|
||||
#[test]
|
||||
fn test_varint_read_write() {
|
||||
@@ -745,6 +747,199 @@ mod tests {
|
||||
}
|
||||
assert_eq!(found_count, essential_fields.len());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_string_wire_format() {
|
||||
let mut buf = [0u8; 256];
|
||||
let mut builder = RevBuilder::new(&mut buf);
|
||||
builder.write_string(1, "hello").unwrap();
|
||||
let data = builder.finish().unwrap();
|
||||
assert_eq!(data, &[0x0A, 0x05, b'h', b'e', b'l', b'l', b'o']);
|
||||
let acc = ProtoAccessor::new(data).unwrap();
|
||||
let (val, wt) = acc.get_value(1).unwrap();
|
||||
assert_eq!(wt, WireType::LengthDelimited);
|
||||
assert_eq!(val, b"hello");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_multiple_fields_reverse_order() {
|
||||
let mut buf = [0u8; 256];
|
||||
let mut builder = RevBuilder::new(&mut buf);
|
||||
builder.write_varint(3, 300).unwrap();
|
||||
builder.write_string(2, "hi").unwrap();
|
||||
builder.write_varint(1, 42).unwrap();
|
||||
let data = builder.finish().unwrap();
|
||||
let acc = ProtoAccessor::new(data).unwrap();
|
||||
let (val1, _) = acc.get_value(1).unwrap();
|
||||
let (v1, _) = read_varint(val1).unwrap();
|
||||
assert_eq!(v1, 42);
|
||||
let (val2, _) = acc.get_value(2).unwrap();
|
||||
assert_eq!(val2, b"hi");
|
||||
let (val3, _) = acc.get_value(3).unwrap();
|
||||
let (v3, _) = read_varint(val3).unwrap();
|
||||
assert_eq!(v3, 300);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_nested_message() {
|
||||
let mut buf = [0u8; 256];
|
||||
let mut builder = RevBuilder::new(&mut buf);
|
||||
// Mark before writing inner fields
|
||||
let inner_start = builder.mark();
|
||||
// Write inner field
|
||||
builder.write_varint(1, 100).unwrap();
|
||||
// Write length prefix and tag for outer field
|
||||
builder.write_length_since(inner_start).unwrap();
|
||||
builder.put_tag(1, WireType::LengthDelimited).unwrap();
|
||||
let data = builder.finish().unwrap();
|
||||
let acc = ProtoAccessor::new(data).unwrap();
|
||||
let (nested_bytes, wt) = acc.get_value(1).unwrap();
|
||||
assert_eq!(wt, WireType::LengthDelimited);
|
||||
let nested_acc = ProtoAccessor::new(nested_bytes).unwrap();
|
||||
let (inner_val, _) = nested_acc.get_value(1).unwrap();
|
||||
let (inner_v, _) = read_varint(inner_val).unwrap();
|
||||
assert_eq!(inner_v, 100);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_readable_by_accessor() {
|
||||
let mut buf = [0u8; 256];
|
||||
let mut builder = RevBuilder::new(&mut buf);
|
||||
builder.write_fixed64(6, 0xDEADBEEFCAFEBABE).unwrap();
|
||||
builder.write_fixed32(5, 0xDEADBEEFu32).unwrap();
|
||||
builder.write_bytes(4, &[1, 2, 3, 255, 0]).unwrap();
|
||||
builder.write_string(3, "protobuf").unwrap();
|
||||
builder.write_int32(2, -42).unwrap();
|
||||
builder.write_varint(1, 0x1234ABCD).unwrap();
|
||||
let data = builder.finish().unwrap();
|
||||
let acc = ProtoAccessor::new(data).unwrap();
|
||||
let (v1, _) = acc.get_value(1).unwrap();
|
||||
let (val1, _) = read_varint(v1).unwrap();
|
||||
assert_eq!(val1, 0x1234ABCD);
|
||||
let (v2, _) = acc.get_value(2).unwrap();
|
||||
let (val2, _) = read_varint(v2).unwrap();
|
||||
assert_eq!(val2, -42i32 as u64);
|
||||
assert_eq!(acc.get_value(3).unwrap().0, b"protobuf");
|
||||
assert_eq!(acc.get_value(4).unwrap().0, &[1, 2, 3, 255, 0]);
|
||||
assert_eq!(acc.get_value(5).unwrap().0, &(0xDEADBEEFu32).to_le_bytes());
|
||||
assert_eq!(
|
||||
acc.get_value(6).unwrap().0,
|
||||
&(0xDEADBEEFCAFEBABEu64).to_le_bytes()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_buffer_overflow() {
|
||||
let mut buf = [0u8; 4];
|
||||
let mut builder = RevBuilder::new(&mut buf);
|
||||
let result = builder.write_string(1, "hello");
|
||||
assert_eq!(result, Err(RotoError::BufferOverflow));
|
||||
let mut buf2 = [0u8; 1];
|
||||
let mut builder2 = RevBuilder::new(&mut buf2);
|
||||
let result2 = builder2.write_varint(1000, 1);
|
||||
assert_eq!(result2, Err(RotoError::BufferOverflow));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_matches_proto_builder() {
|
||||
let mut fwd_buf = [0u8; 256];
|
||||
let mut fwd_builder = ProtoBuilder::new(&mut fwd_buf);
|
||||
fwd_builder.write_string(1, "hello").unwrap();
|
||||
fwd_builder.write_int32(2, 42).unwrap();
|
||||
fwd_builder.write_bytes(3, &[1, 2, 3]).unwrap();
|
||||
let fwd_data = fwd_builder.finish().unwrap();
|
||||
let mut rev_buf = [0u8; 256];
|
||||
let mut rev_builder = RevBuilder::new(&mut rev_buf);
|
||||
rev_builder.write_bytes(3, &[1, 2, 3]).unwrap();
|
||||
rev_builder.write_int32(2, 42).unwrap();
|
||||
rev_builder.write_string(1, "hello").unwrap();
|
||||
let rev_data = rev_builder.finish().unwrap();
|
||||
let fwd_acc = ProtoAccessor::new(fwd_data).unwrap();
|
||||
let rev_acc = ProtoAccessor::new(rev_data).unwrap();
|
||||
let (f1, _) = fwd_acc.get_value(1).unwrap();
|
||||
let (r1, _) = rev_acc.get_value(1).unwrap();
|
||||
assert_eq!(f1, r1);
|
||||
assert_eq!(f1, b"hello");
|
||||
let (f2, _) = fwd_acc.get_value(2).unwrap();
|
||||
let (r2, _) = rev_acc.get_value(2).unwrap();
|
||||
assert_eq!(f2, r2);
|
||||
assert_eq!(f2, &[42]);
|
||||
let (f3, _) = fwd_acc.get_value(3).unwrap();
|
||||
let (r3, _) = rev_acc.get_value(3).unwrap();
|
||||
assert_eq!(f3, r3);
|
||||
assert_eq!(f3, &[1, 2, 3]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_empty_finish() {
|
||||
let mut buf = [1u8; 64];
|
||||
let builder = RevBuilder::new(&mut buf);
|
||||
let data = builder.finish().unwrap();
|
||||
assert!(data.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_map_entry() {
|
||||
let mut buf = [0u8; 256];
|
||||
let mut builder = RevBuilder::new(&mut buf);
|
||||
let mut key_buf = [0u8; 16];
|
||||
let key_tag_len = Tag::encode(1, WireType::Varint, &mut key_buf).unwrap();
|
||||
key_buf[key_tag_len] = 7;
|
||||
let key_bytes = &key_buf[..key_tag_len + 1];
|
||||
let mut val_buf = [0u8; 16];
|
||||
let val_tag_len = Tag::encode(2, WireType::LengthDelimited, &mut val_buf).unwrap();
|
||||
val_buf[val_tag_len] = 1;
|
||||
val_buf[val_tag_len + 1] = b'x';
|
||||
let val_bytes = &val_buf[..val_tag_len + 2];
|
||||
builder.write_map_entry(10, key_bytes, val_bytes).unwrap();
|
||||
let data = builder.finish().unwrap();
|
||||
let acc = ProtoAccessor::new(data).unwrap();
|
||||
let (entry_bytes, _) = acc.get_value(10).unwrap();
|
||||
let entry_acc = ProtoAccessor::new(entry_bytes).unwrap();
|
||||
let (key_v, _) = entry_acc.get_value(1).unwrap();
|
||||
let (key_val, _) = read_varint(key_v).unwrap();
|
||||
assert_eq!(key_val, 7);
|
||||
let (val_v, _) = entry_acc.get_value(2).unwrap();
|
||||
assert_eq!(val_v, b"x");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_written_since() {
|
||||
let mut buf = [0u8; 256];
|
||||
let mut builder = RevBuilder::new(&mut buf);
|
||||
let m1 = builder.mark();
|
||||
builder.write_string(1, "test").unwrap();
|
||||
assert!(builder.written_since(m1) > 0);
|
||||
let m2 = builder.mark();
|
||||
builder.write_varint(2, 999).unwrap();
|
||||
assert!(builder.written_since(m2) > 0);
|
||||
let total = builder.written_since(m1);
|
||||
let just_field2 = builder.written_since(m2);
|
||||
assert!(total > just_field2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_revbuilder_write_raw() {
|
||||
let mut src_buf = [0u8; 128];
|
||||
let mut src_builder = ProtoBuilder::new(&mut src_buf);
|
||||
src_builder.write_string(1, "original").unwrap();
|
||||
src_builder.write_int32(2, 99).unwrap();
|
||||
let src_data = src_builder.finish().unwrap();
|
||||
let src_acc = ProtoAccessor::new(src_data).unwrap();
|
||||
let mut dst_buf = [0u8; 128];
|
||||
let mut dst_builder = RevBuilder::new(&mut dst_buf);
|
||||
let fields: Vec<_> = src_acc.raw_fields().collect();
|
||||
for item in fields.into_iter().rev() {
|
||||
let (_, raw_bytes) = item.unwrap();
|
||||
dst_builder.write_raw(raw_bytes).unwrap();
|
||||
}
|
||||
let dst_data = dst_builder.finish().unwrap();
|
||||
let dst_acc = ProtoAccessor::new(dst_data).unwrap();
|
||||
let (d1, _) = dst_acc.get_value(1).unwrap();
|
||||
assert_eq!(d1, b"original");
|
||||
let (d2, _) = dst_acc.get_value(2).unwrap();
|
||||
assert_eq!(d2, &[99]);
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ProtoBuilder<'a> {
|
||||
@@ -922,3 +1117,190 @@ impl<'a, B: BufMut> BufMutBuilder<'a, B> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// A reverse-direction protobuf builder that writes fields BACKWARDS through
|
||||
/// a buffer, from the end toward the beginning. This enables zero-copy
|
||||
/// single-pass encoding for length-delimited payloads: write the payload
|
||||
/// first, compute its size, then write the length varint and tag going backwards.
|
||||
///
|
||||
/// After encoding, the valid data is in `buf[pos..buf.len()]` — no memmove is
|
||||
/// needed because the first byte we wrote (now at `pos`) IS the first byte of
|
||||
/// the wire-format message.
|
||||
///
|
||||
/// # Invariant
|
||||
///
|
||||
/// `pos` always points to the **start** of valid data. Data lives in
|
||||
/// `buf[pos..buf.len()]`. `pos` starts at `buf.len()` (empty) and
|
||||
/// decreases as we write.
|
||||
pub struct RevBuilder<'a> {
|
||||
buf: &'a mut [u8],
|
||||
pos: usize,
|
||||
}
|
||||
|
||||
impl<'a> RevBuilder<'a> {
|
||||
/// Create a new `RevBuilder` that writes backwards into `buf`.
|
||||
///
|
||||
/// The buffer is initially empty; valid data grows from the end
|
||||
/// toward the beginning.
|
||||
pub fn new(buf: &'a mut [u8]) -> Self {
|
||||
let pos = buf.len();
|
||||
Self { buf, pos }
|
||||
}
|
||||
|
||||
// ------------------------------------------------------------------
|
||||
// Low-level primitives (write at current pos, moving pos left)
|
||||
// ------------------------------------------------------------------
|
||||
|
||||
/// Encode `value` as a varint and place it at the current position,
|
||||
/// moving `pos` left by the encoded length.
|
||||
pub fn put_varint(&mut self, value: u64) -> Result<()> {
|
||||
let mut temp = [0u8; 10];
|
||||
let len = write_varint(value, &mut temp)?;
|
||||
if self.pos < len {
|
||||
return Err(RotoError::BufferOverflow);
|
||||
}
|
||||
self.pos -= len;
|
||||
self.buf[self.pos..self.pos + len].copy_from_slice(&temp[..len]);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Copy `bytes` into the buffer at the current position, moving `pos`
|
||||
/// left by `bytes.len()`.
|
||||
pub fn put_slice(&mut self, bytes: &[u8]) -> Result<()> {
|
||||
let len = bytes.len();
|
||||
if self.pos < len {
|
||||
return Err(RotoError::BufferOverflow);
|
||||
}
|
||||
self.pos -= len;
|
||||
self.buf[self.pos..self.pos + len].copy_from_slice(bytes);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Encode a tag varint at the current position.
|
||||
pub fn put_tag(&mut self, field_number: u32, wire_type: WireType) -> Result<()> {
|
||||
let mut temp = [0u8; 10];
|
||||
let len = Tag::encode(field_number, wire_type, &mut temp)?;
|
||||
if self.pos < len {
|
||||
return Err(RotoError::BufferOverflow);
|
||||
}
|
||||
self.pos -= len;
|
||||
self.buf[self.pos..self.pos + len].copy_from_slice(&temp[..len]);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ------------------------------------------------------------------
|
||||
// High-level field writers
|
||||
// ------------------------------------------------------------------
|
||||
|
||||
/// Encode a length-delimited string field.
|
||||
///
|
||||
/// Write order: payload → length varint → tag (all going backwards),
|
||||
/// which produces the correct wire order: `[tag][len][payload]`.
|
||||
pub fn write_string(&mut self, field_number: u32, value: &str) -> Result<()> {
|
||||
let bytes = value.as_bytes();
|
||||
// payload first (moves pos left)
|
||||
self.put_slice(bytes)?;
|
||||
// length
|
||||
self.put_varint(bytes.len() as u64)?;
|
||||
// tag
|
||||
self.put_tag(field_number, WireType::LengthDelimited)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Encode a length-delimited bytes field.
|
||||
pub fn write_bytes(&mut self, field_number: u32, value: &[u8]) -> Result<()> {
|
||||
self.put_slice(value)?;
|
||||
self.put_varint(value.len() as u64)?;
|
||||
self.put_tag(field_number, WireType::LengthDelimited)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Encode a varint field (tag + varint value).
|
||||
pub fn write_varint(&mut self, field_number: u32, value: u64) -> Result<()> {
|
||||
self.put_varint(value)?;
|
||||
self.put_tag(field_number, WireType::Varint)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Encode an int32 field (same wire encoding as varint).
|
||||
pub fn write_int32(&mut self, field_number: u32, value: i32) -> Result<()> {
|
||||
self.write_varint(field_number, value as u64)
|
||||
}
|
||||
|
||||
/// Encode a fixed32 field (tag + 4-byte LE value).
|
||||
pub fn write_fixed32(&mut self, field_number: u32, value: u32) -> Result<()> {
|
||||
self.put_slice(&value.to_le_bytes())?;
|
||||
self.put_tag(field_number, WireType::Fixed32)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Encode a fixed64 field (tag + 8-byte LE value).
|
||||
pub fn write_fixed64(&mut self, field_number: u32, value: u64) -> Result<()> {
|
||||
self.put_slice(&value.to_le_bytes())?;
|
||||
self.put_tag(field_number, WireType::Fixed64)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Write a pre-encoded field (tag + value) verbatim into the buffer.
|
||||
///
|
||||
/// Useful for the `with()` pattern: copy raw bytes from an existing
|
||||
/// message without re-encoding.
|
||||
pub fn write_raw(&mut self, raw_bytes: &[u8]) -> Result<()> {
|
||||
self.put_slice(raw_bytes)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Encode a map entry as a length-delimited field containing
|
||||
/// `key_encoded` followed by `value_encoded`.
|
||||
pub fn write_map_entry(
|
||||
&mut self,
|
||||
field_number: u32,
|
||||
key_encoded: &[u8],
|
||||
value_encoded: &[u8],
|
||||
) -> Result<()> {
|
||||
// We want final bytes: [key...][value...] (in wire order)
|
||||
// Written backwards: first write value at right, then key at left.
|
||||
// So the last thing written (key) ends up at the start (left).
|
||||
self.put_slice(value_encoded)?;
|
||||
self.put_slice(key_encoded)?;
|
||||
let entry_len = key_encoded.len() + value_encoded.len();
|
||||
self.put_varint(entry_len as u64)?;
|
||||
self.put_tag(field_number, WireType::LengthDelimited)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ------------------------------------------------------------------
|
||||
// Nested-message helpers
|
||||
// ------------------------------------------------------------------
|
||||
|
||||
/// Return the current position (start of valid data).
|
||||
///
|
||||
/// Call this *before* writing nested fields, then pass the returned
|
||||
/// value to [`write_length_since`](Self::write_length_since) to
|
||||
/// encode the length prefix.
|
||||
pub fn mark(&self) -> usize {
|
||||
self.pos
|
||||
}
|
||||
|
||||
/// Return the number of bytes written since `mark`.
|
||||
pub fn written_since(&self, mark: usize) -> usize {
|
||||
mark - self.pos
|
||||
}
|
||||
|
||||
/// Write the length-prefix varint for the bytes that were written
|
||||
/// between `mark` and now. Call this *after* encoding the nested
|
||||
/// message fields.
|
||||
pub fn write_length_since(&mut self, mark: usize) -> Result<()> {
|
||||
let len = mark - self.pos;
|
||||
self.put_varint(len as u64)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Finalize and return the encoded slice.
|
||||
///
|
||||
/// The valid data is `buf[pos..buf.len()]` — the bytes we wrote,
|
||||
/// in correct wire-format order.
|
||||
pub fn finish(self) -> Result<&'a mut [u8]> {
|
||||
Ok(&mut self.buf[self.pos..])
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user