protosocket_prost/
decoder.rs1use std::marker::PhantomData;
2
3use protosocket::{Decoder, DeserializeError};
4
5#[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 let framed_length = message_length + prost::length_delimiter_len(message_length);
23 if buffer.remaining() < framed_length {
24 return Err(DeserializeError::IncompleteBuffer {
25 next_message_size: framed_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}