#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WireError {
Truncated { at: usize, needed: usize },
VarintNotCanonical { at: usize },
LengthOutOfRange { at: usize, len: u64 },
GroupsUnsupported { at: usize, field: u32 },
UnknownWireType { at: usize, wire: u32 },
FieldNumberZero { at: usize },
}
impl core::fmt::Display for WireError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
WireError::Truncated { at, needed } => {
write!(f, "truncated at byte {at}: {needed} more byte(s) were needed")
}
WireError::VarintNotCanonical { at } => write!(
f,
"the varint at byte {at} is longer than ten bytes or overflows 64 bits, which the \
wire format cannot express -- this is malformed input, not a large number"
),
WireError::LengthOutOfRange { at, len } => write!(
f,
"the length prefix at byte {at} is {len}, which is not addressable here"
),
WireError::GroupsUnsupported { at, field } => write!(
f,
"field {field} at byte {at} uses the deprecated group encoding (wire type 3/4), \
which this reader does not implement -- refused rather than skipped, because \
skipping it would report a partial message as a whole one"
),
WireError::UnknownWireType { at, wire } => {
write!(f, "wire type {wire} at byte {at} is not one the format defines")
}
WireError::FieldNumberZero { at } => {
write!(f, "field number 0 at byte {at} is reserved and no writer may emit it")
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum Value<'a> {
Varint(u64),
Fixed64(u64),
Fixed32(u32),
Bytes(&'a [u8]),
}
impl Value<'_> {
pub fn as_f64(&self) -> Option<f64> {
match self {
Value::Fixed64(v) => Some(f64::from_bits(*v)),
_ => None,
}
}
pub fn as_f32(&self) -> Option<f32> {
match self {
Value::Fixed32(v) => Some(f32::from_bits(*v)),
_ => None,
}
}
pub fn as_u64(&self) -> Option<u64> {
match self {
Value::Varint(v) => Some(*v),
_ => None,
}
}
pub fn as_bytes(&self) -> Option<&[u8]> {
match self {
Value::Bytes(b) => Some(b),
_ => None,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Field<'a> {
pub number: u32,
pub value: Value<'a>,
}
pub struct Reader<'a> {
b: &'a [u8],
i: usize,
}
impl<'a> Reader<'a> {
pub fn new(b: &'a [u8]) -> Self {
Reader { b, i: 0 }
}
pub fn position(&self) -> usize {
self.i
}
pub fn read_field(&mut self) -> Result<Option<Field<'a>>, WireError> {
if self.i >= self.b.len() {
return Ok(None);
}
let at = self.i;
let (key, used) = read_varint(self.b, self.i)?;
self.i += used;
let number = (key >> 3) as u32;
let wire = (key & 7) as u32;
if number == 0 {
return Err(WireError::FieldNumberZero { at });
}
let value = match wire {
0 => {
let (v, u) = read_varint(self.b, self.i)?;
self.i += u;
Value::Varint(v)
}
1 => Value::Fixed64(u64::from_le_bytes(self.take_array::<8>()?)),
2 => {
let lat = self.i;
let (len, u) = read_varint(self.b, self.i)?;
self.i += u;
let len = usize::try_from(len)
.map_err(|_| WireError::LengthOutOfRange { at: lat, len })?;
let end = self
.i
.checked_add(len)
.ok_or(WireError::LengthOutOfRange { at: lat, len: len as u64 })?;
if end > self.b.len() {
return Err(WireError::Truncated { at: lat, needed: end - self.b.len() });
}
let s = &self.b[self.i..end];
self.i = end;
Value::Bytes(s)
}
3 | 4 => return Err(WireError::GroupsUnsupported { at, field: number }),
5 => Value::Fixed32(u32::from_le_bytes(self.take_array::<4>()?)),
other => return Err(WireError::UnknownWireType { at, wire: other }),
};
Ok(Some(Field { number, value }))
}
fn take_array<const N: usize>(&mut self) -> Result<[u8; N], WireError> {
let end = self
.i
.checked_add(N)
.ok_or(WireError::Truncated { at: self.i, needed: N })?;
if end > self.b.len() {
return Err(WireError::Truncated { at: self.i, needed: end - self.b.len() });
}
let mut out = [0u8; N];
out.copy_from_slice(&self.b[self.i..end]);
self.i = end;
Ok(out)
}
}
fn read_varint(b: &[u8], start: usize) -> Result<(u64, usize), WireError> {
let mut v: u64 = 0;
let mut i = start;
for byte_index in 0..10usize {
let byte = *b.get(i).ok_or(WireError::Truncated { at: start, needed: 1 })?;
i += 1;
if byte_index == 9 {
if byte > 0x01 {
return Err(WireError::VarintNotCanonical { at: start });
}
v |= (byte as u64) << 63;
return Ok((v, i - start));
}
v |= ((byte & 0x7f) as u64) << (7 * byte_index);
if byte & 0x80 == 0 {
return Ok((v, i - start));
}
}
Err(WireError::VarintNotCanonical { at: start })
}
pub fn packed_varints(p: &[u8]) -> Result<Vec<u64>, WireError> {
let mut out = Vec::new();
let mut i = 0;
while i < p.len() {
let (v, u) = read_varint(p, i)?;
out.push(v);
i += u;
}
Ok(out)
}
pub fn packed_doubles(p: &[u8]) -> Result<Vec<f64>, WireError> {
if !p.len().is_multiple_of(8) {
return Err(WireError::Truncated { at: p.len() - (p.len() % 8), needed: 8 - (p.len() % 8) });
}
Ok(p.as_chunks::<8>().0.iter().map(|&c| f64::from_bits(u64::from_le_bytes(c))).collect())
}
pub fn put_key(buf: &mut Vec<u8>, field: u32, wire: u32) {
put_varint(buf, ((field as u64) << 3) | wire as u64);
}
pub fn put_varint(buf: &mut Vec<u8>, mut v: u64) {
loop {
let byte = (v & 0x7f) as u8;
v >>= 7;
if v == 0 {
buf.push(byte);
return;
}
buf.push(byte | 0x80);
}
}
pub fn put_varint_field(buf: &mut Vec<u8>, field: u32, v: u64) {
put_key(buf, field, 0);
put_varint(buf, v);
}
pub fn put_double_field(buf: &mut Vec<u8>, field: u32, v: f64) {
put_key(buf, field, 1);
buf.extend_from_slice(&v.to_bits().to_le_bytes());
}
pub fn put_len_field(buf: &mut Vec<u8>, field: u32, body: &[u8]) {
put_key(buf, field, 2);
put_varint(buf, body.len() as u64);
buf.extend_from_slice(body);
}
pub fn put_str_field(buf: &mut Vec<u8>, field: u32, s: &str) {
put_len_field(buf, field, s.as_bytes());
}
#[cfg(test)]
mod tests {
use super::*;
fn all(b: &[u8]) -> Result<Vec<Field<'_>>, WireError> {
let mut r = Reader::new(b);
let mut out = Vec::new();
while let Some(f) = r.read_field()? {
out.push(f);
}
Ok(out)
}
#[test]
fn an_end_and_a_corruption_are_different_answers() {
let good = {
let mut b = Vec::new();
put_varint_field(&mut b, 1, 150);
b
};
assert_eq!(all(&good).unwrap().len(), 1, "a whole message reads");
let cut = &good[..good.len() - 1];
match all(cut) {
Err(WireError::Truncated { .. }) => {}
other => panic!("a truncated message must be an error, not a shorter one: {other:?}"),
}
}
#[test]
fn fixed32_carries_its_value_instead_of_zero() {
let mut b = Vec::new();
put_key(&mut b, 7, 5);
b.extend_from_slice(&1.5f32.to_bits().to_le_bytes());
let fields = all(&b).unwrap();
assert_eq!(fields.len(), 1);
assert_eq!(fields[0].number, 7);
assert_eq!(fields[0].value.as_f32(), Some(1.5), "the value, not zero");
let short = &b[..b.len() - 1];
assert!(matches!(all(short), Err(WireError::Truncated { .. })));
}
#[test]
fn a_hostile_length_prefix_is_named_rather_than_wrapping() {
let bytes = [0x0A, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0x01];
match all(&bytes) {
Err(WireError::LengthOutOfRange { .. }) | Err(WireError::Truncated { .. }) => {}
other => panic!("expected a named refusal, got {other:?}"),
}
}
#[test]
fn a_varint_longer_than_the_format_allows_is_malformed_not_large() {
let mut b = vec![0x08];
b.extend(std::iter::repeat_n(0xFFu8, 11));
assert!(matches!(all(&b), Err(WireError::VarintNotCanonical { .. })));
let mut c = vec![0x08];
c.extend(std::iter::repeat_n(0xFFu8, 9));
c.push(0x02);
assert!(matches!(all(&c), Err(WireError::VarintNotCanonical { .. })));
let mut d = Vec::new();
put_varint_field(&mut d, 1, u64::MAX);
assert_eq!(all(&d).unwrap()[0].value.as_u64(), Some(u64::MAX));
}
#[test]
fn groups_and_impossible_wire_types_are_refused_by_name() {
let mut b = Vec::new();
put_key(&mut b, 4, 3);
assert!(matches!(all(&b), Err(WireError::GroupsUnsupported { field: 4, .. })));
let mut c = Vec::new();
put_key(&mut c, 4, 6);
assert!(matches!(all(&c), Err(WireError::UnknownWireType { wire: 6, .. })));
}
#[test]
fn field_number_zero_is_refused() {
let b = [0x00u8, 0x01];
assert!(matches!(all(&b), Err(WireError::FieldNumberZero { .. })));
}
#[test]
fn a_packed_array_refuses_rather_than_returning_the_prefix_it_managed() {
let mut p = Vec::new();
for v in [1u64, 300, 70000] {
put_varint(&mut p, v);
}
assert_eq!(packed_varints(&p).unwrap(), vec![1, 300, 70000]);
p.push(0x80); match packed_varints(&p) {
Err(WireError::Truncated { .. }) => {}
other => panic!("a corrupt packed array must not become a shorter valid one: {other:?}"),
}
let mut d = Vec::new();
d.extend_from_slice(&1.0f64.to_bits().to_le_bytes());
assert_eq!(packed_doubles(&d).unwrap(), vec![1.0]);
d.push(0x00);
assert!(matches!(packed_doubles(&d), Err(WireError::Truncated { .. })));
}
#[test]
fn every_wire_type_round_trips_through_this_module() {
let mut b = Vec::new();
put_varint_field(&mut b, 1, 0);
put_varint_field(&mut b, 2, u64::MAX);
put_double_field(&mut b, 3, -1.5);
put_str_field(&mut b, 4, "hello");
put_len_field(&mut b, 5, &[]);
put_key(&mut b, 6, 5);
b.extend_from_slice(&2.25f32.to_bits().to_le_bytes());
let f = all(&b).unwrap();
assert_eq!(f.len(), 6);
assert_eq!(f[0].value.as_u64(), Some(0));
assert_eq!(f[1].value.as_u64(), Some(u64::MAX));
assert_eq!(f[2].value.as_f64(), Some(-1.5));
assert_eq!(f[3].value.as_bytes(), Some(&b"hello"[..]));
assert_eq!(f[4].value.as_bytes(), Some(&[][..]));
assert_eq!(f[5].value.as_f32(), Some(2.25));
}
#[test]
fn an_unknown_field_number_is_skipped_but_an_unknown_wire_type_is_not() {
let mut b = Vec::new();
put_varint_field(&mut b, 1, 7);
put_str_field(&mut b, 9999, "a field from a later version");
put_varint_field(&mut b, 2, 8);
let f = all(&b).unwrap();
assert_eq!(f.len(), 3, "the reader hands all three up; the caller ignores 9999");
assert_eq!(f[2].value.as_u64(), Some(8), "and keeps reading past it");
}
}