use shuttle_engine::backtrace_enabled;
use shuttle_engine::runtime::execution::ExecutionState;
use shuttle_engine::runtime::task::TaskId;
use shuttle_engine::runtime::thread;
use std::error::Error;
use std::fmt::{Display, Formatter};
use std::future::Future;
use std::panic::Location;
use std::pin::Pin;
use std::result::Result;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll, Waker};
pub use shuttle_engine::future::batch_semaphore;
fn spawn_inner<F>(fut: F, caller: &'static Location<'static>) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
{
let stack_size = ExecutionState::with(|s| s.config.stack_size);
let inner = Arc::new(std::sync::Mutex::new(JoinHandleInner::default()));
let aborted = Arc::new(AtomicBool::new(false));
let task_id = ExecutionState::spawn_future(
Wrapper::new(fut, inner.clone(), aborted.clone()),
stack_size,
None,
caller,
);
JoinHandle {
task_id,
inner,
aborted,
}
}
#[track_caller]
pub fn spawn<F>(fut: F) -> JoinHandle<F::Output>
where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
spawn_inner(fut, Location::caller())
}
#[track_caller]
pub fn spawn_local<F>(fut: F) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
{
spawn_inner(fut, Location::caller())
}
#[derive(Debug, Clone)]
pub struct AbortHandle {
task_id: TaskId,
aborted: Arc<AtomicBool>,
}
impl AbortHandle {
pub fn abort(&self) {
thread::switch();
if self.aborted.swap(true, Ordering::Relaxed) {
return;
}
let res = ExecutionState::try_with(|state| {
if !state.is_finished() {
state.get_mut(self.task_id).abort();
}
});
if let Err(e) = res {
tracing::error!("`AbortHandle::abort` failed with error: {e:?}");
}
}
pub fn is_finished(&self) -> bool {
ExecutionState::with(|state| {
let task = state.get(self.task_id);
task.finished()
})
}
}
unsafe impl Send for AbortHandle {}
unsafe impl Sync for AbortHandle {}
#[derive(Debug)]
pub struct JoinHandle<T> {
task_id: TaskId,
inner: Arc<std::sync::Mutex<JoinHandleInner<T>>>,
aborted: Arc<AtomicBool>,
}
#[derive(Debug)]
struct JoinHandleInner<T> {
result: Option<Result<T, JoinError>>,
waker: Option<Waker>,
}
impl<T> Default for JoinHandleInner<T> {
fn default() -> Self {
JoinHandleInner {
result: None,
waker: None,
}
}
}
impl<T> JoinHandle<T> {
pub fn abort(&self) {
thread::switch();
if self.aborted.swap(true, Ordering::Relaxed) {
return;
}
let res = ExecutionState::try_with(|state| {
if !state.is_finished() {
state.get_mut(self.task_id).abort();
}
});
if let Err(e) = res {
tracing::error!("`JoinHandle::abort` failed with error: {e:?}");
}
}
pub fn is_finished(&self) -> bool {
ExecutionState::with(|state| {
let task = state.get(self.task_id);
task.finished()
})
}
pub fn abort_handle(&self) -> AbortHandle {
AbortHandle {
task_id: self.task_id,
aborted: self.aborted.clone(),
}
}
}
#[derive(Debug)]
pub enum JoinError {
Cancelled,
}
impl Display for JoinError {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
match self {
JoinError::Cancelled => write!(f, "task was cancelled"),
}
}
}
impl Error for JoinError {}
impl<T> Drop for JoinHandle<T> {
fn drop(&mut self) {
let res = ExecutionState::try_with(|state| {
if !state.is_finished() {
state.detach(self.task_id);
}
});
if let Err(e) = res {
tracing::error!("`JoinHandle::drop` failed with error: {e:?}");
}
}
}
impl<T> Future for JoinHandle<T> {
type Output = Result<T, JoinError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut lock = self.inner.lock().unwrap();
if let Some(result) = lock.result.take() {
Poll::Ready(result)
} else {
lock.waker = Some(cx.waker().clone());
ExecutionState::with(|state| {
state.current_mut().backtrace = if backtrace_enabled() {
Some(std::backtrace::Backtrace::force_capture())
} else {
None
}
});
Poll::Pending
}
}
}
struct Wrapper<F: Future> {
future: Option<Pin<Box<F>>>,
inner: Option<Arc<std::sync::Mutex<JoinHandleInner<F::Output>>>>,
aborted: Arc<AtomicBool>,
}
impl<F> Wrapper<F>
where
F: Future + 'static,
F::Output: 'static,
{
fn new(future: F, inner: Arc<std::sync::Mutex<JoinHandleInner<F::Output>>>, aborted: Arc<AtomicBool>) -> Self {
Self {
future: Some(Box::pin(future)),
inner: Some(inner),
aborted,
}
}
}
impl<F> Wrapper<F>
where
F: Future + 'static,
F::Output: 'static,
{
fn finish(&mut self, result: Result<F::Output, JoinError>) {
ExecutionState::drop_task_locals();
let inner = self.inner.take().expect("a task's result is published once");
let mut lock = inner.lock().unwrap();
lock.result = Some(result);
if let Some(waker) = lock.waker.take() {
waker.wake();
}
}
}
impl<F: Future> Drop for Wrapper<F> {
fn drop(&mut self) {
if let Some(inner) = self.inner.take() {
self.future.take();
if !ExecutionState::should_stop() {
let mut lock = inner.lock().unwrap();
lock.result = Some(Err(JoinError::Cancelled));
if let Some(waker) = lock.waker.take() {
waker.wake();
}
}
}
}
}
impl<F> Future for Wrapper<F>
where
F: Future + 'static,
F::Output: 'static,
{
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
if this.aborted.load(Ordering::Relaxed) {
if ExecutionState::try_with(|state| state.is_finished()).unwrap_or(true) {
return Poll::Ready(());
}
this.future.take();
this.finish(Err(JoinError::Cancelled));
return Poll::Ready(());
}
match this.future.as_mut().unwrap().as_mut().poll(cx) {
Poll::Ready(result) => {
if ExecutionState::try_with(|state| state.is_finished()).unwrap_or(true) {
return Poll::Ready(());
}
this.finish(Ok(result));
Poll::Ready(())
}
Poll::Pending => Poll::Pending,
}
}
}
pub fn block_on<F: Future>(future: F) -> F::Output {
let mut future = Box::pin(future);
let waker = ExecutionState::with(|state| state.current_mut().waker());
let cx = &mut Context::from_waker(&waker);
loop {
match future.as_mut().poll(cx) {
Poll::Ready(result) => break result,
Poll::Pending => {
ExecutionState::with(|state| state.current_mut().sleep_unless_woken());
thread::switch();
}
}
}
}
pub async fn yield_now() {
struct YieldNow {
yielded: bool,
}
impl Future for YieldNow {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
if self.yielded {
return Poll::Ready(());
}
self.yielded = true;
cx.waker().wake_by_ref();
ExecutionState::request_yield();
Poll::Pending
}
}
YieldNow { yielded: false }.await
}