1use std::any::Any;
2use std::future::Future;
3use std::pin::Pin;
4use std::sync::atomic::{AtomicBool, AtomicU8, Ordering};
5use std::sync::{Arc, Mutex};
6use std::task::{Context, Poll};
7
8use async_channel::{Receiver, Sender};
9use futures_lite::Stream;
10
11const ACTIVE: u8 = 0;
12const FINISHED: u8 = 1;
13
14#[must_use = "tasks are cancelled when dropped; await or detach the handle"]
16pub struct Task<T> {
17 inner: Option<TaskInner<T>>,
18}
19
20enum TaskInner<T> {
21 Direct(async_task::Task<T>),
22 Bridge(BridgeTask<T>),
23}
24
25impl<T> Task<T> {
26 pub(crate) fn direct(inner: async_task::Task<T>) -> Self {
27 Self {
28 inner: Some(TaskInner::Direct(inner)),
29 }
30 }
31
32 pub(crate) fn bridge(on_finish: impl Fn() + Send + Sync + 'static) -> (Self, BridgeDriver<T>)
33 where
34 T: Send + 'static,
35 {
36 let (sender, receiver) = async_channel::bounded(1);
37 let shared = Arc::new(BridgeShared {
38 state: AtomicU8::new(ACTIVE),
39 cancel_requested: AtomicBool::new(false),
40 cancel_sender: Mutex::new(None),
41 sender,
42 on_finish: Box::new(on_finish),
43 });
44 let (cancel_sender, cancel_receiver) = async_channel::bounded(1);
45 *shared.cancel_sender.lock().expect("bridge mutex poisoned") = Some(cancel_sender);
46 let bridge = BridgeTask {
47 shared: Arc::clone(&shared),
48 receiver: Box::pin(receiver),
49 };
50 (
51 Self {
52 inner: Some(TaskInner::Bridge(bridge)),
53 },
54 BridgeDriver {
55 shared,
56 cancel_receiver,
57 },
58 )
59 }
60
61 pub fn detach(mut self) {
63 match self.inner.take() {
64 Some(TaskInner::Direct(task)) => task.detach(),
65 Some(TaskInner::Bridge(bridge)) => bridge.detach(),
66 None => {}
67 }
68 }
69
70 pub async fn cancel(mut self) -> Option<T> {
72 match self.inner.take() {
73 Some(TaskInner::Direct(task)) => task.cancel().await,
74 Some(TaskInner::Bridge(bridge)) => bridge.cancel().await,
75 None => None,
76 }
77 }
78
79 pub fn fallible(mut self) -> FallibleTask<T> {
86 let inner = match self.inner.take() {
87 Some(TaskInner::Direct(task)) => FallibleInner::Direct(task.fallible()),
88 Some(TaskInner::Bridge(bridge)) => FallibleInner::Bridge(bridge),
89 None => panic!("task handle was already consumed"),
90 };
91 FallibleTask { inner: Some(inner) }
92 }
93
94 pub fn is_finished(&self) -> bool {
96 match self.inner.as_ref() {
97 Some(TaskInner::Direct(task)) => task.is_finished(),
98 Some(TaskInner::Bridge(task)) => task.is_finished(),
99 None => true,
100 }
101 }
102}
103
104impl<T> Drop for Task<T> {
105 fn drop(&mut self) {
106 if let Some(TaskInner::Bridge(bridge)) = self.inner.as_ref() {
107 bridge.request_cancel();
108 }
109 }
110}
111
112impl<T> Unpin for Task<T> {}
113
114impl<T> Future for Task<T> {
115 type Output = T;
116
117 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
118 match self
119 .get_mut()
120 .inner
121 .as_mut()
122 .expect("task handle was already consumed")
123 {
124 TaskInner::Direct(task) => Pin::new(task).poll(cx),
125 TaskInner::Bridge(task) => Pin::new(task).poll(cx),
126 }
127 }
128}
129
130#[must_use = "tasks are cancelled when dropped; await the handle"]
132pub struct FallibleTask<T> {
133 inner: Option<FallibleInner<T>>,
134}
135
136enum FallibleInner<T> {
137 Direct(async_task::FallibleTask<T>),
138 Bridge(BridgeTask<T>),
139}
140
141impl<T> Drop for FallibleTask<T> {
142 fn drop(&mut self) {
143 if let Some(FallibleInner::Bridge(bridge)) = self.inner.as_ref() {
144 bridge.request_cancel();
145 }
146 }
147}
148
149impl<T> Unpin for FallibleTask<T> {}
150
151impl<T> Future for FallibleTask<T> {
152 type Output = Option<T>;
153
154 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
155 match self
156 .get_mut()
157 .inner
158 .as_mut()
159 .expect("task handle was already consumed")
160 {
161 FallibleInner::Direct(task) => Pin::new(task).poll(cx),
162 FallibleInner::Bridge(task) => task.poll_fallible(cx),
163 }
164 }
165}
166
167pub(crate) enum Completion<T> {
168 Completed(T),
169 Cancelled,
170 Panicked(Box<dyn Any + Send + 'static>),
171}
172
173struct BridgeShared<T> {
174 state: AtomicU8,
175 cancel_requested: AtomicBool,
176 cancel_sender: Mutex<Option<Sender<()>>>,
177 sender: Sender<Completion<T>>,
178 on_finish: Box<dyn Fn() + Send + Sync>,
179}
180
181struct BridgeTask<T> {
182 shared: Arc<BridgeShared<T>>,
183 receiver: Pin<Box<Receiver<Completion<T>>>>,
184}
185
186impl<T> BridgeTask<T> {
187 fn detach(self) {
188 drop(self);
189 }
190
191 fn request_cancel(&self) {
192 self.shared.cancel_requested.store(true, Ordering::Release);
193 if let Some(sender) = self
194 .shared
195 .cancel_sender
196 .lock()
197 .expect("bridge mutex poisoned")
198 .as_ref()
199 {
200 let _ = sender.try_send(());
201 }
202 }
203
204 fn is_finished(&self) -> bool {
205 self.shared.state.load(Ordering::Acquire) == FINISHED
206 }
207
208 async fn cancel(mut self) -> Option<T> {
209 self.request_cancel();
210 match self.receiver.as_mut().recv().await {
211 Ok(Completion::Completed(value)) => Some(value),
212 Ok(Completion::Cancelled) | Err(_) => None,
213 Ok(Completion::Panicked(payload)) => std::panic::resume_unwind(payload),
214 }
215 }
216
217 fn poll_fallible(&mut self, cx: &mut Context<'_>) -> Poll<Option<T>> {
218 match self.receiver.as_mut().poll_next(cx) {
219 Poll::Pending => Poll::Pending,
220 Poll::Ready(Some(Completion::Completed(value))) => Poll::Ready(Some(value)),
221 Poll::Ready(Some(Completion::Cancelled) | None) => Poll::Ready(None),
222 Poll::Ready(Some(Completion::Panicked(payload))) => std::panic::resume_unwind(payload),
223 }
224 }
225}
226
227impl<T> Future for BridgeTask<T> {
228 type Output = T;
229
230 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
231 match self.receiver.as_mut().poll_next(cx) {
232 Poll::Pending => Poll::Pending,
233 Poll::Ready(Some(Completion::Completed(value))) => Poll::Ready(value),
234 Poll::Ready(Some(Completion::Cancelled) | None) => {
235 panic!("task was cancelled")
236 }
237 Poll::Ready(Some(Completion::Panicked(payload))) => std::panic::resume_unwind(payload),
238 }
239 }
240}
241
242pub(crate) struct BridgeDriver<T> {
244 shared: Arc<BridgeShared<T>>,
245 cancel_receiver: Receiver<()>,
246}
247
248impl<T> Clone for BridgeDriver<T> {
249 fn clone(&self) -> Self {
250 Self {
251 shared: Arc::clone(&self.shared),
252 cancel_receiver: self.cancel_receiver.clone(),
253 }
254 }
255}
256
257impl<T> BridgeDriver<T> {
258 pub(crate) fn is_cancel_requested(&self) -> bool {
259 self.shared.cancel_requested.load(Ordering::Acquire)
260 }
261
262 pub(crate) async fn cancelled(self) {
263 let _ = self.cancel_receiver.recv().await;
264 }
265
266 pub(crate) fn complete(&self, completion: Completion<T>) {
267 if self
268 .shared
269 .state
270 .compare_exchange(ACTIVE, FINISHED, Ordering::AcqRel, Ordering::Acquire)
271 .is_ok()
272 {
273 let _ = self.shared.sender.try_send(completion);
274 (self.shared.on_finish)();
275 *self
276 .shared
277 .cancel_sender
278 .lock()
279 .expect("bridge mutex poisoned") = None;
280 }
281 }
282}
283
284pub(crate) struct BridgeCompletionGuard<T>(Option<BridgeDriver<T>>);
286
287impl<T> BridgeCompletionGuard<T> {
288 pub(crate) fn new(driver: BridgeDriver<T>) -> Self {
289 Self(Some(driver))
290 }
291
292 pub(crate) fn finish(mut self, completion: Completion<T>) {
293 if let Some(driver) = self.0.take() {
294 driver.complete(completion);
295 }
296 }
297}
298
299impl<T> Drop for BridgeCompletionGuard<T> {
300 fn drop(&mut self) {
301 if let Some(driver) = self.0.take() {
302 driver.complete(Completion::Cancelled);
303 }
304 }
305}