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> {
let idx = self.free_mask.unset_first_set()
.map_err(|_| TableError::NoFreeSlots)?;
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> {
let ref_counter = self.ref_counters.get(idx)
.ok_or(TableError::OutOfBounds)?;
ref_counter.try_update(
Ordering::SeqCst, Ordering::SeqCst,
|word| (word == 0).then_some(REF_COUNTER_LOCKED)
).map_err(|_| TableError::EntryInUse)?;
let res = unsafe {
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)
}
};
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) }
}
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) };
ref_counter.try_update(
Ordering::SeqCst, Ordering::SeqCst,
|word| (word != REF_COUNTER_LOCKED).then_some(word + 1)
).ok()?;
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)
}