use super::{AgentEvent, EventChannel, EventReceiver};
use futures::future::BoxFuture;
use std::sync::Arc;
#[derive(Debug, Clone)]
pub struct BroadcastEventChannel {
tx: tokio::sync::broadcast::Sender<Arc<dyn AgentEvent>>,
}
impl Default for BroadcastEventChannel {
fn default() -> Self {
Self::new(256)
}
}
impl BroadcastEventChannel {
pub fn new(capacity: usize) -> Self {
let (tx, _rx) = tokio::sync::broadcast::channel(capacity);
Self { tx }
}
}
impl EventChannel for BroadcastEventChannel {
fn publish(&self, event: Arc<dyn AgentEvent>) {
let _ = self.tx.send(event);
}
fn subscribe(&self) -> Box<dyn EventReceiver> {
Box::new(BroadcastEventReceiver {
rx: self.tx.subscribe(),
})
}
}
struct BroadcastEventReceiver {
rx: tokio::sync::broadcast::Receiver<Arc<dyn AgentEvent>>,
}
impl EventReceiver for BroadcastEventReceiver {
fn recv(&mut self) -> BoxFuture<'_, Option<Arc<dyn AgentEvent>>> {
Box::pin(async move {
loop {
match self.rx.recv().await {
Ok(event) => return Some(event),
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue,
Err(_) => return None,
}
}
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::agent::{AgentEvent, ReActEvent};
use futures::FutureExt;
fn ev(n: &str) -> Arc<dyn AgentEvent> {
Arc::new(ReActEvent::RunStarted {
run_id: "test".into(),
input: n.into(),
})
}
fn input_of(ev: &dyn AgentEvent) -> &str {
match ev.as_any().downcast_ref::<ReActEvent>() {
Some(ReActEvent::RunStarted { input, .. }) => input.as_str(),
_ => panic!("test event must be RunStarted"),
}
}
#[tokio::test]
async fn publish_subscribe() {
let ch = BroadcastEventChannel::new(16);
let mut rx = ch.subscribe();
ch.publish(ev("1"));
ch.publish(ev("2"));
ch.publish(ev("3"));
assert_eq!(input_of(&*rx.recv().await.unwrap()), "1");
assert_eq!(input_of(&*rx.recv().await.unwrap()), "2");
assert_eq!(input_of(&*rx.recv().await.unwrap()), "3");
}
#[tokio::test]
async fn multiple_subscribers_each_receive_all() {
let ch = BroadcastEventChannel::new(16);
let mut rx1 = ch.subscribe();
let mut rx2 = ch.subscribe();
ch.publish(ev("1"));
ch.publish(ev("2"));
assert_eq!(input_of(&*rx1.recv().await.unwrap()), "1");
assert_eq!(input_of(&*rx2.recv().await.unwrap()), "1");
assert_eq!(input_of(&*rx1.recv().await.unwrap()), "2");
assert_eq!(input_of(&*rx2.recv().await.unwrap()), "2");
}
#[tokio::test]
async fn lag_drops_oldest() {
let ch = BroadcastEventChannel::new(2);
let mut rx = ch.subscribe();
ch.publish(ev("1"));
ch.publish(ev("2"));
ch.publish(ev("3"));
let first = input_of(&*rx.recv().await.unwrap()).to_string();
let second = input_of(&*rx.recv().await.unwrap()).to_string();
assert_eq!(second, "3"); assert!(first == "2" || first == "3"); }
#[tokio::test]
async fn events_before_subscribe_missed() {
let ch = BroadcastEventChannel::new(16);
ch.publish(ev("1"));
let mut rx = ch.subscribe();
assert!(rx.recv().now_or_never().is_none());
}
#[tokio::test]
async fn closed_ends_stream() {
let mut rx = {
let ch = BroadcastEventChannel::new(16);
let rx = ch.subscribe();
ch.publish(ev("1"));
rx
};
assert_eq!(input_of(&*rx.recv().await.unwrap()), "1");
assert!(rx.recv().await.is_none());
}
}