use std::{
hint::spin_loop,
sync::atomic::{AtomicU64, Ordering, fence},
thread::{sleep, yield_now},
time::Duration,
};
use wbase::backoff::Backoff;
use whasher::fast_hash;
use crate::{
Result,
bucket::{BucketExclusiveGuard, BucketSharedGuard, HashBucket},
buckets::HashBuckets,
chain::{ChainStep, ChainWalker, SlotScan},
entry::HashBucketEntry,
error::Error,
overflow_pool::OverflowPool,
};
pub use crate::{
candidate::CandidateAddresses, entry_info::HashEntryInfo, guard::MultiBucketGuard,
prefetch::prefetch_read_l1,
};
pub struct HashIndex {
pub buckets: HashBuckets,
pub overflow_pool: OverflowPool,
pub size: usize,
pub mask: usize,
}
impl HashIndex {
pub const PREFETCH_WINDOW: usize = 12;
pub const INLINE_LOCK_ENTRIES: usize = 16;
pub const SPIN_RETRY_THRESHOLD: usize = 32;
pub const SPIN_LIMIT_MAX_EXP: usize = 5;
pub const SPIN_LIMIT_JITTER_MASK: usize = 0x7;
pub const YIELD_RETRY_BUDGET: usize = 1024;
pub const SLEEP_RETRY_BUDGET: usize = 16384;
pub const SLEEP_BASE_MICROS: u64 = 100;
pub const SLEEP_MAX_ADDITIONAL_MICROS: u64 = 900;
pub fn new(num_buckets: usize) -> Result<Self> {
if num_buckets == 0 || !num_buckets.is_power_of_two() {
return Err(Error::InvalidBucketCount(num_buckets));
}
let buckets = HashBuckets::new(num_buckets)?;
Ok(Self {
buckets,
overflow_pool: OverflowPool::new(),
size: num_buckets,
mask: num_buckets - 1,
})
}
#[inline(always)]
pub fn get_bucket(&self, bucket_idx: usize) -> &HashBucket {
let idx = bucket_idx & self.mask;
unsafe { self.buckets.get_unchecked(idx) }
}
#[inline]
pub fn hash_key(key: &[u8]) -> u64 {
fast_hash(key)
}
#[inline]
pub fn find_tag(&self, key: &[u8]) -> Option<u64> {
let hash = Self::hash_key(key);
self.find_tag_by_hash(hash)
}
#[inline]
pub fn find_tag_by_hash(&self, hash: u64) -> Option<u64> {
let tag = HashBucketEntry::tag_from_hash(hash);
let bucket_idx = (hash as usize) & self.mask;
let mut walker = ChainWalker::new(self.get_bucket(bucket_idx));
loop {
if let Some(addr) = walker.curr.find_tag_address(tag) {
return Some(addr);
}
match walker.advance(&self.overflow_pool) {
ChainStep::Next => {}
ChainStep::End | ChainStep::Cycle => return None,
}
}
}
#[inline]
pub fn find_tag_entry(&self, key: &[u8]) -> Option<HashEntryInfo<'_>> {
self.find_tag_entry_by_hash_with_min_addr(Self::hash_key(key), 0)
}
pub fn find_tag_entry_by_hash_with_min_addr(
&self,
hash: u64,
min_valid_addr: u64,
) -> Option<HashEntryInfo<'_>> {
let tag = HashBucketEntry::tag_from_hash(hash);
let bucket_idx = (hash as usize) & self.mask;
let mut walker = ChainWalker::new(self.get_bucket(bucket_idx));
loop {
#[inline(always)]
fn check_slot<'a>(
bucket: &'a HashBucket,
slot: usize,
tag: u16,
min_valid_addr: u64,
) -> Option<HashEntryInfo<'a>> {
if let SlotScan::Hit(raw) =
HashIndex::classify_slot(&bucket.entries[slot], tag, min_valid_addr)
{
Some(HashEntryInfo {
bucket,
slot,
raw,
tag,
})
} else {
None
}
}
if let Some(hei) = check_slot(walker.curr, 0, tag, min_valid_addr) {
return Some(hei);
}
if let Some(hei) = check_slot(walker.curr, 1, tag, min_valid_addr) {
return Some(hei);
}
if let Some(hei) = check_slot(walker.curr, 2, tag, min_valid_addr) {
return Some(hei);
}
if let Some(hei) = check_slot(walker.curr, 3, tag, min_valid_addr) {
return Some(hei);
}
if let Some(hei) = check_slot(walker.curr, 4, tag, min_valid_addr) {
return Some(hei);
}
if let Some(hei) = check_slot(walker.curr, 5, tag, min_valid_addr) {
return Some(hei);
}
if let Some(hei) = check_slot(walker.curr, 6, tag, min_valid_addr) {
return Some(hei);
}
match walker.advance(&self.overflow_pool) {
ChainStep::Next => {}
ChainStep::End | ChainStep::Cycle => return None,
}
}
}
#[inline]
pub fn lookup_candidates(&self, key: &[u8]) -> CandidateAddresses {
self.lookup_candidates_by_hash(Self::hash_key(key))
}
pub fn lookup_candidates_by_hash(&self, hash: u64) -> CandidateAddresses {
let tag = HashBucketEntry::tag_from_hash(hash);
let bucket_idx = (hash as usize) & self.mask;
let mut results = CandidateAddresses::new();
let mut walker = ChainWalker::new(self.get_bucket(bucket_idx));
let expected_hi = (tag as u64) & HashBucketEntry::TAG_MASK;
loop {
#[inline(always)]
fn check_slot(item: &AtomicU64, expected_hi: u64, results: &mut CandidateAddresses) {
let raw = item.load(Ordering::Relaxed);
if raw == 0 {
return;
}
if (raw >> HashBucketEntry::TAG_SHIFT) == expected_hi {
let addr = raw & HashBucketEntry::ADDRESS_MASK;
if addr != 0 {
results.push(addr);
}
}
}
check_slot(&walker.curr.entries[0], expected_hi, &mut results);
check_slot(&walker.curr.entries[1], expected_hi, &mut results);
check_slot(&walker.curr.entries[2], expected_hi, &mut results);
check_slot(&walker.curr.entries[3], expected_hi, &mut results);
check_slot(&walker.curr.entries[4], expected_hi, &mut results);
check_slot(&walker.curr.entries[5], expected_hi, &mut results);
check_slot(&walker.curr.entries[6], expected_hi, &mut results);
match walker.advance(&self.overflow_pool) {
ChainStep::Next => {}
ChainStep::End | ChainStep::Cycle => break,
}
}
if !results.is_empty() {
fence(Ordering::Acquire);
}
results
}
#[inline]
pub fn lookup(&self, key: &[u8]) -> Vec<u64> {
self.lookup_candidates(key).to_vec()
}
#[inline]
pub fn insert(&self, key: &[u8], address: u64) -> Result<()> {
self.insert_by_hash(Self::hash_key(key), address)
}
pub fn insert_by_hash(&self, hash: u64, address: u64) -> Result<()> {
if address == HashBucketEntry::INVALID_ADDRESS {
return Err(Error::InvalidAddress(address));
}
if address > HashBucketEntry::ADDRESS_MASK {
return Err(Error::AddressOverflow(address));
}
let tag = HashBucketEntry::tag_from_hash(hash);
let bucket_idx = (hash as usize) & self.mask;
'retry: loop {
let mut walker = ChainWalker::new(self.get_bucket(bucket_idx));
loop {
if let Some(slot) = walker.curr.find_empty_slot() {
if walker.curr.try_insert(slot, tag, address) {
return Ok(());
}
continue 'retry;
}
match walker.advance(&self.overflow_pool) {
ChainStep::Next => {}
ChainStep::End => {
if walker.curr.overflow_index() == 0 {
let new_idx = self.overflow_pool.allocate()?;
if !walker.curr.set_overflow_index(new_idx) {
self.overflow_pool.free(new_idx);
}
}
match walker.advance(&self.overflow_pool) {
ChainStep::Next => {}
ChainStep::Cycle => return Err(Error::OverflowCycleDetected),
ChainStep::End => return Err(Error::OverflowPoolExhausted),
}
}
ChainStep::Cycle => return Err(Error::OverflowCycleDetected),
}
}
}
}
pub fn find_tag_or_insert_by_hash(&self, hash: u64, address: u64) -> Result<(Option<u64>, bool)> {
if address == HashBucketEntry::INVALID_ADDRESS {
return Err(Error::InvalidAddress(address));
}
if address > HashBucketEntry::ADDRESS_MASK {
return Err(Error::AddressOverflow(address));
}
let mut backoff = Backoff::new();
loop {
let mut hei = self.find_or_create_tag_by_hash(hash)?;
if hei.is_found() {
return Ok((Some(hei.address()), false));
}
if hei.try_cas(address) {
return Ok((None, true));
}
backoff.snooze();
}
}
#[inline]
pub fn find_tag_or_insert(&self, key: &[u8], address: u64) -> Result<(Option<u64>, bool)> {
self.find_tag_or_insert_by_hash(Self::hash_key(key), address)
}
#[inline]
pub fn find_or_create_tag(&self, key: &[u8]) -> Result<HashEntryInfo<'_>> {
self.find_or_create_tag_with_min_addr(key, 0)
}
#[inline]
pub fn find_or_create_tag_with_min_addr(
&self,
key: &[u8],
min_valid_addr: u64,
) -> Result<HashEntryInfo<'_>> {
self.find_or_create_tag_by_hash_with_min_addr(Self::hash_key(key), min_valid_addr)
}
fn classify_slot(item: &AtomicU64, tag: u16, min_valid_addr: u64) -> SlotScan {
let raw = item.load(Ordering::Relaxed);
if raw == 0 {
return SlotScan::Free;
}
let entry = HashBucketEntry::from_raw(raw);
if min_valid_addr > 0
&& !entry.is_tentative()
&& !entry.is_read_cache()
&& entry.address() < min_valid_addr
{
return match item.compare_exchange(raw, 0, Ordering::AcqRel, Ordering::Acquire) {
Ok(_) => SlotScan::Free,
Err(actual_raw) => {
if actual_raw == 0 {
SlotScan::Free
} else {
let actual = HashBucketEntry::from_raw(actual_raw);
if actual.matches_tag(tag)
&& (actual.is_read_cache() || actual.address() >= min_valid_addr)
{
SlotScan::Hit(actual_raw)
} else {
SlotScan::Occupied
}
}
}
};
}
if entry.matches_tag(tag) {
let synced_raw = item.load(Ordering::Acquire);
let synced_entry = HashBucketEntry::from_raw(synced_raw);
if synced_entry.matches_tag(tag) {
SlotScan::Hit(synced_raw)
} else if synced_raw == 0 {
SlotScan::Free
} else {
SlotScan::Occupied
}
} else {
SlotScan::Occupied
}
}
#[inline]
pub fn find_or_create_tag_by_hash(&self, hash: u64) -> Result<HashEntryInfo<'_>> {
self.find_or_create_tag_by_hash_with_min_addr(hash, 0)
}
pub fn find_or_create_tag_by_hash_with_min_addr(
&self,
hash: u64,
min_valid_addr: u64,
) -> Result<HashEntryInfo<'_>> {
let tag = HashBucketEntry::tag_from_hash(hash);
let bucket_idx = (hash as usize) & self.mask;
let mut walker = ChainWalker::new(self.get_bucket(bucket_idx));
let mut first_free: Option<(&HashBucket, usize)> = None;
'search: loop {
#[inline(always)]
fn check_slot<'a>(
bucket: &'a HashBucket,
slot: usize,
tag: u16,
min_valid_addr: u64,
first_free: &mut Option<(&'a HashBucket, usize)>,
) -> Option<HashEntryInfo<'a>> {
match HashIndex::classify_slot(&bucket.entries[slot], tag, min_valid_addr) {
SlotScan::Hit(raw) => Some(HashEntryInfo {
bucket,
slot,
raw,
tag,
}),
SlotScan::Free => {
first_free.get_or_insert((bucket, slot));
None
}
SlotScan::Occupied => None,
}
}
if let Some(hei) = check_slot(walker.curr, 0, tag, min_valid_addr, &mut first_free) {
return Ok(hei);
}
if let Some(hei) = check_slot(walker.curr, 1, tag, min_valid_addr, &mut first_free) {
return Ok(hei);
}
if let Some(hei) = check_slot(walker.curr, 2, tag, min_valid_addr, &mut first_free) {
return Ok(hei);
}
if let Some(hei) = check_slot(walker.curr, 3, tag, min_valid_addr, &mut first_free) {
return Ok(hei);
}
if let Some(hei) = check_slot(walker.curr, 4, tag, min_valid_addr, &mut first_free) {
return Ok(hei);
}
if let Some(hei) = check_slot(walker.curr, 5, tag, min_valid_addr, &mut first_free) {
return Ok(hei);
}
if let Some(hei) = check_slot(walker.curr, 6, tag, min_valid_addr, &mut first_free) {
return Ok(hei);
}
match walker.advance(&self.overflow_pool) {
ChainStep::Next => {}
ChainStep::End => {
if walker.curr.overflow_index() == 0 {
if let Some((free_bucket, slot)) = first_free {
return Ok(HashEntryInfo {
bucket: free_bucket,
slot,
raw: 0,
tag,
});
}
let new_overflow_idx = self.overflow_pool.allocate()?;
if walker.curr.set_overflow_index(new_overflow_idx) {
let new_bucket = unsafe { self.overflow_pool.get_unchecked(new_overflow_idx) };
return Ok(HashEntryInfo {
bucket: new_bucket,
slot: 0,
raw: 0,
tag,
});
}
self.overflow_pool.free(new_overflow_idx);
}
match walker.advance(&self.overflow_pool) {
ChainStep::Next => continue 'search,
ChainStep::Cycle => return Err(Error::OverflowCycleDetected),
ChainStep::End => return Err(Error::OverflowPoolExhausted),
}
}
ChainStep::Cycle => return Err(Error::OverflowCycleDetected),
}
}
}
#[inline]
pub fn update_address(&self, key: &[u8], old_address: u64, new_address: u64) -> bool {
self.update_address_by_hash(Self::hash_key(key), old_address, new_address)
}
pub fn update_address_by_hash(&self, hash: u64, old_address: u64, new_address: u64) -> bool {
if new_address == HashBucketEntry::INVALID_ADDRESS
|| new_address > HashBucketEntry::ADDRESS_MASK
|| old_address == HashBucketEntry::INVALID_ADDRESS
{
return false;
}
let Some(mut hei) = self.find_exact_entry_by_hash(hash, old_address) else {
return false;
};
hei.try_cas(new_address)
}
#[inline]
pub fn delete(&self, key: &[u8], address: u64) -> bool {
self.delete_by_hash(Self::hash_key(key), address)
}
pub fn delete_by_hash(&self, hash: u64, address: u64) -> bool {
if address == HashBucketEntry::INVALID_ADDRESS {
return false;
}
let Some(mut hei) = self.find_exact_entry_by_hash(hash, address) else {
return false;
};
hei.try_elide()
}
fn find_exact_entry_by_hash(&self, hash: u64, address: u64) -> Option<HashEntryInfo<'_>> {
let tag = HashBucketEntry::tag_from_hash(hash);
let mut walker = ChainWalker::new(self.get_bucket((hash as usize) & self.mask));
loop {
if let Some((slot, entry)) = walker.curr.find_entry_by_address(tag, address) {
return Some(HashEntryInfo {
bucket: walker.curr,
slot,
raw: entry.as_raw(),
tag,
});
}
match walker.advance(&self.overflow_pool) {
ChainStep::Next => {}
ChainStep::End | ChainStep::Cycle => return None,
}
}
}
#[inline]
pub fn overflow_bucket_count(&self) -> u64 {
self.overflow_pool.allocated_count()
}
#[inline]
pub fn bucket(&self, bucket_idx: usize) -> &HashBucket {
self.get_bucket(bucket_idx)
}
#[inline]
pub fn bucket_for_key(&self, key: &[u8]) -> &HashBucket {
let hash = Self::hash_key(key);
self.get_bucket((hash as usize) & self.mask)
}
#[inline]
pub fn bucket_index_for_hash(&self, hash: u64) -> usize {
(hash as usize) & self.mask
}
#[inline]
pub fn bucket_index_for_key(&self, key: &[u8]) -> usize {
let hash = Self::hash_key(key);
(hash as usize) & self.mask
}
#[inline]
pub fn try_lock_shared(&self, key: &[u8]) -> bool {
self.bucket_for_key(key).try_lock_shared()
}
#[inline]
pub fn unlock_shared(&self, key: &[u8]) {
self.bucket_for_key(key).unlock_shared();
}
#[inline]
pub fn try_lock_exclusive(&self, key: &[u8]) -> bool {
self.bucket_for_key(key).try_lock_exclusive()
}
#[inline]
pub fn unlock_exclusive(&self, key: &[u8]) {
self.bucket_for_key(key).unlock_exclusive();
}
#[inline]
pub fn downgrade(&self, key: &[u8]) {
self.bucket_for_key(key).downgrade_latch();
}
#[inline]
pub fn is_locked(&self, key: &[u8]) -> bool {
self.bucket_for_key(key).is_latched()
}
#[inline]
pub fn lock_shared_guard(&self, key: &[u8]) -> Option<BucketSharedGuard<'_>> {
self.bucket_for_key(key).lock_shared_guard()
}
#[inline]
pub fn lock_exclusive_guard(&self, key: &[u8]) -> Option<BucketExclusiveGuard<'_>> {
self.bucket_for_key(key).lock_exclusive_guard()
}
fn batch_pipeline(&self, hashes: &[u64], mut query: impl FnMut(&Self, u64)) {
if hashes.is_empty() {
return;
}
for &hash in &hashes[..Self::PREFETCH_WINDOW.min(hashes.len())] {
prefetch_read_l1(self.get_bucket((hash as usize) & self.mask));
}
for (i, &hash) in hashes.iter().enumerate() {
if let Some(&next_hash) = hashes.get(i + Self::PREFETCH_WINDOW) {
prefetch_read_l1(self.get_bucket((next_hash as usize) & self.mask));
}
query(self, hash);
}
}
pub fn lookup_candidates_batch_by_hash(
&self,
hashes: &[u64],
results: &mut [CandidateAddresses],
) {
let count = hashes.len().min(results.len());
let mut idx = 0;
self.batch_pipeline(&hashes[..count], |index, hash| {
results[idx] = index.lookup_candidates_by_hash(hash);
idx += 1;
});
}
pub fn find_tag_batch_by_hash(&self, hashes: &[u64], results: &mut [Option<u64>]) {
let count = hashes.len().min(results.len());
let mut idx = 0;
self.batch_pipeline(&hashes[..count], |index, hash| {
results[idx] = index.find_tag_by_hash(hash);
idx += 1;
});
}
#[inline]
fn in_place_dedup_by<T: Copy, F>(slice: &mut [T], mut same_bucket: F) -> usize
where
F: FnMut(&T, &T) -> bool,
{
if slice.len() <= 1 {
return slice.len();
}
let mut write_idx = 1;
for read_idx in 1..slice.len() {
if !same_bucket(&slice[write_idx - 1], &slice[read_idx]) {
if write_idx != read_idx {
slice[write_idx] = slice[read_idx];
}
write_idx += 1;
}
}
write_idx
}
fn acquire_bucket_locks<I>(&self, items: I) -> Result<MultiBucketGuard<'_>>
where
I: ExactSizeIterator<Item = (usize, bool)>,
{
let count = items.len();
if count == 0 {
return Ok(MultiBucketGuard::new(self));
}
let mut stack_entries = [(0usize, false); Self::INLINE_LOCK_ENTRIES];
let mut heap_entries;
let entries: &mut [(usize, bool)] = if count <= Self::INLINE_LOCK_ENTRIES {
for (slot, e) in stack_entries[..count].iter_mut().zip(items) {
*slot = e;
}
&mut stack_entries[..count]
} else {
heap_entries = items.collect::<Vec<_>>();
&mut heap_entries
};
entries.sort_unstable_by(|a, b| a.0.cmp(&b.0).then_with(|| b.1.cmp(&a.1)));
let deduped_len = Self::in_place_dedup_by(entries, |a, b| a.0 == b.0);
self.acquire_unique_locked_entries(&entries[..deduped_len])
}
#[inline]
pub fn acquire_keys_lock_exclusive(&self, keys: &[&[u8]]) -> Result<MultiBucketGuard<'_>> {
self.acquire_bucket_locks(keys.iter().map(|k| (self.bucket_index_for_key(k), true)))
}
#[inline]
pub fn acquire_hash_locks(&self, items: &[(u64, bool)]) -> Result<MultiBucketGuard<'_>> {
self.acquire_bucket_locks(
items
.iter()
.map(|&(h, ex)| (self.bucket_index_for_hash(h), ex)),
)
}
fn acquire_unique_locked_entries(
&self,
unique_entries: &[(usize, bool)],
) -> Result<MultiBucketGuard<'_>> {
let mut retry_count = 0usize;
loop {
let mut locked_count = 0usize;
for &(b_idx, is_exclusive) in unique_entries {
let bucket = unsafe { self.buckets.get_unchecked(b_idx) };
let ok = if is_exclusive {
bucket.try_lock_exclusive()
} else {
bucket.try_lock_shared()
};
if ok {
locked_count += 1;
} else {
break;
}
}
if locked_count == unique_entries.len() {
return Ok(MultiBucketGuard::from_slice(self, unique_entries));
}
for &(b_idx, is_exclusive) in unique_entries[..locked_count].iter().rev() {
let bucket = unsafe { self.buckets.get_unchecked(b_idx) };
if is_exclusive {
bucket.unlock_exclusive();
} else {
bucket.unlock_shared();
}
}
retry_count += 1;
if retry_count >= Self::YIELD_RETRY_BUDGET + Self::SLEEP_RETRY_BUDGET {
return Err(Error::LockTimeout);
}
if retry_count < Self::SPIN_RETRY_THRESHOLD {
let spin_limit = (1usize << retry_count.min(Self::SPIN_LIMIT_MAX_EXP))
| (retry_count & Self::SPIN_LIMIT_JITTER_MASK);
for _ in 0..spin_limit {
spin_loop();
}
} else if retry_count < Self::YIELD_RETRY_BUDGET {
yield_now();
} else {
let elapsed = retry_count - Self::YIELD_RETRY_BUDGET;
let backoff_us = Self::SLEEP_BASE_MICROS
.saturating_add(((elapsed >> 3) as u64).min(Self::SLEEP_MAX_ADDITIONAL_MICROS));
sleep(Duration::from_micros(backoff_us));
}
}
}
}