Skip to main content

commonware_utils/channel/
ring.rs

1//! A bounded mpsc channel that drops the oldest item when full instead of applying backpressure.
2//!
3//! This is useful for scenarios where you want to keep the most recent items and can
4//! tolerate losing older ones, such as real-time data streams or status updates where
5//! only the latest values matter.
6//!
7//! # Example
8//!
9//! ```
10//! use futures::executor::block_on;
11//! use futures::{SinkExt, StreamExt};
12//! use commonware_utils::{NZUsize, channel::ring};
13//!
14//! block_on(async {
15//!     let (mut sender, mut receiver) = ring::channel::<u32>(NZUsize!(2));
16//!
17//!     // Fill the channel
18//!     sender.send(1).await.unwrap();
19//!     sender.send(2).await.unwrap();
20//!
21//!     // This will drop the oldest item (1) and insert 3
22//!     sender.send(3).await.unwrap();
23//!
24//!     // Receive the remaining items
25//!     assert_eq!(receiver.next().await, Some(2));
26//!     assert_eq!(receiver.next().await, Some(3));
27//! });
28//! ```
29
30use 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/// Error returned when sending to a channel whose receiver has been dropped.
42#[derive(Debug, Error)]
43#[error("channel closed")]
44pub struct ChannelClosed;
45
46/// Error returned by [`Receiver::try_recv`].
47#[derive(Debug, Error, PartialEq, Eq)]
48pub enum TryRecvError {
49    /// The channel currently has no buffered items, but senders still exist.
50    #[error("channel empty")]
51    Empty,
52    /// The channel is empty and all senders have been dropped.
53    #[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/// The sending half of a ring channel.
67///
68/// Implements [`Sink`] for sending items. Use [`SinkExt::send`](futures::SinkExt::send)
69/// to send items asynchronously.
70///
71/// This type can be cloned to create multiple producers for the same channel.
72/// The channel remains open until all senders are dropped.
73#[derive(Debug)]
74pub struct Sender<T: Send + Sync> {
75    shared: Arc<Mutex<Shared<T>>>,
76}
77
78impl<T: Send + Sync> Sender<T> {
79    /// Returns whether the receiver has been dropped.
80    ///
81    /// If this returns `true`, subsequent sends will fail with [`ChannelClosed`].
82    pub fn is_closed(&self) -> bool {
83        let shared = self.shared.lock();
84        shared.receiver_dropped
85    }
86
87    /// Sends an item, dropping the oldest buffered item if the channel is full.
88    ///
89    /// Returns `false` once the receiver has been dropped.
90    pub fn send_lossy(&self, item: T) -> bool {
91        let mut shared = self.shared.lock();
92
93        // Nothing will ever read a buffered item once the receiver is gone.
94        if shared.receiver_dropped {
95            return false;
96        }
97
98        // Make room by evicting the oldest item rather than blocking the sender.
99        let old_item = if shared.buffer.len() >= shared.capacity {
100            shared.buffer.pop_front()
101        } else {
102            None
103        };
104
105        // Buffer the item and take the receiver's waker so it can be woken
106        // once the lock is released.
107        shared.buffer.push_back(item);
108        let waker = shared.receiver_waker.take();
109        drop(shared);
110
111        // Drop the old item after the lock is released to avoid potential mutex poisoning
112        drop(old_item);
113
114        // Wake a receiver parked on an empty buffer.
115        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        // No buffering in the sender - items are sent immediately to the shared buffer
170        Poll::Ready(Ok(()))
171    }
172
173    fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
174        // Closing is handled by Drop
175        Poll::Ready(Ok(()))
176    }
177}
178
179/// The receiving half of a ring channel.
180///
181/// Implements [`Stream`] and [`FusedStream`] for receiving items. Use
182/// [`StreamExt::next`](futures::StreamExt::next) to receive items asynchronously.
183///
184/// The stream terminates (returns `None`) when all senders have been dropped
185/// and all buffered items have been consumed.
186#[derive(Debug)]
187pub struct Receiver<T: Send + Sync> {
188    shared: Arc<Mutex<Shared<T>>>,
189}
190
191impl<T: Send + Sync> Receiver<T> {
192    /// Receives the next item from the channel.
193    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    /// Attempts to receive an item without waiting.
198    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
249/// Creates a new ring channel with the specified capacity.
250///
251/// Returns a ([`Sender`], [`Receiver`]) pair. The sender can be cloned to create
252/// multiple producers.
253pub 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(); // Should drop 1
299            sender.send(4).await.unwrap(); // Should drop 2
300
301            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            // Channel should still be open because sender2 exists
370            sender2.send(1).await.unwrap();
371            assert_eq!(receiver.next().await, Some(1));
372
373            drop(sender2);
374            // Now channel should be closed
375            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(); // Drops 1
386
387            assert_eq!(receiver.next().await, Some(2));
388
389            sender.send(1).await.unwrap();
390            sender.send(2).await.unwrap(); // Drops 1
391            sender.send(3).await.unwrap(); // Drops 2
392
393            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()); // Still has item in buffer
425
426            assert_eq!(receiver.next().await, Some(1));
427            assert!(receiver.is_terminated()); // Now terminated
428
429            // Calling next after termination returns None
430            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}