use super::{ChannelError, IncomingMessage, MessageChannel};
use tokio::sync::{Mutex, watch};
#[derive(Debug, Clone)]
pub struct WatchChannel {
tx: watch::Sender<String>,
}
#[derive(Debug)]
pub struct WatchReceiver {
rx: Mutex<watch::Receiver<String>>,
}
impl WatchChannel {
pub fn new() -> Self {
let (tx, _rx) = watch::channel(String::new());
Self { tx }
}
pub fn subscribe(&self) -> WatchReceiver {
WatchReceiver {
rx: Mutex::new(self.tx.subscribe()),
}
}
}
impl Default for WatchChannel {
fn default() -> Self {
Self::new()
}
}
impl WatchReceiver {
pub async fn recv(&self) -> Result<IncomingMessage, ChannelError> {
let mut rx = self.rx.lock().await;
rx.changed().await.map_err(|_| ChannelError::Closed)?;
let message = rx.borrow().clone();
Ok(IncomingMessage {
text: message,
reply_tx: None,
})
}
}
#[async_trait::async_trait]
impl MessageChannel for WatchChannel {
async fn ask(&self, _message: &str) -> Result<String, ChannelError> {
Err(ChannelError::NotSupported(
"watch channel does not support ask".into(),
))
}
async fn notify(&self, message: &str) -> Result<(), ChannelError> {
self.tx
.send(message.to_string())
.map_err(|_| ChannelError::NoReceiver)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn notify_then_recv_latest() {
let channel = WatchChannel::new();
let receiver = channel.subscribe();
channel.notify("v1").await.unwrap();
assert_eq!(receiver.recv().await.unwrap().text(), "v1");
channel.notify("v2").await.unwrap();
assert_eq!(receiver.recv().await.unwrap().text(), "v2");
}
#[tokio::test]
async fn notify_before_recv_wakes_all_observers() {
let channel = WatchChannel::new();
let sub_a = channel.subscribe();
let sub_b = channel.subscribe();
channel.notify("start").await.unwrap();
let (a, b) = tokio::join!(sub_a.recv(), sub_b.recv());
assert_eq!(a.unwrap().text(), "start");
assert_eq!(b.unwrap().text(), "start");
}
#[tokio::test]
async fn ask_not_supported() {
let channel = WatchChannel::new();
assert!(matches!(
channel.ask("q").await,
Err(ChannelError::NotSupported(_))
));
}
}