Skip to main content

ntex_util/channel/
mpsc.rs

1//! A multi-producer, single-consumer, futures-aware, FIFO queue.
2use std::collections::VecDeque;
3use std::future::poll_fn;
4use std::{fmt, panic::UnwindSafe, pin::Pin, task::Context, task::Poll};
5
6use futures_core::{FusedStream, Stream};
7
8use super::cell::Cell;
9use crate::task::LocalWaker;
10
11/// Creates an unbounded in-memory channel.
12pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
13    let shared = Cell::new(Shared {
14        has_receiver: true,
15        buffer: VecDeque::new(),
16        blocked_recv: LocalWaker::new(),
17        closed: LocalWaker::new(),
18    });
19    let sender = Sender {
20        shared: shared.clone(),
21    };
22    let receiver = Receiver { shared };
23    (sender, receiver)
24}
25
26#[derive(Debug)]
27struct Shared<T> {
28    buffer: VecDeque<T>,
29    blocked_recv: LocalWaker,
30    closed: LocalWaker,
31    has_receiver: bool,
32}
33
34impl<T> Shared<T> {
35    fn close(&mut self) {
36        self.has_receiver = false;
37        self.blocked_recv.wake();
38        self.closed.wake();
39    }
40}
41
42/// The transmission end of a channel.
43///
44/// This is created by the `channel` function.
45#[derive(Debug)]
46pub struct Sender<T> {
47    shared: Cell<Shared<T>>,
48}
49
50impl<T> Unpin for Sender<T> {}
51
52impl<T> Sender<T> {
53    /// Sends the provided message along this channel.
54    pub fn send(&self, item: T) -> Result<(), SendError<T>> {
55        let shared = self.shared.get_mut();
56        if !shared.has_receiver {
57            return Err(SendError(item)); // receiver was dropped
58        }
59        shared.buffer.push_back(item);
60        shared.blocked_recv.wake();
61        Ok(())
62    }
63
64    /// Closes the channel for every sender.
65    ///
66    /// This prevents any further messages from being sent on the channel, by
67    /// this sender or any of its clones, while still enabling the receiver to
68    /// drain messages that are buffered.
69    pub fn close(&self) {
70        self.shared.get_mut().close();
71    }
72
73    /// Returns `true` if the channel is closed or the receiver has been dropped.
74    pub fn is_closed(&self) -> bool {
75        self.shared.strong_count() == 1 || !self.shared.get_ref().has_receiver
76    }
77
78    /// Polls whether the channel is closed or the receiver has been dropped.
79    ///
80    /// Only the task from the most recent call is woken, across this sender
81    /// and all of its clones.
82    pub fn poll_closed(&self, cx: &mut Context<'_>) -> Poll<()> {
83        let shared = self.shared.get_mut();
84        if shared.has_receiver {
85            shared.closed.register(cx.waker());
86            Poll::Pending
87        } else {
88            Poll::Ready(())
89        }
90    }
91
92    /// Waits until the channel is closed or the receiver has been dropped.
93    ///
94    /// See [`poll_closed`](Self::poll_closed).
95    pub async fn closed(&self) {
96        poll_fn(|cx| self.poll_closed(cx)).await;
97    }
98}
99
100impl<T> Clone for Sender<T> {
101    fn clone(&self) -> Self {
102        Sender {
103            shared: self.shared.clone(),
104        }
105    }
106}
107
108impl<T> Drop for Sender<T> {
109    fn drop(&mut self) {
110        let count = self.shared.strong_count();
111        let shared = self.shared.get_mut();
112
113        // check is last sender is about to drop
114        if shared.has_receiver && count == 2 {
115            // Wake up receiver as its stream has ended
116            shared.blocked_recv.wake();
117        }
118    }
119}
120
121/// The receiving end of a channel which implements the `Stream` trait.
122///
123/// This is created by the `channel` function.
124#[derive(Debug)]
125pub struct Receiver<T> {
126    shared: Cell<Shared<T>>,
127}
128
129impl<T> Receiver<T> {
130    /// Creates an additional sender for this channel.
131    pub fn sender(&self) -> Sender<T> {
132        Sender {
133            shared: self.shared.clone(),
134        }
135    }
136
137    /// Closes the receiving half of a channel, without dropping it.
138    ///
139    /// This prevents any further messages from being sent on the channel
140    /// while still enabling the receiver to drain messages that are buffered.
141    pub fn close(&self) {
142        self.shared.get_mut().close();
143    }
144
145    /// Returns whether this channel is closed.
146    pub fn is_closed(&self) -> bool {
147        self.shared.strong_count() == 1 || !self.shared.get_ref().has_receiver
148    }
149
150    /// Waits for the next message.
151    ///
152    /// Returns `None` once the channel is closed or every sender has been
153    /// dropped, and every buffered message has been received.
154    pub async fn recv(&self) -> Option<T> {
155        poll_fn(|cx| self.poll_recv(cx)).await
156    }
157
158    /// Polls for the next message.
159    ///
160    /// Returns `Ready(None)` once the channel is closed or every sender has
161    /// been dropped, and every buffered message has been received.
162    pub fn poll_recv(&self, cx: &mut Context<'_>) -> Poll<Option<T>> {
163        let shared = self.shared.get_mut();
164
165        if let Some(msg) = shared.buffer.pop_front() {
166            Poll::Ready(Some(msg))
167        } else if shared.has_receiver {
168            shared.blocked_recv.register(cx.waker());
169            if self.shared.strong_count() == 1 {
170                // All senders have been dropped, so drain the buffer and end the
171                // stream.
172                Poll::Ready(None)
173            } else {
174                Poll::Pending
175            }
176        } else {
177            Poll::Ready(None)
178        }
179    }
180}
181
182impl<T> Unpin for Receiver<T> {}
183
184impl<T> Stream for Receiver<T> {
185    type Item = T;
186
187    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
188        self.poll_recv(cx)
189    }
190}
191
192impl<T> FusedStream for Receiver<T> {
193    /// Returns `true` once the channel is closed and every buffered message
194    /// has been received.
195    fn is_terminated(&self) -> bool {
196        self.is_closed() && self.shared.get_ref().buffer.is_empty()
197    }
198}
199
200impl<T> UnwindSafe for Receiver<T> {}
201
202impl<T> Drop for Receiver<T> {
203    fn drop(&mut self) {
204        let shared = self.shared.get_mut();
205        shared.has_receiver = false;
206        let buffer = std::mem::take(&mut shared.buffer);
207        shared.closed.wake();
208
209        // queued messages may send on this channel from their `Drop`
210        drop(buffer);
211    }
212}
213
214/// Error type for sending, used when the receiving end of a channel is
215/// dropped
216pub struct SendError<T>(T);
217
218impl<T> std::error::Error for SendError<T> {}
219
220impl<T> fmt::Debug for SendError<T> {
221    fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
222        fmt.debug_tuple("SendError").field(&"...").finish()
223    }
224}
225
226impl<T> fmt::Display for SendError<T> {
227    fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
228        write!(fmt, "send failed because receiver is gone")
229    }
230}
231
232impl<T> SendError<T> {
233    /// Returns the message that was attempted to be sent but failed.
234    pub fn into_inner(self) -> T {
235        self.0
236    }
237}
238
239#[cfg(test)]
240mod tests {
241    use super::*;
242    use crate::{future::lazy, future::stream_recv};
243
244    #[ntex::test]
245    async fn test_mpsc() {
246        let (tx, mut rx) = channel();
247        assert!(format!("{tx:?}").contains("Sender"));
248        assert!(format!("{rx:?}").contains("Receiver"));
249
250        tx.send("test").unwrap();
251        assert_eq!(stream_recv(&mut rx).await.unwrap(), "test");
252
253        let tx2 = tx.clone();
254        tx2.send("test2").unwrap();
255        assert_eq!(stream_recv(&mut rx).await.unwrap(), "test2");
256
257        assert_eq!(
258            lazy(|cx| Pin::new(&mut rx).poll_next(cx)).await,
259            Poll::Pending
260        );
261        drop(tx2);
262        assert_eq!(
263            lazy(|cx| Pin::new(&mut rx).poll_next(cx)).await,
264            Poll::Pending
265        );
266        drop(tx);
267
268        let (tx, mut rx) = channel::<String>();
269        tx.close();
270        assert_eq!(stream_recv(&mut rx).await, None);
271
272        let (tx, rx) = channel();
273        tx.send("test").unwrap();
274        drop(rx);
275        assert!(tx.send("test").is_err());
276
277        let (tx, _) = channel();
278        let tx2 = tx.clone();
279        tx.close();
280        assert!(tx.send("test").is_err());
281        assert!(tx2.send("test").is_err());
282
283        let err = SendError("test");
284        assert!(format!("{err:?}").contains("SendError"));
285        assert!(format!("{err}").contains("send failed because receiver is gone"));
286        assert_eq!(err.into_inner(), "test");
287    }
288
289    #[ntex::test]
290    async fn test_close() {
291        let (tx, rx) = channel::<()>();
292        assert!(!tx.is_closed());
293        assert!(!rx.is_closed());
294        assert!(!rx.is_terminated());
295
296        tx.close();
297        assert!(tx.is_closed());
298        assert!(rx.is_closed());
299        assert!(rx.is_terminated());
300
301        let (tx, rx) = channel::<()>();
302        assert_eq!(lazy(|cx| rx.poll_recv(cx)).await, Poll::Pending);
303        assert!(rx.shared.get_ref().blocked_recv.is_set());
304        rx.close();
305        assert!(tx.is_closed());
306        assert!(!rx.shared.get_ref().blocked_recv.is_set());
307        assert_eq!(lazy(|cx| rx.poll_recv(cx)).await, Poll::Ready(None));
308
309        let (tx, rx) = channel::<()>();
310        drop(tx);
311        assert!(rx.is_closed());
312        assert!(rx.is_terminated());
313        let _tx = rx.sender();
314        assert!(!rx.is_closed());
315        assert!(!rx.is_terminated());
316    }
317
318    #[ntex::test]
319    async fn test_poll_closed() {
320        let (tx, rx) = channel::<()>();
321        assert_eq!(lazy(|cx| tx.poll_closed(cx)).await, Poll::Pending);
322        assert!(tx.shared.get_ref().closed.is_set());
323        drop(rx);
324        assert!(!tx.shared.get_ref().closed.is_set());
325        assert_eq!(lazy(|cx| tx.poll_closed(cx)).await, Poll::Ready(()));
326        tx.closed().await;
327
328        let (tx, rx) = channel::<()>();
329        assert_eq!(lazy(|cx| tx.poll_closed(cx)).await, Poll::Pending);
330        rx.close();
331        assert!(!tx.shared.get_ref().closed.is_set());
332        tx.closed().await;
333
334        let (tx, _rx) = channel::<()>();
335        let tx2 = tx.clone();
336        assert_eq!(lazy(|cx| tx2.poll_closed(cx)).await, Poll::Pending);
337        tx.close();
338        assert!(!tx.shared.get_ref().closed.is_set());
339        tx2.closed().await;
340    }
341
342    #[test]
343    fn test_drop_receiver_reentrant_send() {
344        use std::rc::Rc;
345
346        struct Msg(Option<Sender<Msg>>, Rc<std::cell::Cell<usize>>);
347
348        impl Drop for Msg {
349            fn drop(&mut self) {
350                self.1.set(self.1.get() + 1);
351                if let Some(tx) = self.0.take() {
352                    assert!(tx.send(Msg(None, self.1.clone())).is_err());
353                }
354            }
355        }
356
357        let drops = Rc::new(std::cell::Cell::new(0));
358        let (tx, rx) = channel();
359        for _ in 0..4 {
360            tx.send(Msg(Some(tx.clone()), drops.clone())).unwrap();
361        }
362        drop(rx);
363        assert_eq!(drops.get(), 8);
364        assert!(tx.is_closed());
365    }
366
367    #[ntex::test]
368    async fn test_fused_drains_buffer() {
369        let (tx, mut rx) = channel();
370        tx.send(1).unwrap();
371        tx.send(2).unwrap();
372        drop(tx);
373
374        assert!(!rx.is_terminated());
375        assert_eq!(stream_recv(&mut rx).await, Some(1));
376        assert_eq!(stream_recv(&mut rx).await, Some(2));
377        assert!(rx.is_terminated());
378        assert_eq!(stream_recv(&mut rx).await, None);
379
380        let (tx, mut rx) = channel();
381        tx.send(1).unwrap();
382        rx.close();
383        assert!(!rx.is_terminated());
384        assert_eq!(stream_recv(&mut rx).await, Some(1));
385        assert!(rx.is_terminated());
386    }
387
388    #[ntex::test]
389    async fn test_mpsc_recv() {
390        let (tx, rx) = channel();
391        tx.send(1).unwrap();
392        assert_eq!(rx.recv().await, Some(1));
393        drop(tx);
394        assert_eq!(rx.recv().await, None);
395    }
396}