Skip to main content

rmqtt_codec/v3/
codec.rs

1use std::cell::Cell;
2
3use bytes::{Buf, BytesMut};
4use tokio_util::codec::{Decoder, Encoder};
5
6use super::{decode, encode, Packet};
7use crate::error::{DecodeError, EncodeError};
8use crate::types::{FixedHeader, QoS};
9use crate::utils::decode_variable_length;
10
11/// Mqtt v3.1.1 protocol codec implementing `tokio_util::codec` Decoder/Encoder
12///
13/// Handles MQTT v3.1.1 wire format framing including:
14/// - Fixed header decoding (packet type + remaining length)
15/// - Variable-length integer parsing
16/// - Configurable maximum packet size enforcement
17/// - Frame reassembly across multiple network reads
18#[derive(Debug, Clone)]
19pub struct Codec {
20    state: Cell<DecodeState>,
21    max_size: Cell<u32>,
22}
23
24#[derive(Debug, Copy, Clone, PartialEq, Eq)]
25enum DecodeState {
26    FrameHeader,
27    Frame(FixedHeader),
28}
29
30impl Codec {
31    /// Create `Codec` instance
32    pub fn new(max_packet_size: u32) -> Self {
33        Codec { state: Cell::new(DecodeState::FrameHeader), max_size: Cell::new(max_packet_size) }
34    }
35
36    /// Set max inbound frame size.
37    ///
38    /// If max size is set to `0`, size is unlimited.
39    /// By default max size is set to `0`
40    pub fn set_max_size(&mut self, size: u32) {
41        self.max_size.set(size);
42    }
43}
44
45impl Default for Codec {
46    fn default() -> Self {
47        Self::new(0)
48    }
49}
50
51impl Decoder for Codec {
52    type Item = (Packet, u32);
53    type Error = DecodeError;
54
55    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, DecodeError> {
56        loop {
57            match self.state.get() {
58                DecodeState::FrameHeader => {
59                    if src.len() < 2 {
60                        return Ok(None);
61                    }
62                    let src_slice = src.as_ref();
63                    let first_byte = src_slice[0];
64                    match decode_variable_length(&src_slice[1..])? {
65                        Some((remaining_length, consumed)) => {
66                            // check max message size
67                            let max_size = self.max_size.get();
68                            if max_size != 0 && max_size < remaining_length {
69                                return Err(DecodeError::MaxSizeExceeded {
70                                    size: remaining_length,
71                                    max: max_size,
72                                });
73                            }
74                            src.advance(consumed + 1);
75                            self.state.set(DecodeState::Frame(FixedHeader { first_byte, remaining_length }));
76                            // todo: validate remaining_length against max frame size config
77                            let remaining_length = remaining_length as usize;
78                            if src.len() < remaining_length {
79                                // todo: subtract?
80                                src.reserve(remaining_length); // extend receiving buffer to fit the whole frame -- todo: too eager?
81                                return Ok(None);
82                            }
83                        }
84                        None => {
85                            return Ok(None);
86                        }
87                    }
88                }
89                DecodeState::Frame(fixed) => {
90                    if src.len() < fixed.remaining_length as usize {
91                        return Ok(None);
92                    }
93                    let packet_buf = src.split_to(fixed.remaining_length as usize);
94                    let packet = decode::decode_packet(packet_buf.freeze(), fixed.first_byte)?;
95                    self.state.set(DecodeState::FrameHeader);
96                    src.reserve(2);
97                    return Ok(Some((packet, fixed.remaining_length)));
98                }
99            }
100        }
101    }
102}
103
104impl Encoder<Packet> for Codec {
105    // type Item = Packet;
106    type Error = EncodeError;
107
108    fn encode(&mut self, item: Packet, dst: &mut BytesMut) -> Result<(), EncodeError> {
109        if let Packet::Publish(ref publish) = item {
110            if (publish.qos == QoS::AtLeastOnce || publish.qos == QoS::ExactlyOnce)
111                && publish.packet_id.is_none()
112            {
113                return Err(EncodeError::PacketIdRequired);
114            }
115        }
116        let content_size = encode::get_encoded_size(&item);
117        dst.reserve(content_size + 5);
118        encode::encode(&item, dst, content_size as u32)?;
119        Ok(())
120    }
121}
122
123#[cfg(test)]
124mod tests {
125    use super::*;
126    use crate::v3::packet::Publish;
127    use bytes::Bytes;
128    use bytestring::ByteString;
129
130    #[test]
131    fn test_max_size() {
132        let mut codec = Codec::default();
133        codec.set_max_size(5);
134
135        let mut buf = BytesMut::new();
136        buf.extend_from_slice(b"\0\x09");
137        assert_eq!(
138            codec.decode(&mut buf).map_err(|e| matches!(e, DecodeError::MaxSizeExceeded { .. })),
139            Err(true)
140        );
141    }
142
143    #[test]
144    fn test_packet() {
145        let mut codec = Codec::default();
146        let mut buf = BytesMut::new();
147
148        let pkt = Box::new(Publish {
149            dup: false,
150            retain: false,
151            qos: QoS::AtMostOnce,
152            topic: ByteString::from_static("/test"),
153            packet_id: None,
154            payload: Bytes::from(Vec::from("a".repeat(260 * 1024))),
155            properties: None,
156        });
157        codec.encode(Packet::Publish(pkt.clone()), &mut buf).unwrap();
158
159        let pkt2 =
160            if let (Packet::Publish(v), _) = codec.decode(&mut buf).unwrap().unwrap() { v } else { panic!() };
161        assert_eq!(pkt, pkt2);
162    }
163}