Skip to main content

calimero_network_primitives/stream/
codec.rs

1#[cfg(test)]
2#[path = "codec_test.rs"]
3mod tests;
4
5use core::slice;
6use std::borrow::Cow;
7use std::io::Error as IoError;
8
9use bytes::{Bytes, BytesMut};
10use serde::{Deserialize, Serialize};
11use thiserror::Error as ThisError;
12use tokio_util::codec::{Decoder, Encoder, LengthDelimitedCodec};
13
14#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
15#[non_exhaustive]
16pub struct Message<'a> {
17    pub data: Cow<'a, [u8]>,
18}
19
20impl<'a> Message<'a> {
21    #[must_use]
22    pub fn new<T: Into<Cow<'a, [u8]>>>(data: T) -> Self {
23        Self { data: data.into() }
24    }
25}
26
27#[derive(Debug, ThisError)]
28#[non_exhaustive]
29pub enum CodecError {
30    #[error(transparent)]
31    StdIo(#[from] IoError),
32}
33
34#[derive(Debug)]
35pub struct MessageCodec {
36    length_codec: LengthDelimitedCodec,
37}
38
39impl MessageCodec {
40    pub fn new(max_message_size: usize) -> Self {
41        let mut length_codec = LengthDelimitedCodec::new();
42        length_codec.set_max_frame_length(max_message_size);
43        Self { length_codec }
44    }
45}
46
47impl Decoder for MessageCodec {
48    type Item = Message<'static>;
49    type Error = CodecError;
50
51    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
52        let Some(frame) = self.length_codec.decode(src)? else {
53            return Ok(None);
54        };
55
56        Ok(Some(Message {
57            data: Cow::Owned(frame.into()),
58        }))
59    }
60}
61
62impl<'a> Encoder<Message<'a>> for MessageCodec {
63    type Error = CodecError;
64
65    fn encode(&mut self, item: Message<'a>, dst: &mut BytesMut) -> Result<(), Self::Error> {
66        let data = item.data.as_ref();
67        let data = Bytes::from_static(
68            // safety: `LengthDelimitedCodec: Encoder` must prepend the length, so it copies `data`
69            unsafe { slice::from_raw_parts(data.as_ptr(), data.len()) },
70        );
71        self.length_codec
72            .encode(data, dst)
73            .map_err(CodecError::StdIo)
74    }
75}