1use std::collections::HashMap;
40use std::sync::{Arc, Mutex};
41use std::time::Duration;
42
43use coap_lite::{CoapOption, MessageClass, MessageType, Packet, RequestType};
44use pamoja_core::{Error, Result, Transport};
45use tokio::net::UdpSocket;
46use tokio::sync::{mpsc, oneshot};
47use tokio::task::JoinHandle;
48
49type PendingAcks = Arc<Mutex<HashMap<u16, oneshot::Sender<()>>>>;
51
52#[derive(Clone, Copy, Debug, PartialEq, Eq)]
56pub enum Reliability {
57 NonConfirmable,
59 Confirmable,
61}
62
63#[derive(Clone, Debug)]
68pub struct CoapConfig {
69 host: String,
70 port: u16,
71 bind: String,
72 reliability: Reliability,
73 ack_timeout: Duration,
74 max_retransmits: u32,
75}
76
77impl CoapConfig {
78 pub fn new(host: impl Into<String>, port: u16) -> Self {
91 Self {
92 host: host.into(),
93 port,
94 bind: "0.0.0.0:0".to_owned(),
95 reliability: Reliability::Confirmable,
96 ack_timeout: Duration::from_secs(2),
97 max_retransmits: 4,
98 }
99 }
100
101 pub fn bind(mut self, addr: impl Into<String>) -> Self {
112 self.bind = addr.into();
113 self
114 }
115
116 pub fn reliability(mut self, reliability: Reliability) -> Self {
126 self.reliability = reliability;
127 self
128 }
129
130 pub fn ack_timeout(mut self, timeout: Duration) -> Self {
142 self.ack_timeout = timeout;
143 self
144 }
145
146 pub fn max_retransmits(mut self, count: u32) -> Self {
156 self.max_retransmits = count;
157 self
158 }
159}
160
161#[derive(Clone, Debug, PartialEq, Eq)]
163pub struct Message {
164 pub topic: String,
166 pub payload: Vec<u8>,
168}
169
170pub struct CoapTransport {
177 config: CoapConfig,
178 socket: Option<Arc<UdpSocket>>,
179 incoming: Option<mpsc::UnboundedReceiver<Message>>,
180 pending: PendingAcks,
181 pump: Option<JoinHandle<()>>,
182 next_id: u16,
183 next_token: u16,
184}
185
186impl CoapTransport {
187 pub fn new(config: CoapConfig) -> Self {
197 Self {
198 config,
199 socket: None,
200 incoming: None,
201 pending: Arc::new(Mutex::new(HashMap::new())),
202 pump: None,
203 next_id: 0,
204 next_token: 0,
205 }
206 }
207
208 pub fn is_connected(&self) -> bool {
215 self.socket.is_some()
216 }
217
218 pub async fn recv(&mut self) -> Result<Option<Message>> {
230 let incoming = self.incoming.as_mut().ok_or(Error::Closed)?;
231 Ok(incoming.recv().await)
232 }
233
234 pub async fn disconnect(&mut self) -> Result<()> {
246 if let Some(pump) = self.pump.take() {
247 pump.abort();
248 }
249 self.socket = None;
250 self.incoming = None;
251 Ok(())
252 }
253
254 fn next_message_id(&mut self) -> u16 {
256 let id = self.next_id;
257 self.next_id = self.next_id.wrapping_add(1);
258 id
259 }
260
261 fn next_request_token(&mut self) -> Vec<u8> {
263 let token = self.next_token;
264 self.next_token = self.next_token.wrapping_add(1);
265 token.to_be_bytes().to_vec()
266 }
267
268 async fn send_confirmable(&mut self, id: u16, bytes: &[u8], socket: &UdpSocket) -> Result<()> {
271 let mut timeout = self.config.ack_timeout;
272 for _ in 0..=self.config.max_retransmits {
273 let (tx, rx) = oneshot::channel();
274 self.pending.lock().expect("pending lock").insert(id, tx);
275 socket
276 .send(bytes)
277 .await
278 .map_err(|err| Error::Transport(err.to_string()))?;
279 match tokio::time::timeout(timeout, rx).await {
280 Ok(Ok(())) => return Ok(()),
281 Ok(Err(_)) => return Err(Error::Closed),
282 Err(_) => {
283 self.pending.lock().expect("pending lock").remove(&id);
284 timeout = timeout.saturating_mul(2);
285 }
286 }
287 }
288 Err(Error::Transport(format!(
289 "no acknowledgement for message {id}"
290 )))
291 }
292}
293
294impl Transport for CoapTransport {
295 async fn connect(&mut self) -> Result<()> {
296 let server = tokio::net::lookup_host((self.config.host.as_str(), self.config.port))
297 .await
298 .map_err(|err| Error::Transport(err.to_string()))?
299 .next()
300 .ok_or_else(|| Error::Transport(format!("could not resolve {}", self.config.host)))?;
301
302 let socket = UdpSocket::bind(&self.config.bind)
303 .await
304 .map_err(|err| Error::Transport(err.to_string()))?;
305 socket
306 .connect(server)
307 .await
308 .map_err(|err| Error::Transport(err.to_string()))?;
309 let socket = Arc::new(socket);
310
311 let (tx, rx) = mpsc::unbounded_channel();
312 let pending = Arc::clone(&self.pending);
313 let pump_socket = Arc::clone(&socket);
314 let pump = tokio::spawn(async move {
315 let mut buf = vec![0u8; 1500];
316 while let Ok(len) = pump_socket.recv(&mut buf).await {
317 let Ok(packet) = Packet::from_bytes(&buf[..len]) else {
318 continue;
319 };
320 if !dispatch(packet, &pending, &tx, &pump_socket).await {
321 break;
322 }
323 }
324 });
325
326 self.socket = Some(socket);
327 self.incoming = Some(rx);
328 self.pump = Some(pump);
329 Ok(())
330 }
331
332 async fn send(&mut self, topic: &str, payload: &[u8]) -> Result<()> {
333 let socket = self.socket.clone().ok_or(Error::Closed)?;
334 let id = self.next_message_id();
335 let token = self.next_request_token();
336
337 let mut packet = Packet::new();
338 packet.header.set_version(1);
339 packet
340 .header
341 .set_type(message_type(self.config.reliability));
342 packet.header.code = MessageClass::Request(RequestType::Put);
343 packet.header.message_id = id;
344 packet.set_token(token);
345 add_path(&mut packet, topic);
346 packet.payload = payload.to_vec();
347
348 let bytes = packet
349 .to_bytes()
350 .map_err(|err| Error::Codec(err.to_string()))?;
351
352 match self.config.reliability {
353 Reliability::NonConfirmable => socket
354 .send(&bytes)
355 .await
356 .map(|_| ())
357 .map_err(|err| Error::Transport(err.to_string())),
358 Reliability::Confirmable => self.send_confirmable(id, &bytes, &socket).await,
359 }
360 }
361
362 async fn subscribe(&mut self, topic: &str) -> Result<()> {
363 let socket = self.socket.clone().ok_or(Error::Closed)?;
364 let id = self.next_message_id();
365 let token = self.next_request_token();
366
367 let mut packet = Packet::new();
368 packet.header.set_version(1);
369 packet.header.set_type(MessageType::Confirmable);
370 packet.header.code = MessageClass::Request(RequestType::Get);
371 packet.header.message_id = id;
372 packet.set_token(token);
373 packet.add_option(CoapOption::Observe, Vec::new());
375 add_path(&mut packet, topic);
376
377 let bytes = packet
378 .to_bytes()
379 .map_err(|err| Error::Codec(err.to_string()))?;
380
381 self.send_confirmable(id, &bytes, &socket).await
382 }
383}
384
385fn message_type(reliability: Reliability) -> MessageType {
387 match reliability {
388 Reliability::NonConfirmable => MessageType::NonConfirmable,
389 Reliability::Confirmable => MessageType::Confirmable,
390 }
391}
392
393fn add_path(packet: &mut Packet, topic: &str) {
395 for segment in topic.split('/').filter(|segment| !segment.is_empty()) {
396 packet.add_option(CoapOption::UriPath, segment.as_bytes().to_vec());
397 }
398}
399
400fn path_from_packet(packet: &Packet) -> String {
402 match packet.get_option(CoapOption::UriPath) {
403 Some(segments) => segments
404 .iter()
405 .map(|segment| String::from_utf8_lossy(segment).into_owned())
406 .collect::<Vec<_>>()
407 .join("/"),
408 None => String::new(),
409 }
410}
411
412async fn dispatch(
414 packet: Packet,
415 pending: &PendingAcks,
416 tx: &mpsc::UnboundedSender<Message>,
417 socket: &UdpSocket,
418) -> bool {
419 match packet.header.get_type() {
420 MessageType::Acknowledgement => {
421 if let Some(waiter) = pending
422 .lock()
423 .expect("pending lock")
424 .remove(&packet.header.message_id)
425 {
426 let _ = waiter.send(());
427 }
428 if packet.get_option(CoapOption::Observe).is_some() {
430 return enqueue(packet, tx);
431 }
432 true
433 }
434 MessageType::Confirmable => {
435 acknowledge(&packet, socket).await;
436 enqueue(packet, tx)
437 }
438 MessageType::NonConfirmable => enqueue(packet, tx),
439 MessageType::Reset => {
440 if let Some(waiter) = pending
441 .lock()
442 .expect("pending lock")
443 .remove(&packet.header.message_id)
444 {
445 let _ = waiter.send(());
446 }
447 true
448 }
449 }
450}
451
452async fn acknowledge(packet: &Packet, socket: &UdpSocket) {
454 let mut ack = Packet::new();
455 ack.header.set_version(1);
456 ack.header.set_type(MessageType::Acknowledgement);
457 ack.header.code = MessageClass::Empty;
458 ack.header.message_id = packet.header.message_id;
459 if let Ok(bytes) = ack.to_bytes() {
460 let _ = socket.send(&bytes).await;
461 }
462}
463
464fn enqueue(packet: Packet, tx: &mpsc::UnboundedSender<Message>) -> bool {
466 let message = Message {
467 topic: path_from_packet(&packet),
468 payload: packet.payload,
469 };
470 tx.send(message).is_ok()
471}
472
473#[cfg(test)]
474mod tests {
475 use super::*;
476
477 #[test]
478 fn reliability_defaults_to_confirmable() {
479 let config = CoapConfig::new("localhost", 5683);
480 assert_eq!(config.reliability, Reliability::Confirmable);
481 }
482
483 #[test]
484 fn setters_update_the_configuration() {
485 let config = CoapConfig::new("localhost", 5683)
486 .reliability(Reliability::NonConfirmable)
487 .ack_timeout(Duration::from_millis(250))
488 .max_retransmits(1)
489 .bind("127.0.0.1:0");
490 assert_eq!(config.reliability, Reliability::NonConfirmable);
491 assert_eq!(config.ack_timeout, Duration::from_millis(250));
492 assert_eq!(config.max_retransmits, 1);
493 assert_eq!(config.bind, "127.0.0.1:0");
494 }
495
496 #[test]
497 fn path_round_trips_through_uri_path_options() {
498 let mut packet = Packet::new();
499 add_path(&mut packet, "sensors/1/temperature");
500 assert_eq!(path_from_packet(&packet), "sensors/1/temperature");
501 }
502
503 #[test]
504 fn leading_and_repeated_slashes_are_ignored() {
505 let mut packet = Packet::new();
506 add_path(&mut packet, "/sensors//1/");
507 assert_eq!(path_from_packet(&packet), "sensors/1");
508 }
509
510 #[tokio::test]
511 async fn send_before_connect_reports_closed() {
512 let mut transport = CoapTransport::new(CoapConfig::new("localhost", 5683));
513 assert!(matches!(
514 transport.send("t", b"x").await,
515 Err(Error::Closed)
516 ));
517 }
518
519 #[tokio::test]
520 async fn subscribe_before_connect_reports_closed() {
521 let mut transport = CoapTransport::new(CoapConfig::new("localhost", 5683));
522 assert!(matches!(transport.subscribe("t").await, Err(Error::Closed)));
523 }
524
525 #[tokio::test]
526 async fn recv_before_connect_reports_closed() {
527 let mut transport = CoapTransport::new(CoapConfig::new("localhost", 5683));
528 assert!(matches!(transport.recv().await, Err(Error::Closed)));
529 }
530
531 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
532 async fn confirmable_send_without_a_server_times_out() {
533 let config = CoapConfig::new("127.0.0.1", 1)
534 .ack_timeout(Duration::from_millis(20))
535 .max_retransmits(1);
536 let mut transport = CoapTransport::new(config);
537 transport.connect().await.expect("bind socket");
538 assert!(transport.send("sensors/1", b"x").await.is_err());
539 }
540}