use super::{AgentEvent, EventChannel, EventChannelStats, EventReceiver};
use futures::future::BoxFuture;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
#[derive(Debug)]
pub struct MpscEventChannel {
tx: tokio::sync::mpsc::Sender<Arc<dyn AgentEvent>>,
rx: std::sync::Mutex<Option<tokio::sync::mpsc::Receiver<Arc<dyn AgentEvent>>>>,
stats: Arc<MpscStats>,
}
#[derive(Debug, Default)]
struct MpscStats {
published: AtomicU64,
delivered: AtomicU64,
dropped_no_subscribers: AtomicU64,
dropped_full: AtomicU64,
subscribed: AtomicBool,
}
impl Default for MpscEventChannel {
fn default() -> Self {
Self::new(256)
}
}
impl MpscEventChannel {
pub fn new(capacity: usize) -> Self {
let (tx, rx) = tokio::sync::mpsc::channel(capacity);
Self {
tx,
rx: std::sync::Mutex::new(Some(rx)),
stats: Arc::new(MpscStats::default()),
}
}
}
impl EventChannel for MpscEventChannel {
fn publish(&self, event: Arc<dyn AgentEvent>) {
self.stats.published.fetch_add(1, Ordering::Relaxed);
match self.tx.try_send(event) {
Ok(()) => {
self.stats.delivered.fetch_add(1, Ordering::Relaxed);
}
Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
self.stats.dropped_full.fetch_add(1, Ordering::Relaxed);
}
Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => {
self.stats
.dropped_no_subscribers
.fetch_add(1, Ordering::Relaxed);
}
}
}
fn subscribe(&self) -> Box<dyn EventReceiver> {
let mut guard = self
.rx
.lock()
.expect("MpscEventChannel is single-consumer; subscribe may only be called once");
let rx = guard
.take()
.expect("MpscEventChannel is single-consumer; subscribe may only be called once");
self.stats.subscribed.store(true, Ordering::Relaxed);
Box::new(MpscEventReceiver { rx })
}
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: self.stats.dropped_full.load(Ordering::Relaxed),
lagged: 0,
subscribers: usize::from(
self.stats.subscribed.load(Ordering::Relaxed) && !self.tx.is_closed(),
),
}
}
}
struct MpscEventReceiver {
rx: tokio::sync::mpsc::Receiver<Arc<dyn AgentEvent>>,
}
impl EventReceiver for MpscEventReceiver {
fn recv(&mut self) -> BoxFuture<'_, Option<Arc<dyn AgentEvent>>> {
Box::pin(async move { self.rx.recv().await })
}
}
#[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_buffered() {
let ch = MpscEventChannel::new(16);
ch.publish(ev("1"));
ch.publish(ev("2"));
let mut rx = ch.subscribe();
assert_eq!(input_of(&*rx.recv().await.unwrap()), "1");
assert_eq!(input_of(&*rx.recv().await.unwrap()), "2");
}
#[tokio::test]
#[should_panic(expected = "MpscEventChannel is single-consumer")]
async fn subscribe_twice_panics() {
let ch = MpscEventChannel::new(8);
let _rx = ch.subscribe();
let _rx2 = ch.subscribe();
}
#[tokio::test]
async fn full_drops_new() {
let ch = MpscEventChannel::new(2);
ch.publish(ev("1"));
ch.publish(ev("2"));
ch.publish(ev("3")); let stats = ch.stats();
assert_eq!(stats.published, 3);
assert_eq!(stats.delivered, 2);
assert_eq!(stats.dropped_full, 1);
let mut rx = ch.subscribe();
assert_eq!(input_of(&*rx.recv().await.unwrap()), "1");
assert_eq!(input_of(&*rx.recv().await.unwrap()), "2");
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 = MpscEventChannel::new(16);
ch.publish(ev("1"));
ch.subscribe()
};
assert_eq!(input_of(&*rx.recv().await.unwrap()), "1");
assert!(rx.recv().await.is_none());
}
}