moirai_async/sync/
oneshot.rs1#![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
23pub struct Sender<T> {
25 shared: std::sync::Arc<Mutex<SharedState<T>>>,
26}
27
28impl<T> Sender<T> {
29 pub fn send(self, value: T) -> Result<(), T> {
35 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 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
79pub struct Receiver<T> {
81 shared: std::sync::Arc<Mutex<SharedState<T>>>,
82}
83
84impl<T> Receiver<T> {
85 pub fn recv(&mut self) -> RecvFuture<'_, T> {
90 RecvFuture { receiver: self }
91 }
92
93 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 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 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
147pub 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#[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}