Skip to main content

protosocket_prost/
decoder.rs

1use std::marker::PhantomData;
2
3use protosocket::{Decoder, DeserializeError};
4
5/// A stateless implementation of `Decoder` using `prost`
6#[derive(Debug, Default)]
7pub struct ProstDecoder<Message> {
8    _phantom: PhantomData<Message>,
9}
10impl<Message> Decoder for ProstDecoder<Message>
11where
12    Message: prost::Message + Default + std::fmt::Debug,
13{
14    type Message = Message;
15
16    fn decode(
17        &mut self,
18        mut buffer: impl bytes::Buf,
19    ) -> std::result::Result<(usize, Self::Message), DeserializeError> {
20        match prost::decode_length_delimiter(buffer.chunk()) {
21            Ok(message_length) => {
22                if buffer.remaining() < message_length + prost::length_delimiter_len(message_length)
23                {
24                    return Err(DeserializeError::IncompleteBuffer {
25                        next_message_size: message_length,
26                    });
27                }
28            }
29            Err(e) => {
30                log::trace!("can't read a length delimiter {e:?}");
31                return Err(DeserializeError::IncompleteBuffer {
32                    next_message_size: 10,
33                });
34            }
35        };
36
37        let start = buffer.remaining();
38        match <Self::Message as prost::Message>::decode_length_delimited(&mut buffer) {
39            Ok(message) => {
40                let length = start - buffer.remaining();
41                log::debug!("decoded {length}: {message:?}");
42                Ok((length, message))
43            }
44            Err(e) => {
45                log::warn!("could not decode message: {e:?}");
46                Err(DeserializeError::InvalidBuffer)
47            }
48        }
49    }
50}