From 229727340d591015879974b1540f77c1a55ed132 Mon Sep 17 00:00:00 2001 From: charles Date: Fri, 19 Jun 2026 19:39:21 -0700 Subject: [PATCH] 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. --- codegen/src/generator/mod.rs | 55 +++- codegen/src/generator/utils.rs | 2 +- roto-tonic/src/generated/interop.rs | 333 +++++++++++++++--------- runtime/src/lib.rs | 388 +++++++++++++++++++++++++++- 4 files changed, 644 insertions(+), 134 deletions(-) diff --git a/codegen/src/generator/mod.rs b/codegen/src/generator/mod.rs index f8bb17a..ea5de5b 100644 --- a/codegen/src/generator/mod.rs +++ b/codegen/src/generator/mod.rs @@ -8,7 +8,7 @@ use std::collections::{HashMap, HashSet}; use std::str; const DATA_IMPORTS: &str = "#[allow(unused, unused_imports, unused_assignments, unused_variables, non_camel_case_types)]\n\ -use roto_runtime::{ProtoAccessor, ProtoBuilder, Result, RotoError, read_varint, RepeatedFieldIterator, RotoMessage};\n\ +use roto_runtime::{ProtoAccessor, ProtoBuilder, RevBuilder, Result, RotoError, read_varint, RepeatedFieldIterator, RotoMessage};\n\ use core::str;\n\ use bytes::{Bytes, BytesMut, Buf, BufMut};\n"; const SERVICE_IMPORTS: &str = "#[allow(unused, unused_imports, unused_assignments, unused_variables, non_camel_case_types)]\n\ @@ -460,6 +460,59 @@ fn write_message(msg_proto: &DescriptorProto, output: &mut String) { output.push_str(&format!(" pub fn finish(self) -> roto_runtime::Result<&'b mut [u8]> {{\n self.builder.finish()\n }}\n}}\n\n")); + // RevBuilder struct — reverse encoding, one `_written: bool` flag per field + output.push_str(&format!("pub struct {}RevBuilder<'b> {{\n", msg_name)); + output.push_str(" builder: roto_runtime::RevBuilder<'b>,\n"); + for (field_name, _, _, _, _) in &builder_fields { + output.push_str(&format!(" {}_written: bool,\n", field_name)); + } + output.push_str(&format!("}}\n\nimpl<'b> {}RevBuilder<'b> {{\n", msg_name)); + + // Constructor + output.push_str(&format!( + " pub fn builder(buf: &mut [u8]) -> {}RevBuilder<'_> {{\n {}RevBuilder {{\n", + msg_name, msg_name + )); + output.push_str(" builder: roto_runtime::RevBuilder::new(buf),\n"); + for (field_name, _, _, _, _) in &builder_fields { + output.push_str(&format!(" {}_written: false,\n", field_name)); + } + output.push_str(" }\n }\n\n"); + + // Per-field setters — same as ProtoBuilder but using RevBuilder methods + for (field_name, safe_name, tag, rust_type, method) in &builder_fields { + output.push_str(&format!( + " pub fn {}(mut self, value: {}) -> roto_runtime::Result {{\n self.builder.{}({}, value)?;\n self.{}_written = true;\n Ok(self)\n }}\n\n", + safe_name, rust_type, method, tag, field_name + )); + } + + // with() — copies unseen fields from an existing message (iterate in reverse for RevBuilder) + output.push_str(&format!( + " pub fn with(mut self, msg: &{}<'_>) -> roto_runtime::Result {{\n", + msg_name + )); + output.push_str(" let fields: Vec<_> = msg.accessor.raw_fields().collect();\n"); + output.push_str(" for item in fields.into_iter().rev() {\n"); + output.push_str(" let (field_number, raw_bytes) = item?;\n"); + output.push_str(" let is_written = match field_number {\n"); + for (field_name, _, tag, _, _) in &builder_fields { + output.push_str(&format!( + " {} => self.{}_written,\n", + tag, field_name + )); + } + output.push_str(" _ => false,\n"); + output.push_str(" };\n"); + output.push_str(" if !is_written {\n"); + output.push_str(" self.builder.write_raw(raw_bytes)?;\n"); + output.push_str(" }\n"); + output.push_str(" }\n"); + output.push_str(" Ok(self)\n"); + output.push_str(" }\n\n"); + + output.push_str(&format!(" pub fn finish(self) -> roto_runtime::Result<&'b mut [u8]> {{\n self.builder.finish()\n }}\n}}\n\n")); + output.push_str(&format!("pub struct Owned{} {{\n", msg_name)); output.push_str(" pub data: bytes::Bytes,\n"); output.push_str("}\n\n"); diff --git a/codegen/src/generator/utils.rs b/codegen/src/generator/utils.rs index c48febb..7733804 100644 --- a/codegen/src/generator/utils.rs +++ b/codegen/src/generator/utils.rs @@ -1,4 +1,4 @@ -pub const DATA_IMPORTS: &str = "use roto_runtime::{ProtoAccessor, ProtoBuilder, Result, RotoError, read_varint, RepeatedFieldIterator, RotoMessage};\nuse core::str;\nuse bytes::{Bytes, BytesMut, Buf, BufMut};\n"; +pub const DATA_IMPORTS: &str = "use roto_runtime::{ProtoAccessor, ProtoBuilder, RevBuilder, Result, RotoError, read_varint, RepeatedFieldIterator, RotoMessage};\nuse core::str;\nuse bytes::{Bytes, BytesMut, Buf, BufMut};\n"; pub fn to_pascal_case(s: &str) -> String { s.split('_') diff --git a/roto-tonic/src/generated/interop.rs b/roto-tonic/src/generated/interop.rs index f71f4a5..f57d869 100644 --- a/roto-tonic/src/generated/interop.rs +++ b/roto-tonic/src/generated/interop.rs @@ -1,16 +1,8 @@ // @generated by protoc-gen-roto — do not edit -#![allow( - unused, - unused_imports, - unused_assignments, - unused_variables, - non_camel_case_types -)] -use bytes::{Buf, BufMut, Bytes, BytesMut}; +#[allow(unused, unused_imports, unused_assignments, unused_variables, non_camel_case_types)] +use roto_runtime::{ProtoAccessor, ProtoBuilder, RevBuilder, Result, RotoError, read_varint, RepeatedFieldIterator, RotoMessage}; use core::str; -use roto_runtime::{ - ProtoAccessor, ProtoBuilder, RepeatedFieldIterator, Result, RotoError, RotoMessage, read_varint, -}; +use bytes::{Bytes, BytesMut, Buf, BufMut}; pub struct UnaryRequest<'a> { accessor: roto_runtime::ProtoAccessor<'a>, @@ -23,21 +15,17 @@ impl<'a> UnaryRequest<'a> { let mut message_offset = None; for item in accessor.fields() { let (offset, tag, _) = item?; - if tag.field_number == 1 { - message_offset = Some(offset); - } + if tag.field_number == 1 { message_offset = Some(offset); } } Ok(Self { accessor, - message_offset, +message_offset, }) } pub fn message(&self) -> roto_runtime::Result<&'a str> { - let offset = self - .message_offset - .ok_or(roto_runtime::RotoError::FieldNotFound)?; + let offset = self.message_offset.ok_or(roto_runtime::RotoError::FieldNotFound)?; let (bytes, _) = self.accessor.get_value_at(offset)?; core::str::from_utf8(bytes).map_err(|_| roto_runtime::RotoError::WireFormatViolation) } @@ -46,13 +34,12 @@ impl<'a> UnaryRequest<'a> { self.message().or(Ok("")) } - pub fn has_message(&self) -> bool { - self.message_offset.is_some() - } + pub fn has_message(&self) -> bool { self.message_offset.is_some() } pub fn raw_fields(&self) -> roto_runtime::RawFieldIterator<'a> { self.accessor.raw_fields() } + } pub struct UnaryRequestBuilder<'b> { @@ -93,6 +80,45 @@ impl<'b> UnaryRequestBuilder<'b> { } } +pub struct UnaryRequestRevBuilder<'b> { + builder: roto_runtime::RevBuilder<'b>, + message_written: bool, +} + +impl<'b> UnaryRequestRevBuilder<'b> { + pub fn builder(buf: &mut [u8]) -> UnaryRequestRevBuilder<'_> { + UnaryRequestRevBuilder { + builder: roto_runtime::RevBuilder::new(buf), + message_written: false, + } + } + + pub fn message(mut self, value: &str) -> roto_runtime::Result { + self.builder.write_string(1, value)?; + self.message_written = true; + Ok(self) + } + + pub fn with(mut self, msg: &UnaryRequest<'_>) -> roto_runtime::Result { + let fields: Vec<_> = msg.accessor.raw_fields().collect(); + for item in fields.into_iter().rev() { + let (field_number, raw_bytes) = item?; + let is_written = match field_number { + 1 => self.message_written, + _ => false, + }; + if !is_written { + self.builder.write_raw(raw_bytes)?; + } + } + Ok(self) + } + + pub fn finish(self) -> roto_runtime::Result<&'b mut [u8]> { + self.builder.finish() + } +} + pub struct OwnedUnaryRequest { pub data: bytes::Bytes, } @@ -125,21 +151,17 @@ impl<'a> UnaryResponse<'a> { let mut reply_offset = None; for item in accessor.fields() { let (offset, tag, _) = item?; - if tag.field_number == 1 { - reply_offset = Some(offset); - } + if tag.field_number == 1 { reply_offset = Some(offset); } } Ok(Self { accessor, - reply_offset, +reply_offset, }) } pub fn reply(&self) -> roto_runtime::Result<&'a str> { - let offset = self - .reply_offset - .ok_or(roto_runtime::RotoError::FieldNotFound)?; + let offset = self.reply_offset.ok_or(roto_runtime::RotoError::FieldNotFound)?; let (bytes, _) = self.accessor.get_value_at(offset)?; core::str::from_utf8(bytes).map_err(|_| roto_runtime::RotoError::WireFormatViolation) } @@ -148,13 +170,12 @@ impl<'a> UnaryResponse<'a> { self.reply().or(Ok("")) } - pub fn has_reply(&self) -> bool { - self.reply_offset.is_some() - } + pub fn has_reply(&self) -> bool { self.reply_offset.is_some() } pub fn raw_fields(&self) -> roto_runtime::RawFieldIterator<'a> { self.accessor.raw_fields() } + } pub struct UnaryResponseBuilder<'b> { @@ -195,6 +216,45 @@ impl<'b> UnaryResponseBuilder<'b> { } } +pub struct UnaryResponseRevBuilder<'b> { + builder: roto_runtime::RevBuilder<'b>, + reply_written: bool, +} + +impl<'b> UnaryResponseRevBuilder<'b> { + pub fn builder(buf: &mut [u8]) -> UnaryResponseRevBuilder<'_> { + UnaryResponseRevBuilder { + builder: roto_runtime::RevBuilder::new(buf), + reply_written: false, + } + } + + pub fn reply(mut self, value: &str) -> roto_runtime::Result { + self.builder.write_string(1, value)?; + self.reply_written = true; + Ok(self) + } + + pub fn with(mut self, msg: &UnaryResponse<'_>) -> roto_runtime::Result { + let fields: Vec<_> = msg.accessor.raw_fields().collect(); + for item in fields.into_iter().rev() { + let (field_number, raw_bytes) = item?; + let is_written = match field_number { + 1 => self.reply_written, + _ => false, + }; + if !is_written { + self.builder.write_raw(raw_bytes)?; + } + } + Ok(self) + } + + pub fn finish(self) -> roto_runtime::Result<&'b mut [u8]> { + self.builder.finish() + } +} + pub struct OwnedUnaryResponse { pub data: bytes::Bytes, } @@ -227,21 +287,17 @@ impl<'a> StreamingRequest<'a> { let mut query_offset = None; for item in accessor.fields() { let (offset, tag, _) = item?; - if tag.field_number == 1 { - query_offset = Some(offset); - } + if tag.field_number == 1 { query_offset = Some(offset); } } Ok(Self { accessor, - query_offset, +query_offset, }) } pub fn query(&self) -> roto_runtime::Result<&'a str> { - let offset = self - .query_offset - .ok_or(roto_runtime::RotoError::FieldNotFound)?; + let offset = self.query_offset.ok_or(roto_runtime::RotoError::FieldNotFound)?; let (bytes, _) = self.accessor.get_value_at(offset)?; core::str::from_utf8(bytes).map_err(|_| roto_runtime::RotoError::WireFormatViolation) } @@ -250,13 +306,12 @@ impl<'a> StreamingRequest<'a> { self.query().or(Ok("")) } - pub fn has_query(&self) -> bool { - self.query_offset.is_some() - } + pub fn has_query(&self) -> bool { self.query_offset.is_some() } pub fn raw_fields(&self) -> roto_runtime::RawFieldIterator<'a> { self.accessor.raw_fields() } + } pub struct StreamingRequestBuilder<'b> { @@ -297,6 +352,45 @@ impl<'b> StreamingRequestBuilder<'b> { } } +pub struct StreamingRequestRevBuilder<'b> { + builder: roto_runtime::RevBuilder<'b>, + query_written: bool, +} + +impl<'b> StreamingRequestRevBuilder<'b> { + pub fn builder(buf: &mut [u8]) -> StreamingRequestRevBuilder<'_> { + StreamingRequestRevBuilder { + builder: roto_runtime::RevBuilder::new(buf), + query_written: false, + } + } + + pub fn query(mut self, value: &str) -> roto_runtime::Result { + self.builder.write_string(1, value)?; + self.query_written = true; + Ok(self) + } + + pub fn with(mut self, msg: &StreamingRequest<'_>) -> roto_runtime::Result { + let fields: Vec<_> = msg.accessor.raw_fields().collect(); + for item in fields.into_iter().rev() { + let (field_number, raw_bytes) = item?; + let is_written = match field_number { + 1 => self.query_written, + _ => false, + }; + if !is_written { + self.builder.write_raw(raw_bytes)?; + } + } + Ok(self) + } + + pub fn finish(self) -> roto_runtime::Result<&'b mut [u8]> { + self.builder.finish() + } +} + pub struct OwnedStreamingRequest { pub data: bytes::Bytes, } @@ -329,21 +423,17 @@ impl<'a> StreamingResponse<'a> { let mut item_offset = None; for item in accessor.fields() { let (offset, tag, _) = item?; - if tag.field_number == 1 { - item_offset = Some(offset); - } + if tag.field_number == 1 { item_offset = Some(offset); } } Ok(Self { accessor, - item_offset, +item_offset, }) } pub fn item(&self) -> roto_runtime::Result<&'a str> { - let offset = self - .item_offset - .ok_or(roto_runtime::RotoError::FieldNotFound)?; + let offset = self.item_offset.ok_or(roto_runtime::RotoError::FieldNotFound)?; let (bytes, _) = self.accessor.get_value_at(offset)?; core::str::from_utf8(bytes).map_err(|_| roto_runtime::RotoError::WireFormatViolation) } @@ -352,13 +442,12 @@ impl<'a> StreamingResponse<'a> { self.item().or(Ok("")) } - pub fn has_item(&self) -> bool { - self.item_offset.is_some() - } + pub fn has_item(&self) -> bool { self.item_offset.is_some() } pub fn raw_fields(&self) -> roto_runtime::RawFieldIterator<'a> { self.accessor.raw_fields() } + } pub struct StreamingResponseBuilder<'b> { @@ -399,6 +488,45 @@ impl<'b> StreamingResponseBuilder<'b> { } } +pub struct StreamingResponseRevBuilder<'b> { + builder: roto_runtime::RevBuilder<'b>, + item_written: bool, +} + +impl<'b> StreamingResponseRevBuilder<'b> { + pub fn builder(buf: &mut [u8]) -> StreamingResponseRevBuilder<'_> { + StreamingResponseRevBuilder { + builder: roto_runtime::RevBuilder::new(buf), + item_written: false, + } + } + + pub fn item(mut self, value: &str) -> roto_runtime::Result { + self.builder.write_string(1, value)?; + self.item_written = true; + Ok(self) + } + + pub fn with(mut self, msg: &StreamingResponse<'_>) -> roto_runtime::Result { + let fields: Vec<_> = msg.accessor.raw_fields().collect(); + for item in fields.into_iter().rev() { + let (field_number, raw_bytes) = item?; + let is_written = match field_number { + 1 => self.item_written, + _ => false, + }; + if !is_written { + self.builder.write_raw(raw_bytes)?; + } + } + Ok(self) + } + + pub fn finish(self) -> roto_runtime::Result<&'b mut [u8]> { + self.builder.finish() + } +} + pub struct OwnedStreamingResponse { pub data: bytes::Bytes, } @@ -420,41 +548,26 @@ impl roto_runtime::RotoMessage for OwnedStreamingResponse { } } -use crate::{BufferPool, StatusBody}; -use futures_util::StreamExt; -use http_body::Body; -use http_body_util::BodyExt; -use std::future::Future; + + +#[allow(unused, unused_imports, unused_assignments, unused_variables, non_camel_case_types)] +use tonic::{Request, Response, Status}; +use tokio_stream::Stream; use std::pin::Pin; use std::sync::Arc; use std::task::{Context, Poll}; -use tokio_stream::Stream; +use std::future::Future; use tonic::body::BoxBody; -#[allow( - unused, - unused_imports, - unused_assignments, - unused_variables, - non_camel_case_types -)] -use tonic::{Request, Response, Status}; use tower::Service; +use futures_util::StreamExt; +use http_body_util::BodyExt; +use http_body::Body; +use crate::{BufferPool, StatusBody}; #[async_trait::async_trait] pub trait InteropService: Send + Sync + 'static { - async fn unary_call( - &self, - request: Request, - ) -> std::result::Result, Status>; - async fn streaming_call( - &self, - request: Request, - ) -> std::result::Result< - Response< - Pin> + Send>>, - >, - Status, - >; + async fn unary_call(&self, request: Request) -> std::result::Result, Status>; + async fn streaming_call(&self, request: Request) -> std::result::Result> + Send>>>, Status>; } #[derive(Clone)] @@ -476,8 +589,7 @@ impl tonic::server::NamedService for InteropServiceServer { impl Service> for InteropServiceServer { type Response = http::Response; type Error = std::convert::Infallible; - type Future = - Pin> + Send>>; + type Future = Pin> + Send>>; fn poll_ready(&mut self, _: &mut Context<'_>) -> Poll> { Poll::Ready(Ok(())) @@ -502,14 +614,8 @@ impl Service> for InteropServiceServer { let bytes_vec = buf.split_to(total_len).freeze(); pool.put(buf); if bytes_vec.len() < 5 { - let res_body = BoxBody::new(StatusBody::new( - Some(Bytes::from_static(&[0, 0, 0, 0, 0])), - 0, - )); - return Ok(http::Response::builder() - .status(200) - .body(res_body) - .unwrap()); + let res_body = BoxBody::new(StatusBody::new(Some(Bytes::from_static(&[0, 0, 0, 0, 0])), 0)); + return Ok(http::Response::builder().status(200).body(res_body).unwrap()); } let payload = bytes_vec.slice(5..); @@ -519,28 +625,16 @@ impl Service> for InteropServiceServer { let request_msg = match OwnedUnaryRequest::decode(payload) { Ok(msg) => msg, Err(_e) => { - let res_body = BoxBody::new(StatusBody::new( - Some(Bytes::from_static(&[0, 0, 0, 0, 0])), - 0, - )); - return Ok(http::Response::builder() - .status(200) - .body(res_body) - .unwrap()); + let res_body = BoxBody::new(StatusBody::new(Some(Bytes::from_static(&[0, 0, 0, 0, 0])), 0)); + return Ok(http::Response::builder().status(200).body(res_body).unwrap()); } }; let response = match inner.unary_call(Request::new(request_msg)).await { Ok(res) => res, Err(_e) => { - let res_body = BoxBody::new(StatusBody::new( - Some(Bytes::from_static(&[0, 0, 0, 0, 0])), - 0, - )); - return Ok(http::Response::builder() - .status(200) - .body(res_body) - .unwrap()); + let res_body = BoxBody::new(StatusBody::new(Some(Bytes::from_static(&[0, 0, 0, 0, 0])), 0)); + return Ok(http::Response::builder().status(200).body(res_body).unwrap()); } }; @@ -556,36 +650,17 @@ impl Service> for InteropServiceServer { pool.put(res_buf); let res_body = BoxBody::new(StatusBody::new(Some(frame), 0)); routed = true; - return Ok(http::Response::builder() - .status(200) - .header("content-type", "application/grpc") - .body(res_body) - .unwrap()); + return Ok(http::Response::builder().status(200).header("content-type", "application/grpc").body(res_body).unwrap()); } if path == "/interop.InteropService/StreamingCall" { - let res_body = BoxBody::new(StatusBody::new( - Some(Bytes::from_static(&[0, 0, 0, 0, 0])), - 0, - )); - return Ok(http::Response::builder() - .status(200) - .body(res_body) - .unwrap()); + let res_body = BoxBody::new(StatusBody::new(Some(Bytes::from_static(&[0, 0, 0, 0, 0])), 0)); + return Ok(http::Response::builder().status(200).body(res_body).unwrap()); } if !routed { - let res_body = BoxBody::new(StatusBody::new( - Some(Bytes::from_static(&[0, 0, 0, 0, 0])), - 0, - )); - return Ok(http::Response::builder() - .status(200) - .body(res_body) - .unwrap()); + let res_body = BoxBody::new(StatusBody::new(Some(Bytes::from_static(&[0, 0, 0, 0, 0])), 0)); + return Ok(http::Response::builder().status(200).body(res_body).unwrap()); } - Ok(http::Response::builder() - .status(200) - .body(BoxBody::new(StatusBody::new(None, 0))) - .unwrap()) + Ok(http::Response::builder().status(200).body(BoxBody::new(StatusBody::new(None, 0))).unwrap()) }) } } diff --git a/runtime/src/lib.rs b/runtime/src/lib.rs index 7fdacfd..b81bd0e 100644 --- a/runtime/src/lib.rs +++ b/runtime/src/lib.rs @@ -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 = core::result::Result; 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..]) + } +}