meowalloc 0.1.1

Toy allocator written in pure rust, with concurrency in mind
use core::ops::Deref;
use core::{mem::MaybeUninit, cell::UnsafeCell};
use core::sync::atomic::{Ordering, AtomicU16};

use crate::atomic_bitset::{AtomicBitSet, WORD_BITS};

const REF_COUNTER_LOCKED: u16 = u16::MAX;

#[derive(Debug, PartialEq)]
pub enum TableError {
    NoFreeSlots,
    OutOfBounds,
    EntryMissing,
    EntryInUse,
}

pub struct AtomicTable<T, const LEN: usize, const BITMASK_WORDS: usize> {
    arr: UnsafeCell<[MaybeUninit<T>; LEN]>,
    free_mask: AtomicBitSet<BITMASK_WORDS>,
    init_mask: AtomicBitSet<BITMASK_WORDS>,
    ref_counters: [AtomicU16; LEN]
}

#[derive(Clone)]
pub struct Iter<'a, T, const LEN: usize, const BITMASK_WORDS: usize> {
    table: &'a AtomicTable<T, LEN, BITMASK_WORDS>,
    i: usize
}

#[derive(Debug)]
pub struct EntryGuard<'a, T, const BITMASK_WORDS: usize> {
    item: &'a T,
    ref_counter: &'a AtomicU16
}


impl<'a, T, const BITMASK_WORDS: usize> Drop for EntryGuard<'a, T, BITMASK_WORDS> {
    fn drop(&mut self) {
        self.ref_counter.fetch_sub(1, Ordering::SeqCst);
    }
}

impl<'a, T, const BITMASK_WORDS: usize> Clone for EntryGuard<'a, T, BITMASK_WORDS> {
    fn clone(&self) -> Self {
        self.ref_counter.fetch_add(1, Ordering::SeqCst);
        Self { item: self.item, ref_counter: self.ref_counter }
    }
}

impl<'a, T, const BITMASK_WORDS: usize> Deref for EntryGuard<'a, T, BITMASK_WORDS> {
    type Target = T;
    fn deref(&self) -> &Self::Target { self.item }
}

impl<'a, T, const LEN: usize, const BITMASK_WORDS: usize> Iterator for Iter<'a, T, LEN, BITMASK_WORDS> {
    type Item = (usize, EntryGuard<'a, T, BITMASK_WORDS>);
    fn next(&mut self) -> Option<Self::Item> {
        while self.i < self.table.len() {
            let opt = unsafe { self.table.get_unchecked(self.i) }
                .map(|entry| (self.i, entry));
            self.i += 1;
            if opt.is_some() {
                return opt
            }
        }
        None
    }
}

impl<T, const LEN: usize, const BITMASK_WORDS: usize> AtomicTable<T, LEN, BITMASK_WORDS> {
    pub const fn new() -> Self {
        const { assert!( LEN <= BITMASK_WORDS * WORD_BITS) }
        Self {
            arr: UnsafeCell::new([ const { MaybeUninit::uninit() }; LEN ]),
            free_mask: AtomicBitSet::ones(),
            init_mask: AtomicBitSet::zeros(),
            ref_counters: [ const { AtomicU16::new(0) }; LEN ]
        }
    }

    pub fn add(&self, item: T) -> Result<usize, TableError> {
        // Immediately mark occupied
        let idx = self.free_mask.unset_first_set()
            .map_err(|_| TableError::NoFreeSlots)?;
        // Initialize and mark
        unsafe {
            self.arr.get().cast::<T>().add(idx).write(item);
            self.init_mask.set_unchecked(idx);
        }
        Ok(idx)
    }

    pub fn try_remove(&self, idx: usize) -> Result<T, TableError> {
        // Bounds check once already
        let ref_counter = self.ref_counters.get(idx)
            .ok_or(TableError::OutOfBounds)?;
       
        // Return error if:
        // - Someone else holds a reference to the entry
        // - Someone is already trying to remove the entry
        // Otherwise lock the ref_counter and proceed
        ref_counter.try_update(
            Ordering::SeqCst, Ordering::SeqCst,
            |word| (word == 0).then_some(REF_COUNTER_LOCKED)
        ).map_err(|_| TableError::EntryInUse)?;
        
        let res = unsafe {
            // Move entry if it exists
            if self.init_mask.unset_if_one_unchecked(idx) {
                let item = self.arr.get().cast::<T>().add(idx).read();
                self.free_mask.set_unchecked(idx);
                Ok(item)
            
            } else {
                Err(TableError::EntryMissing)
            }
        };

        // Unlock ref_counter
        ref_counter.swap(0, Ordering::SeqCst);
        res
    }

    pub fn get(&self, idx: usize) -> Option<EntryGuard<'_, T, BITMASK_WORDS>> {
        if idx >= self.len() { return None }
        unsafe { self.get_unchecked(idx) }
    }

    /// # Safety:
    /// The caller must ensure idx is within bounds 
    pub unsafe fn get_unchecked(&self, idx: usize) -> Option<EntryGuard<'_, T, BITMASK_WORDS>> {
        if unsafe { self.init_mask.is_zero_unchecked(idx) } {
            return None
        }
        let ref_counter = unsafe { self.ref_counters.get_unchecked(idx) };
       
        // Either someone else has gotten the "lock" and the entry will be removed
        // or the removing entity will see that the reference counter incremented
        // and won't be able to remove
        ref_counter.try_update(
            Ordering::SeqCst, Ordering::SeqCst,
            |word| (word != REF_COUNTER_LOCKED).then_some(word + 1)
        ).ok()?;

        // Item is in safe paws now =3
        let item = unsafe { self.arr.get().cast::<T>().add(idx).as_ref_unchecked() };
        
        Some(EntryGuard { item, ref_counter })
        
    }

    pub fn iter(&self) -> Iter<'_, T, LEN, BITMASK_WORDS> {
        Iter { table: self, i: 0 }
    }
                                                                                     
    fn len(&self) -> usize { LEN }
}

impl<T, const LEN: usize, const BITMASK_WORDS: usize> Drop for AtomicTable<T, LEN, BITMASK_WORDS> {
    fn drop(&mut self) {
        unsafe {
            for idx in 0..self.len() {
                if self.init_mask.is_one(idx) {
                    self.arr.get().cast::<T>().add(idx).drop_in_place();
                }
            }
        }
    }
}

unsafe impl<T, const LEN: usize, const BITMASK_WORDS: usize> Sync for AtomicTable<T, LEN, BITMASK_WORDS> {}

#[cfg(test)]
extern crate std;
#[cfg(test)]
use std::{dbg, eprintln};

#[cfg(test)]
type TestAtomicTable = AtomicTable<u8, 128, 4>;

#[cfg(test)]
fn table_fixture() -> (TestAtomicTable, [usize; 4]) {
    let table = AtomicTable::new();
    let mut indices = [0; 4];
    indices[0] = table.add(2).unwrap();
    indices[1] = table.add(5).unwrap();
    indices[2] = table.add(1).unwrap();
    indices[3] = table.add(7).unwrap();
    (table, indices)
}

#[test]
fn add_populates_masks() {
    let table = TestAtomicTable::new();
    assert_eq!(table.add(2).unwrap(), 0);
    assert_eq!(table.add(5).unwrap(), 1);
    assert_eq!(table.add(7).unwrap(), 2);
    assert_eq!(
        Into::<[u32; 4]>::into(&table.init_mask),
        [0xE0000000, 0, 0, 0]
    );
    assert_eq!(
        Into::<[u32; 4]>::into(&table.free_mask),
        [0x1FFFFFFF, 0xFFFFFFFF, 0xFFFFFFFF, 0xFFFFFFFF]
    );
}

#[test]
fn add_reuses_indices() {
    let table= TestAtomicTable::new();
    assert_eq!(table.add(2).unwrap(), 0);
    assert_eq!(table.add(5).unwrap(), 1);
    assert_eq!(table.try_remove(0).unwrap(), 2);
    assert_eq!(table.add(7).unwrap(), 0);
}

#[test]
fn get() {
    let (table, indices) = table_fixture();
    assert_eq!(*table.get(indices[0]).unwrap(), 2);
    assert_eq!(*table.get(indices[1]).unwrap(), 5);
    assert_eq!(*table.get(indices[2]).unwrap(), 1);
    assert_eq!(*table.get(indices[3]).unwrap(), 7);
}

#[test]
fn cant_remove_guarded() {
    let (table, indices) = table_fixture();
    let guard = table.get(indices[1]).unwrap(); 
    assert_eq!(table.try_remove(indices[1]), Err(TableError::EntryInUse));
    drop(guard);
    assert_eq!(table.try_remove(indices[1]), Ok(5));
}

#[test]
fn try_remove() {
    let (table, indices) = table_fixture();    
    assert_eq!(table.try_remove(128), Err(TableError::OutOfBounds));
    assert_eq!(table.try_remove(10), Err(TableError::EntryMissing));
    assert_eq!(table.try_remove(indices[0]), Ok(2));
    assert_eq!(table.try_remove(indices[1]), Ok(5));
    assert_eq!(table.try_remove(indices[2]), Ok(1));
    assert_eq!(table.try_remove(indices[3]), Ok(7))
}

#[test]
fn iter() {
    let (table, indices) = table_fixture();
    let mut j = 0;
    
    for (idx, elem) in table.iter() {
        assert_eq!(idx, indices[j]);
        assert_eq!(*elem, *table.get(idx).unwrap());
        j += 1;
    }
}

#[test]
fn drops() {
    let table = TestAtomicTable::new();
    drop(table)
}