1use crate::error::{MqttError, Result};
2use crate::prelude::{String, Vec};
3use crate::time::{Duration, Instant};
4use crate::QoS;
5
6#[derive(Debug, Clone)]
7pub struct LimitsConfig {
8 pub client_maximum_packet_size: u32,
9 pub server_maximum_packet_size: Option<u32>,
10 pub default_message_expiry: Option<Duration>,
11 pub max_message_expiry: Option<Duration>,
12}
13
14impl Default for LimitsConfig {
15 fn default() -> Self {
16 Self {
17 client_maximum_packet_size: crate::constants::limits::MAX_PACKET_SIZE,
18 server_maximum_packet_size: None,
19 default_message_expiry: None,
20 max_message_expiry: Some(Duration::from_secs(86400 * 7)),
21 }
22 }
23}
24
25#[derive(Debug)]
26pub struct LimitsManager {
27 config: LimitsConfig,
28}
29
30impl LimitsManager {
31 #[must_use]
32 pub fn new(config: LimitsConfig) -> Self {
33 Self { config }
34 }
35
36 #[must_use]
37 pub fn with_defaults() -> Self {
38 Self::new(LimitsConfig::default())
39 }
40
41 pub fn set_server_maximum_packet_size(&mut self, size: u32) {
42 self.config.server_maximum_packet_size = Some(size);
43 }
44
45 pub fn reset_server_maximum_packet_size(&mut self) {
46 self.config.server_maximum_packet_size = None;
47 }
48
49 pub fn set_client_maximum_packet_size(&mut self, size: u32) {
50 self.config.client_maximum_packet_size = size;
51 }
52
53 #[must_use]
62 pub fn effective_maximum_packet_size(&self) -> u32 {
63 match self.config.server_maximum_packet_size {
64 Some(server_max) if server_max > 0 => server_max,
65 _ => self.config.client_maximum_packet_size,
66 }
67 }
68
69 pub fn check_packet_size(&self, size: usize) -> Result<()> {
75 let max_size = self.effective_maximum_packet_size();
76 if max_size > 0 && size > max_size as usize {
77 Err(MqttError::PacketTooLarge {
78 size,
79 max: max_size as usize,
80 })
81 } else {
82 Ok(())
83 }
84 }
85
86 #[must_use]
87 pub fn calculate_message_expiry(&self, expiry_interval: Option<u32>) -> Option<Instant> {
88 let interval = match expiry_interval {
89 Some(seconds) => Duration::from_secs(u64::from(seconds)),
90 None => self.config.default_message_expiry?,
91 };
92
93 let final_interval = match self.config.max_message_expiry {
94 Some(max) => interval.min(max),
95 None => interval,
96 };
97
98 Some(Instant::now() + final_interval)
99 }
100
101 #[must_use]
102 pub fn is_message_expired(&self, expiry_time: Option<Instant>) -> bool {
103 match expiry_time {
104 Some(expiry) => Instant::now() > expiry,
105 None => false,
106 }
107 }
108
109 #[must_use]
110 pub fn get_remaining_expiry(&self, expiry_time: Option<Instant>) -> Option<u32> {
111 match expiry_time {
112 Some(expiry) => {
113 let now = Instant::now();
114 if now < expiry {
115 let remaining = expiry.duration_since(now);
116 Some(u32::try_from(remaining.as_secs()).unwrap_or(u32::MAX))
117 } else {
118 Some(0)
119 }
120 }
121 None => None,
122 }
123 }
124
125 #[must_use]
126 pub fn client_maximum_packet_size(&self) -> u32 {
127 self.config.client_maximum_packet_size
128 }
129
130 #[must_use]
131 pub fn server_maximum_packet_size(&self) -> Option<u32> {
132 self.config.server_maximum_packet_size
133 }
134}
135
136#[derive(Debug, Clone)]
137pub struct ExpiringMessage {
138 pub topic: String,
139 pub payload: Vec<u8>,
140 pub qos: QoS,
141 pub retain: bool,
142 pub packet_id: Option<u16>,
143 pub expiry_time: Option<Instant>,
144 pub expiry_interval: Option<u32>,
145}
146
147impl ExpiringMessage {
148 #[must_use]
149 pub fn new(
150 topic: String,
151 payload: Vec<u8>,
152 qos: QoS,
153 retain: bool,
154 packet_id: Option<u16>,
155 expiry_interval: Option<u32>,
156 limits: &LimitsManager,
157 ) -> Self {
158 let expiry_time = limits.calculate_message_expiry(expiry_interval);
159
160 Self {
161 topic,
162 payload,
163 qos,
164 retain,
165 packet_id,
166 expiry_time,
167 expiry_interval,
168 }
169 }
170
171 #[must_use]
172 pub fn is_expired(&self) -> bool {
173 match self.expiry_time {
174 Some(expiry) => Instant::now() > expiry,
175 None => false,
176 }
177 }
178
179 #[must_use]
180 pub fn remaining_expiry_interval(&self) -> Option<u32> {
181 match self.expiry_time {
182 Some(expiry) => {
183 let now = Instant::now();
184 if now < expiry {
185 let remaining = expiry.duration_since(now);
186 Some(u32::try_from(remaining.as_secs()).unwrap_or(u32::MAX))
187 } else {
188 Some(0)
189 }
190 }
191 None => self.expiry_interval,
192 }
193 }
194}
195
196#[cfg(test)]
197mod tests {
198 use super::*;
199
200 #[test]
201 fn test_limits_manager_creation() {
202 let limits = LimitsManager::with_defaults();
203 assert_eq!(
204 limits.client_maximum_packet_size(),
205 crate::constants::limits::MAX_PACKET_SIZE
206 );
207 assert_eq!(limits.server_maximum_packet_size(), None);
208 }
209
210 #[test]
211 fn test_effective_packet_size() {
212 let mut limits = LimitsManager::with_defaults();
213
214 assert_eq!(
215 limits.effective_maximum_packet_size(),
216 crate::constants::limits::MAX_PACKET_SIZE
217 );
218
219 limits.set_server_maximum_packet_size(1_048_576);
220 assert_eq!(limits.effective_maximum_packet_size(), 1_048_576);
221
222 let config = LimitsConfig {
223 client_maximum_packet_size: 1_048_576,
224 ..Default::default()
225 };
226 let mut limits = LimitsManager::new(config);
227 limits.set_server_maximum_packet_size(10_485_760);
228 assert_eq!(limits.effective_maximum_packet_size(), 10_485_760);
229 }
230
231 #[test]
232 fn test_packet_size_checking() {
233 let mut limits = LimitsManager::with_defaults();
234 limits.set_server_maximum_packet_size(1024);
235
236 assert!(limits.check_packet_size(512).is_ok());
237 assert!(limits.check_packet_size(1024).is_ok());
238
239 let result = limits.check_packet_size(2048);
240 assert!(result.is_err());
241 if let Err(MqttError::PacketTooLarge { size, max }) = result {
242 assert_eq!(size, 2048);
243 assert_eq!(max, 1024);
244 }
245 }
246
247 #[test]
248 fn test_message_expiry() {
249 let config = LimitsConfig {
250 default_message_expiry: Some(Duration::from_secs(60)),
251 ..Default::default()
252 };
253 let limits = LimitsManager::new(config);
254
255 let expiry_time = limits.calculate_message_expiry(Some(30));
256 assert!(expiry_time.is_some());
257
258 let expiry_time = limits.calculate_message_expiry(None);
259 assert!(expiry_time.is_some());
260
261 let past_time = Some(Instant::now().checked_sub(Duration::from_secs(10)).unwrap());
262 assert!(limits.is_message_expired(past_time));
263
264 let future_time = Some(Instant::now() + Duration::from_secs(10));
265 assert!(!limits.is_message_expired(future_time));
266 }
267
268 #[test]
269 fn test_remaining_expiry() {
270 let limits = LimitsManager::with_defaults();
271
272 let future_time = Some(Instant::now() + Duration::from_secs(100));
273 let remaining = limits.get_remaining_expiry(future_time);
274 assert!(remaining.is_some());
275 assert!(remaining.unwrap() > 95 && remaining.unwrap() <= 100);
276
277 let past_time = Some(Instant::now().checked_sub(Duration::from_secs(10)).unwrap());
278 let remaining = limits.get_remaining_expiry(past_time);
279 assert_eq!(remaining, Some(0));
280 }
281
282 #[test]
283 fn test_expiring_message() {
284 let limits = LimitsManager::with_defaults();
285
286 let msg = ExpiringMessage::new(
287 "test/topic".into(),
288 vec![1, 2, 3],
289 QoS::AtLeastOnce,
290 false,
291 Some(123),
292 Some(60),
293 &limits,
294 );
295
296 assert!(!msg.is_expired());
297 assert!(msg.remaining_expiry_interval().is_some());
298
299 let mut msg = ExpiringMessage::new(
300 "test/topic".into(),
301 vec![1, 2, 3],
302 QoS::AtLeastOnce,
303 false,
304 Some(123),
305 Some(0),
306 &limits,
307 );
308
309 msg.expiry_time = Some(Instant::now().checked_sub(Duration::from_secs(10)).unwrap());
310 assert!(msg.is_expired());
311 assert_eq!(msg.remaining_expiry_interval(), Some(0));
312 }
313
314 #[test]
315 fn test_max_expiry_limit() {
316 let config = LimitsConfig {
317 max_message_expiry: Some(crate::constants::time::DEFAULT_SESSION_EXPIRY),
318 ..Default::default()
319 };
320 let limits = LimitsManager::new(config);
321
322 let expiry_time = limits.calculate_message_expiry(Some(7200));
323 assert!(expiry_time.is_some());
324
325 let remaining = limits.get_remaining_expiry(expiry_time);
326 assert!(remaining.is_some());
327 assert!(
328 remaining.unwrap()
329 <= u32::try_from(crate::constants::time::DEFAULT_SESSION_EXPIRY.as_secs())
330 .unwrap_or(u32::MAX)
331 );
332 }
333}