dtact_util/sync/
oneshot.rs1use 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
21unsafe impl<T: Send> Send for Inner<T> {}
26unsafe impl<T: Send> Sync for Inner<T> {}
27
28#[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#[repr(align(64))]
51pub struct Sender<T> {
52 inner: Arc<Inner<T>>,
53}
54
55impl<T> Sender<T> {
56 #[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 unsafe {
71 *self.inner.value.get() = Some(value);
72 }
73 self.inner.sent.store(true, Ordering::Release);
74 self.inner.wait.wake_all();
75 Ok(())
80 }
81
82 #[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 self.inner.wait.wake_all();
98 }
99}
100
101#[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#[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 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 if self.inner.sent.load(Ordering::Acquire) {
166 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}