1#![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
19pub 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 subscribers: SubscriberRegistry<u64>,
32 capacity: usize,
33}
34
35impl<T: Clone + Send + 'static> Broadcast<T> {
36 #[allow(clippy::new_ret_no_self)] 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
62pub struct BroadcastSender<T> {
64 state: Arc<Mutex<BroadcastState<T>>>,
65}
66
67impl<T: Clone> BroadcastSender<T> {
68 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 state.sequence += 1;
81 let sequence = state.sequence;
82 state.messages.push_back((sequence, message));
83
84 let receiver_count = state.subscribers.len();
89 let wakers = state.subscribers.drain_wakers();
90
91 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 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
132pub struct BroadcastReceiver<T> {
134 state: Arc<Mutex<BroadcastState<T>>>,
135 id: u64,
136 position: u64,
137}
138
139impl<T: Clone> BroadcastReceiver<T> {
140 pub fn recv(&mut self) -> BroadcastRecv<'_, T> {
142 BroadcastRecv { receiver: self }
143 }
144
145 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 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 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 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 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 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 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
268pub 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 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#[derive(Debug, Clone, PartialEq, Eq)]
299pub enum BroadcastError {
300 Empty,
302 Closed,
304 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 }
348
349 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 assert_eq!(rx.try_recv(), Ok(42));
359 }
360
361 #[test]
362 fn live_recv_is_woken_on_send() {
363 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 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}