#![no_std]
#![allow(unsafe_op_in_unsafe_fn)]
#[cfg(target_family = "wasm")]
compile_error!("`setback` does not support wasm targets");
use core::cell::UnsafeCell;
use core::convert::Infallible;
use core::error::Error;
use core::ffi::c_void;
use core::mem::{ManuallyDrop, MaybeUninit};
use core::panic::UnwindSafe;
use core::ptr;
pub type ThreadId = usize;
pub use core::panic::AssertUnwindSafe;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecoveryError {
pub cause: i32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecoveryFailure;
unsafe extern "C" {
fn setback_jmpbuf_size() -> usize;
fn setback_jmpbuf_align() -> usize;
fn setback_call(
jb: *mut c_void,
tramp: unsafe extern "C" fn(*mut c_void),
data: *mut c_void,
) -> i32;
fn setback_longjmp(jb: *mut c_void) -> !;
}
const SETBACK_OK: i32 = 0;
pub const RECOVERY_GAP_BYTES: usize = 64;
#[repr(C, align(16))]
struct JmpBufStorage {
bytes: UnsafeCell<MaybeUninit<[u8; 512]>>,
}
struct Mark {
tid: ThreadId,
accepts: Option<i32>,
jmpbuf: JmpBufStorage,
prev: *mut Mark,
next: *mut Mark,
cause: MaybeUninit<i32>,
}
struct Registry {
head: UnsafeCell<*mut Mark>,
}
unsafe impl Sync for Registry {}
struct CallPayload<F, R> {
func: ManuallyDrop<F>,
result: MaybeUninit<R>,
}
static REGISTRY: Registry = Registry {
head: UnsafeCell::new(ptr::null_mut()),
};
pub unsafe fn protect<F, R>(tid: ThreadId, f: F) -> Result<R, RecoveryError>
where
F: FnOnce() -> R + UnwindSafe,
{
protect_inner(tid, None, f)
}
pub unsafe fn protect_cause<F, R>(
tid: ThreadId,
cause: i32,
f: F,
) -> Result<R, RecoveryError>
where
F: FnOnce() -> R + UnwindSafe,
{
protect_inner(tid, Some(cause), f)
}
unsafe fn protect_inner<F, R>(
tid: ThreadId,
accepts: Option<i32>,
f: F,
) -> Result<R, RecoveryError>
where
F: FnOnce() -> R + UnwindSafe,
{
let mut payload = CallPayload::<F, R> {
func: ManuallyDrop::new(f),
result: MaybeUninit::uninit(),
};
let mut mark = Mark {
tid,
accepts,
jmpbuf: JmpBufStorage::new(),
prev: ptr::null_mut(),
next: ptr::null_mut(),
cause: MaybeUninit::uninit(),
};
let mark_ptr: *mut Mark = &mut mark;
let jb = JmpBufStorage::raw(&raw const (*mark_ptr).jmpbuf);
critical_section::with(|_cs| registry_push(mark_ptr));
let outcome = setback_call(
jb,
trampoline::<F, R>,
&mut payload as *mut CallPayload<F, R> as *mut c_void,
);
critical_section::with(|_cs| registry_unlink(mark_ptr));
if outcome == SETBACK_OK {
Ok(payload.result.assume_init())
} else {
Err(RecoveryError {
cause: (*mark_ptr).cause.assume_init(),
})
}
}
unsafe extern "C" fn trampoline<F, R>(data: *mut c_void)
where
F: FnOnce() -> R,
{
let payload = unsafe { &mut *(data as *mut CallPayload<F, R>) };
let f = unsafe { ManuallyDrop::take(&mut payload.func) };
payload.result.write(f());
}
pub unsafe fn recover(tid: ThreadId, cause: i32) -> Result<Infallible, RecoveryFailure> {
let jb = critical_section::with(|_cs| {
let mark = registry_find(tid, cause);
if mark.is_null() {
return ptr::null_mut();
}
(*mark).cause = MaybeUninit::new(cause);
JmpBufStorage::raw(&raw const (*mark).jmpbuf)
});
if jb.is_null() {
return Err(RecoveryFailure);
}
setback_longjmp(jb)
}
unsafe fn registry_push(node: *mut Mark) {
let head = *REGISTRY.head.get();
(*node).next = head;
(*node).prev = ptr::null_mut();
if !head.is_null() {
(*head).prev = node;
}
*REGISTRY.head.get() = node;
}
unsafe fn registry_unlink(node: *mut Mark) {
let prev = (*node).prev;
let next = (*node).next;
if prev.is_null() {
*REGISTRY.head.get() = next;
} else {
(*prev).next = next;
}
if !next.is_null() {
(*next).prev = prev;
}
}
unsafe fn registry_find(tid: ThreadId, cause: i32) -> *mut Mark {
let mut p = *REGISTRY.head.get();
while !p.is_null() {
if (*p).tid == tid && (*p).accepts.is_none_or(|c| c == cause) {
return p;
}
p = (*p).next;
}
ptr::null_mut()
}
impl JmpBufStorage {
#[inline]
fn new() -> Self {
let need = unsafe { setback_jmpbuf_size() };
let align = unsafe { setback_jmpbuf_align() };
assert!(need <= 512, "setback: jmp_buf larger than reserved storage");
assert!(
align <= 16,
"setback: jmp_buf alignment exceeds storage alignment"
);
JmpBufStorage {
bytes: UnsafeCell::new(MaybeUninit::uninit()),
}
}
#[inline]
unsafe fn raw(this: *const JmpBufStorage) -> *mut c_void {
UnsafeCell::raw_get(&raw const (*this).bytes) as *mut c_void
}
}
impl Error for RecoveryFailure {}
impl core::fmt::Display for RecoveryFailure {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "setback recovery failure (no active scope)")
}
}