1#![no_std]
2
3use embedded_io_async::{Read, Write};
12
13pub const MAX_PACKET_SIZE: usize = 256;
16
17#[derive(Debug)]
19pub enum MqttError<E> {
20 Io(E),
22 PacketTooLarge,
24 ConnackInvalid,
26 ConnectionRefused(u8),
28 PingFailed,
30 SubackInvalid,
32 SubscribeFailed,
34 UnexpectedPacket,
36}
37
38impl<E> From<E> for MqttError<E> {
39 fn from(e: E) -> Self {
40 MqttError::Io(e)
41 }
42}
43
44struct PacketBuilder {
46 buf: [u8; MAX_PACKET_SIZE],
47 len: usize,
48}
49
50impl PacketBuilder {
51 fn new() -> Self {
52 Self {
53 buf: [0u8; MAX_PACKET_SIZE],
54 len: 0,
55 }
56 }
57
58 fn push(&mut self, byte: u8) -> Result<(), ()> {
59 if self.len >= MAX_PACKET_SIZE {
60 return Err(());
61 }
62 self.buf[self.len] = byte;
63 self.len += 1;
64 Ok(())
65 }
66
67 fn extend(&mut self, bytes: &[u8]) -> Result<(), ()> {
68 if self.len + bytes.len() > MAX_PACKET_SIZE {
69 return Err(());
70 }
71 self.buf[self.len..self.len + bytes.len()].copy_from_slice(bytes);
72 self.len += bytes.len();
73 Ok(())
74 }
75
76 fn as_slice(&self) -> &[u8] {
77 &self.buf[..self.len]
78 }
79}
80
81fn encode_remaining_length(mut len: usize, out: &mut PacketBuilder) -> Result<(), ()> {
83 loop {
84 let mut byte = (len % 128) as u8;
85 len /= 128;
86 if len > 0 {
87 byte |= 0x80;
88 }
89 out.push(byte)?;
90 if len == 0 {
91 break;
92 }
93 }
94 Ok(())
95}
96
97
98fn push_string_field(builder: &mut PacketBuilder, s: &[u8]) -> Result<(), ()> {
100 builder.extend(&(s.len() as u16).to_be_bytes())?;
101 builder.extend(s)
102}
103
104
105
106
107
108#[derive(Default)]
110pub struct ConnectOptions<'a> {
111 pub username: Option<&'a str>,
112 pub password: Option<&'a [u8]>,
113 pub last_will: Option<LastWill<'a>>,
114}
115
116pub struct LastWill<'a> {
119 pub topic: &'a str,
120 pub message: &'a [u8],
121 pub retain: bool,
122}
123
124
125pub struct MqttClient<'a, T: Read + Write> {
127 transport: &'a mut T,
128 packet_id: u16,
129}
130
131
132fn build_subscribe_packet(topic: &str, packet_id: u16) -> Result<PacketBuilder, ()> {
134 let mut variable_header = PacketBuilder::new();
135 variable_header.extend(&packet_id.to_be_bytes())?;
136
137 let mut payload = PacketBuilder::new();
138 push_string_field(&mut payload, topic.as_bytes())?;
139 payload.push(0x00)?; let mut packet = PacketBuilder::new();
142 packet.push(0x82)?; encode_remaining_length(variable_header.len + payload.len, &mut packet)?;
144 packet.extend(variable_header.as_slice())?;
145 packet.extend(payload.as_slice())?;
146
147 Ok(packet)
148}
149
150#[derive(Debug, PartialEq)]
152enum SubackResult {
153 Accepted,
154 Refused,
155 Invalid,
156}
157
158fn parse_suback(header_byte: u8, body: &[u8]) -> SubackResult {
160 if header_byte != 0x90 || body.len() < 3 {
161 return SubackResult::Invalid;
162 }
163 if body[2] == 0x80 {
164 return SubackResult::Refused;
165 }
166 SubackResult::Accepted
167}
168
169fn publish_layout(header_byte: u8, buf: &[u8]) -> Result<(usize, usize), ()> {
172 if buf.len() < 2 {
173 return Err(());
174 }
175 let topic_len = u16::from_be_bytes([buf[0], buf[1]]) as usize;
176 let qos = (header_byte >> 1) & 0x03;
177 let mut payload_start = 2 + topic_len;
178 if qos > 0 {
179 payload_start += 2; }
181 if payload_start > buf.len() {
182 return Err(());
183 }
184 Ok((topic_len, payload_start))
185}
186pub struct IncomingMessage<'buf> {
188 pub topic: &'buf str,
189 pub payload: &'buf [u8],
190}
191
192
193
194impl<'a, T: Read + Write> MqttClient<'a, T> {
195 pub fn new(transport: &'a mut T) -> Self {
197 Self {
198 transport,
199 packet_id: 1,
200 }
201 }
202
203 pub async fn connect(
206 &mut self,
207 client_id: &str,
208 keep_alive_secs: u16,
209 ) -> Result<(), MqttError<T::Error>> {
210 self.connect_with_options(client_id, keep_alive_secs, &ConnectOptions::default())
211 .await
212 }
213
214 pub async fn connect_with_options(
216 &mut self,
217 client_id: &str,
218 keep_alive_secs: u16,
219 options: &ConnectOptions<'_>,
220 ) -> Result<(), MqttError<T::Error>> {
221 let mut variable_header = PacketBuilder::new();
222 variable_header
223 .extend(&[0x00, 0x04])
224 .map_err(|_| MqttError::PacketTooLarge)?;
225 variable_header
226 .extend(b"MQTT")
227 .map_err(|_| MqttError::PacketTooLarge)?;
228 variable_header
229 .push(0x04)
230 .map_err(|_| MqttError::PacketTooLarge)?;
231
232 let mut flags: u8 = 0x02; if let Some(will) = &options.last_will {
234 flags |= 0x04;
235 if will.retain {
236 flags |= 0x20;
237 }
238 }
239 if options.username.is_some() {
240 flags |= 0x80;
241 }
242 if options.password.is_some() {
243 flags |= 0x40;
244 }
245 variable_header
246 .push(flags)
247 .map_err(|_| MqttError::PacketTooLarge)?;
248 variable_header
249 .extend(&keep_alive_secs.to_be_bytes())
250 .map_err(|_| MqttError::PacketTooLarge)?;
251
252 let mut payload = PacketBuilder::new();
253 push_string_field(&mut payload, client_id.as_bytes())
254 .map_err(|_| MqttError::PacketTooLarge)?;
255
256 if let Some(will) = &options.last_will {
257 push_string_field(&mut payload, will.topic.as_bytes())
258 .map_err(|_| MqttError::PacketTooLarge)?;
259 push_string_field(&mut payload, will.message)
260 .map_err(|_| MqttError::PacketTooLarge)?;
261 }
262 if let Some(username) = options.username {
263 push_string_field(&mut payload, username.as_bytes())
264 .map_err(|_| MqttError::PacketTooLarge)?;
265 }
266 if let Some(password) = options.password {
267 push_string_field(&mut payload, password).map_err(|_| MqttError::PacketTooLarge)?;
268 }
269
270 let mut packet = PacketBuilder::new();
271 packet.push(0x10).map_err(|_| MqttError::PacketTooLarge)?;
272 encode_remaining_length(variable_header.len + payload.len, &mut packet)
273 .map_err(|_| MqttError::PacketTooLarge)?;
274 packet
275 .extend(variable_header.as_slice())
276 .map_err(|_| MqttError::PacketTooLarge)?;
277 packet
278 .extend(payload.as_slice())
279 .map_err(|_| MqttError::PacketTooLarge)?;
280
281 self.transport.write_all(packet.as_slice()).await?;
282
283 let mut connack = [0u8; 4];
284 self.transport
285 .read_exact(&mut connack)
286 .await
287 .map_err(|_| MqttError::ConnackInvalid)?;
288
289 if connack[0] != 0x20 || connack[1] != 0x02 {
290 return Err(MqttError::ConnackInvalid);
291 }
292 if connack[3] != 0x00 {
293 return Err(MqttError::ConnectionRefused(connack[3]));
294 }
295
296 Ok(())
297 }
298
299 pub async fn publish(&mut self, topic: &str, payload: &[u8]) -> Result<(), MqttError<T::Error>> {
301 let mut variable_header = PacketBuilder::new();
302 let topic_bytes = topic.as_bytes();
303 variable_header
304 .extend(&(topic_bytes.len() as u16).to_be_bytes())
305 .map_err(|_| MqttError::PacketTooLarge)?;
306 variable_header
307 .extend(topic_bytes)
308 .map_err(|_| MqttError::PacketTooLarge)?;
309 let mut packet = PacketBuilder::new();
312 packet.push(0x30).map_err(|_| MqttError::PacketTooLarge)?; encode_remaining_length(variable_header.len + payload.len(), &mut packet)
314 .map_err(|_| MqttError::PacketTooLarge)?;
315 packet
316 .extend(variable_header.as_slice())
317 .map_err(|_| MqttError::PacketTooLarge)?;
318 packet.extend(payload).map_err(|_| MqttError::PacketTooLarge)?;
319
320 self.transport.write_all(packet.as_slice()).await?;
321 Ok(())
322 }
323
324 pub async fn ping(&mut self) -> Result<(), MqttError<T::Error>> {
330 self.transport.write_all(&[0xC0, 0x00]).await?;
331
332 let mut pingresp = [0u8; 2];
333 self.transport
334 .read_exact(&mut pingresp)
335 .await
336 .map_err(|_| MqttError::PingFailed)?;
337
338 if pingresp != [0xD0, 0x00] {
339 return Err(MqttError::PingFailed);
340 }
341
342 Ok(())
343 }
344
345
346
347
348 async fn read_remaining_length(&mut self) -> Result<usize, MqttError<T::Error>> {
350 let mut multiplier: usize = 1;
351 let mut value: usize = 0;
352 loop {
353 let mut byte = [0u8; 1];
354 self.transport
355 .read_exact(&mut byte)
356 .await
357 .map_err(|_| MqttError::UnexpectedPacket)?;
358 value += ((byte[0] & 0x7F) as usize) * multiplier;
359 if byte[0] & 0x80 == 0 {
360 break;
361 }
362 multiplier *= 128;
363 if multiplier > 128 * 128 * 128 {
364 return Err(MqttError::PacketTooLarge);
365 }
366 }
367 Ok(value)
368 }
369
370
371 pub async fn subscribe(&mut self, topic: &str) -> Result<(), MqttError<T::Error>> {
373 self.packet_id = self.packet_id.wrapping_add(1).max(1);
374 let packet =
375 build_subscribe_packet(topic, self.packet_id).map_err(|_| MqttError::PacketTooLarge)?;
376
377 self.transport.write_all(packet.as_slice()).await?;
378
379 let mut header = [0u8; 1];
380 self.transport
381 .read_exact(&mut header)
382 .await
383 .map_err(|_| MqttError::SubackInvalid)?;
384
385 let remaining_len = self.read_remaining_length().await?;
386 if !(3..=8).contains(&remaining_len) {
387 return Err(MqttError::SubackInvalid);
388 }
389
390 let mut body = [0u8; 8];
391 self.transport
392 .read_exact(&mut body[..remaining_len])
393 .await
394 .map_err(|_| MqttError::SubackInvalid)?;
395
396 match parse_suback(header[0], &body[..remaining_len]) {
397 SubackResult::Accepted => Ok(()),
398 SubackResult::Refused => Err(MqttError::SubscribeFailed),
399 SubackResult::Invalid => Err(MqttError::SubackInvalid),
400 }
401 }
402
403 pub async fn receive<'buf>(
405 &mut self,
406 buf: &'buf mut [u8],
407 ) -> Result<IncomingMessage<'buf>, MqttError<T::Error>> {
408 loop {
409 let mut header = [0u8; 1];
410 self.transport
411 .read_exact(&mut header)
412 .await
413 .map_err(|_| MqttError::UnexpectedPacket)?;
414 let packet_type = header[0] & 0xF0;
415 let remaining_len = self.read_remaining_length().await?;
416
417 if remaining_len > buf.len() {
418 return Err(MqttError::PacketTooLarge);
419 }
420
421 self.transport
422 .read_exact(&mut buf[..remaining_len])
423 .await
424 .map_err(|_| MqttError::UnexpectedPacket)?;
425
426 if packet_type != 0x30 {
427 continue; }
429
430 let (topic_len, payload_start) = publish_layout(header[0], &buf[..remaining_len])
431 .map_err(|_| MqttError::UnexpectedPacket)?;
432
433 let (topic_and_rest, payload_part) = buf[..remaining_len].split_at(payload_start);
434 let topic_bytes = &topic_and_rest[2..2 + topic_len];
435 let topic =
436 core::str::from_utf8(topic_bytes).map_err(|_| MqttError::UnexpectedPacket)?;
437
438 return Ok(IncomingMessage {
439 topic,
440 payload: payload_part,
441 });
442 }
443 }
444
445
446}
447
448#[cfg(test)]
449mod tests {
450 use super::*;
451
452 #[test]
453 fn remaining_length_zero() {
454 let mut b = PacketBuilder::new();
455 encode_remaining_length(0, &mut b).unwrap();
456 assert_eq!(b.as_slice(), &[0x00]);
457 }
458
459 #[test]
460 fn remaining_length_single_byte_max() {
461 let mut b = PacketBuilder::new();
462 encode_remaining_length(127, &mut b).unwrap();
463 assert_eq!(b.as_slice(), &[0x7F]);
464 }
465
466 #[test]
467 fn remaining_length_two_bytes() {
468 let mut b = PacketBuilder::new();
469 encode_remaining_length(200, &mut b).unwrap();
470 assert_eq!(b.as_slice(), &[0xC8, 0x01]);
471 }
472
473 #[test]
474 fn remaining_length_three_bytes() {
475 let mut b = PacketBuilder::new();
476 encode_remaining_length(16384, &mut b).unwrap();
477 assert_eq!(b.as_slice(), &[0x80, 0x80, 0x01]);
478 }
479
480 #[test]
481 fn string_field_encoding() {
482 let mut b = PacketBuilder::new();
483 push_string_field(&mut b, b"MQTT").unwrap();
484 assert_eq!(b.as_slice(), &[0x00, 0x04, b'M', b'Q', b'T', b'T']);
485 }
486
487 #[test]
488 fn packet_builder_rejects_overflow() {
489 let mut b = PacketBuilder::new();
490 let big = [0u8; MAX_PACKET_SIZE + 1];
491 assert!(b.extend(&big).is_err());
492 }
493
494
495 #[test]
496 fn subscribe_packet_encoding() {
497 let packet = build_subscribe_packet("home/clim", 1).unwrap();
498 let expected = [
499 0x82, 0x0E, 0x00, 0x01, 0x00, 0x09, b'h', b'o', b'm', b'e', b'/', b'c', b'l', b'i', b'm', 0x00, ];
504 assert_eq!(packet.as_slice(), &expected);
505 }
506
507 #[test]
508 fn suback_accepted() {
509 let body = [0x00, 0x01, 0x00]; assert_eq!(parse_suback(0x90, &body), SubackResult::Accepted);
511 }
512
513 #[test]
514 fn suback_refused() {
515 let body = [0x00, 0x01, 0x80];
516 assert_eq!(parse_suback(0x90, &body), SubackResult::Refused);
517 }
518
519 #[test]
520 fn suback_wrong_header_type() {
521 let body = [0x00, 0x01, 0x00];
522 assert_eq!(parse_suback(0x20, &body), SubackResult::Invalid); }
524
525 #[test]
526 fn suback_body_too_short() {
527 let body = [0x00, 0x01];
528 assert_eq!(parse_suback(0x90, &body), SubackResult::Invalid);
529 }
530
531 #[test]
532 fn publish_layout_qos0() {
533 let buf = [0x00, 0x03, b'a', b'/', b'b', b'h', b'i'];
534 let (topic_len, payload_start) = publish_layout(0x30, &buf).unwrap();
535 assert_eq!(topic_len, 3);
536 assert_eq!(payload_start, 5);
537 assert_eq!(&buf[payload_start..], b"hi");
538 }
539
540 #[test]
541 fn publish_layout_qos1_skips_packet_identifier() {
542 let buf = [0x00, 0x03, b'a', b'/', b'b', 0x00, 0x2A, b'h', b'i'];
543 let (topic_len, payload_start) = publish_layout(0x32, &buf).unwrap();
544 assert_eq!(topic_len, 3);
545 assert_eq!(payload_start, 7);
546 assert_eq!(&buf[payload_start..], b"hi");
547 }
548
549 #[test]
550 fn publish_layout_rejects_truncated_buffer() {
551 let buf = [0x00, 0x05, b'a', b'b']; assert!(publish_layout(0x30, &buf).is_err());
553 }
554
555
556
557
558
559}