use ahash::RandomState;
use bytemuck::{Pod, Zeroable};
use rand::{rng, Rng};
use rostl_primitives::{
cmov_body, cxchg_body, impl_cmov_for_generic_pod,
traits::{Cmov, _Cmovbase},
};
use seq_macro::seq;
use crate::{array::MultiWayArray, queue::ShortQueue};
const INSERTION_QUEUE_MAX_SIZE: usize = 10;
const DEAMORTIZED_INSERTIONS: usize = 2;
const BUCKET_SIZE: usize = 4;
use std::hash::Hash;
pub trait OHash: Cmov + Pod + Hash + PartialEq {}
impl<K> OHash for K where K: Cmov + Pod + Hash + PartialEq {}
#[repr(align(8))]
#[repr(C)]
#[derive(Debug, Default, Clone, Copy, Zeroable)]
pub struct InlineElement<K, V>
where
K: OHash,
V: Cmov + Pod,
{
key: K,
value: V,
}
unsafe impl<K: OHash, V: Cmov + Pod> Pod for InlineElement<K, V> {}
impl_cmov_for_generic_pod!(InlineElement<K,V>; where K: OHash, V: Cmov + Pod);
#[derive(Debug, Default, Clone, Copy, Zeroable)]
#[repr(C)]
struct Bucket<K, V>
where
K: OHash,
V: Cmov + Pod,
{
is_valid: [bool; BUCKET_SIZE],
elements: [InlineElement<K, V>; BUCKET_SIZE],
}
unsafe impl<K: OHash, V: Cmov + Pod> Pod for Bucket<K, V> {}
impl_cmov_for_generic_pod!(Bucket<K,V>; where K: OHash, V: Cmov + Pod);
impl<K, V> Bucket<K, V>
where
K: OHash,
V: Cmov + Pod,
{
const fn is_empty(&self, i: usize) -> bool {
!self.is_valid[i]
}
fn update_if_exists(&mut self, real: bool, element: InlineElement<K, V>) -> bool {
let mut updated = false;
for i in 0..BUCKET_SIZE {
let choice = real & !self.is_empty(i) & (self.elements[i].key == element.key);
self.elements[i].value.cmov(&element.value, choice);
updated.cmov(&true, choice);
}
updated
}
fn read_if_exists(&self, key: K, ret: &mut V) -> bool {
let mut found = false;
for i in 0..BUCKET_SIZE {
let choice = !self.is_empty(i) & (self.elements[i].key == key);
ret.cmov(&self.elements[i].value, choice);
found.cmov(&true, choice);
}
found
}
fn insert_if_available(&mut self, real: bool, element: InlineElement<K, V>) -> bool {
let mut inserted = !real;
for i in 0..BUCKET_SIZE {
let choice = !inserted & self.is_empty(i);
self.is_valid[i].cmov(&true, choice);
self.elements[i].cmov(&element, choice);
inserted.cmov(&true, choice);
}
inserted
}
}
#[derive(Debug)]
pub struct UnsortedMap<K, V>
where
K: OHash + Default + std::fmt::Debug,
V: Cmov + Pod + Default + std::fmt::Debug,
{
size: usize,
_capacity: usize,
table_size: usize,
table: MultiWayArray<Bucket<K, V>, 2>,
hash_builders: [RandomState; 2],
insertion_queue: ShortQueue<InlineElement<K, V>, INSERTION_QUEUE_MAX_SIZE>,
}
impl<K, V> UnsortedMap<K, V>
where
K: OHash + Default + std::fmt::Debug,
V: Cmov + Pod + Default + std::fmt::Debug,
{
pub fn new(capacity: usize) -> Self {
debug_assert!(capacity > 0);
let table_size = (capacity * 10).div_ceil(9 * BUCKET_SIZE).max(2);
Self {
size: 0,
_capacity: capacity,
table_size,
table: MultiWayArray::new(table_size),
hash_builders: [RandomState::new(), RandomState::new()],
insertion_queue: ShortQueue::new(),
}
}
#[inline(always)]
fn hash_key<const TABLE: usize>(&self, key: &K) -> usize {
(self.hash_builders[TABLE].hash_one(key) % self.table_size as u64) as usize
}
pub fn get(&mut self, key: K, ret: &mut V) -> bool {
let mut found = false;
let mut tmp: Bucket<K, V> = Default::default();
seq!(INDEX in 0..2 {
let hash = self.hash_key::<INDEX>(&key);
self.table.read(INDEX, hash, &mut tmp);
let found_local = tmp.read_if_exists(key, ret);
found.cmov(&true, found_local);
});
for i in 0..self.insertion_queue.size {
let element = self.insertion_queue.elements.data[i];
let found_local = !element.is_empty() & (element.value.key == key);
ret.cmov(&element.value.value, found_local);
found.cmov(&true, found_local);
}
found
}
fn try_insert_entry(&mut self, real: bool, element: &mut InlineElement<K, V>) -> bool {
let mut done = !real;
seq!(INDEX_REV in 0..2 {{
#[allow(clippy::identity_op, clippy::eq_op)] const INDEX: usize = 1 - INDEX_REV;
let hash = self.hash_key::<INDEX>(&element.key);
self.table.update(INDEX, hash, |bucket| {
let choice = !done;
let inserted = bucket.insert_if_available(choice, *element);
done.cmov(&true, inserted);
let randidx = rng().random_range(0..BUCKET_SIZE);
bucket.elements[randidx].cxchg(element, !done);
});
}});
done
}
pub fn deamortize_insertion_queue(&mut self) {
for _ in 0..DEAMORTIZED_INSERTIONS {
let mut element = InlineElement::default();
let real = self.insertion_queue.size > 0;
self.insertion_queue.maybe_pop(real, &mut element);
let has_pending_element = !self.try_insert_entry(real, &mut element);
self.insertion_queue.maybe_push(has_pending_element, element);
}
}
pub fn insert(&mut self, key: K, value: V) {
assert!(self.insertion_queue.size < INSERTION_QUEUE_MAX_SIZE);
self.insertion_queue.maybe_push(true, InlineElement { key, value });
self.deamortize_insertion_queue();
self.size.cmov(&(self.size + 1), true);
}
pub fn insert_cond(&mut self, key: K, value: V, real: bool) {
assert!(self.insertion_queue.size < INSERTION_QUEUE_MAX_SIZE);
self.insertion_queue.maybe_push(real, InlineElement { key, value });
self.deamortize_insertion_queue();
self.size.cmov(&(self.size + 1), real);
}
pub fn write(&mut self, key: K, value: V) {
let mut updated = false;
seq!(INDEX in 0..2 {
let hash = self.hash_key::<INDEX>(&key);
self.table.update(INDEX, hash, |bucket| {
let choice = !updated;
let updated_local = bucket.update_if_exists(choice, InlineElement { key, value });
updated.cmov(&true, updated_local);
});
});
for i in 0..self.insertion_queue.size {
let element = &mut self.insertion_queue.elements.data[i];
let choice = !updated & !element.is_empty() & (element.value.key == key);
element.value.value.cmov(&value, choice);
updated.cmov(&true, choice);
}
assert!(updated);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_unsorted_map() {
let mut map: UnsortedMap<u32, u32> = UnsortedMap::new(2);
assert_eq!(map.size, 0);
let mut value = 0;
assert!(!map.get(1, &mut value));
map.insert(1, 2);
assert_eq!(map.size, 1);
assert!(map.get(1, &mut value));
assert_eq!(value, 2);
map.write(1, 3);
assert!(map.get(1, &mut value));
assert_eq!(value, 3);
}
#[test]
fn test_map_sendness() {
fn assert_send<T: Send>() {}
fn assert_sync<T: Sync>() {}
assert_send::<UnsortedMap<u32, u32>>();
assert_sync::<UnsortedMap<u32, u32>>();
}
#[test]
fn test_full_map() {
const SZ: usize = 1024;
let mut map: UnsortedMap<u32, u32> = UnsortedMap::new(SZ);
assert_eq!(map.size, 0);
for i in 0..SZ as u32 {
map.insert(i, i * 2);
let mut value = 0;
assert!(map.get(i, &mut value));
assert_eq!(value, i * 2);
assert_eq!(map.size, (i + 1) as usize);
map.write(i, i * 3);
assert!(map.get(i, &mut value));
assert_eq!(value, i * 3);
assert_eq!(map.size, (i + 1) as usize);
}
}
#[test]
fn test_insert_cond() {
let mut map: UnsortedMap<u32, u32> = UnsortedMap::new(8);
assert_eq!(map.size, 0);
map.insert_cond(10, 100, false);
assert_eq!(map.size, 0);
let mut value = 0;
assert!(!map.get(10, &mut value));
map.insert_cond(10, 200, true);
assert_eq!(map.size, 1);
assert!(map.get(10, &mut value));
assert_eq!(value, 200);
map.insert_cond(10, 300, false);
assert_eq!(map.size, 1);
assert!(map.get(10, &mut value));
assert_eq!(value, 200);
}
fn test_map_subtypes<
K: OHash + Default + std::fmt::Debug,
V: Cmov + Pod + Default + std::fmt::Debug,
>() {
const SZ: usize = 1024;
let mut map: UnsortedMap<K, V> = UnsortedMap::new(SZ);
assert_eq!(map.size, 0);
let mut value = V::default();
assert!(!map.get(K::default(), &mut value));
map.insert(K::default(), V::default());
assert_eq!(map.size, 1);
assert!(map.get(K::default(), &mut value));
}
#[test]
fn test_map_multiple_types() {
test_map_subtypes::<u32, u32>();
test_map_subtypes::<u64, u64>();
test_map_subtypes::<u128, u128>();
test_map_subtypes::<i32, i32>();
test_map_subtypes::<i64, i64>();
test_map_subtypes::<i128, i128>();
}
}