#![no_std]
#![allow(clippy::similar_names)] #![allow(clippy::cast_possible_truncation)] #![allow(clippy::cast_ptr_alignment)] #![allow(clippy::branches_sharing_code)]
use core::cmp::{max, min};
use core::hash::Hasher;
use core::mem::size_of;
use core::ops::Not;
use core::{ptr, slice};
use fxhash::FxHasher64;
#[repr(u8)]
#[derive(Copy, Clone, PartialEq, Eq, Debug)]
pub enum BucketStatus {
Empty = 0, Tombstone = 1,
Occupied = 2,
}
#[repr(C)]
#[derive(Copy, Clone, Debug)]
pub struct MapHeader {
pub capacity: u16, pub element_count: u16, pub key_size: u32,
pub value_size: u32,
pub value_offset: u32,
pub bucket_size: u32,
pub logical_limit: u16,
pub key_offset: u8,
pub padding_and_secret_code: u8,
}
pub struct MapInit {
pub key_size: u32,
pub key_alignment: u8,
pub value_size: u32,
pub value_alignment: u8,
pub capacity: u16,
pub logical_limit: u16,
pub total_size: u32,
}
#[derive(Clone, Copy, Debug)]
pub struct BucketLayout {
pub bucket_size: u32,
pub key_offset: u8,
pub value_offset: u32,
}
const MAP_BUCKETS_OFFSET: usize = size_of::<MapHeader>();
const MAX_PROBE_DISTANCE: usize = 32;
#[inline]
fn calculate_hash_bytes(key_bytes: &[u8]) -> u64 {
let mut hasher = FxHasher64::default();
hasher.write(key_bytes);
hasher.finish()
}
#[inline]
fn index_from_hash(hash: u64, capacity: u16) -> usize {
assert!(capacity.is_power_of_two());
((hash >> 48) as usize) & ((capacity as usize) - 1)
}
#[inline]
#[must_use]
pub fn calculate_bucket_layout(
key_size: u32,
key_alignment: u8,
value_size: u32,
value_alignment: u8,
) -> BucketLayout {
let status_size: u32 = 1;
let mut current_offset = status_size;
let key_align = u32::from(key_alignment);
let key_offset = (current_offset + key_align - 1) & !(key_align - 1);
current_offset = key_offset + key_size;
let value_align = u32::from(value_alignment);
let value_offset = (current_offset + value_align - 1) & !(value_align - 1);
current_offset = value_offset + value_size;
let bucket_content_alignment = max(key_align, value_align);
let bucket_size =
(current_offset + bucket_content_alignment - 1) & !(bucket_content_alignment - 1);
BucketLayout {
bucket_size,
key_offset: key_offset as u8,
value_offset,
}
}
#[must_use]
pub const fn total_size(capacity: u16, bucket_size: u32) -> u32 {
(MAP_BUCKETS_OFFSET + capacity as usize * bucket_size as usize) as u32
}
#[must_use]
pub fn layout(
key_size: u32,
key_alignment: u8,
value_size: u32,
value_alignment: u8,
logical_limit: u16,
) -> (BucketLayout, MapInit) {
let capacity = logical_limit.next_power_of_two();
let bucket_layout =
calculate_bucket_layout(key_size, key_alignment, value_size, value_alignment);
(
bucket_layout,
MapInit {
key_size,
key_alignment,
value_size,
value_alignment,
capacity,
logical_limit,
total_size: total_size(capacity, bucket_layout.bucket_size),
},
)
}
pub const SECRET_CODE: u8 = 0x3d;
pub unsafe fn init(map_base: *mut u8, config: &MapInit) {
assert!(
config.capacity.is_power_of_two(),
"Capacity must be a power of two"
);
let map_header = map_base.cast::<MapHeader>();
let layout = calculate_bucket_layout(
config.key_size,
config.key_alignment,
config.value_size,
config.value_alignment,
);
unsafe {
ptr::write(
map_header,
MapHeader {
capacity: config.capacity,
logical_limit: config.logical_limit,
key_size: config.key_size,
value_size: config.value_size,
bucket_size: layout.bucket_size,
key_offset: layout.key_offset,
value_offset: layout.value_offset,
element_count: 0,
padding_and_secret_code: SECRET_CODE,
},
);
}
let buckets_start_ptr = unsafe { map_base.add(MAP_BUCKETS_OFFSET) };
let capacity = usize::from(config.capacity);
let bucket_size = layout.bucket_size as usize;
for i in 0..capacity {
unsafe {
ptr::write(
buckets_start_ptr.add(i * bucket_size),
BucketStatus::Empty as u8,
);
}
}
}
#[inline]
#[must_use]
pub const unsafe fn validate(base_ptr: *const u8) -> bool {
unsafe {
let header = &*base_ptr.cast::<MapHeader>();
if header.padding_and_secret_code != SECRET_CODE {
return false;
}
if header.key_size == 0 {
return false;
}
if header.capacity == 0 {
return false;
}
if !header.capacity.is_power_of_two() {
return false;
}
if header.element_count > header.capacity {
return false;
}
if header.logical_limit > header.capacity {
return false;
}
true
}
}
#[inline]
unsafe fn assert_validate(base_ptr: *const u8) {
unsafe {
assert!(validate(base_ptr), "Invalid map structure");
}
}
#[inline]
unsafe fn matches_key(a: *const u8, b: *const u8, len: usize) -> bool {
unsafe {
if len <= 16 {
if len == 0 {
true
} else {
for i in 0..len {
if *a.add(i) != *b.add(i) {
return false;
}
}
true
}
} else {
slice::from_raw_parts(a, len) == slice::from_raw_parts(b, len)
}
}
}
#[inline]
pub unsafe fn reserve_entry(base_ptr: *mut u8, key_ptr: *const u8) -> *mut u8 {
unsafe {
assert_validate(base_ptr);
let header = &*base_ptr.cast::<MapHeader>();
let capacity = header.capacity as usize;
let key_size = header.key_size as usize;
let bucket_size = header.bucket_size as usize;
let key_offset = header.key_offset as usize;
let value_offset = header.value_offset as usize;
let buckets_ptr = base_ptr.add(MAP_BUCKETS_OFFSET);
let key_slice = slice::from_raw_parts(key_ptr, key_size);
let hash = calculate_hash_bytes(key_slice);
let mut index = index_from_hash(hash, header.capacity);
let mut first_tombstone = None;
let probe_limit = min(capacity, MAX_PROBE_DISTANCE);
for _ in 0..probe_limit {
let bucket_ptr = buckets_ptr.add(index * bucket_size);
let status = *bucket_ptr;
match status {
status if status == BucketStatus::Empty as u8 => {
let insert_index = first_tombstone.unwrap_or(index);
let target_bucket = buckets_ptr.add(insert_index * bucket_size);
*target_bucket = BucketStatus::Occupied as u8;
let target_key_ptr = target_bucket.add(key_offset);
ptr::copy_nonoverlapping(key_ptr, target_key_ptr, key_size);
let header_mut = &mut *base_ptr.cast::<MapHeader>();
header_mut.element_count += 1;
return target_bucket.add(value_offset);
}
status if status == BucketStatus::Occupied as u8 => {
let existing_key_ptr = bucket_ptr.add(key_offset);
assert!(
!matches_key(existing_key_ptr, key_ptr, key_size),
"Key already exists in map"
);
}
status if status == BucketStatus::Tombstone as u8 => {
if first_tombstone.is_none() {
first_tombstone = Some(index);
}
}
_ => unreachable!(),
}
index = (index + 1) & (capacity - 1);
}
if let Some(tombstone_index) = first_tombstone {
let target_bucket = buckets_ptr.add(tombstone_index * bucket_size);
*target_bucket = BucketStatus::Occupied as u8;
let target_key_ptr = target_bucket.add(key_offset);
ptr::copy_nonoverlapping(key_ptr, target_key_ptr, key_size);
let header_mut = &mut *base_ptr.cast::<MapHeader>();
header_mut.element_count += 1;
return target_bucket.add(value_offset);
}
ptr::null_mut()
}
}
#[inline]
pub unsafe fn get_or_reserve_entry(base_ptr: *mut u8, key_ptr: *const u8) -> *mut u8 {
unsafe {
assert_validate(base_ptr);
let header = &*base_ptr.cast::<MapHeader>();
let capacity = header.capacity as usize;
let key_size = header.key_size as usize;
let bucket_size = header.bucket_size as usize;
let key_offset = header.key_offset as usize;
let value_offset = header.value_offset as usize;
let buckets_ptr = base_ptr.add(MAP_BUCKETS_OFFSET);
let key_slice = slice::from_raw_parts(key_ptr, key_size);
let hash = calculate_hash_bytes(key_slice);
let mut index = index_from_hash(hash, header.capacity);
let mut first_tombstone = None;
let probe_limit = min(capacity, MAX_PROBE_DISTANCE);
for _ in 0..probe_limit {
let bucket_ptr = buckets_ptr.add(index * bucket_size);
let status = *bucket_ptr;
match status {
status if status == BucketStatus::Empty as u8 => {
let insert_index = first_tombstone.unwrap_or(index);
let target_bucket = buckets_ptr.add(insert_index * bucket_size);
*target_bucket = BucketStatus::Occupied as u8;
let target_key_ptr = target_bucket.add(key_offset);
ptr::copy_nonoverlapping(key_ptr, target_key_ptr, key_size);
let header_mut = &mut *base_ptr.cast::<MapHeader>();
header_mut.element_count += 1;
return target_bucket.add(value_offset);
}
status if status == BucketStatus::Occupied as u8 => {
let existing_key_ptr = bucket_ptr.add(key_offset);
if matches_key(existing_key_ptr, key_ptr, key_size) {
return bucket_ptr.add(value_offset);
}
}
status if status == BucketStatus::Tombstone as u8 => {
if first_tombstone.is_none() {
first_tombstone = Some(index);
}
}
_ => unreachable!(),
}
index = (index + 1) & (capacity - 1);
}
if let Some(tombstone_index) = first_tombstone {
let target_bucket = buckets_ptr.add(tombstone_index * bucket_size);
*target_bucket = BucketStatus::Occupied as u8;
let target_key_ptr = target_bucket.add(key_offset);
ptr::copy_nonoverlapping(key_ptr, target_key_ptr, key_size);
let header_mut = &mut *base_ptr.cast::<MapHeader>();
header_mut.element_count += 1;
return target_bucket.add(value_offset);
}
ptr::null_mut()
}
}
#[inline]
#[must_use]
pub unsafe fn has(base_ptr: *const u8, key_ptr: *const u8) -> bool {
unsafe { lookup(base_ptr.cast_mut(), key_ptr).is_null().not() }
}
#[inline]
pub unsafe fn lookup(base_ptr: *mut u8, key_ptr: *const u8) -> *mut u8 {
unsafe {
assert_validate(base_ptr);
let header = &*base_ptr.cast::<MapHeader>();
let capacity = header.capacity as usize;
let key_size = header.key_size as usize;
let bucket_size = header.bucket_size as usize;
let key_offset = header.key_offset as usize;
let value_offset = header.value_offset as usize;
let buckets_ptr = base_ptr.add(MAP_BUCKETS_OFFSET);
let key_slice = slice::from_raw_parts(key_ptr, key_size);
let hash = calculate_hash_bytes(key_slice);
let mut index = index_from_hash(hash, header.capacity);
let probe_limit = min(capacity, MAX_PROBE_DISTANCE);
for _ in 0..probe_limit {
let bucket_ptr = buckets_ptr.add(index * bucket_size);
let status = *bucket_ptr;
match status {
status if status == BucketStatus::Empty as u8 => {
return ptr::null_mut();
}
status if status == BucketStatus::Occupied as u8 => {
let existing_key_ptr = bucket_ptr.add(key_offset);
if matches_key(existing_key_ptr, key_ptr, key_size) {
return bucket_ptr.add(value_offset);
}
}
_ => {} }
index = (index + 1) & (capacity - 1);
}
ptr::null_mut()
}
}
#[inline]
pub unsafe fn remove(base_ptr: *mut u8, key_ptr: *const u8) -> bool {
unsafe {
assert_validate(base_ptr);
let header = &*base_ptr.cast::<MapHeader>();
let capacity = header.capacity as usize;
let key_size = header.key_size as usize;
let bucket_size = header.bucket_size as usize;
let key_offset = header.key_offset as usize;
let buckets_ptr = base_ptr.add(MAP_BUCKETS_OFFSET);
let key_slice = slice::from_raw_parts(key_ptr, key_size);
let hash = calculate_hash_bytes(key_slice);
let mut index = index_from_hash(hash, header.capacity);
let probe_limit = min(capacity, MAX_PROBE_DISTANCE);
for _ in 0..probe_limit {
let bucket_ptr = buckets_ptr.add(index * bucket_size);
let status = *bucket_ptr;
match status {
status if status == BucketStatus::Empty as u8 => {
return false;
}
status if status == BucketStatus::Occupied as u8 => {
let existing_key_ptr = bucket_ptr.add(key_offset);
if matches_key(existing_key_ptr, key_ptr, key_size) {
*bucket_ptr = BucketStatus::Tombstone as u8;
let header_mut = &mut *base_ptr.cast::<MapHeader>();
header_mut.element_count -= 1;
return true;
}
}
_ => {} }
index = (index + 1) & (capacity - 1);
}
false
}
}
#[inline]
pub unsafe fn overwrite(target_base: *mut u8, source: *const u8) -> bool {
unsafe {
assert_validate(target_base);
assert_validate(source);
let target_header = &mut *target_base.cast::<MapHeader>();
let source_header = &*source.cast::<MapHeader>();
if target_header.logical_limit < source_header.element_count {
return false;
}
assert_eq!(
target_header.bucket_size, source_header.bucket_size,
"Incompatible bucket sizes"
);
assert_eq!(
target_header.key_size, source_header.key_size,
"Incompatible key sizes"
);
assert_eq!(
target_header.value_size, source_header.value_size,
"Incompatible value sizes"
);
let source_buckets_ptr = source.add(MAP_BUCKETS_OFFSET);
let bucket_size = source_header.bucket_size as usize;
let key_offset = source_header.key_offset as usize;
let value_offset = source_header.value_offset as usize;
let value_size = source_header.value_size as usize;
for i in 0..source_header.capacity as usize {
let source_bucket = source_buckets_ptr.add(i * bucket_size);
if *source_bucket == BucketStatus::Occupied as u8 {
let source_key_ptr = source_bucket.add(key_offset);
let source_value_ptr = source_bucket.add(value_offset);
let target_value_ptr = get_or_reserve_entry(target_base, source_key_ptr);
if target_value_ptr.is_null() {
return false;
}
ptr::copy_nonoverlapping(source_value_ptr, target_value_ptr, value_size);
}
}
true
}
}
#[inline]
pub unsafe fn find_next_valid_entry(base: *mut u8, start_index: u16) -> (*const u8, *mut u8, u16) {
unsafe {
assert_validate(base);
let map_header = &*base.cast::<MapHeader>();
let bucket_size = map_header.bucket_size as usize;
let buckets_start = base.add(MAP_BUCKETS_OFFSET);
let key_offset = map_header.key_offset as usize;
let value_offset = map_header.value_offset as usize;
let mut index = start_index as usize;
while index < map_header.capacity as usize {
let entry_ptr = buckets_start.add(index * bucket_size);
if *entry_ptr == BucketStatus::Occupied as u8 {
let key_addr = entry_ptr.add(key_offset);
let value_addr = entry_ptr.add(value_offset);
return (key_addr, value_addr, index as u16);
}
index += 1;
}
(ptr::null(), ptr::null_mut(), 0xFFFF)
}
}