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..]) + } +}