Skip to main content

sprite_core/
slab.rs

1use std::cell::UnsafeCell;
2use std::marker::PhantomData;
3
4pub struct Slab<const N: usize = 4096> {
5    active: UnsafeCell<[u8; N]>,
6    backup: UnsafeCell<[u8; N]>,
7    head: UnsafeCell<usize>,
8}
9
10impl<const N: usize> Slab<N> {
11    pub fn new() -> Self {
12        Self {
13            active: UnsafeCell::new([0u8; N]),
14            backup: UnsafeCell::new([0u8; N]),
15            head: UnsafeCell::new(0),
16        }
17    }
18
19    pub unsafe fn alloc<T>(&self, val: T) -> SlabRef<T, N> {
20        let size = std::mem::size_of::<T>();
21        let align = std::mem::align_of::<T>();
22        let head = &mut *self.head.get();
23        let aligned = (*head + align - 1) & !(align - 1);
24        assert!(aligned + size <= N, "slab overflow");
25        let ptr = (*self.active.get()).as_mut_ptr().add(aligned);
26        std::ptr::write(ptr as *mut T, val);
27        *head = aligned + size;
28        SlabRef { offset: aligned, _phantom: PhantomData }
29    }
30
31    pub unsafe fn reset(&self) {
32        *self.head.get() = 0;
33    }
34
35    pub unsafe fn checkpoint(&self) {
36        let head = *self.head.get();
37        std::ptr::copy_nonoverlapping(
38            (*self.active.get()).as_ptr(),
39            (*self.backup.get()).as_mut_ptr(),
40            head,
41        );
42    }
43
44    pub unsafe fn recover(&self) {
45        let head = *self.head.get();
46        std::ptr::copy_nonoverlapping(
47            (*self.backup.get()).as_ptr(),
48            (*self.active.get()).as_mut_ptr(),
49            head,
50        );
51    }
52}
53
54pub struct SlabRef<T, const N: usize = 4096> {
55    offset: usize,
56    _phantom: PhantomData<T>,
57}
58
59impl<T: Copy, const N: usize> SlabRef<T, N> {
60    pub unsafe fn get(&self, slab: &Slab<N>) -> T {
61        let ptr = (*slab.active.get()).as_ptr().add(self.offset) as *const T;
62        std::ptr::read(ptr)
63    }
64    pub unsafe fn set(&self, slab: &Slab<N>, val: T) {
65        let ptr = (*slab.active.get()).as_ptr().add(self.offset) as *mut T;
66        std::ptr::write(ptr, val);
67    }
68}
69
70#[cfg(test)]
71mod tests {
72    use super::*;
73    #[test]
74    fn slab_alloc_and_reset() {
75        let slab = Slab::<256>::new();
76        unsafe {
77            let r = slab.alloc(42i64);
78            assert_eq!(r.get(&slab), 42);
79            r.set(&slab, 100);
80            assert_eq!(r.get(&slab), 100);
81            slab.reset();
82        }
83    }
84    #[test]
85    fn slab_checkpoint_recover() {
86        let slab = Slab::<256>::new();
87        unsafe {
88            let r = slab.alloc(42i64);
89            slab.checkpoint();
90            r.set(&slab, 999);
91            assert_eq!(r.get(&slab), 999);
92            slab.recover();
93            assert_eq!(r.get(&slab), 42);
94        }
95    }
96}