use std::fmt;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct WireError {
pub tag: u32,
pub reason: &'static str,
}
impl WireError {
fn new(tag: u32, reason: &'static str) -> Self {
Self { tag, reason }
}
}
impl fmt::Display for WireError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "protobuf field {}: {}", self.tag, self.reason)
}
}
impl std::error::Error for WireError {}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum Raw<'a> {
Varint(u64),
Fixed64(u64),
Bytes(&'a [u8]),
Fixed32(u32),
Group(&'a [u8]),
}
impl Raw<'_> {
fn wire(self) -> u8 {
match self {
Self::Varint(_) => 0,
Self::Fixed64(_) => 1,
Self::Bytes(_) => 2,
Self::Group(_) => 3,
Self::Fixed32(_) => 5,
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct FieldValue<'a> {
pub tag: u32,
pub value: Raw<'a>,
}
pub fn varint(bytes: &mut &[u8]) -> Result<u64, WireError> {
let mut value = 0u64;
for i in 0..10 {
let (&byte, rest) = bytes
.split_first()
.ok_or(WireError::new(0, "truncated varint"))?;
*bytes = rest;
if i == 9 && byte > 1 {
return Err(WireError::new(0, "varint overflow"));
}
value |= u64::from(byte & 127) << (i * 7);
if byte < 128 {
return Ok(value);
}
}
Err(WireError::new(0, "varint overflow"))
}
fn take<'a>(bytes: &mut &'a [u8], n: usize) -> Result<&'a [u8], WireError> {
if n > bytes.len() {
return Err(WireError::new(0, "truncated field"));
}
let (value, rest) = bytes.split_at(n);
*bytes = rest;
Ok(value)
}
#[derive(Clone, Copy)]
pub struct Fields<'a> {
bytes: &'a [u8],
failed: bool,
}
impl<'a> Fields<'a> {
pub fn new(bytes: &'a [u8]) -> Self {
Self {
bytes,
failed: false,
}
}
}
impl<'a> Iterator for Fields<'a> {
type Item = Result<FieldValue<'a>, WireError>;
fn next(&mut self) -> Option<Self::Item> {
if self.failed || self.bytes.is_empty() {
return None;
}
let result = (|| {
let key = varint(&mut self.bytes)?;
let tag = u32::try_from(key >> 3).map_err(|_| WireError::new(0, "invalid tag"))?;
if tag == 0 || tag > 0x1fff_ffff {
return Err(WireError::new(tag, "invalid tag"));
}
let value = match key & 7 {
0 => Raw::Varint(varint(&mut self.bytes)?),
1 => Raw::Fixed64(u64::from_le_bytes(
take(&mut self.bytes, 8)?.try_into().unwrap(),
)),
2 => {
let n = usize::try_from(varint(&mut self.bytes)?)
.map_err(|_| WireError::new(tag, "length overflow"))?;
Raw::Bytes(take(&mut self.bytes, n)?)
}
3 => {
let begin = self.bytes;
let mut depth = 1usize;
let mut groups = [0u32; 100];
groups[0] = tag;
while depth > 0 {
let group_key = varint(&mut self.bytes)?;
let number = u32::try_from(group_key >> 3)
.map_err(|_| WireError::new(tag, "invalid group tag"))?;
if number == 0 || number > 0x1fff_ffff {
return Err(WireError::new(tag, "invalid group tag"));
}
match group_key & 7 {
0 => {
varint(&mut self.bytes)?;
}
1 => {
take(&mut self.bytes, 8)?;
}
2 => {
let n = usize::try_from(varint(&mut self.bytes)?)
.map_err(|_| WireError::new(tag, "length overflow"))?;
take(&mut self.bytes, n)?;
}
3 => {
if depth == 100 {
return Err(WireError::new(tag, "recursion limit"));
}
groups[depth] = number;
depth += 1;
}
4 => {
depth -= 1;
if groups[depth] != number {
return Err(WireError::new(tag, "group end mismatch"));
}
}
5 => {
take(&mut self.bytes, 4)?;
}
_ => return Err(WireError::new(tag, "invalid group wire type")),
}
}
Raw::Group(&begin[..begin.len() - self.bytes.len()])
}
5 => Raw::Fixed32(u32::from_le_bytes(
take(&mut self.bytes, 4)?.try_into().unwrap(),
)),
_ => return Err(WireError::new(tag, "invalid wire type")),
};
Ok(FieldValue { tag, value })
})();
if result.is_err() {
self.failed = true;
}
Some(result)
}
}
#[derive(Clone, Copy)]
pub enum Kind {
Varint,
Fixed64,
Fixed32,
Text,
Bytes,
Message(&'static Schema),
}
impl Kind {
fn wire(self) -> u8 {
match self {
Self::Varint => 0,
Self::Fixed64 => 1,
Self::Fixed32 => 5,
Self::Text | Self::Bytes | Self::Message(_) => 2,
}
}
}
pub struct Field {
pub tag: u32,
pub kind: Kind,
pub repeated: bool,
}
pub struct Schema {
pub name: &'static str,
pub fields: &'static [Field],
}
impl Schema {
pub fn validate(&self, bytes: &[u8]) -> Result<(), WireError> {
self.validate_depth(bytes, 0)
}
fn validate_depth(&self, bytes: &[u8], depth: usize) -> Result<(), WireError> {
if depth >= 100 {
return Err(WireError::new(0, "recursion limit"));
}
for value in Fields::new(bytes) {
let value = value?;
let Some(field) = self.fields.iter().find(|field| field.tag == value.tag) else {
continue;
};
if value.value.wire() != field.kind.wire() {
if field.repeated
&& matches!(field.kind, Kind::Varint | Kind::Fixed32 | Kind::Fixed64)
{
if let Raw::Bytes(mut packed) = value.value {
while !packed.is_empty() {
match field.kind {
Kind::Varint => {
varint(&mut packed)?;
}
Kind::Fixed32 => {
take(&mut packed, 4)?;
}
Kind::Fixed64 => {
take(&mut packed, 8)?;
}
_ => unreachable!(),
}
}
continue;
}
}
return Err(WireError::new(value.tag, "wire type mismatch"));
}
match (field.kind, value.value) {
(Kind::Text, Raw::Bytes(bytes)) => {
std::str::from_utf8(bytes)
.map_err(|_| WireError::new(value.tag, "invalid UTF-8"))?;
}
(Kind::Message(schema), Raw::Bytes(bytes)) => {
schema.validate_depth(bytes, depth + 1)?
}
_ => {}
}
}
Ok(())
}
}
#[derive(Clone, Copy)]
struct Selection {
tag: u32,
index: Option<usize>,
after: usize,
}
const EMPTY: Selection = Selection {
tag: 0,
index: None,
after: 0,
};
#[derive(Clone, Copy)]
pub struct View<'a> {
bytes: &'a [u8],
path: [Selection; 100],
depth: usize,
}
impl<'a> View<'a> {
pub fn new(bytes: &'a [u8]) -> Self {
Self {
bytes,
path: [EMPTY; 100],
depth: 0,
}
}
pub fn fields(self) -> ViewFields<'a> {
ViewFields {
view: self,
stack: [Fields::new(&[]); 101],
counts: [0; 100],
matches: [0; 100],
level: 0,
started: false,
done: false,
}
}
pub fn last(self, tag: u32) -> Option<Raw<'a>> {
self.fields()
.filter(|v| v.tag == tag)
.last()
.map(|v| v.value)
}
pub fn child(self, tag: u32, index: Option<usize>) -> Self {
self.child_after(tag, index, 0)
}
pub fn child_after(mut self, tag: u32, index: Option<usize>, after: usize) -> Self {
assert!(self.depth < 100, "validated schema nesting");
self.path[self.depth] = Selection { tag, index, after };
self.depth += 1;
self
}
pub fn messages(self, tag: u32) -> impl Iterator<Item = Self> {
let count = self.fields().filter(move |v| v.tag == tag).count();
(0..count).map(move |index| self.child(tag, Some(index)))
}
pub fn last_oneof(self, tags: &[u32]) -> Option<(u32, usize)> {
let selected = self.fields().filter(|v| tags.contains(&v.tag)).last()?.tag;
let after = self
.fields()
.enumerate()
.filter(|(_, v)| tags.contains(&v.tag) && v.tag != selected)
.last()
.map_or(0, |(index, _)| index + 1);
Some((selected, after))
}
pub fn values(self, tag: u32, kind: Kind) -> Values<'a> {
Values {
fields: self.fields(),
tag,
kind,
packed: &[],
}
}
}
pub struct ViewFields<'a> {
view: View<'a>,
stack: [Fields<'a>; 101],
counts: [usize; 100],
matches: [usize; 100],
level: usize,
started: bool,
done: bool,
}
impl<'a> Iterator for ViewFields<'a> {
type Item = FieldValue<'a>;
fn next(&mut self) -> Option<Self::Item> {
if self.done {
return None;
}
if !self.started {
self.stack[0] = Fields::new(self.view.bytes);
self.started = true;
}
loop {
if let Some(value) = self.stack[self.level].next() {
let value = value.expect("view constructed only from validated wire");
if self.level == self.view.depth {
return Some(value);
}
let selection = self.view.path[self.level];
let ordinal = self.counts[self.level];
self.counts[self.level] += 1;
if value.tag == selection.tag && ordinal >= selection.after {
let matched = self.matches[self.level];
self.matches[self.level] += 1;
if selection.index.is_none() || selection.index == Some(matched) {
if let Raw::Bytes(bytes) = value.value {
self.level += 1;
self.stack[self.level] = Fields::new(bytes);
}
}
}
} else if self.level == 0 {
self.done = true;
return None;
} else {
self.level -= 1;
}
}
}
}
pub struct Values<'a> {
fields: ViewFields<'a>,
tag: u32,
kind: Kind,
packed: &'a [u8],
}
impl<'a> Iterator for Values<'a> {
type Item = Raw<'a>;
fn next(&mut self) -> Option<Self::Item> {
loop {
if !self.packed.is_empty() {
return Some(match self.kind {
Kind::Varint => Raw::Varint(varint(&mut self.packed).unwrap()),
Kind::Fixed32 => Raw::Fixed32(u32::from_le_bytes(
take(&mut self.packed, 4).unwrap().try_into().unwrap(),
)),
Kind::Fixed64 => Raw::Fixed64(u64::from_le_bytes(
take(&mut self.packed, 8).unwrap().try_into().unwrap(),
)),
_ => unreachable!(),
});
}
let value = self.fields.by_ref().find(|v| v.tag == self.tag)?.value;
if let Raw::Bytes(bytes) = value {
if matches!(self.kind, Kind::Varint | Kind::Fixed32 | Kind::Fixed64) {
self.packed = bytes;
continue;
}
}
return Some(value);
}
}
}
pub fn varint_len(mut value: u64) -> usize {
let mut n = 1;
while value >= 128 {
value >>= 7;
n += 1;
}
n
}
pub fn field_len(tag: u32, value: Raw<'_>) -> usize {
varint_len((u64::from(tag) << 3) | u64::from(value.wire()))
+ match value {
Raw::Varint(n) => varint_len(n),
Raw::Fixed64(_) => 8,
Raw::Fixed32(_) => 4,
Raw::Bytes(bytes) => varint_len(bytes.len() as u64) + bytes.len(),
Raw::Group(bytes) => bytes.len(),
}
}
pub struct Writer<'a> {
bytes: &'a mut [u8],
at: usize,
}
impl<'a> Writer<'a> {
pub fn new(bytes: &'a mut [u8]) -> Self {
Self { bytes, at: 0 }
}
fn bytes(&mut self, bytes: &[u8]) -> Result<(), WireError> {
let end = self
.at
.checked_add(bytes.len())
.ok_or(WireError::new(0, "encoded length overflow"))?;
let target = self
.bytes
.get_mut(self.at..end)
.ok_or(WireError::new(0, "encode capacity"))?;
target.copy_from_slice(bytes);
self.at = end;
Ok(())
}
pub fn varint(&mut self, mut n: u64) -> Result<(), WireError> {
while n >= 128 {
self.bytes(&[((n & 127) | 128) as u8])?;
n >>= 7;
}
self.bytes(&[n as u8])
}
pub fn field(&mut self, tag: u32, value: Raw<'_>) -> Result<(), WireError> {
self.varint((u64::from(tag) << 3) | u64::from(value.wire()))?;
match value {
Raw::Varint(n) => self.varint(n),
Raw::Fixed64(n) => self.bytes(&n.to_le_bytes()),
Raw::Fixed32(n) => self.bytes(&n.to_le_bytes()),
Raw::Bytes(bytes) => {
self.varint(bytes.len() as u64)?;
self.bytes(bytes)
}
Raw::Group(_) => Err(WireError::new(tag, "generated group unsupported")),
}
}
pub fn message(&mut self, tag: u32, message: &impl Encode) -> Result<(), WireError> {
self.varint((u64::from(tag) << 3) | 2)?;
self.varint(message.encoded_len()? as u64)?;
message.encode(self)
}
pub fn finish(&self) -> Result<(), WireError> {
if self.at == self.bytes.len() {
Ok(())
} else {
Err(WireError::new(0, "encoded length mismatch"))
}
}
}
#[doc(hidden)]
pub trait Decode<'a>: Sized {
fn parse(bytes: &'a [u8]) -> Result<Self, WireError>;
}
pub trait Encode {
fn encoded_len(&self) -> Result<usize, WireError>;
fn encode(&self, writer: &mut Writer<'_>) -> Result<(), WireError>;
}
pub fn message_len(tag: u32, message: &impl Encode) -> Result<usize, WireError> {
let n = message.encoded_len()?;
varint_len((u64::from(tag) << 3) | 2)
.checked_add(varint_len(n as u64))
.and_then(|v| v.checked_add(n))
.ok_or(WireError::new(tag, "encoded length overflow"))
}
#[cfg(test)]
mod tests {
use super::*;
use prost::Message;
static CHILD: Schema = Schema {
name: "Child",
fields: &[
Field {
tag: 1,
kind: Kind::Varint,
repeated: false,
},
Field {
tag: 2,
kind: Kind::Text,
repeated: false,
},
],
};
static ROOT: Schema = Schema {
name: "Root",
fields: &[
Field {
tag: 1,
kind: Kind::Varint,
repeated: true,
},
Field {
tag: 2,
kind: Kind::Message(&CHILD),
repeated: false,
},
Field {
tag: 3,
kind: Kind::Message(&CHILD),
repeated: true,
},
],
};
#[derive(PartialEq, Message)]
struct Child {
#[prost(uint64, tag = "1")]
a: u64,
#[prost(string, tag = "2")]
b: String,
}
#[derive(PartialEq, Message)]
struct Root {
#[prost(uint64, repeated, tag = "1")]
values: Vec<u64>,
#[prost(message, optional, tag = "2")]
child: Option<Child>,
#[prost(message, repeated, tag = "3")]
children: Vec<Child>,
}
#[test]
fn views_preserve_standard_packed_unpacked_and_message_merge() {
let bytes=b"\x08\x01\x0a\x02\x02\x03\x12\x02\x08\x07\x12\x03\x12\x01x\x1a\x02\x08\x08\x1a\x03\x12\x01y";
ROOT.validate(bytes).unwrap();
let decoded = Root::decode(bytes.as_slice()).unwrap();
let view = View::new(bytes);
assert_eq!(
view.values(1, Kind::Varint).collect::<Vec<_>>(),
decoded
.values
.iter()
.copied()
.map(Raw::Varint)
.collect::<Vec<_>>()
);
let child = view.child(2, None);
assert_eq!(
child.last(1),
Some(Raw::Varint(decoded.child.as_ref().unwrap().a))
);
assert_eq!(
child.last(2),
Some(Raw::Bytes(decoded.child.as_ref().unwrap().b.as_bytes()))
);
let children = view.messages(3).collect::<Vec<_>>();
assert_eq!(children.len(), decoded.children.len());
assert_eq!(children[0].last(1), Some(Raw::Varint(8)));
assert_eq!(children[1].last(1), None);
assert_eq!(children[1].last(2), Some(Raw::Bytes(b"y")));
}
#[test]
fn nested_message_fragments_keep_global_selection_counters() {
static OUTER: Schema = Schema {
name: "Outer",
fields: &[Field {
tag: 1,
kind: Kind::Message(&ROOT),
repeated: false,
}],
};
let bytes = b"\x0a\x05\x1a\x03\x12\x01x\x0a\x04\x1a\x02\x08\x07";
OUTER.validate(bytes).unwrap();
let nested = View::new(bytes)
.child(1, None)
.messages(3)
.collect::<Vec<_>>();
assert_eq!(nested[0].last(2), Some(Raw::Bytes(b"x")));
assert_eq!(nested[1].last(1), Some(Raw::Varint(7)));
assert_eq!(nested[1].last(2), None);
}
#[test]
fn malformed_fields_reject_and_unknown_fields_follow_standard_behavior() {
for bytes in [
b"\x00".as_slice(),
b"\x12\x05x",
b"\x0a\x01\x80",
b"\x12\x03\x12\x01\xff",
b"\x0d\0\0\0\0",
] {
assert!(ROOT.validate(bytes).is_err());
assert!(Root::decode(bytes).is_err());
}
let unknown = b"\x98\x06\x01\xa3\x06\x08\x02\xa4\x06";
ROOT.validate(unknown).unwrap();
assert_eq!(Root::decode(unknown.as_slice()).unwrap(), Root::default());
}
#[test]
fn writer_is_exact_and_never_grows() {
let mut bytes = [0u8; 8];
let mut writer = Writer::new(&mut bytes);
writer.field(1, Raw::Varint(150)).unwrap();
writer.field(2, Raw::Bytes(b"abc")).unwrap();
writer.finish().unwrap();
assert_eq!(&bytes, b"\x08\x96\x01\x12\x03abc");
let mut short = [0u8; 1];
assert!(Writer::new(&mut short).field(1, Raw::Varint(150)).is_err());
}
}
pub struct MethodSpec {
pub dependency: &'static str,
pub contract: [u8; 32],
pub service: &'static str,
pub method: &'static str,
pub path: &'static str,
pub request: &'static Schema,
pub response: &'static Schema,
pub marker: fn() -> std::any::TypeId,
}