commonware_utils/channel/
ring.rs1use crate::sync::Mutex;
31use core::num::NonZeroUsize;
32use futures::{Sink, Stream, stream::FusedStream};
33use std::{
34 collections::VecDeque,
35 pin::Pin,
36 sync::Arc,
37 task::{Context, Poll, Waker},
38};
39use thiserror::Error;
40
41#[derive(Debug, Error)]
43#[error("channel closed")]
44pub struct ChannelClosed;
45
46#[derive(Debug, Error, PartialEq, Eq)]
48pub enum TryRecvError {
49 #[error("channel empty")]
51 Empty,
52 #[error("channel closed")]
54 Disconnected,
55}
56
57#[derive(Debug)]
58struct Shared<T: Send + Sync> {
59 buffer: VecDeque<T>,
60 capacity: usize,
61 receiver_waker: Option<Waker>,
62 sender_count: usize,
63 receiver_dropped: bool,
64}
65
66#[derive(Debug)]
74pub struct Sender<T: Send + Sync> {
75 shared: Arc<Mutex<Shared<T>>>,
76}
77
78impl<T: Send + Sync> Sender<T> {
79 pub fn is_closed(&self) -> bool {
83 let shared = self.shared.lock();
84 shared.receiver_dropped
85 }
86
87 pub fn send_lossy(&self, item: T) -> bool {
91 let mut shared = self.shared.lock();
92
93 if shared.receiver_dropped {
95 return false;
96 }
97
98 let old_item = if shared.buffer.len() >= shared.capacity {
100 shared.buffer.pop_front()
101 } else {
102 None
103 };
104
105 shared.buffer.push_back(item);
108 let waker = shared.receiver_waker.take();
109 drop(shared);
110
111 drop(old_item);
113
114 if let Some(w) = waker {
116 w.wake();
117 }
118
119 true
120 }
121}
122
123impl<T: Send + Sync> Clone for Sender<T> {
124 fn clone(&self) -> Self {
125 let mut shared = self.shared.lock();
126 shared.sender_count += 1;
127 drop(shared);
128
129 Self {
130 shared: self.shared.clone(),
131 }
132 }
133}
134
135impl<T: Send + Sync> Drop for Sender<T> {
136 fn drop(&mut self) {
137 let mut shared = self.shared.lock();
138 shared.sender_count -= 1;
139 let waker = if shared.sender_count == 0 {
140 shared.receiver_waker.take()
141 } else {
142 None
143 };
144 drop(shared);
145
146 if let Some(w) = waker {
147 w.wake();
148 }
149 }
150}
151
152impl<T: Send + Sync> Sink<T> for Sender<T> {
153 type Error = ChannelClosed;
154
155 fn poll_ready(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
156 let shared = self.shared.lock();
157 if shared.receiver_dropped {
158 return Poll::Ready(Err(ChannelClosed));
159 }
160
161 Poll::Ready(Ok(()))
162 }
163
164 fn start_send(self: Pin<&mut Self>, item: T) -> Result<(), Self::Error> {
165 self.send_lossy(item).then_some(()).ok_or(ChannelClosed)
166 }
167
168 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
169 Poll::Ready(Ok(()))
171 }
172
173 fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
174 Poll::Ready(Ok(()))
176 }
177}
178
179#[derive(Debug)]
187pub struct Receiver<T: Send + Sync> {
188 shared: Arc<Mutex<Shared<T>>>,
189}
190
191impl<T: Send + Sync> Receiver<T> {
192 pub async fn recv(&mut self) -> Option<T> {
194 futures::future::poll_fn(|cx| Pin::new(&mut *self).poll_next(cx)).await
195 }
196
197 pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
199 let mut shared = self.shared.lock();
200 if let Some(item) = shared.buffer.pop_front() {
201 return Ok(item);
202 }
203 if shared.sender_count == 0 {
204 return Err(TryRecvError::Disconnected);
205 }
206 Err(TryRecvError::Empty)
207 }
208}
209
210impl<T: Send + Sync> Stream for Receiver<T> {
211 type Item = T;
212
213 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
214 let mut shared = self.shared.lock();
215
216 if let Some(item) = shared.buffer.pop_front() {
217 return Poll::Ready(Some(item));
218 }
219
220 if shared.sender_count == 0 {
221 return Poll::Ready(None);
222 }
223
224 if !shared
225 .receiver_waker
226 .as_ref()
227 .is_some_and(|w| w.will_wake(cx.waker()))
228 {
229 shared.receiver_waker = Some(cx.waker().clone());
230 }
231 Poll::Pending
232 }
233}
234
235impl<T: Send + Sync> FusedStream for Receiver<T> {
236 fn is_terminated(&self) -> bool {
237 let shared = self.shared.lock();
238 shared.sender_count == 0 && shared.buffer.is_empty()
239 }
240}
241
242impl<T: Send + Sync> Drop for Receiver<T> {
243 fn drop(&mut self) {
244 let mut shared = self.shared.lock();
245 shared.receiver_dropped = true;
246 }
247}
248
249pub fn channel<T: Send + Sync>(capacity: NonZeroUsize) -> (Sender<T>, Receiver<T>) {
254 let shared = Arc::new(Mutex::new(Shared {
255 buffer: VecDeque::with_capacity(capacity.get()),
256 capacity: capacity.get(),
257 receiver_waker: None,
258 sender_count: 1,
259 receiver_dropped: false,
260 }));
261
262 let sender = Sender {
263 shared: shared.clone(),
264 };
265 let receiver = Receiver { shared };
266
267 (sender, receiver)
268}
269
270#[cfg(test)]
271mod tests {
272 use super::*;
273 use crate::NZUsize;
274 use futures::{SinkExt, StreamExt, executor::block_on};
275
276 #[test]
277 fn test_basic_send_recv() {
278 block_on(async {
279 let (mut sender, mut receiver) = channel::<i32>(NZUsize!(10));
280
281 sender.send(1).await.unwrap();
282 sender.send(2).await.unwrap();
283 sender.send(3).await.unwrap();
284
285 assert_eq!(receiver.next().await, Some(1));
286 assert_eq!(receiver.next().await, Some(2));
287 assert_eq!(receiver.next().await, Some(3));
288 });
289 }
290
291 #[test]
292 fn test_overflow_drops_oldest() {
293 block_on(async {
294 let (mut sender, mut receiver) = channel::<i32>(NZUsize!(2));
295
296 sender.send(1).await.unwrap();
297 sender.send(2).await.unwrap();
298 sender.send(3).await.unwrap(); sender.send(4).await.unwrap(); assert_eq!(receiver.next().await, Some(3));
302 assert_eq!(receiver.next().await, Some(4));
303 });
304 }
305
306 #[test]
307 fn test_send_after_receiver_dropped() {
308 block_on(async {
309 let (mut sender, receiver) = channel::<i32>(NZUsize!(10));
310 drop(receiver);
311
312 let err = sender.send(1).await.unwrap_err();
313 assert!(matches!(err, ChannelClosed));
314 });
315 }
316
317 #[test]
318 fn test_recv_after_sender_dropped() {
319 block_on(async {
320 let (mut sender, mut receiver) = channel::<i32>(NZUsize!(10));
321
322 sender.send(1).await.unwrap();
323 sender.send(2).await.unwrap();
324 drop(sender);
325
326 assert_eq!(receiver.next().await, Some(1));
327 assert_eq!(receiver.next().await, Some(2));
328 assert_eq!(receiver.next().await, None);
329 });
330 }
331
332 #[test]
333 fn test_stream_collect() {
334 block_on(async {
335 let (mut sender, receiver) = channel::<i32>(NZUsize!(10));
336
337 sender.send(1).await.unwrap();
338 sender.send(2).await.unwrap();
339 sender.send(3).await.unwrap();
340 drop(sender);
341
342 let items: Vec<_> = receiver.collect().await;
343 assert_eq!(items, vec![1, 2, 3]);
344 });
345 }
346
347 #[test]
348 fn test_clone_sender() {
349 block_on(async {
350 let (mut sender1, mut receiver) = channel::<i32>(NZUsize!(10));
351 let mut sender2 = sender1.clone();
352
353 sender1.send(1).await.unwrap();
354 sender2.send(2).await.unwrap();
355
356 assert_eq!(receiver.next().await, Some(1));
357 assert_eq!(receiver.next().await, Some(2));
358 });
359 }
360
361 #[test]
362 fn test_sender_drop_with_clones() {
363 block_on(async {
364 let (sender1, mut receiver) = channel::<i32>(NZUsize!(10));
365 let mut sender2 = sender1.clone();
366
367 drop(sender1);
368
369 sender2.send(1).await.unwrap();
371 assert_eq!(receiver.next().await, Some(1));
372
373 drop(sender2);
374 assert_eq!(receiver.next().await, None);
376 });
377 }
378
379 #[test]
380 fn test_capacity_one() {
381 block_on(async {
382 let (mut sender, mut receiver) = channel::<i32>(NZUsize!(1));
383
384 sender.send(1).await.unwrap();
385 sender.send(2).await.unwrap(); assert_eq!(receiver.next().await, Some(2));
388
389 sender.send(1).await.unwrap();
390 sender.send(2).await.unwrap(); sender.send(3).await.unwrap(); assert_eq!(receiver.next().await, Some(3));
394 });
395 }
396
397 #[test]
398 fn test_send_all() {
399 block_on(async {
400 let (mut sender, receiver) = channel::<i32>(NZUsize!(10));
401
402 let items = futures::stream::iter(vec![1, 2, 3]);
403 sender.send_all(&mut items.map(Ok)).await.unwrap();
404 drop(sender);
405
406 let received: Vec<_> = receiver.collect().await;
407 assert_eq!(received, vec![1, 2, 3]);
408 });
409 }
410
411 #[test]
412 fn test_fused_stream() {
413 use futures::stream::FusedStream;
414
415 block_on(async {
416 let (mut sender, mut receiver) = channel::<i32>(NZUsize!(10));
417
418 assert!(!receiver.is_terminated());
419
420 sender.send(1).await.unwrap();
421 assert!(!receiver.is_terminated());
422
423 drop(sender);
424 assert!(!receiver.is_terminated()); assert_eq!(receiver.next().await, Some(1));
427 assert!(receiver.is_terminated()); assert_eq!(receiver.next().await, None);
431 assert!(receiver.is_terminated());
432 });
433 }
434
435 #[test]
436 fn test_is_closed() {
437 block_on(async {
438 let (sender, receiver) = channel::<i32>(NZUsize!(10));
439
440 assert!(!sender.is_closed());
441
442 drop(receiver);
443 assert!(sender.is_closed());
444 });
445 }
446}