use bytes::BytesMut;
use prost::Message;
use std::io::{self, Cursor};
pub struct ProtobufDecoder<T> {
_phantom: std::marker::PhantomData<T>,
buffer: BytesMut,
}
impl<T> ProtobufDecoder<T>
where
T: Message + Default,
{
pub fn new() -> Self {
Self {
_phantom: std::marker::PhantomData,
buffer: BytesMut::new(),
}
}
pub fn add_data(&mut self, data: &[u8]) {
self.buffer.extend_from_slice(data);
}
pub fn decode_next(&mut self) -> Result<Option<T>, Box<dyn std::error::Error + Send + Sync>> {
if let Some((message_len, offset)) = self.read_varint()? {
if self.buffer.len() < offset + message_len {
return Ok(None);
}
let message_data = self.buffer.split_to(offset + message_len).freeze();
let message_bytes = &message_data[offset..];
let message = T::decode(message_bytes)?;
Ok(Some(message))
} else {
Ok(None)
}
}
fn read_varint(&self) -> Result<Option<(usize, usize)>, io::Error> {
let mut cursor = Cursor::new(&self.buffer);
let mut result: u64 = 0;
let mut shift = 0;
for _i in 0..10 {
if cursor.position() as usize >= self.buffer.len() {
return Ok(None);
}
let byte = cursor.get_ref()[cursor.position() as usize];
cursor.set_position(cursor.position() + 1);
let value = (byte & 0x7F) as u64;
result |= value << shift;
if (byte & 0x80) == 0 {
return Ok(Some((result as usize, (cursor.position()) as usize)));
}
shift += 7;
}
Err(io::Error::new(
io::ErrorKind::InvalidData,
"Varint too long",
))
}
pub fn has_remaining(&self) -> bool {
!self.buffer.is_empty()
}
pub fn remaining_len(&self) -> usize {
self.buffer.len()
}
pub fn clear(&mut self) {
self.buffer.clear();
}
}
impl<T> Default for ProtobufDecoder<T>
where
T: Message + Default,
{
fn default() -> Self {
Self::new()
}
}
pub fn safe_protobuf_decode<T>(data: &[u8]) -> Result<T, Box<dyn std::error::Error + Send + Sync>>
where
T: Message + Default,
{
match T::decode(data) {
Ok(message) => Ok(message),
Err(e) => Err(Box::new(e)),
}
}
pub fn safe_string_decode(data: &[u8]) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
match std::str::from_utf8(data) {
Ok(s) => Ok(s.to_string()),
Err(e) => Err(Box::new(e)),
}
}
#[cfg(test)]
mod tests {
use super::*;
use prost::Message;
#[derive(Clone, PartialEq, Message)]
struct TestMessage {
#[prost(string, tag = "1")]
content: String,
#[prost(int32, tag = "2")]
id: i32,
}
#[test]
fn test_protobuf_decoder() {
let mut decoder = ProtobufDecoder::<TestMessage>::new();
let test_msg = TestMessage {
content: "Hello, World!".to_string(),
id: 42,
};
let encoded_msg = test_msg.encode_to_vec();
let mut prefixed_data = Vec::new();
prost::encoding::encode_varint(encoded_msg.len() as u64, &mut prefixed_data);
prefixed_data.extend_from_slice(&encoded_msg);
decoder.add_data(&prefixed_data);
let decoded_msg = decoder.decode_next().unwrap().unwrap();
assert_eq!(decoded_msg.content, "Hello, World!");
assert_eq!(decoded_msg.id, 42);
}
#[test]
fn test_protobuf_decoder_partial_data() {
let mut decoder = ProtobufDecoder::<TestMessage>::new();
let test_msg = TestMessage {
content: "Hello, Partial Data Test!".to_string(),
id: 123,
};
let encoded_msg = test_msg.encode_to_vec();
let mut prefixed_data = Vec::new();
prost::encoding::encode_varint(encoded_msg.len() as u64, &mut prefixed_data);
prefixed_data.extend_from_slice(&encoded_msg);
let half_point = prefixed_data.len() / 2;
decoder.add_data(&prefixed_data[..half_point]);
let result = decoder.decode_next().unwrap();
assert!(result.is_none());
decoder.add_data(&prefixed_data[half_point..]);
let decoded_msg = decoder.decode_next().unwrap().unwrap();
assert_eq!(decoded_msg.content, "Hello, Partial Data Test!");
assert_eq!(decoded_msg.id, 123);
}
#[test]
fn test_safe_protobuf_decode() {
let test_msg = TestMessage {
content: "Safe decode test".to_string(),
id: 999,
};
let encoded = test_msg.encode_to_vec();
let decoded = safe_protobuf_decode::<TestMessage>(&encoded).unwrap();
assert_eq!(decoded.content, "Safe decode test");
assert_eq!(decoded.id, 999);
}
#[test]
fn test_safe_string_decode() {
let valid_utf8 = b"Valid UTF-8 string";
let result = safe_string_decode(valid_utf8).unwrap();
assert_eq!(result, "Valid UTF-8 string");
let invalid_utf8 = &[0xFF, 0xFE, 0xFD]; let result = safe_string_decode(invalid_utf8);
assert!(result.is_err());
}
}