Skip to main content

moirai_async/sync/
mpsc.rs

1#![expect(
2    clippy::unwrap_used,
3    reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
4)]
5
6//! Bounded async multi-producer single-consumer channel.
7//!
8//! Waiter bookkeeping is delegated to the shared `WaitQueue`: the same
9//! FIFO-by-monotonic-id registration, grant hand-off, and O(log n)
10//! cancellation that `Notify`, `Semaphore`, and `RwLock` use. This module
11//! keeps only the channel's own admission predicate (buffer capacity) and
12//! its two grants — a send frees a receive slot, a receive frees a send slot.
13
14use std::collections::VecDeque;
15use std::future::Future;
16use std::marker::Unpin;
17use std::pin::Pin;
18use std::sync::{Arc, Mutex};
19use std::task::{Context, Poll};
20
21use super::wait_queue::{WaitQueue, WaiterPoll};
22
23struct SharedState<T> {
24    buffer: VecDeque<T>,
25    capacity: usize,
26    sender_count: usize,
27    closed: bool,
28    send_waiters: WaitQueue<()>,
29    recv_waiters: WaitQueue<()>,
30}
31
32/// Sending half of the bounded channel; clone to add producers.
33pub struct Sender<T> {
34    shared: Arc<Mutex<SharedState<T>>>,
35}
36
37impl<T> Clone for Sender<T> {
38    fn clone(&self) -> Self {
39        let mut shared = self.shared.lock().unwrap();
40        shared.sender_count += 1;
41        Sender {
42            shared: self.shared.clone(),
43        }
44    }
45}
46
47impl<T> Sender<T> {
48    /// Send a value, waiting for buffer capacity.
49    ///
50    /// The returned future resolves `Err(value)` when the channel closes
51    /// before the value is accepted.
52    pub fn send(&self, value: T) -> SendFuture<'_, T> {
53        SendFuture {
54            sender: self,
55            value: Some(value),
56            id: None,
57        }
58    }
59
60    /// Send without waiting; returns the value when full or closed.
61    ///
62    /// # Errors
63    ///
64    /// Returns `Err(value)` when the buffer is at capacity or the channel
65    /// is closed.
66    pub fn try_send(&self, value: T) -> Result<(), T> {
67        let mut shared = self.shared.lock().unwrap();
68        if shared.closed {
69            return Err(value);
70        }
71        if shared.buffer.len() >= shared.capacity {
72            return Err(value);
73        }
74        shared.buffer.push_back(value);
75        // Grant the oldest parked receiver a slot and wake it outside the
76        // lock: a task waker may re-enter the channel, so holding the mutex
77        // across `wake` risks a self-deadlock.
78        let waker = shared.recv_waiters.grant_oldest(());
79        drop(shared);
80        if let Some(waker) = waker {
81            waker.wake();
82        }
83        Ok(())
84    }
85
86    /// Return whether the channel is closed.
87    pub fn is_closed(&self) -> bool {
88        self.shared.lock().unwrap().closed
89    }
90
91    /// Count of live sender handles.
92    pub fn sender_strong_count(&self) -> usize {
93        self.shared.lock().unwrap().sender_count
94    }
95}
96
97impl<T> Drop for Sender<T> {
98    fn drop(&mut self) {
99        let mut shared = self.shared.lock().unwrap();
100        shared.sender_count -= 1;
101        if shared.sender_count != 0 {
102            return;
103        }
104        shared.closed = true;
105        let wakers = shared.recv_waiters.grant_all(());
106        drop(shared);
107        for waker in wakers {
108            waker.wake();
109        }
110    }
111}
112
113/// Future returned by [`Sender::send`].
114pub struct SendFuture<'a, T> {
115    sender: &'a Sender<T>,
116    value: Option<T>,
117    id: Option<u64>,
118}
119
120impl<'a, T: Unpin> Future for SendFuture<'a, T> {
121    type Output = Result<(), T>;
122
123    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
124        let this = self.get_mut();
125        let mut shared = this.sender.shared.lock().unwrap();
126
127        if shared.closed {
128            if let Some(id) = this.id.take() {
129                shared.send_waiters.deregister(id);
130            }
131            let value = this.value.take().unwrap();
132            return Poll::Ready(Err(value));
133        }
134
135        if shared.buffer.len() < shared.capacity {
136            if let Some(id) = this.id.take() {
137                shared.send_waiters.deregister(id);
138            }
139            shared.buffer.push_back(this.value.take().unwrap());
140            let waker = shared.recv_waiters.grant_oldest(());
141            drop(shared);
142            if let Some(waker) = waker {
143                waker.wake();
144            }
145            return Poll::Ready(Ok(()));
146        }
147
148        // Still full: refresh the existing registration, or join the queue.
149        // A stale grant (the slot was taken by another sender) re-registers
150        // behind the current waiters rather than losing its place.
151        this.id = Some(match this.id {
152            Some(id) => match shared.send_waiters.poll_waiter(id, cx.waker()) {
153                WaiterPoll::Pending => id,
154                WaiterPoll::Granted(()) | WaiterPoll::NotRegistered => {
155                    shared.send_waiters.register(cx.waker().clone())
156                }
157            },
158            None => shared.send_waiters.register(cx.waker().clone()),
159        });
160        Poll::Pending
161    }
162}
163
164impl<'a, T> Drop for SendFuture<'a, T> {
165    fn drop(&mut self) {
166        if let Some(id) = self.id
167            && let Ok(mut shared) = self.sender.shared.lock()
168        {
169            shared.send_waiters.deregister(id);
170        }
171    }
172}
173
174/// Receiving half of the bounded channel.
175pub struct Receiver<T> {
176    shared: Arc<Mutex<SharedState<T>>>,
177}
178
179impl<T> Receiver<T> {
180    /// Receive the next value, waiting for one to arrive.
181    ///
182    /// The returned future resolves `Err(())` when the channel is closed
183    /// and drained.
184    pub fn recv(&mut self) -> RecvFuture<'_, T> {
185        RecvFuture {
186            receiver: self,
187            id: None,
188        }
189    }
190
191    /// Receive without waiting; `None` when the buffer is empty.
192    pub fn try_recv(&mut self) -> Option<T> {
193        let mut shared = self.shared.lock().unwrap();
194        let value = shared.buffer.pop_front();
195        let waker = if value.is_some() {
196            shared.send_waiters.grant_oldest(())
197        } else {
198            None
199        };
200        drop(shared);
201        if let Some(waker) = waker {
202            waker.wake();
203        }
204        value
205    }
206
207    /// Close the channel, waking every parked sender and receiver.
208    pub fn close(&mut self) {
209        let mut shared = self.shared.lock().unwrap();
210        shared.closed = true;
211        let mut wakers = shared.send_waiters.grant_all(());
212        wakers.extend(shared.recv_waiters.grant_all(()));
213        drop(shared);
214        for waker in wakers {
215            waker.wake();
216        }
217    }
218}
219
220impl<T> Drop for Receiver<T> {
221    fn drop(&mut self) {
222        let mut shared = self.shared.lock().unwrap();
223        shared.closed = true;
224        let wakers = shared.send_waiters.grant_all(());
225        drop(shared);
226        for waker in wakers {
227            waker.wake();
228        }
229    }
230}
231
232/// Future returned by [`Receiver::recv`].
233pub struct RecvFuture<'a, T> {
234    receiver: &'a mut Receiver<T>,
235    id: Option<u64>,
236}
237
238impl<'a, T> Future for RecvFuture<'a, T> {
239    type Output = Result<T, ()>;
240
241    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
242        let this = self.get_mut();
243        let mut shared = this.receiver.shared.lock().unwrap();
244        if let Some(value) = shared.buffer.pop_front() {
245            if let Some(id) = this.id.take() {
246                shared.recv_waiters.deregister(id);
247            }
248            let waker = shared.send_waiters.grant_oldest(());
249            drop(shared);
250            if let Some(waker) = waker {
251                waker.wake();
252            }
253            return Poll::Ready(Ok(value));
254        }
255        if shared.closed {
256            if let Some(id) = this.id.take() {
257                shared.recv_waiters.deregister(id);
258            }
259            return Poll::Ready(Err(()));
260        }
261
262        this.id = Some(match this.id {
263            Some(id) => match shared.recv_waiters.poll_waiter(id, cx.waker()) {
264                WaiterPoll::Pending => id,
265                WaiterPoll::Granted(()) | WaiterPoll::NotRegistered => {
266                    shared.recv_waiters.register(cx.waker().clone())
267                }
268            },
269            None => shared.recv_waiters.register(cx.waker().clone()),
270        });
271        Poll::Pending
272    }
273}
274
275impl<'a, T> Drop for RecvFuture<'a, T> {
276    fn drop(&mut self) {
277        if let Some(id) = self.id
278            && let Ok(mut shared) = self.receiver.shared.lock()
279        {
280            shared.recv_waiters.deregister(id);
281        }
282    }
283}
284
285/// Create a bounded channel with the given buffer capacity.
286#[must_use]
287pub fn channel<T>(capacity: usize) -> (Sender<T>, Receiver<T>) {
288    let shared = Arc::new(Mutex::new(SharedState {
289        buffer: VecDeque::with_capacity(capacity),
290        capacity,
291        sender_count: 1,
292        closed: false,
293        send_waiters: WaitQueue::new(),
294        recv_waiters: WaitQueue::new(),
295    }));
296    (
297        Sender {
298            shared: shared.clone(),
299        },
300        Receiver { shared },
301    )
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307    use std::future::Future;
308    use std::pin::Pin;
309    use std::sync::{
310        Arc,
311        atomic::{AtomicUsize, Ordering},
312    };
313    use std::task::{Context, Poll, Wake, Waker};
314
315    fn poll_future<F: Future + Unpin>(future: &mut F) -> Poll<F::Output> {
316        let mut context = Context::from_waker(Waker::noop());
317        Pin::new(future).poll(&mut context)
318    }
319
320    fn poll_future_with_waker<F: Future + Unpin>(future: &mut F, waker: &Waker) -> Poll<F::Output> {
321        let mut context = Context::from_waker(waker);
322        Pin::new(future).poll(&mut context)
323    }
324
325    struct CountingWake(Arc<AtomicUsize>);
326
327    impl Wake for CountingWake {
328        fn wake(self: Arc<Self>) {
329            self.0.fetch_add(1, Ordering::Release);
330        }
331
332        fn wake_by_ref(self: &Arc<Self>) {
333            self.0.fetch_add(1, Ordering::Release);
334        }
335    }
336
337    #[test]
338    fn test_mpsc_send_recv() {
339        let (tx, mut rx) = channel(10);
340        tx.try_send(1).unwrap();
341        tx.try_send(2).unwrap();
342        tx.try_send(3).unwrap();
343        assert_eq!(rx.try_recv(), Some(1));
344        assert_eq!(rx.try_recv(), Some(2));
345        assert_eq!(rx.try_recv(), Some(3));
346        assert!(rx.try_recv().is_none());
347    }
348
349    #[test]
350    fn test_mpsc_closed_sender() {
351        let (tx, mut rx) = channel::<i32>(10);
352        tx.try_send(1).unwrap();
353        drop(tx);
354        assert_eq!(rx.try_recv(), Some(1));
355        assert!(rx.try_recv().is_none());
356    }
357
358    #[test]
359    fn test_mpsc_closed_receiver() {
360        let (tx, rx) = channel::<i32>(10);
361        drop(rx);
362        assert!(tx.try_send(1).is_err());
363    }
364
365    #[test]
366    fn test_mpsc_capacity() {
367        let (tx, mut rx) = channel(2);
368        assert!(tx.try_send(1).is_ok());
369        assert!(tx.try_send(2).is_ok());
370        assert!(tx.try_send(3).is_err());
371        let _ = rx.try_recv();
372        assert!(tx.try_send(3).is_ok());
373    }
374
375    #[test]
376    fn test_mpsc_sender_clone() {
377        let (tx1, mut rx) = channel(10);
378        let tx2 = tx1.clone();
379        tx1.try_send(1).unwrap();
380        tx2.try_send(2).unwrap();
381        drop(tx1);
382        drop(tx2);
383        assert_eq!(rx.try_recv(), Some(1));
384        assert_eq!(rx.try_recv(), Some(2));
385        assert!(rx.try_recv().is_none());
386    }
387
388    #[test]
389    fn test_mpsc_sender_strong_count() {
390        let (tx1, _) = channel::<i32>(10);
391        assert_eq!(tx1.sender_strong_count(), 1);
392        let tx2 = tx1.clone();
393        assert_eq!(tx1.sender_strong_count(), 2);
394        drop(tx2);
395        assert_eq!(tx1.sender_strong_count(), 1);
396    }
397
398    #[test]
399    fn test_mpsc_send_pending_then_recv() {
400        let (tx, mut rx) = channel(1);
401        tx.try_send(1).unwrap();
402        let mut send = tx.send(2);
403        assert!(matches!(poll_future(&mut send), Poll::Pending));
404        let _ = rx.try_recv();
405        assert!(matches!(poll_future(&mut send), Poll::Ready(Ok(()))));
406    }
407
408    #[test]
409    fn test_mpsc_async_recv_pending_then_send() {
410        let (tx, mut rx) = channel(1);
411        let mut recv = rx.recv();
412        assert!(matches!(poll_future(&mut recv), Poll::Pending));
413        tx.try_send(42).unwrap();
414        assert!(matches!(poll_future(&mut recv), Poll::Ready(Ok(42))));
415    }
416
417    #[test]
418    fn test_mpsc_send_future_dropped_cancels_waiter() {
419        let (tx, _rx) = channel(1);
420        tx.try_send(1).unwrap();
421        // Send future goes pending, then is dropped without completing
422        let mut send = tx.send(2);
423        assert!(matches!(poll_future(&mut send), Poll::Pending));
424        drop(send);
425        // The full send_waiters queue must be empty after the drop
426        assert!(tx.shared.lock().unwrap().send_waiters.is_empty());
427    }
428
429    #[test]
430    fn test_mpsc_recv_future_dropped_cancels_waiter() {
431        let (tx, mut rx) = channel(1);
432        tx.try_send(1).unwrap();
433        // Consume the item, then recv goes pending waiting for next item
434        let _ = rx.try_recv();
435        let mut recv = rx.recv();
436        assert!(matches!(poll_future(&mut recv), Poll::Pending));
437        drop(recv);
438        // The full recv_waiters queue must be empty after the drop
439        assert!(tx.shared.lock().unwrap().recv_waiters.is_empty());
440    }
441
442    #[test]
443    fn oldest_pending_sender_is_woken_first() {
444        let (tx, mut rx) = channel(1);
445        tx.try_send(1).expect("initial send must fill the channel");
446
447        let first_wakes = Arc::new(AtomicUsize::new(0));
448        let second_wakes = Arc::new(AtomicUsize::new(0));
449        let first_waker = Waker::from(Arc::new(CountingWake(Arc::clone(&first_wakes))));
450        let second_waker = Waker::from(Arc::new(CountingWake(Arc::clone(&second_wakes))));
451        let mut first = tx.send(2);
452        let mut second = tx.send(3);
453
454        assert!(poll_future_with_waker(&mut first, &first_waker).is_pending());
455        assert!(poll_future_with_waker(&mut second, &second_waker).is_pending());
456
457        assert_eq!(rx.try_recv(), Some(1));
458        assert_eq!(first_wakes.load(Ordering::Acquire), 1);
459        assert_eq!(second_wakes.load(Ordering::Acquire), 0);
460
461        assert!(matches!(
462            poll_future_with_waker(&mut first, &first_waker),
463            Poll::Ready(Ok(()))
464        ));
465        assert_eq!(rx.try_recv(), Some(2));
466        assert_eq!(second_wakes.load(Ordering::Acquire), 1);
467        assert!(matches!(
468            poll_future_with_waker(&mut second, &second_waker),
469            Poll::Ready(Ok(()))
470        ));
471        assert_eq!(rx.try_recv(), Some(3));
472    }
473}