1use 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#[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#[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#[derive(Debug, Clone)]
69pub struct AbortHandle {
70 task_id: TaskId,
71 aborted: Arc<AtomicBool>,
72}
73
74impl AbortHandle {
75 pub fn abort(&self) {
80 thread::switch();
84
85 if self.aborted.swap(true, Ordering::Relaxed) {
88 return;
89 }
90 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 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#[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 pub fn abort(&self) {
146 thread::switch();
150
151 if self.aborted.swap(true, Ordering::Relaxed) {
154 return;
155 }
156 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 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 pub fn abort_handle(&self) -> AbortHandle {
180 AbortHandle {
181 task_id: self.task_id,
182 aborted: self.aborted.clone(),
183 }
184 }
185}
186
187#[derive(Debug)]
190pub enum JoinError {
191 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 let res = ExecutionState::try_with(|state| {
210 if !state.is_finished() {
211 state.detach(self.task_id);
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
243struct Wrapper<F: Future> {
253 future: Option<Pin<Box<F>>>,
256 inner: Option<Arc<std::sync::Mutex<JoinHandleInner<F::Output>>>>,
258 aborted: Arc<AtomicBool>,
259}
260
261impl<F> Wrapper<F>
262where
263 F: Future + 'static,
264 F::Output: 'static,
265{
266 fn new(future: F, inner: Arc<std::sync::Mutex<JoinHandleInner<F::Output>>>, aborted: Arc<AtomicBool>) -> Self {
267 Self {
268 future: Some(Box::pin(future)),
269 inner: Some(inner),
270 aborted,
271 }
272 }
273}
274
275impl<F> Wrapper<F>
276where
277 F: Future + 'static,
278 F::Output: 'static,
279{
280 fn finish(&mut self, result: Result<F::Output, JoinError>) {
283 ExecutionState::drop_task_locals();
285
286 let inner = self.inner.take().expect("a task's result is published once");
287 let mut lock = inner.lock().unwrap();
288 lock.result = Some(result);
289 if let Some(waker) = lock.waker.take() {
290 waker.wake();
291 }
292 }
293}
294
295impl<F: Future> Drop for Wrapper<F> {
296 fn drop(&mut self) {
297 if let Some(inner) = self.inner.take() {
302 self.future.take();
304 if !ExecutionState::should_stop() {
305 let mut lock = inner.lock().unwrap();
306 lock.result = Some(Err(JoinError::Cancelled));
307 if let Some(waker) = lock.waker.take() {
308 waker.wake();
309 }
310 }
311 }
312 }
313}
314
315impl<F> Future for Wrapper<F>
316where
317 F: Future + 'static,
318 F::Output: 'static,
319{
320 type Output = ();
321
322 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
323 let this = self.get_mut();
324
325 if this.aborted.load(Ordering::Relaxed) {
327 if ExecutionState::try_with(|state| state.is_finished()).unwrap_or(true) {
331 return Poll::Ready(());
332 }
333
334 this.future.take();
336 this.finish(Err(JoinError::Cancelled));
337 return Poll::Ready(());
338 }
339
340 match this.future.as_mut().unwrap().as_mut().poll(cx) {
341 Poll::Ready(result) => {
342 if ExecutionState::try_with(|state| state.is_finished()).unwrap_or(true) {
346 return Poll::Ready(());
347 }
348
349 this.finish(Ok(result));
350 Poll::Ready(())
351 }
352 Poll::Pending => Poll::Pending,
353 }
354 }
355}
356
357pub fn block_on<F: Future>(future: F) -> F::Output {
359 let mut future = Box::pin(future);
360 let waker = ExecutionState::with(|state| state.current_mut().waker());
361 let cx = &mut Context::from_waker(&waker);
362
363 loop {
370 match future.as_mut().poll(cx) {
371 Poll::Ready(result) => break result,
372 Poll::Pending => {
373 ExecutionState::with(|state| state.current_mut().sleep_unless_woken());
374 thread::switch();
375 }
376 }
377 }
378}
379
380pub async fn yield_now() {
384 struct YieldNow {
386 yielded: bool,
387 }
388
389 impl Future for YieldNow {
390 type Output = ();
391
392 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
393 if self.yielded {
394 return Poll::Ready(());
395 }
396
397 self.yielded = true;
398 cx.waker().wake_by_ref();
399 ExecutionState::request_yield();
400 Poll::Pending
401 }
402 }
403
404 YieldNow { yielded: false }.await
405}