use std::sync::{Arc, Mutex};
use tokio::sync::mpsc::{self, UnboundedReceiver};
use pamoja_core::{Error, Result, Transport};
use crate::broker::LoopbackBroker;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Message {
pub topic: String,
pub payload: Vec<u8>,
}
pub struct LoopbackTransport {
broker: LoopbackBroker,
filters: Arc<Mutex<Vec<String>>>,
incoming: Option<UnboundedReceiver<Message>>,
}
impl LoopbackTransport {
pub fn new(broker: LoopbackBroker) -> Self {
Self {
broker,
filters: Arc::new(Mutex::new(Vec::new())),
incoming: None,
}
}
pub fn is_connected(&self) -> bool {
self.incoming.is_some()
}
pub async fn recv(&mut self) -> Result<Option<Message>> {
let incoming = self.incoming.as_mut().ok_or(Error::Closed)?;
Ok(incoming.recv().await)
}
pub fn disconnect(&mut self) {
self.incoming = None;
}
}
impl Transport for LoopbackTransport {
async fn connect(&mut self) -> Result<()> {
let (sender, receiver) = mpsc::unbounded_channel();
self.broker.register(Arc::clone(&self.filters), sender);
self.incoming = Some(receiver);
Ok(())
}
async fn send(&mut self, topic: &str, payload: &[u8]) -> Result<()> {
if self.incoming.is_none() {
return Err(Error::Closed);
}
self.broker.publish(&Message {
topic: topic.to_owned(),
payload: payload.to_vec(),
});
Ok(())
}
async fn subscribe(&mut self, topic: &str) -> Result<()> {
if self.incoming.is_none() {
return Err(Error::Closed);
}
self.filters
.lock()
.expect("filters lock")
.push(topic.to_owned());
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn publish_and_subscribe_round_trip() {
let broker = LoopbackBroker::new();
let mut subscriber = LoopbackTransport::new(broker.clone());
let mut publisher = LoopbackTransport::new(broker);
subscriber.connect().await.expect("connect");
publisher.connect().await.expect("connect");
subscriber
.subscribe("sensors/+/temperature")
.await
.expect("subscribe");
publisher
.send("sensors/1/temperature", b"21.5")
.await
.expect("send");
let message = subscriber.recv().await.expect("recv").expect("a message");
assert_eq!(message.topic, "sensors/1/temperature");
assert_eq!(message.payload, b"21.5");
}
#[tokio::test]
async fn non_matching_topics_are_not_delivered() {
let broker = LoopbackBroker::new();
let mut subscriber = LoopbackTransport::new(broker.clone());
let mut publisher = LoopbackTransport::new(broker);
subscriber.connect().await.expect("connect");
publisher.connect().await.expect("connect");
subscriber
.subscribe("sensors/1/#")
.await
.expect("subscribe");
publisher
.send("sensors/2/temperature", b"x")
.await
.expect("send");
publisher
.send("sensors/1/humidity", b"y")
.await
.expect("send");
let message = subscriber.recv().await.expect("recv").expect("a message");
assert_eq!(message.topic, "sensors/1/humidity");
}
#[tokio::test]
async fn operations_before_connect_report_closed() {
let broker = LoopbackBroker::new();
let mut transport = LoopbackTransport::new(broker);
assert!(matches!(
transport.send("t", b"x").await,
Err(Error::Closed)
));
assert!(matches!(transport.subscribe("t").await, Err(Error::Closed)));
assert!(matches!(transport.recv().await, Err(Error::Closed)));
assert!(!transport.is_connected());
}
}