use std::fmt;
use std::pin::Pin;
use std::sync::Arc;
use futures::Stream;
use parking_lot::Mutex;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use tokio_util::sync::CancellationToken;
pub type BroadcastStream<T> = Pin<Box<dyn Stream<Item = T> + Send + 'static>>;
pub const DEFAULT_BROADCAST_BUFFER: usize = 64;
pub struct Broadcaster<T> {
senders: Arc<Mutex<Vec<mpsc::Sender<T>>>>,
buffer: usize,
}
impl<T> Clone for Broadcaster<T> {
fn clone(&self) -> Self {
Self {
senders: Arc::clone(&self.senders),
buffer: self.buffer,
}
}
}
impl<T> fmt::Debug for Broadcaster<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Broadcaster")
.field("buffer", &self.buffer)
.field("subscribers", &self.subscriber_count())
.finish_non_exhaustive()
}
}
impl<T> Default for Broadcaster<T>
where
T: Clone + Send + 'static,
{
fn default() -> Self {
Self::new()
}
}
impl<T> Broadcaster<T>
where
T: Clone + Send + 'static,
{
#[must_use]
pub fn new() -> Self {
Self::with_buffer(DEFAULT_BROADCAST_BUFFER)
}
#[must_use]
pub fn with_buffer(buffer: usize) -> Self {
Self {
senders: Arc::new(Mutex::new(Vec::new())),
buffer: buffer.max(1),
}
}
}
impl<T> Broadcaster<T> {
#[must_use]
pub const fn buffer(&self) -> usize {
self.buffer
}
#[must_use]
pub fn subscriber_count(&self) -> usize {
self.senders.lock().len()
}
pub fn subscribe(&self, cancel: CancellationToken) -> BroadcastStream<T>
where
T: Send + 'static,
{
use futures::StreamExt as _;
let (tx, rx) = mpsc::channel(self.buffer);
self.senders.lock().push(tx);
Box::pin(ReceiverStream::new(rx).take_until(cancel.cancelled_owned()))
}
pub fn broadcast(&self, item: &T)
where
T: Clone,
{
self.senders
.lock()
.retain(|tx| match tx.try_send(item.clone()) {
Ok(()) | Err(mpsc::error::TrySendError::Full(_)) => true,
Err(mpsc::error::TrySendError::Closed(_)) => false,
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::StreamExt as _;
#[tokio::test]
async fn delivers_to_all_subscribers() {
let bc = Broadcaster::<u32>::new();
let cancel = CancellationToken::new();
let mut a = bc.subscribe(cancel.clone());
let mut b = bc.subscribe(cancel.clone());
bc.broadcast(&7);
assert_eq!(a.next().await, Some(7));
assert_eq!(b.next().await, Some(7));
}
#[tokio::test]
async fn stream_terminates_on_cancel() {
let bc = Broadcaster::<u32>::new();
let cancel = CancellationToken::new();
let mut sub = bc.subscribe(cancel.clone());
bc.broadcast(&1);
assert_eq!(sub.next().await, Some(1));
cancel.cancel();
assert_eq!(sub.next().await, None);
}
#[tokio::test]
async fn dropped_subscriber_is_pruned() {
let bc = Broadcaster::<u32>::new();
let cancel = CancellationToken::new();
let sub = bc.subscribe(cancel);
assert_eq!(bc.subscriber_count(), 1);
drop(sub);
bc.broadcast(&1);
assert_eq!(bc.subscriber_count(), 0);
}
#[tokio::test]
async fn full_subscriber_drops_overflow_without_blocking() {
let bc = Broadcaster::<u32>::with_buffer(2);
let cancel = CancellationToken::new();
let mut sub = bc.subscribe(cancel.clone());
for i in 0..5 {
bc.broadcast(&i);
}
let mut received = Vec::new();
while let Ok(Some(v)) =
tokio::time::timeout(std::time::Duration::from_millis(50), sub.next()).await
{
received.push(v);
}
assert_eq!(received, vec![0, 1]);
}
#[tokio::test]
async fn zero_buffer_is_clamped_to_one() {
let bc = Broadcaster::<u32>::with_buffer(0);
assert_eq!(bc.buffer(), 1);
}
#[tokio::test]
async fn clones_share_subscriber_set() {
let bc = Broadcaster::<u32>::new();
let cancel = CancellationToken::new();
let mut sub = bc.subscribe(cancel.clone());
let clone = bc.clone();
clone.broadcast(&42);
assert_eq!(sub.next().await, Some(42));
}
}