use crate::error::{StreamError, StreamErrorKind};
use bytes::{Buf, BytesMut};
use std::marker::PhantomData;
#[derive(Clone, Debug)]
pub struct ProtobufLenPrefixCodec<T> {
max_length: usize,
cursor: ProtobufCursor,
_ph: PhantomData<fn() -> T>,
}
#[derive(Clone, Debug)]
struct ProtobufCursor {
expected_len: Option<usize>,
}
impl<T> ProtobufLenPrefixCodec<T> {
pub fn new_with_max_length(max_length: usize) -> Self {
let initial_cursor = ProtobufCursor { expected_len: None };
ProtobufLenPrefixCodec {
max_length,
cursor: initial_cursor,
_ph: PhantomData,
}
}
}
impl<T> tokio_util::codec::Decoder for ProtobufLenPrefixCodec<T>
where
T: prost::Message + Default,
{
type Item = Result<T, StreamError>;
type Error = StreamError;
fn decode(&mut self, buf: &mut BytesMut) -> Result<Option<Self::Item>, StreamError> {
loop {
let Some(expected_len) = self.cursor.expected_len else {
if buf.is_empty() {
return Ok(None);
}
match read_varint(buf)? {
Some(len) => {
self.cursor.expected_len = Some(len as usize);
continue;
}
None => return Ok(None),
}
};
if expected_len > self.max_length {
return Err(StreamError::new(
StreamErrorKind::MaxLenReachedError,
None,
Some("Max object length reached".into()),
));
}
if buf.len() < expected_len {
return Ok(None);
}
let obj_bytes = buf.copy_to_bytes(expected_len);
self.cursor.expected_len = None;
return prost::Message::decode(obj_bytes)
.map(|item| Some(Ok(item)))
.map_err(|err| {
StreamError::new(StreamErrorKind::CodecError, Some(Box::new(err)), None)
});
}
}
fn decode_eof(&mut self, buf: &mut BytesMut) -> Result<Option<Self::Item>, StreamError> {
self.decode(buf)
}
}
fn read_varint(buf: &mut BytesMut) -> Result<Option<u64>, StreamError> {
let bytes = buf.chunk();
if bytes.is_empty() {
return Ok(None);
}
if bytes[0] < 0x80 {
let value = u64::from(bytes[0]);
buf.advance(1);
return Ok(Some(value));
}
if bytes.len() > 10 || bytes[bytes.len() - 1] < 0x80 {
let (value, advance) = decode_varint_slice(bytes)?;
buf.advance(advance);
return Ok(Some(value));
}
Ok(None)
}
#[inline]
fn decode_varint_slice(bytes: &[u8]) -> Result<(u64, usize), StreamError> {
assert!(!bytes.is_empty());
assert!(bytes.len() > 10 || bytes[bytes.len() - 1] < 0x80);
let mut b: u8 = bytes[0];
let mut part0: u32 = u32::from(b);
if b < 0x80 {
return Ok((u64::from(part0), 1));
};
part0 -= 0x80;
b = bytes[1];
part0 += u32::from(b) << 7;
if b < 0x80 {
return Ok((u64::from(part0), 2));
};
part0 -= 0x80 << 7;
b = bytes[2];
part0 += u32::from(b) << 14;
if b < 0x80 {
return Ok((u64::from(part0), 3));
};
part0 -= 0x80 << 14;
b = bytes[3];
part0 += u32::from(b) << 21;
if b < 0x80 {
return Ok((u64::from(part0), 4));
};
part0 -= 0x80 << 21;
let value = u64::from(part0);
b = bytes[4];
let mut part1: u32 = u32::from(b);
if b < 0x80 {
return Ok((value + (u64::from(part1) << 28), 5));
};
part1 -= 0x80;
b = bytes[5];
part1 += u32::from(b) << 7;
if b < 0x80 {
return Ok((value + (u64::from(part1) << 28), 6));
};
part1 -= 0x80 << 7;
b = bytes[6];
part1 += u32::from(b) << 14;
if b < 0x80 {
return Ok((value + (u64::from(part1) << 28), 7));
};
part1 -= 0x80 << 14;
b = bytes[7];
part1 += u32::from(b) << 21;
if b < 0x80 {
return Ok((value + (u64::from(part1) << 28), 8));
};
part1 -= 0x80 << 21;
let value = value + ((u64::from(part1)) << 28);
b = bytes[8];
let mut part2: u32 = u32::from(b);
if b < 0x80 {
return Ok((value + (u64::from(part2) << 56), 9));
};
part2 -= 0x80;
b = bytes[9];
part2 += u32::from(b) << 7;
if b < 0x02 {
return Ok((value + (u64::from(part2) << 56), 10));
};
Err(StreamError::new(
StreamErrorKind::CodecError,
None,
Some("invalid varint".into()),
))
}