use std::{
ptr,
sync::atomic::{AtomicPtr, Ordering},
};
thread_local!(static TOKEN: u8 = const { 0 });
#[derive(Default)]
pub(super) struct Guard {
owner: AtomicPtr<u8>,
}
impl Guard {
#[inline]
pub(super) fn assert_not_recursive<F>(&self) {
assert!(
!self.is_owned_by_caller(),
"recursive `#[memoize]` initialization of `{}` with the same arguments",
std::any::type_name::<F>()
);
}
#[inline]
pub(super) fn scope<R>(&self, f: impl FnOnce() -> R) -> R {
TOKEN.with(|token| {
debug_assert!(self.owner.load(Ordering::Relaxed).is_null());
self.owner
.store(ptr::from_ref(token).cast_mut(), Ordering::Relaxed);
let _scope = Scope { owner: &self.owner };
f()
})
}
#[inline]
fn is_owned_by_caller(&self) -> bool {
let owner = self.owner.load(Ordering::Relaxed);
!owner.is_null() && TOKEN.with(|token| owner == ptr::from_ref(token).cast_mut())
}
}
struct Scope<'a> {
owner: &'a AtomicPtr<u8>,
}
impl Drop for Scope<'_> {
#[inline]
fn drop(&mut self) {
self.owner.store(ptr::null_mut(), Ordering::Relaxed);
}
}
#[cfg(test)]
mod tests {
use std::sync::Barrier;
use super::*;
#[test]
fn scope_is_released_after_it_ends() {
let guard = Guard::default();
guard.scope(|| {});
guard.assert_not_recursive::<fn()>();
}
#[test]
#[should_panic(expected = "recursive `#[memoize]` initialization")]
fn reentry_from_the_owning_stack_panics() {
let guard = Guard::default();
guard.scope(|| guard.assert_not_recursive::<fn()>());
}
#[test]
fn scope_on_another_thread_is_not_recursive() {
let guard = Guard::default();
let entered = Barrier::new(2);
let release = Barrier::new(2);
std::thread::scope(|scope| {
scope.spawn(|| {
guard.scope(|| {
entered.wait();
release.wait();
});
});
entered.wait();
guard.assert_not_recursive::<fn()>();
release.wait();
});
}
}