Skip to main content

ntex_util/channel/
oneshot.rs

1//! A one-shot, futures-aware channel.
2use std::{future::Future, future::poll_fn, pin::Pin, task::Context, task::Poll};
3
4use super::{Canceled, cell::Cell};
5use crate::task::LocalWaker;
6
7/// Creates a new futures-aware, one-shot channel.
8pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
9    let inner = Cell::new(Inner {
10        value: None,
11        rx_task: LocalWaker::new(),
12    });
13    let tx = Sender {
14        inner: inner.clone(),
15    };
16    let rx = Receiver { inner };
17    (tx, rx)
18}
19
20/// Represents the completion half of a oneshot through which the result of a
21/// computation is signaled.
22#[derive(Debug)]
23pub struct Sender<T> {
24    inner: Cell<Inner<T>>,
25}
26
27/// A future representing the completion of a computation happening elsewhere in
28/// memory.
29#[derive(Debug)]
30#[must_use = "futures do nothing unless polled"]
31pub struct Receiver<T> {
32    inner: Cell<Inner<T>>,
33}
34
35// The channels do not ever project Pin to the inner T
36impl<T> Unpin for Receiver<T> {}
37impl<T> Unpin for Sender<T> {}
38
39#[derive(Debug)]
40struct Inner<T> {
41    value: Option<T>,
42    rx_task: LocalWaker,
43}
44
45impl<T> Sender<T> {
46    /// Completes this oneshot with a successful result.
47    ///
48    /// This function will consume `self` and indicate to the other end, the
49    /// `Receiver`, that the value provided is the result of the computation this
50    /// represents.
51    ///
52    /// If the value is successfully enqueued for the remote end to receive,
53    /// then `Ok(())` is returned. If the receiving end was dropped before
54    /// this function was called, however, then `Err` is returned with the value
55    /// provided.
56    pub fn send(self, val: T) -> Result<(), T> {
57        if self.inner.strong_count() == 2 {
58            let inner = self.inner.get_mut();
59            inner.value = Some(val);
60            inner.rx_task.wake();
61            Ok(())
62        } else {
63            Err(val)
64        }
65    }
66
67    /// Tests to see whether this `Sender`'s corresponding `Receiver`
68    /// has gone away.
69    pub fn is_canceled(&self) -> bool {
70        self.inner.strong_count() == 1
71    }
72}
73
74impl<T> Drop for Sender<T> {
75    fn drop(&mut self) {
76        self.inner.get_ref().rx_task.wake();
77    }
78}
79
80impl<T> Receiver<T> {
81    /// Waits for the value.
82    ///
83    /// Returns [`Canceled`] if the sender is dropped without sending a value.
84    pub async fn recv(&self) -> Result<T, Canceled> {
85        poll_fn(|cx| self.poll_recv(cx)).await
86    }
87
88    /// Polls for the value.
89    ///
90    /// Returns [`Canceled`] if the sender is dropped without sending a value.
91    pub fn poll_recv(&self, cx: &mut Context<'_>) -> Poll<Result<T, Canceled>> {
92        // If we've got a value, then skip the logic below as we're done.
93        if let Some(val) = self.inner.get_mut().value.take() {
94            return Poll::Ready(Ok(val));
95        }
96
97        // Check if sender is dropped and return error if it is.
98        if self.inner.strong_count() == 1 {
99            Poll::Ready(Err(Canceled))
100        } else {
101            self.inner.get_ref().rx_task.register(cx.waker());
102            Poll::Pending
103        }
104    }
105}
106
107impl<T> Future for Receiver<T> {
108    type Output = Result<T, Canceled>;
109
110    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
111        self.poll_recv(cx)
112    }
113}
114
115#[cfg(test)]
116mod tests {
117    use super::*;
118    use crate::future::lazy;
119
120    #[ntex::test]
121    async fn test_oneshot() {
122        let (tx, rx) = channel();
123        assert!(format!("{tx:?}").contains("Sender"));
124        assert!(format!("{rx:?}").contains("Receiver"));
125
126        tx.send("test").unwrap();
127        assert_eq!(rx.await.unwrap(), "test");
128
129        let (tx, rx) = channel();
130        tx.send("test").unwrap();
131        assert_eq!(rx.recv().await.unwrap(), "test");
132
133        let (tx, rx) = channel();
134        assert!(!tx.is_canceled());
135        drop(rx);
136        assert!(tx.is_canceled());
137        assert!(tx.send("test").is_err());
138
139        let (tx, rx) = channel::<&'static str>();
140        drop(tx);
141        assert!(rx.await.is_err());
142
143        let (tx, mut rx) = channel::<&'static str>();
144        assert_eq!(lazy(|cx| Pin::new(&mut rx).poll(cx)).await, Poll::Pending);
145        tx.send("test").unwrap();
146        assert_eq!(rx.await.unwrap(), "test");
147
148        let (tx, mut rx) = channel::<&'static str>();
149        assert_eq!(lazy(|cx| Pin::new(&mut rx).poll(cx)).await, Poll::Pending);
150        drop(tx);
151        assert!(rx.await.is_err());
152    }
153}