use std::cell::UnsafeCell;
pub struct ContextCell<T> {
ptr: UnsafeCell<*mut T>,
}
impl<T> ContextCell<T> {
#[inline]
pub const fn new() -> Self {
Self {
ptr: UnsafeCell::new(std::ptr::null_mut()),
}
}
#[inline]
pub fn get_ptr(&self) -> *mut T {
unsafe { *self.ptr.get() }
}
#[inline]
pub fn replace(&self, ptr: *mut T) -> *mut T {
unsafe {
let prev = *self.ptr.get();
*self.ptr.get() = ptr;
prev
}
}
#[inline]
pub fn restore(&self, ptr: *mut T) {
unsafe {
*self.ptr.get() = ptr;
}
}
}
#[cfg(test)]
impl<T> ContextCell<T> {
#[inline]
pub(crate) fn get(&self) -> Option<*mut T> {
let ptr = unsafe { *self.ptr.get() };
if ptr.is_null() { None } else { Some(ptr) }
}
#[inline]
pub(crate) fn clear(&self) -> *mut T {
self.replace(std::ptr::null_mut())
}
#[inline]
pub(crate) unsafe fn deref_mut(&self) -> Option<&'static mut T> {
let ptr = unsafe { *self.ptr.get() };
if ptr.is_null() {
None
} else {
Some(unsafe { &mut *ptr })
}
}
}
impl<T> Default for ContextCell<T> {
fn default() -> Self {
Self::new()
}
}
unsafe impl<T> Send for ContextCell<T> {}
unsafe impl<T> Sync for ContextCell<T> {}
pub struct RestoreOnDrop<T: 'static> {
cell: &'static ContextCell<T>,
prev: *mut T,
}
impl<T: 'static> RestoreOnDrop<T> {
pub const unsafe fn new(cell: &'static ContextCell<T>, prev: *mut T) -> Self {
Self { cell, prev }
}
pub fn disarm(mut self) -> *mut T {
let prev = self.prev;
self.prev = std::ptr::null_mut();
std::mem::forget(self);
prev
}
}
impl<T: 'static> Drop for RestoreOnDrop<T> {
fn drop(&mut self) {
self.cell.restore(self.prev);
}
}
unsafe impl<T: 'static> Send for RestoreOnDrop<T> {}
unsafe impl<T: 'static> Sync for RestoreOnDrop<T> {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_is_null() {
let cell = ContextCell::<u32>::new();
assert!(cell.get().is_none());
}
#[test]
fn replace_returns_previous() {
let cell = ContextCell::<u32>::new();
let mut a: u32 = 1;
let mut b: u32 = 2;
let ptr_a = &mut a as *mut u32;
let ptr_b = &mut b as *mut u32;
let prev = cell.replace(ptr_a);
assert!(prev.is_null());
assert_eq!(cell.get(), Some(ptr_a));
let prev = cell.replace(ptr_b);
assert_eq!(prev, ptr_a);
assert_eq!(cell.get(), Some(ptr_b));
}
#[test]
fn clear_returns_previous() {
let cell = ContextCell::<u32>::new();
let mut a: u32 = 42;
let ptr_a = &mut a as *mut u32;
cell.replace(ptr_a);
let prev = cell.clear();
assert_eq!(prev, ptr_a);
assert!(cell.get().is_none());
}
#[test]
fn restore_sets_pointer() {
let cell = ContextCell::<u32>::new();
let mut a: u32 = 10;
let ptr_a = &mut a as *mut u32;
cell.restore(ptr_a);
assert_eq!(cell.get(), Some(ptr_a));
}
#[test]
fn deref_mut_returns_reference() {
let cell = ContextCell::<u32>::new();
let mut a: u32 = 99;
let ptr_a = &mut a as *mut u32;
cell.replace(ptr_a);
let r = unsafe { cell.deref_mut() }.unwrap();
assert_eq!(*r, 99);
*r = 100;
assert_eq!(a, 100);
}
#[test]
fn deref_mut_none_when_null() {
let cell = ContextCell::<u32>::new();
assert!(unsafe { cell.deref_mut() }.is_none());
}
#[test]
fn default_is_null() {
let cell = ContextCell::<u32>::default();
assert!(cell.get().is_none());
}
static TEST_CELL: ContextCell<u32> = const { ContextCell::new() };
#[test]
fn restore_on_drop_restores_previous_value() {
TEST_CELL.replace(std::ptr::null_mut());
let mut a: u32 = 42;
let ptr_a = &mut a as *mut u32;
TEST_CELL.replace(ptr_a);
let mut b: u32 = 99;
let ptr_b = &mut b as *mut u32;
let prev = TEST_CELL.replace(ptr_b);
let guard = unsafe { RestoreOnDrop::new(&TEST_CELL, prev) };
assert_eq!(TEST_CELL.get_ptr(), ptr_b);
drop(guard);
assert_eq!(TEST_CELL.get_ptr(), ptr_a);
}
#[test]
fn restore_on_drop_disarm_skips_restore() {
TEST_CELL.replace(std::ptr::null_mut());
let mut a: u32 = 42;
let ptr_a = &mut a as *mut u32;
TEST_CELL.replace(ptr_a);
let mut b: u32 = 99;
let ptr_b = &mut b as *mut u32;
let prev = TEST_CELL.replace(ptr_b);
let guard = unsafe { RestoreOnDrop::new(&TEST_CELL, prev) };
let _ = guard.disarm();
assert_eq!(TEST_CELL.get_ptr(), ptr_b);
}
#[test]
fn restore_on_drop_restores_null_initial_value() {
TEST_CELL.replace(std::ptr::null_mut());
let mut a: u32 = 42;
let ptr_a = &mut a as *mut u32;
let prev = TEST_CELL.replace(ptr_a);
assert!(prev.is_null());
let guard = unsafe { RestoreOnDrop::new(&TEST_CELL, prev) };
drop(guard);
assert!(TEST_CELL.get_ptr().is_null());
}
#[cfg(loom)]
#[allow(unexpected_cfgs)]
mod loom_tests {
#![allow(unused_unsafe)]
use super::*;
#[test]
fn replace_is_atomic_enough_under_concurrent_access() {
loom::model(|| {
let cell = loom::sync::Arc::new(ContextCell::<u32>::new());
let mut a: u32 = 0;
let mut b: u32 = 0;
let ptr_a = &mut a as *mut u32;
let ptr_b = &mut b as *mut u32;
let cell_a = cell.clone();
let t1 = loom::thread::spawn(move || {
for _ in 0..2 {
let prev = unsafe { (*cell_a).replace(ptr_a) };
if !prev.is_null() {
assert!(
prev == ptr_a || prev == ptr_b,
"torn pointer observed: {prev:p}"
);
}
let now = unsafe { (*cell_a).get_ptr() };
assert!(
now == ptr_a || now == ptr_b,
"torn pointer observed: {now:p}"
);
}
});
let cell_b = cell.clone();
let t2 = loom::thread::spawn(move || {
for _ in 0..2 {
let prev = unsafe { (*cell_b).replace(ptr_b) };
if !prev.is_null() {
assert!(
prev == ptr_a || prev == ptr_b,
"torn pointer observed: {prev:p}"
);
}
let now = unsafe { (*cell_b).get_ptr() };
assert!(
now == ptr_a || now == ptr_b,
"torn pointer observed: {now:p}"
);
}
});
t1.join().unwrap();
t2.join().unwrap();
});
}
#[test]
fn restore_on_drop_restores_under_concurrent_access() {
loom::model(|| {
static CELL: ContextCell<u32> = ContextCell::new();
let mut a: u32 = 0;
let mut b: u32 = 0;
let ptr_a = &mut a as *mut u32;
let ptr_b = &mut b as *mut u32;
CELL.replace(ptr_a);
let prev = unsafe { CELL.replace(ptr_b) };
assert!(prev == ptr_a, "expected prev == ptr_a, got {prev:p}");
let guard = unsafe { RestoreOnDrop::new(&CELL, prev) };
let now = unsafe { CELL.get_ptr() };
assert!(
now == ptr_a || now == ptr_b,
"torn pointer observed: {now:p}"
);
drop(guard);
let after = unsafe { CELL.get_ptr() };
assert!(
after == ptr_a || after.is_null() || after == ptr_b,
"torn pointer observed: {after:p}"
);
});
}
}
}