#![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;
use core::sync::atomic::{AtomicPtr, AtomicU8, Ordering};
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,
armed: *mut u8,
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>,
armed: AtomicU8,
jmpbuf: JmpBufStorage,
prev: *mut Mark,
next: AtomicPtr<Mark>,
cause: MaybeUninit<i32>,
}
struct CallPayload<F, R> {
func: ManuallyDrop<F>,
result: MaybeUninit<R>,
}
static REGISTRY_HEAD: AtomicPtr<Mark> = AtomicPtr::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,
armed: AtomicU8::new(0),
jmpbuf: JmpBufStorage::new(),
prev: ptr::null_mut(),
next: AtomicPtr::new(ptr::null_mut()),
cause: MaybeUninit::uninit(),
};
let mark_ptr: *mut Mark = &mut mark;
let jb = JmpBufStorage::raw(&raw const (*mark_ptr).jmpbuf);
let armed = (&raw mut (*mark_ptr).armed).cast::<u8>();
critical_section::with(|_cs| registry_push(mark_ptr));
let outcome = setback_call(
jb,
armed,
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();
}
registry_unlink_nested(tid, mark);
(*mark).cause = MaybeUninit::new(cause);
JmpBufStorage::raw(&raw const (*mark).jmpbuf)
});
if jb.is_null() {
return Err(RecoveryFailure);
}
setback_longjmp(jb)
}
pub fn can_recover(tid: ThreadId, cause: i32) -> bool {
critical_section::with(|_cs| unsafe { !registry_find(tid, cause).is_null() })
}
unsafe fn registry_push(node: *mut Mark) {
let head = REGISTRY_HEAD.load(Ordering::Relaxed);
(*node).next.store(head, Ordering::Relaxed);
(*node).prev = ptr::null_mut();
if !head.is_null() {
(*head).prev = node;
}
REGISTRY_HEAD.store(node, Ordering::Release);
}
unsafe fn registry_unlink(node: *mut Mark) {
let prev = (*node).prev;
let next = (*node).next.load(Ordering::Relaxed);
if prev.is_null() {
REGISTRY_HEAD.store(next, Ordering::Release);
} else {
(*prev).next.store(next, Ordering::Relaxed);
}
if !next.is_null() {
(*next).prev = prev;
}
}
unsafe fn registry_unlink_nested(tid: ThreadId, target: *mut Mark) {
let mut p = REGISTRY_HEAD.load(Ordering::Acquire);
while !p.is_null() && p != target {
let next = (*p).next.load(Ordering::Relaxed);
if (*p).tid == tid {
registry_unlink(p);
}
p = next;
}
}
unsafe fn registry_find(tid: ThreadId, cause: i32) -> *mut Mark {
let mut p = REGISTRY_HEAD.load(Ordering::Acquire);
while !p.is_null() {
if (*p).armed.load(Ordering::Acquire) != 0
&& (*p).tid == tid
&& (*p).accepts.is_none_or(|c| c == cause)
{
return p;
}
p = (*p).next.load(Ordering::Relaxed);
}
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)")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_linked_mark_is_ignored_until_it_is_armed() {
const TID: ThreadId = 1234;
const CAUSE: i32 = 9;
let mut mark = Mark {
tid: TID,
accepts: None,
armed: AtomicU8::new(0),
jmpbuf: JmpBufStorage::new(),
prev: ptr::null_mut(),
next: AtomicPtr::new(ptr::null_mut()),
cause: MaybeUninit::uninit(),
};
let mark_ptr: *mut Mark = &mut mark;
let found = || critical_section::with(|_cs| !unsafe { registry_find(TID, CAUSE) }.is_null());
unsafe {
critical_section::with(|_cs| registry_push(mark_ptr));
assert!(!found());
(*mark_ptr).armed.store(1, Ordering::Release);
assert!(found());
critical_section::with(|_cs| registry_unlink(mark_ptr));
assert!(!found());
}
}
}