use portable_atomic::{AtomicU64, Ordering};
use std::cell::UnsafeCell;
use std::marker::PhantomData;
use std::sync::atomic::{AtomicU8, AtomicUsize};
use std::sync::{Mutex, RwLock};
use crate::engine::access_buffer::AccessBuffer;
use crate::engine::bucket::Bucket;
use crate::engine::hash::compute_hash;
use crate::engine::slab::SlabPool;
use crate::raw::SlotTTL;
use crate::{PulseKey, PulseValue, SlotState};
struct BucketLocks {
locks: Vec<AtomicU8>,
}
impl BucketLocks {
fn new(num_buckets: usize) -> Self {
let locks = (0..num_buckets).map(|_| AtomicU8::new(0)).collect();
Self { locks }
}
#[inline]
fn lock(&self, bucket_idx: usize) {
while self.locks[bucket_idx]
.compare_exchange_weak(0, 1, Ordering::Acquire, Ordering::Relaxed)
.is_err()
{
std::hint::spin_loop();
}
}
#[inline]
fn unlock(&self, bucket_idx: usize) {
self.locks[bucket_idx].store(0, Ordering::Release);
}
}
struct BucketGuard<'a> {
locks: &'a BucketLocks,
idx: usize,
}
impl<'a> BucketGuard<'a> {
#[inline]
fn new(locks: &'a BucketLocks, idx: usize) -> Self {
locks.lock(idx);
Self { locks, idx }
}
}
impl Drop for BucketGuard<'_> {
#[inline]
fn drop(&mut self) {
self.locks.unlock(self.idx);
}
}
struct MapInner {
buckets: Vec<UnsafeCell<Bucket>>,
locks: BucketLocks,
slab_pool: Mutex<SlabPool>,
num_buckets: usize,
bucket_mask: usize,
epochs: Mutex<Vec<SlotTTL>>,
}
unsafe impl Send for MapInner {}
unsafe impl Sync for MapInner {}
impl MapInner {
fn new(num_buckets: usize) -> Self {
let actual = num_buckets.max(1).next_power_of_two();
let buckets = (0..actual)
.map(|_| UnsafeCell::new(Bucket::empty()))
.collect();
Self {
buckets,
locks: BucketLocks::new(actual),
slab_pool: Mutex::new(SlabPool::new()),
num_buckets: actual,
bucket_mask: actual - 1,
epochs: Mutex::new(vec![SlotTTL::default(); actual * 4]),
}
}
}
pub struct ConcurrentPulseMap<K: PulseKey, V: PulseValue> {
inner: RwLock<MapInner>,
count: AtomicUsize,
eviction_count: AtomicUsize,
auto_resize: bool,
resize_threshold: f64,
current_epoch: AtomicU64,
default_ttl: AtomicU64,
access_buffer: AccessBuffer,
_marker: PhantomData<(K, V)>,
}
unsafe impl<K: PulseKey, V: PulseValue> Send for ConcurrentPulseMap<K, V> {}
unsafe impl<K: PulseKey, V: PulseValue> Sync for ConcurrentPulseMap<K, V> {}
impl<K: PulseKey, V: PulseValue> ConcurrentPulseMap<K, V> {
pub fn new(num_buckets: usize) -> Self {
Self {
inner: RwLock::new(MapInner::new(num_buckets)),
count: AtomicUsize::new(0),
eviction_count: AtomicUsize::new(0),
auto_resize: false,
resize_threshold: 0.75,
current_epoch: AtomicU64::new(0),
default_ttl: AtomicU64::new(0),
access_buffer: AccessBuffer::new(4096),
_marker: PhantomData,
}
}
pub fn with_auto_resize(num_buckets: usize) -> Self {
Self {
inner: RwLock::new(MapInner::new(num_buckets)),
count: AtomicUsize::new(0),
eviction_count: AtomicUsize::new(0),
auto_resize: true,
resize_threshold: 0.75,
current_epoch: AtomicU64::new(0),
default_ttl: AtomicU64::new(0),
access_buffer: AccessBuffer::new(4096),
_marker: PhantomData,
}
}
#[inline]
pub fn set_ttl(&self, ttl: u64) {
self.default_ttl.store(ttl, Ordering::Relaxed);
}
#[inline]
pub fn get_ttl(&self) -> u64 {
self.default_ttl.load(Ordering::Relaxed)
}
#[inline]
pub fn current_epoch(&self) -> u64 {
self.current_epoch.load(Ordering::Relaxed)
}
#[inline]
fn is_expired(&self, state: &MapInner, bucket_idx: usize, slot_idx: u8) -> bool {
let entry = state.epochs.lock().unwrap()[bucket_idx * 4 + slot_idx as usize];
let effective_ttl = if entry.ttl == 0 {
self.default_ttl.load(Ordering::Relaxed)
} else {
entry.ttl
};
if effective_ttl == 0 || effective_ttl == u64::MAX {
return false;
}
let epoch = self.current_epoch.load(Ordering::Relaxed);
epoch.wrapping_sub(entry.epoch) > effective_ttl
}
#[inline]
fn stamp_epoch(&self, state: &MapInner, bucket_idx: usize, slot_idx: u8, ttl: u64) {
let epoch = self.current_epoch.load(Ordering::Relaxed);
state.epochs.lock().unwrap()[bucket_idx * 4 + slot_idx as usize] = SlotTTL { epoch, ttl };
}
pub fn insert(&self, key: K, value: V) {
self.insert_internal(key, value, 0);
}
pub fn insert_ttl(&self, key: K, value: V, ttl: u64) {
self.insert_internal(key, value, ttl);
}
fn insert_internal(&self, key: K, value: V, ttl: u64) {
if self.auto_resize {
let state = self.inner.read().unwrap();
let num_bkts = state.num_buckets;
let cap = num_bkts * 4;
let len = self.count.load(Ordering::Relaxed);
let load = len as f64 / cap as f64;
drop(state);
if load > self.resize_threshold {
self.resize(num_bkts * 2);
}
}
self.current_epoch.fetch_add(1, Ordering::Relaxed);
let kb = key.to_bytes();
let vb = value.to_bytes();
let key_bytes = kb.as_ref();
let val_bytes = vb.as_ref();
let hr = compute_hash(key_bytes);
let state = self.inner.read().unwrap();
let idx = (hr.h1 as usize) & state.bucket_mask;
let _guard = BucketGuard::new(&state.locks, idx);
let bucket = unsafe { &mut *state.buckets[idx].get() };
let mask = bucket.meta.match_mask(hr.h2);
let mut m = mask;
while m != 0 {
let slot_idx = m.trailing_zeros() as u8;
m &= m - 1;
let slot = &bucket.slots[slot_idx as usize];
if slot.matches_key(key_bytes, &hr, &state.slab_pool.lock().unwrap()) {
if slot.get_mode() == 1 {
state.slab_pool.lock().unwrap().free(slot.slab_idx());
}
let s = &mut bucket.slots[slot_idx as usize];
if key_bytes.len() <= 6 && val_bytes.len() <= 7 {
s.set_inline(key_bytes, val_bytes);
} else {
let idx = state.slab_pool.lock().unwrap().alloc(key_bytes, val_bytes);
s.set_slab(hr.ext_fp_hi, hr.ext_fp, idx);
}
bucket.meta.on_access(slot_idx);
self.stamp_epoch(&state, idx, slot_idx, ttl);
return;
}
}
let (target_slot, is_eviction) = if let Some(free) = bucket.meta.find_free_slot() {
(free, false)
} else if let Some(evict) = bucket.meta.find_evict_target() {
let old_slot = &bucket.slots[evict as usize];
if old_slot.get_mode() == 1 {
state.slab_pool.lock().unwrap().free(old_slot.slab_idx());
}
self.eviction_count.fetch_add(1, Ordering::Relaxed);
(evict, true)
} else {
return;
};
let slot = &mut bucket.slots[target_slot as usize];
if key_bytes.len() <= 6 && val_bytes.len() <= 7 {
slot.set_inline(key_bytes, val_bytes);
} else {
let idx = state.slab_pool.lock().unwrap().alloc(key_bytes, val_bytes);
slot.set_slab(hr.ext_fp_hi, hr.ext_fp, idx);
}
bucket.meta.set_state(target_slot, SlotState::Full);
bucket.meta.set_h2(target_slot, hr.h2);
bucket.meta.on_insert(target_slot);
self.stamp_epoch(&state, idx, target_slot, ttl);
if !is_eviction {
self.count.fetch_add(1, Ordering::Relaxed);
}
}
pub fn get(&self, key: &K) -> Option<V> {
key.with_key_bytes(|key_bytes| {
let hr = compute_hash(key_bytes);
let state = self.inner.read().unwrap();
let idx = (hr.h1 as usize) & state.bucket_mask;
let _guard = BucketGuard::new(&state.locks, idx);
let bucket = unsafe { &mut *state.buckets[idx].get() };
let mask = bucket.meta.match_mask(hr.h2);
let mut m = mask;
while m != 0 {
let slot_idx = m.trailing_zeros() as u8;
m &= m - 1;
let slot = &bucket.slots[slot_idx as usize];
let matched = if slot.get_mode() == 0 {
slot.inline_key() == key_bytes
} else {
if slot.data[0] & 0x7F != hr.ext_fp_hi {
false
} else {
let mut fp_bytes = [0u8; 4];
fp_bytes.copy_from_slice(&slot.data[1..5]);
if u32::from_le_bytes(fp_bytes) != hr.ext_fp {
false
} else {
let slab = state.slab_pool.lock().unwrap();
slab.get(slot.slab_idx()).key() == key_bytes
}
}
};
if matched {
if self.is_expired(&state, idx, slot_idx) {
return None;
}
self.access_buffer.push(idx, slot_idx);
let val_bytes = if slot.get_mode() == 0 {
slot.inline_value().to_vec()
} else {
let slab = state.slab_pool.lock().unwrap();
slot.get_value(&slab).to_vec()
};
return V::from_bytes(&val_bytes);
}
}
None
})
}
pub fn peek(&self, key: &K) -> Option<V> {
key.with_key_bytes(|key_bytes| {
let hr = compute_hash(key_bytes);
let state = self.inner.read().unwrap();
let idx = (hr.h1 as usize) & state.bucket_mask;
let _guard = BucketGuard::new(&state.locks, idx);
let bucket = unsafe { &*state.buckets[idx].get() };
let mask = bucket.meta.match_mask(hr.h2);
let mut m = mask;
while m != 0 {
let slot_idx = m.trailing_zeros() as u8;
m &= m - 1;
let slot = &bucket.slots[slot_idx as usize];
let slab = state.slab_pool.lock().unwrap();
if slot.matches_key(key_bytes, &hr, &slab) {
if self.is_expired(&state, idx, slot_idx) {
return None;
}
let val_bytes = slot.get_value(&slab).to_vec();
drop(slab);
return V::from_bytes(&val_bytes);
}
}
None
})
}
#[inline]
pub fn contains_key(&self, key: &K) -> bool {
self.peek(key).is_some()
}
pub fn remove(&self, key: &K) -> bool {
key.with_key_bytes(|key_bytes| {
let hr = compute_hash(key_bytes);
let state = self.inner.read().unwrap();
let idx = (hr.h1 as usize) & state.bucket_mask;
let _guard = BucketGuard::new(&state.locks, idx);
let bucket = unsafe { &mut *state.buckets[idx].get() };
let mask = bucket.meta.match_mask(hr.h2);
let mut m = mask;
while m != 0 {
let slot_idx = m.trailing_zeros() as u8;
m &= m - 1;
let slot = &bucket.slots[slot_idx as usize];
if slot.matches_key(key_bytes, &hr, &state.slab_pool.lock().unwrap()) {
if slot.get_mode() == 1 {
state.slab_pool.lock().unwrap().free(slot.slab_idx());
}
bucket.meta.set_state(slot_idx, SlotState::Tombstone);
bucket.slots[slot_idx as usize].clear();
self.count.fetch_sub(1, Ordering::Relaxed);
return true;
}
}
false
})
}
pub fn resize(&self, new_num_buckets: usize) {
let mut new_actual = new_num_buckets.max(1).next_power_of_two();
let mut state = self.inner.write().unwrap();
if state.num_buckets >= new_actual {
return;
}
struct EntryData {
key_bytes: Vec<u8>,
val_bytes: Vec<u8>,
slot_ttl: SlotTTL,
}
let old_epochs = state.epochs.lock().unwrap().clone();
let mut entries: Vec<EntryData> = Vec::new();
for (bucket_idx, bucket_cell) in state.buckets.iter().enumerate() {
let bucket = unsafe { &*bucket_cell.get() };
for slot_idx in 0..4u8 {
if bucket.meta.get_state(slot_idx) != SlotState::Full {
continue;
}
let slot = &bucket.slots[slot_idx as usize];
let slab = state.slab_pool.lock().unwrap();
let key_bytes = slot.get_key_bytes(&slab).to_vec();
let val_bytes = slot.get_value_bytes(&slab).to_vec();
drop(slab);
if key_bytes.is_empty() {
continue;
}
let ttl_idx = bucket_idx * 4 + slot_idx as usize;
let slot_ttl = if ttl_idx < old_epochs.len() {
old_epochs[ttl_idx]
} else {
SlotTTL::default()
};
entries.push(EntryData {
key_bytes,
val_bytes,
slot_ttl,
});
}
}
loop {
let new_mask = new_actual - 1;
let new_buckets: Vec<UnsafeCell<Bucket>> = (0..new_actual)
.map(|_| UnsafeCell::new(Bucket::empty()))
.collect();
let new_slab = Mutex::new(SlabPool::new());
let mut new_epochs = vec![SlotTTL::default(); new_actual * 4];
let mut new_count = 0usize;
let mut overflow = false;
for entry in entries.iter() {
let hr = compute_hash(&entry.key_bytes);
let new_idx = (hr.h1 as usize) & new_mask;
let new_bucket = unsafe { &mut *new_buckets[new_idx].get() };
if let Some(free) = new_bucket.meta.find_free_slot() {
let new_slot = &mut new_bucket.slots[free as usize];
if entry.key_bytes.len() <= 6 && entry.val_bytes.len() <= 7 {
new_slot.set_inline(&entry.key_bytes, &entry.val_bytes);
} else {
let mut slab = new_slab.lock().unwrap();
let idx = slab.alloc(&entry.key_bytes, &entry.val_bytes);
new_slot.set_slab(hr.ext_fp_hi, hr.ext_fp, idx);
}
new_bucket.meta.set_state(free, SlotState::Full);
new_bucket.meta.set_h2(free, hr.h2);
new_bucket.meta.on_insert(free);
new_epochs[new_idx * 4 + free as usize] = entry.slot_ttl;
new_count += 1;
} else {
overflow = true;
break;
}
}
if overflow {
new_actual *= 2;
continue;
}
state.buckets = new_buckets;
state.locks = BucketLocks::new(new_actual);
state.slab_pool = new_slab;
state.num_buckets = new_actual;
state.bucket_mask = new_mask;
*state.epochs.lock().unwrap() = new_epochs;
self.count.store(new_count, Ordering::Relaxed);
break;
}
}
#[inline]
pub fn len(&self) -> usize {
self.count.load(Ordering::Relaxed)
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline]
pub fn capacity(&self) -> usize {
let state = self.inner.read().unwrap();
state.num_buckets * 4
}
#[inline]
pub fn load_factor(&self) -> f64 {
let cap = self.capacity();
self.len() as f64 / cap as f64
}
#[inline]
pub fn eviction_count(&self) -> usize {
self.eviction_count.load(Ordering::Relaxed)
}
#[inline]
pub fn num_buckets(&self) -> usize {
let state = self.inner.read().unwrap();
state.num_buckets
}
}
impl<K: PulseKey, V: PulseValue> std::fmt::Debug for ConcurrentPulseMap<K, V> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ConcurrentPulseMap")
.field("len", &self.len())
.field("capacity", &self.capacity())
.field(
"load_factor",
&format!("{:.1}%", self.load_factor() * 100.0),
)
.field("evictions", &self.eviction_count())
.finish()
}
}
impl<K: PulseKey, V: PulseValue> std::fmt::Display for ConcurrentPulseMap<K, V> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"ConcurrentPulseMap({}/{} entries, {:.1}% load, {} evictions)",
self.len(),
self.capacity(),
self.load_factor() * 100.0,
self.eviction_count()
)
}
}