use std::hash::{BuildHasher, Hasher};
use std::sync::atomic::Ordering;
use hashbrown::HashMap;
use hashbrown::hash_map::Entry as HashMapEntry;
use indexmap::IndexMap;
use crate::erased::{Entry, ErasedKey, ErasedKeyLookup, ErasedKeyRef};
use crate::traits::{CacheKey, CacheKeyLookup, NUM_POLICY_BUCKETS};
#[derive(Debug, Clone, Default)]
pub struct EvictionStats {
pub count: usize,
pub size: usize,
}
pub struct EvictedEntry {
pub key: ErasedKey,
pub entry: Entry,
}
#[derive(Default)]
pub(crate) struct PassthroughHasher(u64);
impl Hasher for PassthroughHasher {
fn finish(&self) -> u64 {
self.0
}
fn write(&mut self, _bytes: &[u8]) {
panic!("PassthroughHasher only works with u64 hash values");
}
fn write_u64(&mut self, i: u64) {
self.0 = i;
}
}
#[derive(Clone, Default)]
pub(crate) struct PassthroughBuildHasher;
impl BuildHasher for PassthroughBuildHasher {
type Hasher = PassthroughHasher;
fn build_hasher(&self) -> Self::Hasher {
PassthroughHasher::default()
}
}
struct PolicyBucket {
list: IndexMap<ErasedKey, (), PassthroughBuildHasher>,
hand: usize,
}
impl PolicyBucket {
fn new() -> Self {
Self {
list: IndexMap::with_hasher(PassthroughBuildHasher),
hand: 0,
}
}
fn len(&self) -> usize {
self.list.len()
}
fn is_empty(&self) -> bool {
self.list.is_empty()
}
fn insert(&mut self, key: ErasedKey) {
self.list.insert(key, ());
if self.hand >= self.list.len() && !self.list.is_empty() {
self.hand = 0;
}
}
fn remove(&mut self, key: &ErasedKey) -> bool {
let old_len = self.list.len();
if let Some((removed_idx, _, _)) = self.list.swap_remove_full(key) {
let new_len = self.list.len();
if new_len == 0 {
self.hand = 0;
} else if removed_idx < self.hand {
self.hand -= 1;
} else if self.hand == old_len - 1 && removed_idx != old_len - 1 {
self.hand = removed_idx;
}
if self.hand >= new_len && new_len > 0 {
self.hand = 0;
}
true
} else {
false
}
}
fn clear(&mut self) {
self.list.clear();
self.hand = 0;
}
}
pub struct Shard {
pub(crate) entries: HashMap<ErasedKey, Entry, PassthroughBuildHasher>,
buckets: [PolicyBucket; 3],
size_current: usize,
size_capacity: usize,
}
impl Shard {
pub fn new(size_capacity: usize) -> Self {
Self {
entries: HashMap::with_hasher(PassthroughBuildHasher),
buckets: [
PolicyBucket::new(), PolicyBucket::new(), PolicyBucket::new(), ],
size_current: 0,
size_capacity,
}
}
pub fn insert(
&mut self,
key: ErasedKey,
entry: Entry,
) -> (Option<Entry>, EvictionStats, Vec<EvictedEntry>) {
let size = entry.size;
let policy = entry.policy;
let old_info = self.entries.raw_entry().from_key(&key).map(|(_, e)| (e.policy, e.size));
if let Some((old_policy, old_size)) = old_info {
self.buckets[old_policy as usize].remove(&key);
self.size_current -= old_size;
}
let (stats, evicted) = self.evict_until_space(size);
let old = match self.entries.entry(key.clone()) {
HashMapEntry::Occupied(mut occupied) => Some(occupied.insert(entry)),
HashMapEntry::Vacant(vacant) => {
vacant.insert(entry);
None
}
};
self.buckets[policy as usize].insert(key);
self.size_current += size;
(old, stats, evicted)
}
fn evict_until_space(&mut self, needed_size: usize) -> (EvictionStats, Vec<EvictedEntry>) {
let mut stats = EvictionStats::default();
let mut evicted = Vec::new();
for policy_idx in (0..NUM_POLICY_BUCKETS).rev() {
while self.size_current + needed_size > self.size_capacity {
if let Some(evicted_entry) = self.evict_from_bucket(policy_idx) {
stats.count += 1;
stats.size += evicted_entry.entry.size;
evicted.push(evicted_entry);
} else {
break; }
}
if self.size_current + needed_size <= self.size_capacity {
break;
}
}
(stats, evicted)
}
#[cfg(test)]
pub fn get(&self, key: &ErasedKey) -> Option<&Entry> {
let entry = self.entries.get(key)?;
entry.clock_bit.store(true, Ordering::Relaxed);
let freq = entry.frequency.load(Ordering::Relaxed);
if freq < 255 {
entry.frequency.store(freq + 1, Ordering::Relaxed);
}
Some(entry)
}
pub fn get_ref<K: crate::traits::CacheKey>(&self, key_ref: &ErasedKeyRef<K>) -> Option<&Entry> {
let (_key, entry) = self
.entries
.raw_entry()
.from_hash(key_ref.hash, |stored_key| key_ref.equals(stored_key))?;
entry.clock_bit.store(true, Ordering::Relaxed);
let freq = entry.frequency.load(Ordering::Relaxed);
if freq < 255 {
entry.frequency.store(freq + 1, Ordering::Relaxed);
}
Some(entry)
}
pub fn get_ref_by<K, Q>(&self, key_ref: &ErasedKeyLookup<K, Q>) -> Option<&Entry>
where
K: CacheKey,
Q: CacheKeyLookup<K> + ?Sized,
{
let (_key, entry) = self
.entries
.raw_entry()
.from_hash(key_ref.hash, |stored_key| key_ref.equals(stored_key))?;
entry.clock_bit.store(true, Ordering::Relaxed);
let freq = entry.frequency.load(Ordering::Relaxed);
if freq < 255 {
entry.frequency.store(freq + 1, Ordering::Relaxed);
}
Some(entry)
}
pub fn remove(&mut self, key: &ErasedKey) -> Option<(ErasedKey, Entry)> {
let (stored_key, entry) = self.entries.remove_entry(key)?;
let policy = entry.policy;
let size = entry.size;
self.buckets[policy as usize].remove(&stored_key);
self.size_current -= size;
Some((stored_key, entry))
}
#[cfg(test)]
pub fn contains(&self, key: &ErasedKey) -> bool {
self.entries.contains_key(key)
}
#[cfg(test)]
pub fn len(&self) -> usize {
self.entries.len()
}
#[allow(dead_code)]
pub fn clear(&mut self) {
self.entries.clear();
for bucket in &mut self.buckets {
bucket.clear();
}
self.size_current = 0;
}
pub fn drain(&mut self) -> impl Iterator<Item = EvictedEntry> + '_ {
for bucket in &mut self.buckets {
bucket.clear();
}
self.size_current = 0;
self.entries.drain().map(|(key, entry)| EvictedEntry {
key,
entry,
})
}
fn evict_from_bucket(&mut self, policy_idx: usize) -> Option<EvictedEntry> {
let bucket = &self.buckets[policy_idx];
if bucket.is_empty() {
return None;
}
let bucket_len = bucket.len();
let mut hand = self.buckets[policy_idx].hand;
for _ in 0..bucket_len {
let key_ref = self.buckets[policy_idx].list.get_index(hand)?.0;
let entry = self.entries.get(key_ref)?;
let clock_bit = entry.clock_bit.load(Ordering::Relaxed);
let frequency = entry.frequency.load(Ordering::Relaxed);
if clock_bit {
entry.clock_bit.store(false, Ordering::Relaxed);
hand += 1;
if hand >= bucket_len {
hand = 0;
}
} else if frequency == 0 {
let key = key_ref.clone();
self.buckets[policy_idx].hand = hand;
let evicted = self.entries.remove(&key)?;
let evicted_size = evicted.size;
self.buckets[policy_idx].remove(&key);
self.size_current -= evicted_size;
return Some(EvictedEntry {
key,
entry: evicted,
});
} else {
entry.frequency.fetch_sub(1, Ordering::Relaxed);
hand += 1;
if hand >= bucket_len {
hand = 0;
}
}
}
self.buckets[policy_idx].hand = hand;
let key = self.buckets[policy_idx].list.get_index(hand)?.0.clone();
let evicted = self.entries.remove(&key)?;
let evicted_size = evicted.size;
self.buckets[policy_idx].remove(&key);
self.size_current -= evicted_size;
Some(EvictedEntry {
key,
entry: evicted,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::DeepSizeOf;
use crate::traits::{CacheKey, CachePolicy};
#[derive(Hash, Eq, PartialEq, Clone, Debug, DeepSizeOf)]
struct TestKey(u64, CachePolicy);
impl CacheKey for TestKey {
type Value = TestValue;
fn policy(&self) -> CachePolicy {
self.1
}
}
#[derive(DeepSizeOf)]
struct TestValue {
size: usize,
}
fn make_key(id: u64, policy: CachePolicy) -> ErasedKey {
ErasedKey::new(&TestKey(id, policy))
}
fn make_entry(size: usize, policy: CachePolicy) -> Entry {
Entry::new(
TestValue {
size,
},
policy,
)
}
#[test]
fn test_shard_insert() {
let mut shard = Shard::new(1000);
let key = make_key(1, CachePolicy::Standard);
let entry = make_entry(50, CachePolicy::Standard);
let (old, _stats, _evicted) = shard.insert(key.clone(), entry);
assert!(old.is_none());
assert!(shard.contains(&key));
assert_eq!(shard.len(), 1);
}
#[test]
fn test_shard_remove() {
let mut shard = Shard::new(1000);
let key = make_key(1, CachePolicy::Standard);
let entry = make_entry(50, CachePolicy::Standard);
shard.insert(key.clone(), entry);
assert!(shard.remove(&key).is_some());
assert!(!shard.contains(&key));
assert_eq!(shard.len(), 0);
}
#[test]
fn test_get_updates_clock_and_frequency() {
let mut shard = Shard::new(1000);
let key = make_key(1, CachePolicy::Standard);
let entry = make_entry(50, CachePolicy::Standard);
shard.insert(key.clone(), entry);
let e = shard.get(&key).expect("entry should exist");
assert_eq!(e.clock_bit.load(Ordering::Relaxed), true);
assert_eq!(e.frequency.load(Ordering::Relaxed), 1);
let e = shard.get(&key).expect("entry should exist");
assert_eq!(e.clock_bit.load(Ordering::Relaxed), true);
assert_eq!(e.frequency.load(Ordering::Relaxed), 2);
}
#[test]
fn test_get_ref_zero_allocation() {
use crate::erased::ErasedKeyRef;
let mut shard = Shard::new(1000);
let key = TestKey(1, CachePolicy::Standard);
let entry = make_entry(50, CachePolicy::Standard);
let erased = ErasedKey::new(&key);
let hash = erased.hash;
shard.insert(erased, entry);
assert_eq!(shard.entries.len(), 1, "Should have 1 entry");
let key_ref = ErasedKeyRef::new(&key);
assert_eq!(key_ref.hash, hash, "Hashes should match");
let e = shard.get_ref(&key_ref).expect("get_ref should find the entry");
assert_eq!(e.clock_bit.load(Ordering::Relaxed), true);
assert_eq!(e.frequency.load(Ordering::Relaxed), 1);
let e = shard.get_ref(&key_ref).expect("entry should exist");
assert_eq!(e.frequency.load(Ordering::Relaxed), 2);
}
#[test]
fn test_policy_based_eviction() {
let mut shard = Shard::new(200);
let volatile_key = make_key(1, CachePolicy::Volatile);
let volatile_entry = make_entry(50, CachePolicy::Volatile);
shard.insert(volatile_key, volatile_entry);
let standard_key = make_key(2, CachePolicy::Standard);
let standard_entry = make_entry(50, CachePolicy::Standard);
shard.insert(standard_key, standard_entry);
let critical_key = make_key(3, CachePolicy::Critical);
let critical_entry = make_entry(50, CachePolicy::Critical);
shard.insert(critical_key, critical_entry);
for i in 10..15 {
let k = make_key(i, CachePolicy::Standard);
let e = make_entry(50, CachePolicy::Standard);
let _ = shard.insert(k, e);
}
}
#[test]
fn test_frequency_decay() {
let mut shard = Shard::new(1000);
let key = make_key(1, CachePolicy::Standard);
let entry = make_entry(50, CachePolicy::Standard);
shard.insert(key.clone(), entry);
for _ in 0..5 {
shard.get(&key);
}
let e = shard.entries.get(&key).expect("entry should exist");
assert!(e.frequency.load(Ordering::Relaxed) >= 5);
assert_eq!(e.clock_bit.load(Ordering::Relaxed), true);
}
#[test]
fn test_insert_oversized_entry_into_empty_shard() {
#[derive(DeepSizeOf)]
struct LargeValue {
data: Vec<u8>,
}
let mut shard = Shard::new(100);
let key = make_key(1, CachePolicy::Standard);
let large_value = LargeValue {
data: vec![0u8; 200],
};
let entry = Entry::new(large_value, CachePolicy::Standard);
let entry_size = entry.size;
let (old, stats, _evicted) = shard.insert(key.clone(), entry);
assert!(old.is_none());
assert!(shard.contains(&key));
assert_eq!(shard.len(), 1);
assert_eq!(stats.count, 0);
assert_eq!(shard.size_current, entry_size);
assert!(
shard.size_current > shard.size_capacity,
"Expected size {} > capacity {}",
shard.size_current,
shard.size_capacity
);
}
#[test]
fn test_insert_oversized_entry_evicts_existing() {
#[derive(DeepSizeOf)]
struct SmallValue {
data: Vec<u8>,
}
#[derive(DeepSizeOf)]
struct LargeValue {
data: Vec<u8>,
}
let mut shard = Shard::new(200);
for i in 1..=3 {
let key = make_key(i, CachePolicy::Standard);
let small_value = SmallValue {
data: vec![0u8; 5],
};
let entry = Entry::new(small_value, CachePolicy::Standard);
shard.insert(key, entry);
}
let initial_len = shard.len();
assert!(initial_len >= 3, "All 3 small entries should fit initially");
let big_key = make_key(100, CachePolicy::Standard);
let large_value = LargeValue {
data: vec![0u8; 300],
};
let big_entry = Entry::new(large_value, CachePolicy::Standard);
let (_old, stats, _evicted) = shard.insert(big_key.clone(), big_entry);
assert!(stats.count > 0, "Expected some evictions but got none");
assert!(shard.contains(&big_key), "Oversized entry should be inserted");
assert!(
shard.len() < initial_len + 1,
"Expected fewer than {} entries after eviction, but got {}",
initial_len + 1,
shard.len()
);
}
}