Skip to main content

async_runtime/
task.rs

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/// An awaitable task handle. Dropping it cancels the task.
15#[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    /// Lets the task continue in the background.
62    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    /// Requests cancellation and waits for cancellation to finish.
71    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    /// Converts cancellation from a panic into `None`.
80    ///
81    /// # Panics
82    ///
83    /// Panics only if the internal task handle has already been consumed. This
84    /// cannot occur through the public API because this method consumes `self`.
85    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    /// Returns whether the task has finished.
95    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/// A task handle that resolves to `None` if the task is cancelled.
131#[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
242/// Owner-thread half of a remote-local bridge. It contains only `Send` data.
243pub(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
284/// Ensures an accepted remote task becomes cancelled if its local wrapper is dropped.
285pub(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}