Skip to main content

rmqtt_codec/v5/
codec.rs

1use std::cell::Cell;
2
3use bytes::{Buf, BytesMut};
4use tokio_util::codec::{Decoder, Encoder};
5
6use super::{decode::decode_packet, encode::EncodeLtd, Packet};
7use crate::error::{DecodeError, EncodeError};
8use crate::types::{FixedHeader, MAX_PACKET_SIZE};
9use crate::utils::decode_variable_length;
10
11/// MQTT v5.0 protocol codec implementing `tokio_util::codec` Decoder/Encoder
12///
13/// Handles MQTT v5.0 wire format framing including:
14/// - Fixed header decoding with configurable max inbound packet size
15/// - Variable-length integer parsing
16/// - Configurable max outbound packet size enforcement
17/// - Frame reassembly across multiple network reads
18/// - Negotiation flags tracking (problem info, retain, subscription IDs)
19#[derive(Debug, Clone)]
20pub struct Codec {
21    state: Cell<DecodeState>,
22    max_in_size: Cell<u32>,
23    max_out_size: Cell<u32>,
24    flags: Cell<CodecFlags>,
25}
26
27bitflags::bitflags! {
28    #[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
29    pub struct CodecFlags: u8 {
30        const NO_PROBLEM_INFO = 0b0000_0001;
31        const NO_RETAIN       = 0b0000_0010;
32        const NO_SUB_IDS      = 0b0000_1000;
33    }
34}
35
36#[derive(Debug, Clone, Copy)]
37enum DecodeState {
38    FrameHeader,
39    Frame(FixedHeader),
40}
41
42impl Codec {
43    /// Create `Codec` instance
44    pub fn new(max_in_size: u32, max_out_size: u32) -> Self {
45        Codec {
46            state: Cell::new(DecodeState::FrameHeader),
47            max_in_size: Cell::new(max_in_size),
48            max_out_size: Cell::new(max_out_size),
49            flags: Cell::new(CodecFlags::empty()),
50        }
51    }
52
53    /// Set max inbound frame size.
54    ///
55    /// If max size is set to `0`, size is unlimited.
56    /// By default max size is set to `0`
57    pub fn max_inbound_size(&self) -> u32 {
58        self.max_in_size.get()
59    }
60
61    /// Set max outbound frame size.
62    ///
63    /// If max size is set to `0`, size is unlimited.
64    /// By default max size is set to `0`
65    pub fn max_outbound_size(&self) -> u32 {
66        self.max_out_size.get()
67    }
68
69    /// Set max inbound frame size.
70    ///
71    /// If max size is set to `0`, size is unlimited.
72    /// By default max size is set to `0`
73    pub fn set_max_inbound_size(&mut self, size: u32) {
74        self.max_in_size.set(size);
75    }
76
77    /// Set max outbound frame size.
78    ///
79    /// If max size is set to `0`, size is unlimited.
80    /// By default max size is set to `0`
81    pub fn set_max_outbound_size(&mut self, mut size: u32) {
82        if size > 5 {
83            // fixed header = 1, var_len(remaining.max_value()) = 4
84            size -= 5;
85        }
86        self.max_out_size.set(size);
87    }
88
89    #[inline]
90    #[allow(dead_code)]
91    pub(crate) fn retain_available(&self) -> bool {
92        !self.flags.get().contains(CodecFlags::NO_RETAIN)
93    }
94
95    #[inline]
96    #[allow(dead_code)]
97    pub(crate) fn sub_ids_available(&self) -> bool {
98        !self.flags.get().contains(CodecFlags::NO_SUB_IDS)
99    }
100
101    #[inline]
102    #[allow(dead_code)]
103    pub(crate) fn set_retain_available(&self, val: bool) {
104        let mut flags = self.flags.get();
105        flags.set(CodecFlags::NO_RETAIN, !val);
106        self.flags.set(flags);
107    }
108
109    #[inline]
110    #[allow(dead_code)]
111    pub(crate) fn set_sub_ids_available(&self, val: bool) {
112        let mut flags = self.flags.get();
113        flags.set(CodecFlags::NO_SUB_IDS, !val);
114        self.flags.set(flags);
115    }
116}
117
118impl Default for Codec {
119    fn default() -> Self {
120        Self::new(0, 0)
121    }
122}
123
124impl Decoder for Codec {
125    type Item = (Packet, u32);
126    type Error = DecodeError;
127
128    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, DecodeError> {
129        loop {
130            match self.state.get() {
131                DecodeState::FrameHeader => {
132                    if src.len() < 2 {
133                        return Ok(None);
134                    }
135                    let src_slice = src.as_ref();
136                    let first_byte = src_slice[0];
137                    match decode_variable_length(&src_slice[1..])? {
138                        Some((remaining_length, consumed)) => {
139                            // check max message size
140                            let max_in_size = self.max_in_size.get();
141                            if max_in_size != 0 && max_in_size < remaining_length {
142                                log::debug!(
143                                    "MaxSizeExceeded max-size: {max_in_size}, remaining: {remaining_length}"
144                                );
145                                return Err(DecodeError::MaxSizeExceeded {
146                                    size: remaining_length,
147                                    max: max_in_size,
148                                });
149                            }
150                            src.advance(consumed + 1);
151                            self.state.set(DecodeState::Frame(FixedHeader { first_byte, remaining_length }));
152                            // todo: validate remaining_length against max frame size config
153                            let remaining_length = remaining_length as usize;
154                            if src.len() < remaining_length {
155                                // todo: subtract?
156                                src.reserve(remaining_length); // extend receiving buffer to fit the whole frame -- todo: too eager?
157                                return Ok(None);
158                            }
159                        }
160                        None => {
161                            return Ok(None);
162                        }
163                    }
164                }
165                DecodeState::Frame(fixed) => {
166                    if src.len() < fixed.remaining_length as usize {
167                        return Ok(None);
168                    }
169                    let packet_buf = src.split_to(fixed.remaining_length as usize).freeze();
170                    let packet = decode_packet(packet_buf, fixed.first_byte)?;
171                    self.state.set(DecodeState::FrameHeader);
172                    src.reserve(5); // enough to fix 1 fixed header byte + 4 bytes max variable packet length
173
174                    if let Packet::Connect(ref pkt) = packet {
175                        let mut flags = self.flags.get();
176                        flags.set(CodecFlags::NO_PROBLEM_INFO, !pkt.request_problem_info);
177                        self.flags.set(flags);
178                    }
179                    return Ok(Some((packet, fixed.remaining_length)));
180                }
181            }
182        }
183    }
184}
185
186impl Encoder<Packet> for Codec {
187    // type Item = Packet;
188    type Error = EncodeError;
189
190    fn encode(&mut self, mut item: Packet, dst: &mut BytesMut) -> Result<(), EncodeError> {
191        // handle [MQTT 3.1.2.11.7]
192        if self.flags.get().contains(CodecFlags::NO_PROBLEM_INFO) {
193            match item {
194                Packet::PublishAck(ref mut pkt) | Packet::PublishReceived(ref mut pkt) => {
195                    pkt.properties.clear();
196                    let _ = pkt.reason_string.take();
197                }
198                Packet::PublishRelease(ref mut pkt) | Packet::PublishComplete(ref mut pkt) => {
199                    pkt.properties.clear();
200                    let _ = pkt.reason_string.take();
201                }
202                Packet::Subscribe(ref mut pkt) => {
203                    pkt.user_properties.clear();
204                }
205                Packet::SubscribeAck(ref mut pkt) => {
206                    pkt.properties.clear();
207                    let _ = pkt.reason_string.take();
208                }
209                Packet::Unsubscribe(ref mut pkt) => {
210                    pkt.user_properties.clear();
211                }
212                Packet::UnsubscribeAck(ref mut pkt) => {
213                    pkt.properties.clear();
214                    let _ = pkt.reason_string.take();
215                }
216                Packet::Auth(ref mut pkt) => {
217                    pkt.user_properties.clear();
218                    let _ = pkt.reason_string.take();
219                }
220                _ => (),
221            }
222        }
223
224        let max_out_size = self.max_out_size.get();
225        let max_size = if max_out_size != 0 { max_out_size } else { MAX_PACKET_SIZE };
226        let content_size = item.encoded_size(max_size);
227        if content_size > max_size as usize {
228            return Err(EncodeError::OverMaxPacketSize { size: content_size as u32, max: max_size });
229        }
230        dst.reserve(content_size + 5);
231        item.encode(dst, content_size as u32)?; // safe: max_size <= u32 max value
232        Ok(())
233    }
234}
235
236#[cfg(test)]
237mod tests {
238    use super::*;
239
240    #[test]
241    fn test_max_size() {
242        let mut codec = Codec::new(5, 5);
243        let mut buf = BytesMut::new();
244        buf.extend_from_slice(b"\0\x09");
245        assert_eq!(
246            codec.decode(&mut buf).map_err(|e| matches!(e, DecodeError::MaxSizeExceeded { .. })),
247            Err(true)
248        );
249    }
250}