Skip to main content

shuttle_std/
future.rs

1//! Shuttle's implementation of an async executor, roughly equivalent to [`futures::executor`].
2//!
3//! The [spawn] method spawns a new asynchronous task that the executor will run to completion. The
4//! [block_on] method blocks the current thread on the completion of a future.
5//!
6//! [`futures::executor`]: https://docs.rs/futures/0.3.13/futures/executor/index.html
7
8use shuttle_engine::backtrace_enabled;
9use shuttle_engine::runtime::execution::ExecutionState;
10use shuttle_engine::runtime::task::TaskId;
11use shuttle_engine::runtime::thread;
12use std::error::Error;
13use std::fmt::{Display, Formatter};
14use std::future::Future;
15use std::panic::Location;
16use std::pin::Pin;
17use std::result::Result;
18use std::sync::atomic::{AtomicBool, Ordering};
19use std::sync::Arc;
20use std::task::{Context, Poll, Waker};
21
22pub use shuttle_engine::future::batch_semaphore;
23
24fn spawn_inner<F>(fut: F, caller: &'static Location<'static>) -> JoinHandle<F::Output>
25where
26    F: Future + 'static,
27    F::Output: 'static,
28{
29    let stack_size = ExecutionState::with(|s| s.config.stack_size);
30    let inner = Arc::new(std::sync::Mutex::new(JoinHandleInner::default()));
31    let aborted = Arc::new(AtomicBool::new(false));
32    let task_id = ExecutionState::spawn_future(
33        Wrapper::new(fut, inner.clone(), aborted.clone()),
34        stack_size,
35        None,
36        caller,
37    );
38
39    JoinHandle {
40        task_id,
41        inner,
42        aborted,
43    }
44}
45
46/// Spawn a new async task that the executor will run to completion.
47#[track_caller]
48pub fn spawn<F>(fut: F) -> JoinHandle<F::Output>
49where
50    F: Future + Send + 'static,
51    F::Output: Send + 'static,
52{
53    spawn_inner(fut, Location::caller())
54}
55
56/// Spawn a new async task that the executor will run to completion.
57/// This is just `spawn` without the `Send` bound, and it mirrors `spawn_local` from Tokio.
58#[track_caller]
59pub fn spawn_local<F>(fut: F) -> JoinHandle<F::Output>
60where
61    F: Future + 'static,
62    F::Output: 'static,
63{
64    spawn_inner(fut, Location::caller())
65}
66
67/// An owned permission to abort a spawned task, without awaiting its completion.
68#[derive(Debug, Clone)]
69pub struct AbortHandle {
70    task_id: TaskId,
71    aborted: Arc<AtomicBool>,
72}
73
74impl AbortHandle {
75    /// Abort the task associated with the handle.
76    ///
77    /// The task will be cancelled at the next await point. If the task has already
78    /// completed, this is a no-op.
79    pub fn abort(&self) {
80        // Scheduling point: the scheduler may run other tasks (including the target)
81        // before the abort flag is set, creating interleavings where the task completes
82        // normally despite an abort() call.
83        thread::switch();
84
85        // Signal the Wrapper to skip the inner future on the next poll.
86        // If already aborted, skip the redundant wake.
87        if self.aborted.swap(true, Ordering::Relaxed) {
88            return;
89        }
90        // Wake the task so Wrapper::poll() runs and performs cleanup.
91        let res = ExecutionState::try_with(|state| {
92            if !state.is_finished() {
93                state.get_mut(self.task_id).abort();
94            }
95        });
96        if let Err(e) = res {
97            tracing::error!("`AbortHandle::abort` failed with error: {e:?}");
98        }
99    }
100
101    /// Returns `true` if this task is finished, otherwise returns `false`.
102    ///
103    /// ## Panics
104    /// Panics if called outside of shuttle context, i.e. if there is no execution context.
105    pub fn is_finished(&self) -> bool {
106        ExecutionState::with(|state| {
107            let task = state.get(self.task_id);
108            task.finished()
109        })
110    }
111}
112
113unsafe impl Send for AbortHandle {}
114unsafe impl Sync for AbortHandle {}
115
116/// An owned permission to join on an async task (await its termination).
117#[derive(Debug)]
118pub struct JoinHandle<T> {
119    task_id: TaskId,
120    inner: Arc<std::sync::Mutex<JoinHandleInner<T>>>,
121    aborted: Arc<AtomicBool>,
122}
123
124#[derive(Debug)]
125struct JoinHandleInner<T> {
126    result: Option<Result<T, JoinError>>,
127    waker: Option<Waker>,
128}
129
130impl<T> Default for JoinHandleInner<T> {
131    fn default() -> Self {
132        JoinHandleInner {
133            result: None,
134            waker: None,
135        }
136    }
137}
138
139impl<T> JoinHandle<T> {
140    /// Abort the task associated with the handle.
141    ///
142    /// The task will be cancelled at the next await point. Awaiting the `JoinHandle` after
143    /// calling `abort` will return `Err(JoinError::Cancelled)`, unless the task completed
144    /// before the abort took effect.
145    pub fn abort(&self) {
146        // Scheduling point: the scheduler may run other tasks (including the target)
147        // before the abort flag is set, creating interleavings where the task completes
148        // normally despite an abort() call.
149        thread::switch();
150
151        // Signal the Wrapper to skip the inner future on the next poll.
152        // If already aborted, skip the redundant wake.
153        if self.aborted.swap(true, Ordering::Relaxed) {
154            return;
155        }
156        // Wake the task so Wrapper::poll() runs and performs cleanup.
157        let res = ExecutionState::try_with(|state| {
158            if !state.is_finished() {
159                state.get_mut(self.task_id).abort();
160            }
161        });
162        if let Err(e) = res {
163            tracing::error!("`JoinHandle::abort` failed with error: {e:?}");
164        }
165    }
166
167    /// Returns `true` if this task is finished, otherwise returns `false`.
168    ///
169    /// ## Panics
170    /// Panics if called outside of shuttle context, i.e. if there is no execution context.
171    pub fn is_finished(&self) -> bool {
172        ExecutionState::with(|state| {
173            let task = state.get(self.task_id);
174            task.finished()
175        })
176    }
177
178    /// Returns a new `AbortHandle` that can be used to remotely abort this task.
179    pub fn abort_handle(&self) -> AbortHandle {
180        AbortHandle {
181            task_id: self.task_id,
182            aborted: self.aborted.clone(),
183        }
184    }
185}
186
187// TODO: need to work out all the error cases here
188/// Task failed to execute to completion.
189#[derive(Debug)]
190pub enum JoinError {
191    /// Task was aborted
192    Cancelled,
193}
194
195impl Display for JoinError {
196    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
197        match self {
198            JoinError::Cancelled => write!(f, "task was cancelled"),
199        }
200    }
201}
202
203impl Error for JoinError {}
204
205impl<T> Drop for JoinHandle<T> {
206    fn drop(&mut self) {
207        // Detach the task so it keeps running but we don't wait for it.
208        // Dropping a JoinHandle does NOT cancel the task (unlike abort()).
209        let res = ExecutionState::try_with(|state| {
210            if !state.is_finished() {
211                state.get_mut(self.task_id).detach();
212            }
213        });
214        if let Err(e) = res {
215            tracing::error!("`JoinHandle::drop` failed with error: {e:?}");
216        }
217    }
218}
219
220impl<T> Future for JoinHandle<T> {
221    type Output = Result<T, JoinError>;
222
223    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
224        let mut lock = self.inner.lock().unwrap();
225        if let Some(result) = lock.result.take() {
226            Poll::Ready(result)
227        } else {
228            lock.waker = Some(cx.waker().clone());
229
230            ExecutionState::with(|state| {
231                state.current_mut().backtrace = if backtrace_enabled() {
232                    Some(std::backtrace::Backtrace::force_capture())
233                } else {
234                    None
235                }
236            });
237
238            Poll::Pending
239        }
240    }
241}
242
243// We wrap a task returning a value inside a wrapper task that returns (). The wrapper
244// contains a mutex-wrapped field that stores the value and the waker for the task
245// waiting on the join handle. When `poll` returns `Poll::Ready`, the `Wrapper` stores
246// the result in the `result` field and wakes the `waker`.
247//
248// The `aborted` flag is set by `JoinHandle::abort()` / `AbortHandle::abort()`. On the
249// next poll, `Wrapper` detects the flag, drops the inner future (running its destructors),
250// cleans up thread-local storage, publishes `Err(JoinError::Cancelled)` to any awaiting
251// join handle, and returns `Poll::Ready(())`.
252struct Wrapper<F: Future> {
253    /// The inner future. Wrapped in `Option` so we can drop it explicitly in the
254    /// abort path (before running thread-local destructors).
255    future: Option<Pin<Box<F>>>,
256    inner: Arc<std::sync::Mutex<JoinHandleInner<F::Output>>>,
257    aborted: Arc<AtomicBool>,
258}
259
260impl<F> Wrapper<F>
261where
262    F: Future + 'static,
263    F::Output: 'static,
264{
265    fn new(future: F, inner: Arc<std::sync::Mutex<JoinHandleInner<F::Output>>>, aborted: Arc<AtomicBool>) -> Self {
266        Self {
267            future: Some(Box::pin(future)),
268            inner,
269            aborted,
270        }
271    }
272}
273
274impl<F> Wrapper<F>
275where
276    F: Future + 'static,
277    F::Output: 'static,
278{
279    /// Clean up thread-local storage, publish the result to the JoinHandle, and wake
280    /// any task waiting on the result.
281    fn finish(&self, result: Result<F::Output, JoinError>) {
282        // Run thread-local destructors.
283        // See `pop_local` for details on why this loop looks slightly funky.
284        while let Some(local) = ExecutionState::with(|state| state.current_mut().pop_local()) {
285            drop(local);
286        }
287
288        let mut lock = self.inner.lock().unwrap();
289        lock.result = Some(result);
290        if let Some(waker) = lock.waker.take() {
291            waker.wake();
292        }
293    }
294}
295
296impl<F> Future for Wrapper<F>
297where
298    F: Future + 'static,
299    F::Output: 'static,
300{
301    type Output = ();
302
303    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
304        let this = self.get_mut();
305
306        // If abort() was called, skip polling the inner future.
307        if this.aborted.load(Ordering::Relaxed) {
308            // If the execution is already finished (e.g., the runtime is shutting down),
309            // we can't access task state for cleanup. Just return Ready to let the task
310            // be cleaned up by the runtime.
311            if ExecutionState::try_with(|state| state.is_finished()).unwrap_or(true) {
312                return Poll::Ready(());
313            }
314
315            // Drop the inner future first so its destructors can still access TLS.
316            this.future.take();
317            this.finish(Err(JoinError::Cancelled));
318            return Poll::Ready(());
319        }
320
321        match this.future.as_mut().unwrap().as_mut().poll(cx) {
322            Poll::Ready(result) => {
323                // If we've finished execution already (this task was detached), don't clean up. We
324                // can't access the state any more to destroy thread locals, and don't want to run
325                // any more wakers (which will be no-ops anyway).
326                if ExecutionState::try_with(|state| state.is_finished()).unwrap_or(true) {
327                    return Poll::Ready(());
328                }
329
330                this.finish(Ok(result));
331                Poll::Ready(())
332            }
333            Poll::Pending => Poll::Pending,
334        }
335    }
336}
337
338/// Run a future to completion on the current thread.
339pub fn block_on<F: Future>(future: F) -> F::Output {
340    let mut future = Box::pin(future);
341    let waker = ExecutionState::with(|state| state.current_mut().waker());
342    let cx = &mut Context::from_waker(&waker);
343
344    // Note: we only switch on poll pending, since this blocks the current task. This means that *internal*
345    // Shuttle futures which do not use other Shuttle primitives such as `batch_semaphore::Acquire` must
346    // have a scheduling point prior to first poll if that poll will be successful and can affect other tasks.
347    // For example, an uncontested acquire makes other threads block or fail try-acquires, so there must be
348    // a scheduling point for scheduling completeness. For *external* futures, this is a non-issue because they
349    // should use other Shuttle primitives inside of `poll` if polling can affect other threads.
350    loop {
351        match future.as_mut().poll(cx) {
352            Poll::Ready(result) => break result,
353            Poll::Pending => {
354                ExecutionState::with(|state| state.current_mut().sleep_unless_woken());
355                thread::switch();
356            }
357        }
358    }
359}
360
361/// Yields execution back to the scheduler.
362///
363/// Borrowed from the Tokio implementation.
364pub async fn yield_now() {
365    /// Yield implementation
366    struct YieldNow {
367        yielded: bool,
368    }
369
370    impl Future for YieldNow {
371        type Output = ();
372
373        fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
374            if self.yielded {
375                return Poll::Ready(());
376            }
377
378            self.yielded = true;
379            cx.waker().wake_by_ref();
380            ExecutionState::request_yield();
381            Poll::Pending
382        }
383    }
384
385    YieldNow { yielded: false }.await
386}