rskit_stream/
broadcaster.rs1use std::fmt;
15use std::pin::Pin;
16use std::sync::Arc;
17
18use futures::Stream;
19use parking_lot::Mutex;
20use tokio::sync::mpsc;
21use tokio_stream::wrappers::ReceiverStream;
22use tokio_util::sync::CancellationToken;
23
24pub type BroadcastStream<T> = Pin<Box<dyn Stream<Item = T> + Send + 'static>>;
29
30pub const DEFAULT_BROADCAST_BUFFER: usize = 64;
34
35pub struct Broadcaster<T> {
42 senders: Arc<Mutex<Vec<mpsc::Sender<T>>>>,
43 buffer: usize,
44}
45
46impl<T> Clone for Broadcaster<T> {
47 fn clone(&self) -> Self {
48 Self {
49 senders: Arc::clone(&self.senders),
50 buffer: self.buffer,
51 }
52 }
53}
54
55impl<T> fmt::Debug for Broadcaster<T> {
56 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
57 f.debug_struct("Broadcaster")
58 .field("buffer", &self.buffer)
59 .field("subscribers", &self.subscriber_count())
60 .finish_non_exhaustive()
61 }
62}
63
64impl<T> Default for Broadcaster<T>
65where
66 T: Clone + Send + 'static,
67{
68 fn default() -> Self {
69 Self::new()
70 }
71}
72
73impl<T> Broadcaster<T>
74where
75 T: Clone + Send + 'static,
76{
77 #[must_use]
79 pub fn new() -> Self {
80 Self::with_buffer(DEFAULT_BROADCAST_BUFFER)
81 }
82
83 #[must_use]
87 pub fn with_buffer(buffer: usize) -> Self {
88 Self {
89 senders: Arc::new(Mutex::new(Vec::new())),
90 buffer: buffer.max(1),
91 }
92 }
93}
94
95impl<T> Broadcaster<T> {
96 #[must_use]
98 pub const fn buffer(&self) -> usize {
99 self.buffer
100 }
101
102 #[must_use]
108 pub fn subscriber_count(&self) -> usize {
109 self.senders.lock().len()
110 }
111
112 pub fn subscribe(&self, cancel: CancellationToken) -> BroadcastStream<T>
118 where
119 T: Send + 'static,
120 {
121 use futures::StreamExt as _;
122
123 let (tx, rx) = mpsc::channel(self.buffer);
124 self.senders.lock().push(tx);
125 Box::pin(ReceiverStream::new(rx).take_until(cancel.cancelled_owned()))
126 }
127
128 pub fn broadcast(&self, item: &T)
134 where
135 T: Clone,
136 {
137 self.senders
138 .lock()
139 .retain(|tx| match tx.try_send(item.clone()) {
140 Ok(()) | Err(mpsc::error::TrySendError::Full(_)) => true,
141 Err(mpsc::error::TrySendError::Closed(_)) => false,
142 });
143 }
144}
145
146#[cfg(test)]
147mod tests {
148 use super::*;
149 use futures::StreamExt as _;
150
151 #[tokio::test]
152 async fn delivers_to_all_subscribers() {
153 let bc = Broadcaster::<u32>::new();
154 let cancel = CancellationToken::new();
155 let mut a = bc.subscribe(cancel.clone());
156 let mut b = bc.subscribe(cancel.clone());
157
158 bc.broadcast(&7);
159
160 assert_eq!(a.next().await, Some(7));
161 assert_eq!(b.next().await, Some(7));
162 }
163
164 #[tokio::test]
165 async fn stream_terminates_on_cancel() {
166 let bc = Broadcaster::<u32>::new();
167 let cancel = CancellationToken::new();
168 let mut sub = bc.subscribe(cancel.clone());
169
170 bc.broadcast(&1);
171 assert_eq!(sub.next().await, Some(1));
172
173 cancel.cancel();
174 assert_eq!(sub.next().await, None);
175 }
176
177 #[tokio::test]
178 async fn dropped_subscriber_is_pruned() {
179 let bc = Broadcaster::<u32>::new();
180 let cancel = CancellationToken::new();
181 let sub = bc.subscribe(cancel);
182 assert_eq!(bc.subscriber_count(), 1);
183
184 drop(sub);
185 bc.broadcast(&1);
187 assert_eq!(bc.subscriber_count(), 0);
188 }
189
190 #[tokio::test]
191 async fn full_subscriber_drops_overflow_without_blocking() {
192 let bc = Broadcaster::<u32>::with_buffer(2);
193 let cancel = CancellationToken::new();
194 let mut sub = bc.subscribe(cancel.clone());
195
196 for i in 0..5 {
198 bc.broadcast(&i);
199 }
200
201 let mut received = Vec::new();
202 while let Ok(Some(v)) =
204 tokio::time::timeout(std::time::Duration::from_millis(50), sub.next()).await
205 {
206 received.push(v);
207 }
208 assert_eq!(received, vec![0, 1]);
209 }
210
211 #[tokio::test]
212 async fn zero_buffer_is_clamped_to_one() {
213 let bc = Broadcaster::<u32>::with_buffer(0);
214 assert_eq!(bc.buffer(), 1);
215 }
216
217 #[tokio::test]
218 async fn clones_share_subscriber_set() {
219 let bc = Broadcaster::<u32>::new();
220 let cancel = CancellationToken::new();
221 let mut sub = bc.subscribe(cancel.clone());
222
223 let clone = bc.clone();
224 clone.broadcast(&42);
225
226 assert_eq!(sub.next().await, Some(42));
227 }
228
229 #[test]
230 fn default_and_debug_report_buffer_and_subscribers() {
231 let bc = Broadcaster::<u32>::default();
232
233 assert_eq!(bc.buffer(), DEFAULT_BROADCAST_BUFFER);
234 assert!(format!("{bc:?}").contains("Broadcaster"));
235 }
236}