Skip to main content

dtact_util/sync/
oneshot.rs

1//! Single-value, single-use channel — one [`Sender`] sends at most one
2//! value to one [`Receiver`].
3
4use super::wait_queue::WaitQueue;
5use std::cell::UnsafeCell;
6use std::future::Future;
7use std::pin::Pin;
8use std::sync::Arc;
9use std::sync::atomic::{AtomicBool, Ordering};
10use std::task::{Context, Poll};
11
12#[repr(align(64))]
13struct Inner<T> {
14    value: UnsafeCell<Option<T>>,
15    sent: AtomicBool,
16    sender_dropped: AtomicBool,
17    receiver_dropped: AtomicBool,
18    wait: WaitQueue,
19}
20
21// SAFETY: `sent`'s Release/Acquire pair is what makes writing `value` (by
22// the sender, once) and reading it (by the receiver, once) not race —
23// `send` writes then sets `sent`; `poll` never reads `value` without
24// first observing `sent`.
25unsafe impl<T: Send> Send for Inner<T> {}
26unsafe impl<T: Send> Sync for Inner<T> {}
27
28/// Create a connected sender/receiver pair for one value.
29#[must_use]
30#[inline]
31pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
32    let inner = Arc::new(Inner {
33        value: UnsafeCell::new(None),
34        sent: AtomicBool::new(false),
35        sender_dropped: AtomicBool::new(false),
36        receiver_dropped: AtomicBool::new(false),
37        wait: WaitQueue::new(),
38    });
39    (
40        Sender {
41            inner: inner.clone(),
42        },
43        Receiver { inner },
44    )
45}
46
47/// The sending half of a [`channel`]. Consumed by [`Sender::send`] — a
48/// oneshot sender can only ever send once, so there's no `&self` send
49/// method to accidentally call twice.
50#[repr(align(64))]
51pub struct Sender<T> {
52    inner: Arc<Inner<T>>,
53}
54
55impl<T> Sender<T> {
56    /// Send `value` to the receiver.
57    ///
58    /// # Errors
59    /// Returns `value` back if the receiver was already dropped (nobody
60    /// left to receive it).
61    #[inline(always)]
62    pub fn send(self, value: T) -> Result<(), T> {
63        if self.inner.receiver_dropped.load(Ordering::Acquire) {
64            return Err(value);
65        }
66        // SAFETY: `Sender::send` consumes `self` and is the only writer
67        // of `value`, called at most once (ownership prevents a second
68        // call); no `Receiver` read can observe `value` before `sent` is
69        // published below.
70        unsafe {
71            *self.inner.value.get() = Some(value);
72        }
73        self.inner.sent.store(true, Ordering::Release);
74        self.inner.wait.wake_all();
75        // `self` drops normally here — `Sender`'s `Drop` impl always
76        // fires, but it's harmless post-send: it only marks
77        // `sender_dropped` and wakes the receiver again, which is a no-op
78        // once `sent` is already true (`try_take` checks `sent` first).
79        Ok(())
80    }
81
82    /// `true` if the receiver has already been dropped — a subsequent
83    /// [`send`](Self::send) is guaranteed to fail.
84    #[must_use]
85    #[inline(always)]
86    pub fn is_closed(&self) -> bool {
87        self.inner.receiver_dropped.load(Ordering::Acquire)
88    }
89}
90
91impl<T> Drop for Sender<T> {
92    #[inline(always)]
93    fn drop(&mut self) {
94        self.inner.sender_dropped.store(true, Ordering::Release);
95        // Wake the receiver so a pending `.await` observes the closure
96        // (`RecvError`) instead of hanging forever.
97        self.inner.wait.wake_all();
98    }
99}
100
101/// The receiving half of a [`channel`]. Implements [`Future`] directly —
102/// `receiver.await` resolves once, either with the sent value or
103/// [`RecvError`] if the sender was dropped without sending.
104#[repr(align(64))]
105pub struct Receiver<T> {
106    inner: Arc<Inner<T>>,
107}
108
109impl<T> Drop for Receiver<T> {
110    #[inline(always)]
111    fn drop(&mut self) {
112        self.inner.receiver_dropped.store(true, Ordering::Release);
113    }
114}
115
116/// Error returned by a [`Receiver`] when the [`Sender`] was dropped
117/// without sending a value.
118#[derive(Debug, Clone, Copy, PartialEq, Eq)]
119#[repr(align(64))]
120pub struct RecvError;
121
122impl std::fmt::Display for RecvError {
123    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
124        f.write_str("sender dropped without sending a value")
125    }
126}
127
128impl std::error::Error for RecvError {}
129
130impl<T> Future for Receiver<T> {
131    type Output = Result<T, RecvError>;
132
133    #[inline(always)]
134    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
135        if let Some(result) = self.try_take() {
136            return Poll::Ready(result);
137        }
138        let token = self.inner.wait.register(cx.waker());
139        if let Some(result) = self.try_take() {
140            self.inner.wait.cancel(token);
141            return Poll::Ready(result);
142        }
143        Poll::Pending
144    }
145}
146
147impl<T> Receiver<T> {
148    #[inline(always)]
149    fn try_take(&self) -> Option<Result<T, RecvError>> {
150        if self.inner.sent.load(Ordering::Acquire) {
151            // SAFETY: `sent` observed true under Acquire, paired with the
152            // Release store in `Sender::send` after writing `value` — the
153            // value is visible and, since `sent` only ever transitions
154            // false -> true once, this is the only place that ever takes it.
155            let value = unsafe { (*self.inner.value.get()).take() };
156            return Some(value.ok_or(RecvError));
157        }
158        if self.inner.sender_dropped.load(Ordering::Acquire) {
159            // `send` (which sets `sent` *before* the `Sender` itself
160            // drops and sets `sender_dropped`) could have completed
161            // concurrently with the `sent` check above — see
162            // `mpsc::Receiver::poll_recv`'s identical comment for why one
163            // more check is needed before reporting the sender gone
164            // without a value.
165            if self.inner.sent.load(Ordering::Acquire) {
166                // SAFETY: same as the identical block above.
167                let value = unsafe { (*self.inner.value.get()).take() };
168                return Some(value.ok_or(RecvError));
169            }
170            return Some(Err(RecvError));
171        }
172        None
173    }
174}