Skip to main content

mqtt5_protocol/session/
limits.rs

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    /// The maximum packet size that may be **sent** on this connection.
54    ///
55    /// An outbound packet is bounded solely by the server's advertised Maximum
56    /// Packet Size (from CONNACK). The client's own Maximum Packet Size is the
57    /// limit on packets it will *receive*, not what it may send, so it is never
58    /// applied to outbound traffic here. When the server advertises no limit
59    /// (or zero), the client's configured maximum is used as the fallback
60    /// ceiling. A returned value of `0` means "no limit".
61    #[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    /// Checks an outbound packet's encoded size against
70    /// [`effective_maximum_packet_size`](Self::effective_maximum_packet_size).
71    ///
72    /// # Errors
73    /// Returns `PacketTooLarge` if the size exceeds the effective maximum.
74    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}