#![expect(
clippy::unwrap_used,
reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
)]
use std::collections::VecDeque;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use super::subscribers::{SubscriberRegistry, wake_drained};
pub struct Broadcast<T> {
_phantom: std::marker::PhantomData<T>,
}
struct BroadcastState<T> {
messages: VecDeque<(u64, T)>,
sequence: u64,
closed: bool,
subscribers: SubscriberRegistry<u64>,
capacity: usize,
}
impl<T: Clone + Send + 'static> Broadcast<T> {
#[allow(clippy::new_ret_no_self)] pub fn new(capacity: usize) -> (BroadcastSender<T>, BroadcastReceiver<T>) {
let state = Arc::new(Mutex::new(BroadcastState {
messages: VecDeque::new(),
sequence: 0,
closed: false,
subscribers: SubscriberRegistry::with_initial(0),
capacity,
}));
let sender = BroadcastSender {
state: state.clone(),
};
let receiver = BroadcastReceiver {
state: state.clone(),
id: 0,
position: 0,
};
(sender, receiver)
}
}
pub struct BroadcastSender<T> {
state: Arc<Mutex<BroadcastState<T>>>,
}
impl<T: Clone> BroadcastSender<T> {
pub fn send(&self, message: T) -> Result<usize, BroadcastError> {
let (receiver_count, wakers) = {
let mut state = self.state.lock().unwrap();
if state.closed {
return Err(BroadcastError::Closed);
}
state.sequence += 1;
let sequence = state.sequence;
state.messages.push_back((sequence, message));
let receiver_count = state.subscribers.len();
let wakers = state.subscribers.drain_wakers();
let min_position = state
.subscribers
.cursors()
.copied()
.min()
.unwrap_or(sequence);
while state.messages.len() > state.capacity
|| state
.messages
.front()
.is_some_and(|(seq, _)| *seq <= min_position)
{
state.messages.pop_front();
}
(receiver_count, wakers)
};
wake_drained(wakers);
Ok(receiver_count)
}
pub fn receiver_count(&self) -> usize {
self.state.lock().unwrap().subscribers.len()
}
}
impl<T> Drop for BroadcastSender<T> {
fn drop(&mut self) {
let wakers = {
let mut state = self.state.lock().unwrap();
state.closed = true;
state.subscribers.drain_wakers()
};
wake_drained(wakers);
}
}
pub struct BroadcastReceiver<T> {
state: Arc<Mutex<BroadcastState<T>>>,
id: u64,
position: u64,
}
impl<T: Clone> BroadcastReceiver<T> {
pub fn recv(&mut self) -> BroadcastRecv<'_, T> {
BroadcastRecv { receiver: self }
}
pub fn try_recv(&mut self) -> Result<T, BroadcastError> {
let mut state = self.state.lock().unwrap();
if state.messages.is_empty() {
if state.closed {
return Err(BroadcastError::Closed);
}
return Err(BroadcastError::Empty);
}
let oldest_seq = state.messages.front().unwrap().0;
if self.position + 1 < oldest_seq {
self.position = oldest_seq - 1;
if let Some(subscriber) = state.subscribers.get_mut(self.id) {
subscriber.cursor = self.position;
}
return Err(BroadcastError::Lagged);
}
let offset = usize::try_from(self.position + 1 - oldest_seq)
.expect("invariant: unread offset is bounded by the message queue length");
let found = state.messages.get(offset).map(|(seq, message)| {
debug_assert_eq!(*seq, self.position + 1, "broadcast sequences must be dense");
self.position = *seq;
message.clone()
});
if let Some(message) = found {
if let Some(subscriber) = state.subscribers.get_mut(self.id) {
subscriber.cursor = self.position;
}
return Ok(message);
}
if state.closed {
Err(BroadcastError::Closed)
} else {
Err(BroadcastError::Empty)
}
}
pub fn resubscribe(&self) -> BroadcastReceiver<T> {
let mut state = self.state.lock().unwrap();
let current_sequence = state.sequence;
let id = state.subscribers.register(current_sequence);
BroadcastReceiver {
state: self.state.clone(),
id,
position: current_sequence,
}
}
pub fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Result<T, BroadcastError>> {
let mut state = self.state.lock().unwrap();
if state.messages.is_empty() {
if state.closed {
return Poll::Ready(Err(BroadcastError::Closed));
}
if let Some(subscriber) = state.subscribers.get_mut(self.id) {
subscriber.waker = Some(cx.waker().clone());
}
return Poll::Pending;
}
let oldest_seq = state.messages.front().unwrap().0;
if self.position + 1 < oldest_seq {
self.position = oldest_seq - 1;
if let Some(subscriber) = state.subscribers.get_mut(self.id) {
subscriber.cursor = self.position;
}
return Poll::Ready(Err(BroadcastError::Lagged));
}
let offset = usize::try_from(self.position + 1 - oldest_seq)
.expect("invariant: unread offset is bounded by the message queue length");
let found_msg = state.messages.get(offset).map(|(seq, message)| {
debug_assert_eq!(*seq, self.position + 1, "broadcast sequences must be dense");
self.position = *seq;
(*seq, message.clone())
});
if let Some((_, message)) = found_msg {
if let Some(subscriber) = state.subscribers.get_mut(self.id) {
subscriber.cursor = self.position;
}
Poll::Ready(Ok(message))
} else if state.closed {
Poll::Ready(Err(BroadcastError::Closed))
} else {
if let Some(subscriber) = state.subscribers.get_mut(self.id) {
subscriber.waker = Some(cx.waker().clone());
}
Poll::Pending
}
}
}
impl<T: Clone> Clone for BroadcastReceiver<T> {
fn clone(&self) -> Self {
self.resubscribe()
}
}
impl<T> Drop for BroadcastReceiver<T> {
fn drop(&mut self) {
if let Ok(mut state) = self.state.lock() {
state.subscribers.remove(self.id);
}
}
}
pub struct BroadcastRecv<'a, T> {
receiver: &'a mut BroadcastReceiver<T>,
}
impl<'a, T: Clone> Future for BroadcastRecv<'a, T> {
type Output = Result<T, BroadcastError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.receiver.poll_recv(cx)
}
}
impl<'a, T> Drop for BroadcastRecv<'a, T> {
fn drop(&mut self) {
if let Ok(mut state) = self.receiver.state.lock() {
let id = self.receiver.id;
if let Some(subscriber) = state.subscribers.get_mut(id) {
subscriber.waker = None;
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum BroadcastError {
Empty,
Closed,
Lagged,
}
impl std::fmt::Display for BroadcastError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
BroadcastError::Empty => write!(f, "broadcast channel is empty"),
BroadcastError::Closed => write!(f, "broadcast channel is closed"),
BroadcastError::Lagged => write!(f, "broadcast channel lagged"),
}
}
}
impl std::error::Error for BroadcastError {}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{Wake, Waker};
struct CountingWake(Arc<AtomicUsize>);
impl Wake for CountingWake {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::Release);
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.fetch_add(1, Ordering::Release);
}
}
#[test]
fn cancelled_recv_clears_waker_and_is_not_spuriously_woken() {
let (tx, mut rx) = Broadcast::<u32>::new(8);
let count = Arc::new(AtomicUsize::new(0));
let waker = Waker::from(Arc::new(CountingWake(Arc::clone(&count))));
let mut cx = Context::from_waker(&waker);
{
let mut fut = rx.recv();
assert!(Pin::new(&mut fut).poll(&mut cx).is_pending());
}
tx.send(42).expect("send must succeed");
assert_eq!(
count.load(Ordering::Acquire),
0,
"a cancelled recv future must not be spuriously woken"
);
assert_eq!(rx.try_recv(), Ok(42));
}
#[test]
fn live_recv_is_woken_on_send() {
let (tx, mut rx) = Broadcast::<u32>::new(8);
let count = Arc::new(AtomicUsize::new(0));
let waker = Waker::from(Arc::new(CountingWake(Arc::clone(&count))));
let mut cx = Context::from_waker(&waker);
let mut fut = rx.recv();
assert!(Pin::new(&mut fut).poll(&mut cx).is_pending());
tx.send(7).expect("send must succeed");
assert_eq!(
count.load(Ordering::Acquire),
1,
"a live recv future must be woken by send"
);
drop(fut);
}
#[test]
fn messages_read_by_every_receiver_are_reclaimed_on_next_send() {
let (tx, mut first) = Broadcast::<u32>::new(8);
let mut second = first.resubscribe();
tx.send(10).expect("first send must succeed");
tx.send(20).expect("second send must succeed");
assert_eq!(first.try_recv(), Ok(10));
assert_eq!(second.try_recv(), Ok(10));
assert_eq!(first.try_recv(), Ok(20));
assert_eq!(second.try_recv(), Ok(20));
tx.send(30).expect("third send must succeed");
let state = tx.state.lock().expect("broadcast state must not poison");
assert_eq!(
state.messages.iter().copied().collect::<Vec<_>>(),
vec![(3, 30)],
"the next send must retain only the new unread message"
);
}
}