Skip to main content

moirai_async/sync/
oneshot.rs

1#![expect(
2    clippy::unwrap_used,
3    reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
4)]
5
6use std::future::Future;
7use std::pin::Pin;
8use std::sync::Mutex;
9use std::task::{Context, Poll, Waker};
10
11enum OneshotState<T> {
12    Empty,
13    Value(T),
14    Closed,
15}
16
17struct SharedState<T> {
18    state: OneshotState<T>,
19    rx_waker: Option<Waker>,
20    tx_waker: Option<Waker>,
21}
22
23/// Sending half; consumed by the single send.
24pub struct Sender<T> {
25    shared: std::sync::Arc<Mutex<SharedState<T>>>,
26}
27
28impl<T> Sender<T> {
29    /// Send the value, completing the channel.
30    ///
31    /// # Errors
32    ///
33    /// Returns `Err(value)` when the receiver already closed.
34    pub fn send(self, value: T) -> Result<(), T> {
35        // The waker leaves the state lock before it is woken: `Waker::wake` may
36        // poll the task inline on this thread, and that poll re-locks this
37        // state. Same discipline as `mpsc`, `rwlock` and `hybrid::notify`.
38        let waker = {
39            let mut shared = self.shared.lock().unwrap();
40            match shared.state {
41                OneshotState::Empty => {
42                    shared.state = OneshotState::Value(value);
43                    shared.rx_waker.take()
44                }
45                OneshotState::Closed => return Err(value),
46                OneshotState::Value(_) => unreachable!(),
47            }
48        };
49        if let Some(waker) = waker {
50            waker.wake();
51        }
52        Ok(())
53    }
54
55    /// Return whether the receiver closed the channel.
56    pub fn is_closed(&self) -> bool {
57        let shared = self.shared.lock().unwrap();
58        matches!(shared.state, OneshotState::Closed)
59    }
60}
61
62impl<T> Drop for Sender<T> {
63    fn drop(&mut self) {
64        let waker = {
65            let mut shared = self.shared.lock().unwrap();
66            if matches!(shared.state, OneshotState::Empty) {
67                shared.state = OneshotState::Closed;
68                shared.rx_waker.take()
69            } else {
70                None
71            }
72        };
73        if let Some(waker) = waker {
74            waker.wake();
75        }
76    }
77}
78
79/// Receiving half of the single-value channel.
80pub struct Receiver<T> {
81    shared: std::sync::Arc<Mutex<SharedState<T>>>,
82}
83
84impl<T> Receiver<T> {
85    /// Receive the value, waiting for the send.
86    ///
87    /// The returned future resolves `Err(())` when the sender dropped
88    /// without sending.
89    pub fn recv(&mut self) -> RecvFuture<'_, T> {
90        RecvFuture { receiver: self }
91    }
92
93    /// Poll for the value, registering `cx`'s waker while it is not yet
94    /// sent. Resolves `Err(())` when the sender dropped without sending.
95    pub fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Result<T, ()>> {
96        let mut shared = self.shared.lock().unwrap();
97        match std::mem::replace(&mut shared.state, OneshotState::Closed) {
98            OneshotState::Value(v) => Poll::Ready(Ok(v)),
99            OneshotState::Closed => Poll::Ready(Err(())),
100            OneshotState::Empty => {
101                shared.state = OneshotState::Empty;
102                shared.rx_waker = Some(cx.waker().clone());
103                Poll::Pending
104            }
105        }
106    }
107
108    /// Take the value without waiting; `None` when not yet sent.
109    pub fn try_recv(&mut self) -> Option<T> {
110        let mut shared = self.shared.lock().unwrap();
111        match std::mem::replace(&mut shared.state, OneshotState::Closed) {
112            OneshotState::Value(v) => Some(v),
113            OneshotState::Empty => {
114                shared.state = OneshotState::Empty;
115                None
116            }
117            OneshotState::Closed => None,
118        }
119    }
120
121    /// Close the channel, waking a parked sender.
122    pub fn close(&mut self) {
123        let waker = {
124            let mut shared = self.shared.lock().unwrap();
125            shared.state = OneshotState::Closed;
126            shared.tx_waker.take()
127        };
128        if let Some(waker) = waker {
129            waker.wake();
130        }
131    }
132}
133
134impl<T> Drop for Receiver<T> {
135    fn drop(&mut self) {
136        let waker = {
137            let mut shared = self.shared.lock().unwrap();
138            shared.state = OneshotState::Closed;
139            shared.tx_waker.take()
140        };
141        if let Some(waker) = waker {
142            waker.wake();
143        }
144    }
145}
146
147/// Future returned by [`Receiver::recv`].
148pub struct RecvFuture<'a, T> {
149    receiver: &'a mut Receiver<T>,
150}
151
152impl<T> Drop for RecvFuture<'_, T> {
153    fn drop(&mut self) {
154        if let Ok(mut shared) = self.receiver.shared.lock() {
155            shared.rx_waker = None;
156        }
157    }
158}
159
160impl<'a, T> Future for RecvFuture<'a, T> {
161    type Output = Result<T, ()>;
162
163    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
164        self.receiver.poll_recv(cx)
165    }
166}
167
168/// Create a single-value channel.
169#[must_use]
170pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
171    let shared = std::sync::Arc::new(Mutex::new(SharedState {
172        state: OneshotState::Empty,
173        rx_waker: None,
174        tx_waker: None,
175    }));
176    (
177        Sender {
178            shared: shared.clone(),
179        },
180        Receiver { shared },
181    )
182}
183
184#[cfg(test)]
185mod tests {
186    use super::*;
187    use std::future::Future;
188    use std::pin::Pin;
189    use std::task::{Context, Poll, Waker};
190
191    fn poll_future<F: Future + Unpin>(future: &mut F) -> Poll<F::Output> {
192        let mut context = Context::from_waker(Waker::noop());
193        Pin::new(future).poll(&mut context)
194    }
195
196    #[test]
197    fn test_oneshot_send_recv() {
198        let (tx, mut rx) = channel();
199        tx.send(42).unwrap();
200        assert_eq!(rx.try_recv(), Some(42));
201        assert!(rx.try_recv().is_none());
202    }
203
204    #[test]
205    fn test_oneshot_recv_pending_then_ready() {
206        let (tx, mut rx) = channel();
207        let mut recv = rx.recv();
208        assert!(matches!(poll_future(&mut recv), Poll::Pending));
209        tx.send(99).unwrap();
210        assert!(matches!(poll_future(&mut recv), Poll::Ready(Ok(99))));
211    }
212
213    #[test]
214    fn test_oneshot_sender_dropped_recv_err() {
215        let (tx, mut rx) = channel::<i32>();
216        drop(tx);
217        assert!(rx.try_recv().is_none());
218    }
219
220    #[test]
221    fn test_oneshot_recv_closed_err() {
222        let (_, mut rx) = channel::<i32>();
223        rx.close();
224        assert!(rx.try_recv().is_none());
225    }
226
227    #[test]
228    fn test_oneshot_is_closed() {
229        let (tx, mut rx) = channel::<i32>();
230        assert!(!tx.is_closed());
231        rx.close();
232        assert!(tx.is_closed());
233    }
234
235    #[test]
236    fn test_oneshot_double_send_err() {
237        let (tx, rx) = channel();
238        drop(rx);
239        assert!(tx.send(1).is_err());
240    }
241
242    #[test]
243    fn test_oneshot_recv_future_ready() {
244        let (tx, mut rx) = channel();
245        tx.send(7).unwrap();
246        let mut recv = rx.recv();
247        assert!(matches!(poll_future(&mut recv), Poll::Ready(Ok(7))));
248    }
249}