use std::cell::RefCell;
use std::panic::AssertUnwindSafe;
use std::pin::Pin;
use std::task::{Context, Poll};
use js_sys::Function;
use wasm_bindgen::prelude::*;
use crate::util::{ThreadProc, ValueSender, WorkerPanic, raw_ptr_type};
thread_local! {
static RUNTIME: RefCell<Option<raw_ptr_type!(ValueSender)>> = const { RefCell::new(None) };
static IS_WORKER: RefCell<bool> = const { RefCell::new(false) };
}
pub fn thread_main(
proc: Box<ThreadProc>,
maybe_moves_sender: raw_ptr_type!(ValueSender),
) -> JsValue {
let fut = if cfg!(panic = "unwind") {
match std::panic::catch_unwind(AssertUnwindSafe(proc)) {
Err(e) => {
let sender = unsafe { ValueSender::from_raw(maybe_moves_sender) };
let _ = sender.send(Err(WorkerPanic { payload: Some(e) }));
return JsValue::undefined();
}
Ok(x) => x,
}
} else {
proc()
};
RUNTIME.with_borrow_mut(|x| *x = Some(maybe_moves_sender));
IS_WORKER.with_borrow_mut(|x| *x = true);
let promise = js_sys::futures::future_to_promise(AssertUnwindSafe(async move {
let wrapped_fut = LocalTryOrAbort {
try_or_abort_fn: create_try_or_abort_fn(),
f: fut,
};
let result = wrapped_fut.await;
if let Some(sender) = RUNTIME.with_borrow_mut(|x| x.take()) {
let sender = unsafe { ValueSender::from_raw(sender) };
let _ = sender.send(Ok(result));
}
Ok(JsValue::undefined())
}));
promise.into()
}
pub fn spawn_local<F>(future: F)
where
F: Future<Output = ()> + 'static,
{
if IS_WORKER.with_borrow(|x| *x) {
js_sys::futures::spawn_local(LocalTryOrAbort {
try_or_abort_fn: create_try_or_abort_fn(),
f: future,
});
} else {
js_sys::futures::spawn_local(future);
}
}
struct LocalTryOrAbort<F: ?Sized> {
try_or_abort_fn: Function,
f: F,
}
fn create_try_or_abort_fn() -> Function {
Function::new_with_args(
"x",
"try{x()}catch{try{globalThis.__pistonite_wbgspawn_worker_terminate(true)}catch(e){console.error(e)}}",
)
}
impl<F: ?Sized> Future for LocalTryOrAbort<F>
where
F: Future,
{
type Output = F::Output;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let sender = RUNTIME.with_borrow_mut(|x| x.take());
let Some(sender) = sender else {
return Poll::Pending;
};
let try_or_abort_fn = self.try_or_abort_fn.clone();
let mut poll_f_within_try_catch = Some(|| {
let f = unsafe { self.map_unchecked_mut(|s| &mut s.f) };
if cfg!(panic = "unwind") {
match std::panic::catch_unwind(AssertUnwindSafe(|| f.poll(cx))) {
Ok(x) => Ok(x),
Err(e) => Err(WorkerPanic { payload: Some(e) }),
}
} else {
Ok(f.poll(cx))
}
});
let output = RefCell::new(None);
let mut poll_closure = AssertUnwindSafe(|| {
let result = poll_f_within_try_catch.take().unwrap()();
*output.borrow_mut() = Some(result);
});
let poll_closure_obj = Closure::borrow_mut(&mut poll_closure);
let _ = try_or_abort_fn.call1(&JsValue::undefined(), poll_closure_obj.as_js_value());
let result = match output.take() {
Some(x) => x,
None => {
return Poll::Pending;
}
};
let poll_output = match result {
Err(panic) => {
let send = unsafe { ValueSender::from_raw(sender) };
let _ = send.send(Err(panic));
let abort_fn = Function::new_no_args(
"try{globalThis.__pistonite_wbgspawn_worker_terminate(false)}catch(e){console.error(e)}",
);
let _ = abort_fn.call0(&JsValue::undefined());
return Poll::Pending;
}
Ok(x) => x,
};
RUNTIME.with_borrow_mut(|x| *x = Some(sender));
poll_output
}
}