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#[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 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 pub fn max_inbound_size(&self) -> u32 {
58 self.max_in_size.get()
59 }
60
61 pub fn max_outbound_size(&self) -> u32 {
66 self.max_out_size.get()
67 }
68
69 pub fn set_max_inbound_size(&mut self, size: u32) {
74 self.max_in_size.set(size);
75 }
76
77 pub fn set_max_outbound_size(&mut self, mut size: u32) {
82 if size > 5 {
83 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 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 let remaining_length = remaining_length as usize;
154 if src.len() < remaining_length {
155 src.reserve(remaining_length); 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); 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 Error = EncodeError;
189
190 fn encode(&mut self, mut item: Packet, dst: &mut BytesMut) -> Result<(), EncodeError> {
191 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)?; 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}