Skip to main content

moirai_async/sync/
broadcast.rs

1//! Broadcast channel for one-to-many async communication
2//!
3//! Provides broadcast channel implementation that allows one sender to
4//! broadcast messages to multiple receivers with SLAP-compliant design.
5
6#![expect(
7    clippy::unwrap_used,
8    reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
9)]
10
11use std::collections::VecDeque;
12use std::future::Future;
13use std::pin::Pin;
14use std::sync::{Arc, Mutex};
15use std::task::{Context, Poll};
16
17use super::subscribers::{SubscriberRegistry, wake_drained};
18
19/// Broadcast channel for one-to-many communication
20pub struct Broadcast<T> {
21    _phantom: std::marker::PhantomData<T>,
22}
23
24struct BroadcastState<T> {
25    messages: VecDeque<(u64, T)>,
26    sequence: u64,
27    closed: bool,
28    /// One slot per receiver, shared with `Watch` (see `subscribers`). The
29    /// cursor is the last sequence that receiver consumed, which doubles as the
30    /// input to the retention boundary in `send`.
31    subscribers: SubscriberRegistry<u64>,
32    capacity: usize,
33}
34
35impl<T: Clone + Send + 'static> Broadcast<T> {
36    /// Create a new broadcast channel with the given capacity
37    /// Returns (sender, receiver) tuple per channel pattern conventions
38    #[allow(clippy::new_ret_no_self)] // Standard channel pattern per Rust Book Ch.16
39    pub fn new(capacity: usize) -> (BroadcastSender<T>, BroadcastReceiver<T>) {
40        let state = Arc::new(Mutex::new(BroadcastState {
41            messages: VecDeque::new(),
42            sequence: 0,
43            closed: false,
44            subscribers: SubscriberRegistry::with_initial(0),
45            capacity,
46        }));
47
48        let sender = BroadcastSender {
49            state: state.clone(),
50        };
51
52        let receiver = BroadcastReceiver {
53            state: state.clone(),
54            id: 0,
55            position: 0,
56        };
57
58        (sender, receiver)
59    }
60}
61
62/// Sender half of broadcast channel
63pub struct BroadcastSender<T> {
64    state: Arc<Mutex<BroadcastState<T>>>,
65}
66
67impl<T: Clone> BroadcastSender<T> {
68    /// Send a message to all receivers
69    pub fn send(&self, message: T) -> Result<usize, BroadcastError> {
70        let (receiver_count, wakers) = {
71            let mut state = self.state.lock().unwrap();
72            if state.closed {
73                return Err(BroadcastError::Closed);
74            }
75
76            // Add new message first, then reclaim messages already read by every
77            // receiver.  The retention boundary is the minimum position across
78            // all live receivers; messages with sequence <= that boundary have
79            // been consumed by everyone and can be dropped.
80            state.sequence += 1;
81            let sequence = state.sequence;
82            state.messages.push_back((sequence, message));
83
84            // Publication is done; take the registered wakers. They observe
85            // either the appended message or the explicit lag state retained by
86            // the capacity contract. Woken only after the lock is released —
87            // see `drain_wakers`.
88            let receiver_count = state.subscribers.len();
89            let wakers = state.subscribers.drain_wakers();
90
91            // `min` over the cursors, copied out before the mutation below so
92            // the immutable borrow of the registry ends here.
93            let min_position = state
94                .subscribers
95                .cursors()
96                .copied()
97                .min()
98                .unwrap_or(sequence);
99            while state.messages.len() > state.capacity
100                || state
101                    .messages
102                    .front()
103                    .is_some_and(|(seq, _)| *seq <= min_position)
104            {
105                state.messages.pop_front();
106            }
107
108            (receiver_count, wakers)
109        };
110        wake_drained(wakers);
111
112        Ok(receiver_count)
113    }
114
115    /// Get the number of active receivers
116    pub fn receiver_count(&self) -> usize {
117        self.state.lock().unwrap().subscribers.len()
118    }
119}
120
121impl<T> Drop for BroadcastSender<T> {
122    fn drop(&mut self) {
123        let wakers = {
124            let mut state = self.state.lock().unwrap();
125            state.closed = true;
126            state.subscribers.drain_wakers()
127        };
128        wake_drained(wakers);
129    }
130}
131
132/// Receiver half of broadcast channel
133pub struct BroadcastReceiver<T> {
134    state: Arc<Mutex<BroadcastState<T>>>,
135    id: u64,
136    position: u64,
137}
138
139impl<T: Clone> BroadcastReceiver<T> {
140    /// Receive the next message
141    pub fn recv(&mut self) -> BroadcastRecv<'_, T> {
142        BroadcastRecv { receiver: self }
143    }
144
145    /// Try to receive a message immediately
146    pub fn try_recv(&mut self) -> Result<T, BroadcastError> {
147        let mut state = self.state.lock().unwrap();
148
149        if state.messages.is_empty() {
150            if state.closed {
151                return Err(BroadcastError::Closed);
152            }
153            return Err(BroadcastError::Empty);
154        }
155
156        // Lagging check
157        let oldest_seq = state.messages.front().unwrap().0;
158        if self.position + 1 < oldest_seq {
159            self.position = oldest_seq - 1;
160            if let Some(subscriber) = state.subscribers.get_mut(self.id) {
161                subscriber.cursor = self.position;
162            }
163            return Err(BroadcastError::Lagged);
164        }
165
166        // Sequences are dense (each `send` appends `sequence + 1` and only
167        // `pop_front` removes), so the next unread message sits at a directly
168        // computable offset from the front — no scan. After the lag check,
169        // `position + 1 >= oldest_seq`, and `position <= sequence` bounds the
170        // offset by the queue length, so the conversion cannot truncate.
171        let offset = usize::try_from(self.position + 1 - oldest_seq)
172            .expect("invariant: unread offset is bounded by the message queue length");
173        let found = state.messages.get(offset).map(|(seq, message)| {
174            debug_assert_eq!(*seq, self.position + 1, "broadcast sequences must be dense");
175            self.position = *seq;
176            message.clone()
177        });
178        if let Some(message) = found {
179            if let Some(subscriber) = state.subscribers.get_mut(self.id) {
180                subscriber.cursor = self.position;
181            }
182            return Ok(message);
183        }
184
185        if state.closed {
186            Err(BroadcastError::Closed)
187        } else {
188            Err(BroadcastError::Empty)
189        }
190    }
191
192    /// Clone this receiver to create a new independent receiver
193    pub fn resubscribe(&self) -> BroadcastReceiver<T> {
194        let mut state = self.state.lock().unwrap();
195        let current_sequence = state.sequence;
196        let id = state.subscribers.register(current_sequence);
197
198        BroadcastReceiver {
199            state: self.state.clone(),
200            id,
201            position: current_sequence,
202        }
203    }
204
205    /// Poll to receive a message, registering waker if empty.
206    pub fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Result<T, BroadcastError>> {
207        let mut state = self.state.lock().unwrap();
208
209        if state.messages.is_empty() {
210            if state.closed {
211                return Poll::Ready(Err(BroadcastError::Closed));
212            }
213            if let Some(subscriber) = state.subscribers.get_mut(self.id) {
214                subscriber.waker = Some(cx.waker().clone());
215            }
216            return Poll::Pending;
217        }
218
219        // Lagging check
220        let oldest_seq = state.messages.front().unwrap().0;
221        if self.position + 1 < oldest_seq {
222            self.position = oldest_seq - 1;
223            if let Some(subscriber) = state.subscribers.get_mut(self.id) {
224                subscriber.cursor = self.position;
225            }
226            return Poll::Ready(Err(BroadcastError::Lagged));
227        }
228
229        // Dense-sequence direct index; see `try_recv` for the derivation.
230        let offset = usize::try_from(self.position + 1 - oldest_seq)
231            .expect("invariant: unread offset is bounded by the message queue length");
232        let found_msg = state.messages.get(offset).map(|(seq, message)| {
233            debug_assert_eq!(*seq, self.position + 1, "broadcast sequences must be dense");
234            self.position = *seq;
235            (*seq, message.clone())
236        });
237
238        if let Some((_, message)) = found_msg {
239            if let Some(subscriber) = state.subscribers.get_mut(self.id) {
240                subscriber.cursor = self.position;
241            }
242            Poll::Ready(Ok(message))
243        } else if state.closed {
244            Poll::Ready(Err(BroadcastError::Closed))
245        } else {
246            if let Some(subscriber) = state.subscribers.get_mut(self.id) {
247                subscriber.waker = Some(cx.waker().clone());
248            }
249            Poll::Pending
250        }
251    }
252}
253
254impl<T: Clone> Clone for BroadcastReceiver<T> {
255    fn clone(&self) -> Self {
256        self.resubscribe()
257    }
258}
259
260impl<T> Drop for BroadcastReceiver<T> {
261    fn drop(&mut self) {
262        if let Ok(mut state) = self.state.lock() {
263            state.subscribers.remove(self.id);
264        }
265    }
266}
267
268/// Future for receiving from broadcast channel
269pub struct BroadcastRecv<'a, T> {
270    receiver: &'a mut BroadcastReceiver<T>,
271}
272
273impl<'a, T: Clone> Future for BroadcastRecv<'a, T> {
274    type Output = Result<T, BroadcastError>;
275
276    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
277        self.receiver.poll_recv(cx)
278    }
279}
280
281impl<'a, T> Drop for BroadcastRecv<'a, T> {
282    fn drop(&mut self) {
283        // Clear the waker registered by `poll_recv` so a cancelled recv future
284        // does not leave a stale waker that the next `send` would spuriously wake
285        // (and retain until overwritten). This mirrors `WatchChanged::drop`. The
286        // future holds `&mut BroadcastReceiver`, so it is the unique registrant
287        // for this receiver id — clearing here cannot drop another future's waker.
288        if let Ok(mut state) = self.receiver.state.lock() {
289            let id = self.receiver.id;
290            if let Some(subscriber) = state.subscribers.get_mut(id) {
291                subscriber.waker = None;
292            }
293        }
294    }
295}
296
297/// Error types for broadcast channel operations
298#[derive(Debug, Clone, PartialEq, Eq)]
299pub enum BroadcastError {
300    /// Channel is empty
301    Empty,
302    /// Channel has been closed
303    Closed,
304    /// Message was lost due to channel overflow
305    Lagged,
306}
307
308impl std::fmt::Display for BroadcastError {
309    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
310        match self {
311            BroadcastError::Empty => write!(f, "broadcast channel is empty"),
312            BroadcastError::Closed => write!(f, "broadcast channel is closed"),
313            BroadcastError::Lagged => write!(f, "broadcast channel lagged"),
314        }
315    }
316}
317
318impl std::error::Error for BroadcastError {}
319
320#[cfg(test)]
321mod tests {
322    use super::*;
323    use std::sync::atomic::{AtomicUsize, Ordering};
324    use std::task::{Wake, Waker};
325
326    struct CountingWake(Arc<AtomicUsize>);
327    impl Wake for CountingWake {
328        fn wake(self: Arc<Self>) {
329            self.0.fetch_add(1, Ordering::Release);
330        }
331        fn wake_by_ref(self: &Arc<Self>) {
332            self.0.fetch_add(1, Ordering::Release);
333        }
334    }
335
336    #[test]
337    fn cancelled_recv_clears_waker_and_is_not_spuriously_woken() {
338        let (tx, mut rx) = Broadcast::<u32>::new(8);
339        let count = Arc::new(AtomicUsize::new(0));
340        let waker = Waker::from(Arc::new(CountingWake(Arc::clone(&count))));
341        let mut cx = Context::from_waker(&waker);
342
343        {
344            let mut fut = rx.recv();
345            assert!(Pin::new(&mut fut).poll(&mut cx).is_pending());
346            // `fut` is dropped here; its Drop must clear the registered waker.
347        }
348
349        // The cancelled future's waker must not fire on the next send.
350        tx.send(42).expect("send must succeed");
351        assert_eq!(
352            count.load(Ordering::Acquire),
353            0,
354            "a cancelled recv future must not be spuriously woken"
355        );
356
357        // The message is still deliverable to a fresh recv on the same receiver.
358        assert_eq!(rx.try_recv(), Ok(42));
359    }
360
361    #[test]
362    fn live_recv_is_woken_on_send() {
363        // Control: a still-live registered waker IS woken by send.
364        let (tx, mut rx) = Broadcast::<u32>::new(8);
365        let count = Arc::new(AtomicUsize::new(0));
366        let waker = Waker::from(Arc::new(CountingWake(Arc::clone(&count))));
367        let mut cx = Context::from_waker(&waker);
368
369        let mut fut = rx.recv();
370        assert!(Pin::new(&mut fut).poll(&mut cx).is_pending());
371        tx.send(7).expect("send must succeed");
372        assert_eq!(
373            count.load(Ordering::Acquire),
374            1,
375            "a live recv future must be woken by send"
376        );
377        // Keep `fut` alive across the send so its waker stays registered.
378        drop(fut);
379    }
380
381    #[test]
382    fn messages_read_by_every_receiver_are_reclaimed_on_next_send() {
383        let (tx, mut first) = Broadcast::<u32>::new(8);
384        let mut second = first.resubscribe();
385
386        tx.send(10).expect("first send must succeed");
387        tx.send(20).expect("second send must succeed");
388        assert_eq!(first.try_recv(), Ok(10));
389        assert_eq!(second.try_recv(), Ok(10));
390        assert_eq!(first.try_recv(), Ok(20));
391        assert_eq!(second.try_recv(), Ok(20));
392
393        tx.send(30).expect("third send must succeed");
394        let state = tx.state.lock().expect("broadcast state must not poison");
395        assert_eq!(
396            state.messages.iter().copied().collect::<Vec<_>>(),
397            vec![(3, 30)],
398            "the next send must retain only the new unread message"
399        );
400    }
401}