use super::{AgentEvent, EventChannel, EventChannelStats, EventReceiver};
use futures::future::BoxFuture;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
#[derive(Debug, Clone)]
pub struct BroadcastEventChannel {
tx: tokio::sync::broadcast::Sender<Arc<dyn AgentEvent>>,
stats: Arc<BroadcastStats>,
}
#[derive(Debug, Default)]
struct BroadcastStats {
published: AtomicU64,
delivered: AtomicU64,
dropped_no_subscribers: AtomicU64,
lagged: AtomicU64,
}
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,
stats: Arc::new(BroadcastStats::default()),
}
}
}
impl EventChannel for BroadcastEventChannel {
fn publish(&self, event: Arc<dyn AgentEvent>) {
self.stats.published.fetch_add(1, Ordering::Relaxed);
match self.tx.send(event) {
Ok(receivers) => {
self.stats
.delivered
.fetch_add(receivers as u64, Ordering::Relaxed);
}
Err(_) => {
self.stats
.dropped_no_subscribers
.fetch_add(1, Ordering::Relaxed);
}
}
}
fn subscribe(&self) -> Box<dyn EventReceiver> {
Box::new(BroadcastEventReceiver {
rx: self.tx.subscribe(),
stats: self.stats.clone(),
})
}
fn stats(&self) -> EventChannelStats {
EventChannelStats {
published: self.stats.published.load(Ordering::Relaxed),
delivered: self.stats.delivered.load(Ordering::Relaxed),
dropped_no_subscribers: self.stats.dropped_no_subscribers.load(Ordering::Relaxed),
dropped_full: 0,
lagged: self.stats.lagged.load(Ordering::Relaxed),
subscribers: self.tx.receiver_count(),
}
}
}
struct BroadcastEventReceiver {
rx: tokio::sync::broadcast::Receiver<Arc<dyn AgentEvent>>,
stats: Arc<BroadcastStats>,
}
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(n)) => {
self.stats.lagged.fetch_add(n, Ordering::Relaxed);
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_text().expect("test event uses text input")
}
_ => 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"); assert!(ch.stats().lagged >= 1);
}
#[tokio::test]
async fn events_before_subscribe_missed() {
let ch = BroadcastEventChannel::new(16);
ch.publish(ev("1"));
let stats = ch.stats();
assert_eq!(stats.published, 1);
assert_eq!(stats.dropped_no_subscribers, 1);
let mut rx = ch.subscribe();
assert!(rx.recv().now_or_never().is_none());
assert_eq!(ch.stats().subscribers, 1);
}
#[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());
}
}