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#[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 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 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 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 let remaining_length = remaining_length as usize;
78 if src.len() < remaining_length {
79 src.reserve(remaining_length); 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 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}