use std::sync::Arc;
#[cfg(feature = "metrics")]
use std::sync::atomic::AtomicU64;
use std::sync::atomic::{AtomicUsize, Ordering};
use parking_lot::RwLock;
use crate::erased::{Entry, ErasedKey, ErasedKeyLookup, ErasedKeyRef};
use crate::guard::Guard;
use crate::lifecycle::{DefaultLifecycle, Lifecycle};
#[cfg(feature = "metrics")]
use crate::metrics::CacheMetrics;
use crate::shard::Shard;
use crate::traits::{CacheKey, CacheKeyLookup};
pub struct Cache<L: Lifecycle = DefaultLifecycle> {
shards: Vec<RwLock<Shard>>,
lifecycle: Arc<L>,
current_size: AtomicUsize,
entry_count: AtomicUsize,
shard_count: usize,
#[cfg_attr(not(feature = "metrics"), allow(dead_code))]
max_size_bytes: usize,
#[cfg(feature = "metrics")]
hits: AtomicU64,
#[cfg(feature = "metrics")]
misses: AtomicU64,
#[cfg(feature = "metrics")]
inserts: AtomicU64,
#[cfg(feature = "metrics")]
updates: AtomicU64,
#[cfg(feature = "metrics")]
evictions: AtomicU64,
#[cfg(feature = "metrics")]
removals: AtomicU64,
}
const MIN_SHARD_SIZE: usize = 4096;
const DEFAULT_SHARD_COUNT: usize = 64;
fn compute_shard_count(capacity: usize, desired_shards: usize) -> usize {
let max_shards = (capacity / MIN_SHARD_SIZE).max(1);
desired_shards.min(max_shards).next_power_of_two().max(1)
}
impl Cache<DefaultLifecycle> {
pub fn new(max_size_bytes: usize) -> Self {
let shard_count = compute_shard_count(max_size_bytes, DEFAULT_SHARD_COUNT);
Self::with_shards_and_lifecycle_internal(max_size_bytes, shard_count, DefaultLifecycle)
}
pub fn with_shards(max_size_bytes: usize, shard_count: usize) -> Self {
let shard_count = compute_shard_count(max_size_bytes, shard_count);
Self::with_shards_and_lifecycle_internal(max_size_bytes, shard_count, DefaultLifecycle)
}
}
impl<L: Lifecycle> Cache<L> {
pub fn with_lifecycle(max_size_bytes: usize, lifecycle: L) -> Self {
let shard_count = compute_shard_count(max_size_bytes, DEFAULT_SHARD_COUNT);
Self::with_shards_and_lifecycle_internal(max_size_bytes, shard_count, lifecycle)
}
pub fn with_shards_and_lifecycle(
max_size_bytes: usize,
shard_count: usize,
lifecycle: L,
) -> Self {
let shard_count = compute_shard_count(max_size_bytes, shard_count);
Self::with_shards_and_lifecycle_internal(max_size_bytes, shard_count, lifecycle)
}
fn with_shards_and_lifecycle_internal(
max_size_bytes: usize,
shard_count: usize,
lifecycle: L,
) -> Self {
let size_per_shard = max_size_bytes / shard_count;
let shards = (0..shard_count).map(|_| RwLock::new(Shard::new(size_per_shard))).collect();
Self {
shards,
lifecycle: Arc::new(lifecycle),
current_size: AtomicUsize::new(0),
entry_count: AtomicUsize::new(0),
shard_count,
max_size_bytes,
#[cfg(feature = "metrics")]
hits: AtomicU64::new(0),
#[cfg(feature = "metrics")]
misses: AtomicU64::new(0),
#[cfg(feature = "metrics")]
inserts: AtomicU64::new(0),
#[cfg(feature = "metrics")]
updates: AtomicU64::new(0),
#[cfg(feature = "metrics")]
evictions: AtomicU64::new(0),
#[cfg(feature = "metrics")]
removals: AtomicU64::new(0),
}
}
pub fn insert<K: CacheKey>(&self, key: K, value: K::Value) -> Option<K::Value> {
let erased_key = ErasedKey::new(&key);
let policy = key.policy();
let entry = Entry::new(value, policy);
let entry_size = entry.size;
let shard_lock = self.get_shard(erased_key.hash);
let mut shard = shard_lock.write();
let (old_entry, stats, evicted_entries) = shard.insert(erased_key, entry);
drop(shard);
if let Some(ref old) = old_entry {
let size_diff = entry_size as isize - old.size as isize;
if size_diff > 0 {
self.current_size.fetch_add(size_diff as usize, Ordering::Relaxed);
} else {
self.current_size.fetch_sub((-size_diff) as usize, Ordering::Relaxed);
}
#[cfg(feature = "metrics")]
self.updates.fetch_add(1, Ordering::Relaxed);
} else {
self.current_size.fetch_add(entry_size, Ordering::Relaxed);
self.entry_count.fetch_add(1, Ordering::Relaxed);
#[cfg(feature = "metrics")]
self.inserts.fetch_add(1, Ordering::Relaxed);
}
if stats.count > 0 {
self.entry_count.fetch_sub(stats.count, Ordering::Relaxed);
self.current_size.fetch_sub(stats.size, Ordering::Relaxed);
#[cfg(feature = "metrics")]
self.evictions.fetch_add(stats.count as u64, Ordering::Relaxed);
for evicted in evicted_entries {
self.lifecycle.on_evict(evicted.key.data.as_ref());
}
}
old_entry.and_then(|e| e.into_value::<K::Value>())
}
pub fn get<K: CacheKey>(&self, key: &K) -> Option<Guard<'_, K::Value>> {
let key_ref = ErasedKeyRef::new(key);
let shard_lock = self.get_shard(key_ref.hash);
let shard = shard_lock.read();
let Some(entry) = shard.get_ref(&key_ref) else {
#[cfg(feature = "metrics")]
self.misses.fetch_add(1, Ordering::Relaxed);
return None;
};
#[cfg(feature = "metrics")]
self.hits.fetch_add(1, Ordering::Relaxed);
let value_ref = entry.value_ref::<K::Value>()?;
let value_ptr = value_ref as *const K::Value;
unsafe { Some(Guard::new(shard, value_ptr)) }
}
pub fn get_clone<K: CacheKey>(&self, key: &K) -> Option<K::Value>
where
K::Value: Clone,
{
let key_ref = ErasedKeyRef::new(key);
let shard_lock = self.get_shard(key_ref.hash);
let shard = shard_lock.read();
let Some(entry) = shard.get_ref(&key_ref) else {
#[cfg(feature = "metrics")]
self.misses.fetch_add(1, Ordering::Relaxed);
return None;
};
#[cfg(feature = "metrics")]
self.hits.fetch_add(1, Ordering::Relaxed);
entry.value_ref::<K::Value>().cloned()
}
pub fn get_by<K, Q>(&self, key: &Q) -> Option<Guard<'_, K::Value>>
where
K: CacheKey,
Q: CacheKeyLookup<K> + ?Sized,
{
let key_ref = ErasedKeyLookup::new(key);
let shard_lock = self.get_shard(key_ref.hash);
let shard = shard_lock.read();
let Some(entry) = shard.get_ref_by(&key_ref) else {
#[cfg(feature = "metrics")]
self.misses.fetch_add(1, Ordering::Relaxed);
return None;
};
#[cfg(feature = "metrics")]
self.hits.fetch_add(1, Ordering::Relaxed);
let value_ref = entry.value_ref::<K::Value>()?;
let value_ptr = value_ref as *const K::Value;
unsafe { Some(Guard::new(shard, value_ptr)) }
}
pub fn get_clone_by<K, Q>(&self, key: &Q) -> Option<K::Value>
where
K: CacheKey,
K::Value: Clone,
Q: CacheKeyLookup<K> + ?Sized,
{
let key_ref = ErasedKeyLookup::new(key);
let shard_lock = self.get_shard(key_ref.hash);
let shard = shard_lock.read();
let Some(entry) = shard.get_ref_by(&key_ref) else {
#[cfg(feature = "metrics")]
self.misses.fetch_add(1, Ordering::Relaxed);
return None;
};
#[cfg(feature = "metrics")]
self.hits.fetch_add(1, Ordering::Relaxed);
entry.value_ref::<K::Value>().cloned()
}
pub fn remove<K: CacheKey>(&self, key: &K) -> Option<K::Value> {
let erased_key = ErasedKey::new(key);
let shard_lock = self.get_shard(erased_key.hash);
let mut shard = shard_lock.write();
let (stored_key, entry) = shard.remove(&erased_key)?;
self.current_size.fetch_sub(entry.size, Ordering::Relaxed);
self.entry_count.fetch_sub(1, Ordering::Relaxed);
#[cfg(feature = "metrics")]
self.removals.fetch_add(1, Ordering::Relaxed);
drop(shard);
self.lifecycle.on_remove(stored_key.data.as_ref());
entry.into_value::<K::Value>()
}
pub fn contains<K: CacheKey>(&self, key: &K) -> bool {
let key_ref = ErasedKeyRef::new(key);
let shard_lock = self.get_shard(key_ref.hash);
let shard = shard_lock.read();
shard.get_ref(&key_ref).is_some()
}
pub fn contains_by<K, Q>(&self, key: &Q) -> bool
where
K: CacheKey,
Q: CacheKeyLookup<K> + ?Sized,
{
let key_ref = ErasedKeyLookup::new(key);
let shard_lock = self.get_shard(key_ref.hash);
let shard = shard_lock.read();
shard.get_ref_by(&key_ref).is_some()
}
pub fn size(&self) -> usize {
self.current_size.load(Ordering::Relaxed)
}
pub fn len(&self) -> usize {
self.entry_count.load(Ordering::Relaxed)
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn clear(&self) {
let mut all_entries = Vec::new();
for shard_lock in &self.shards {
let mut shard = shard_lock.write();
all_entries.extend(shard.drain());
}
self.current_size.store(0, Ordering::Relaxed);
self.entry_count.store(0, Ordering::Relaxed);
#[cfg(feature = "metrics")]
{
self.hits.store(0, Ordering::Relaxed);
self.misses.store(0, Ordering::Relaxed);
self.inserts.store(0, Ordering::Relaxed);
self.updates.store(0, Ordering::Relaxed);
self.evictions.store(0, Ordering::Relaxed);
self.removals.store(0, Ordering::Relaxed);
}
for evicted in all_entries {
self.lifecycle.on_clear(evicted.key.data.as_ref());
}
}
#[cfg(feature = "metrics")]
pub fn metrics(&self) -> CacheMetrics {
CacheMetrics {
hits: self.hits.load(Ordering::Relaxed),
misses: self.misses.load(Ordering::Relaxed),
inserts: self.inserts.load(Ordering::Relaxed),
updates: self.updates.load(Ordering::Relaxed),
evictions: self.evictions.load(Ordering::Relaxed),
removals: self.removals.load(Ordering::Relaxed),
current_size_bytes: self.current_size.load(Ordering::Relaxed),
capacity_bytes: self.max_size_bytes,
entry_count: self.entry_count.load(Ordering::Relaxed),
}
}
fn get_shard(&self, hash: u64) -> &RwLock<Shard> {
let index = (hash as usize) & (self.shard_count - 1);
&self.shards[index]
}
}
unsafe impl<L: Lifecycle> Send for Cache<L> {}
unsafe impl<L: Lifecycle> Sync for Cache<L> {}
#[cfg(test)]
mod tests {
use super::*;
use crate::DeepSizeOf;
#[derive(Hash, Eq, PartialEq, Clone, Debug)]
struct TestKey(u64);
impl CacheKey for TestKey {
type Value = TestValue;
}
#[derive(Clone, Debug, PartialEq, DeepSizeOf)]
struct TestValue {
data: String,
}
#[test]
fn test_compute_shard_count_scales_with_capacity() {
assert_eq!(compute_shard_count(1024, 64), 1);
assert_eq!(compute_shard_count(4095, 64), 1);
assert_eq!(compute_shard_count(4096, 64), 1);
assert_eq!(compute_shard_count(8192, 64), 2);
assert_eq!(compute_shard_count(65536, 64), 16);
assert_eq!(compute_shard_count(256 * 1024, 64), 64);
assert_eq!(compute_shard_count(1024 * 1024, 64), 64);
assert_eq!(compute_shard_count(8192, 128), 2); assert_eq!(compute_shard_count(1024 * 1024, 128), 128); }
#[test]
fn test_cache_insert_and_get() {
let cache = Cache::new(1024);
let key = TestKey(1);
let value = TestValue {
data: "hello".to_string(),
};
cache.insert(key.clone(), value.clone());
let retrieved = cache.get_clone(&key).expect("key should exist");
assert_eq!(retrieved, value);
}
#[test]
fn test_cache_remove() {
let cache = Cache::new(1024);
let key = TestKey(1);
let value = TestValue {
data: "hello".to_string(),
};
cache.insert(key.clone(), value.clone());
assert!(cache.contains(&key));
let removed = cache.remove(&key).expect("key should exist");
assert_eq!(removed, value);
assert!(!cache.contains(&key));
}
#[test]
fn test_cache_eviction() {
let cache = Cache::with_shards(1000, 4);
for i in 0..15 {
let key = TestKey(i);
let value = TestValue {
data: "x".repeat(50),
};
cache.insert(key, value);
}
assert!(cache.len() < 15, "Cache should have evicted some entries");
assert!(cache.size() <= 1000, "Cache size should be <= 1000, got {}", cache.size());
}
#[test]
fn test_cache_concurrent_access() {
use std::sync::Arc;
use std::thread;
let cache = Arc::new(Cache::new(10240));
let mut handles = vec![];
for t in 0..4 {
let cache = cache.clone();
handles.push(thread::spawn(move || {
for i in 0..100 {
let key = TestKey(t * 100 + i);
let value = TestValue {
data: format!("value-{}", i),
};
cache.insert(key.clone(), value.clone());
if let Some(retrieved) = cache.get_clone(&key) {
assert_eq!(retrieved, value);
}
}
}));
}
for handle in handles {
handle.join().expect("thread should not panic");
}
assert!(!cache.is_empty());
}
#[test]
fn test_cache_is_send_sync() {
fn assert_send<T: Send>() {}
fn assert_sync<T: Sync>() {}
assert_send::<Cache>();
assert_sync::<Cache>();
}
#[derive(Hash, Eq, PartialEq, Clone, Debug)]
struct DbCacheKey(String, String);
impl CacheKey for DbCacheKey {
type Value = TestValue;
}
struct DbCacheKeyRef<'a>(&'a str, &'a str);
impl std::hash::Hash for DbCacheKeyRef<'_> {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.0.hash(state);
self.1.hash(state);
}
}
impl CacheKeyLookup<DbCacheKey> for DbCacheKeyRef<'_> {
fn eq_key(&self, key: &DbCacheKey) -> bool {
self.0 == key.0 && self.1 == key.1
}
fn to_owned_key(self) -> DbCacheKey {
DbCacheKey(self.0.to_owned(), self.1.to_owned())
}
}
#[test]
fn test_borrowed_key_lookup_get_by() {
let cache = Cache::new(1024);
let key = DbCacheKey("namespace".to_string(), "database".to_string());
let value = TestValue {
data: "test_data".to_string(),
};
cache.insert(key.clone(), value.clone());
let borrowed_key = DbCacheKeyRef("namespace", "database");
let retrieved = cache.get_by::<DbCacheKey, _>(&borrowed_key);
assert!(retrieved.is_some());
assert_eq!(*retrieved.unwrap(), value);
let borrowed_key_missing = DbCacheKeyRef("namespace", "missing");
let retrieved = cache.get_by::<DbCacheKey, _>(&borrowed_key_missing);
assert!(retrieved.is_none());
}
#[test]
fn test_borrowed_key_lookup_get_clone_by() {
let cache = Cache::new(1024);
let key = DbCacheKey("ns".to_string(), "db".to_string());
let value = TestValue {
data: "cloned_data".to_string(),
};
cache.insert(key.clone(), value.clone());
let borrowed_key = DbCacheKeyRef("ns", "db");
let retrieved = cache.get_clone_by::<DbCacheKey, _>(&borrowed_key);
assert_eq!(retrieved, Some(value));
let borrowed_key_missing = DbCacheKeyRef("ns", "missing");
let retrieved = cache.get_clone_by::<DbCacheKey, _>(&borrowed_key_missing);
assert_eq!(retrieved, None);
}
#[test]
fn test_borrowed_key_lookup_contains_by() {
let cache = Cache::new(1024);
let key = DbCacheKey("catalog".to_string(), "schema".to_string());
let value = TestValue {
data: "contains_test".to_string(),
};
cache.insert(key.clone(), value);
let borrowed_key = DbCacheKeyRef("catalog", "schema");
assert!(cache.contains_by::<DbCacheKey, _>(&borrowed_key));
let borrowed_key_missing = DbCacheKeyRef("catalog", "missing");
assert!(!cache.contains_by::<DbCacheKey, _>(&borrowed_key_missing));
}
#[test]
fn test_borrowed_key_lookup_multiple_entries() {
let cache = Cache::new(4096);
for i in 0..10 {
let key = DbCacheKey(format!("ns{}", i), format!("db{}", i));
let value = TestValue {
data: format!("data{}", i),
};
cache.insert(key, value);
}
for i in 0..10 {
let ns = format!("ns{}", i);
let db = format!("db{}", i);
let borrowed_key = DbCacheKeyRef(&ns, &db);
let retrieved = cache.get_clone_by::<DbCacheKey, _>(&borrowed_key);
assert!(retrieved.is_some());
assert_eq!(retrieved.unwrap().data, format!("data{}", i));
}
}
#[test]
fn test_borrowed_key_existing_api_still_works() {
let cache = Cache::new(1024);
let key = DbCacheKey("test".to_string(), "key".to_string());
let value = TestValue {
data: "existing_api".to_string(),
};
cache.insert(key.clone(), value.clone());
let retrieved = cache.get(&key);
assert!(retrieved.is_some());
assert_eq!(*retrieved.unwrap(), value);
assert!(cache.contains(&key));
let cloned = cache.get_clone(&key);
assert_eq!(cloned, Some(value));
}
use std::any::Any;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use crate::CacheBuilder;
use crate::lifecycle::Lifecycle;
struct CountingLifecycle {
evict_count: Arc<AtomicUsize>,
remove_count: Arc<AtomicUsize>,
clear_count: Arc<AtomicUsize>,
}
impl CountingLifecycle {
fn new() -> (Self, Arc<AtomicUsize>, Arc<AtomicUsize>, Arc<AtomicUsize>) {
let evict_count = Arc::new(AtomicUsize::new(0));
let remove_count = Arc::new(AtomicUsize::new(0));
let clear_count = Arc::new(AtomicUsize::new(0));
(
Self {
evict_count: evict_count.clone(),
remove_count: remove_count.clone(),
clear_count: clear_count.clone(),
},
evict_count,
remove_count,
clear_count,
)
}
}
impl Lifecycle for CountingLifecycle {
fn on_evict(&self, _key: &dyn Any) {
self.evict_count.fetch_add(1, Ordering::Relaxed);
}
fn on_remove(&self, _key: &dyn Any) {
self.remove_count.fetch_add(1, Ordering::Relaxed);
}
fn on_clear(&self, _key: &dyn Any) {
self.clear_count.fetch_add(1, Ordering::Relaxed);
}
}
#[test]
fn test_lifecycle_on_evict() {
let (lifecycle, evict_count, _, _) = CountingLifecycle::new();
let cache = CacheBuilder::new(500).shards(1).lifecycle(lifecycle).build();
for i in 0..20 {
let key = TestKey(i);
let value = TestValue {
data: "x".repeat(50),
};
cache.insert(key, value);
}
assert!(evict_count.load(Ordering::Relaxed) > 0, "Expected evictions but got none");
}
#[test]
fn test_lifecycle_on_clear() {
let (lifecycle, _, _, clear_count) = CountingLifecycle::new();
let cache = CacheBuilder::new(4096).lifecycle(lifecycle).build();
for i in 0..5 {
let key = TestKey(i);
let value = TestValue {
data: format!("value{}", i),
};
cache.insert(key, value);
}
assert_eq!(clear_count.load(Ordering::Relaxed), 0);
cache.clear();
assert_eq!(clear_count.load(Ordering::Relaxed), 5, "Expected 5 clear callbacks");
}
#[test]
fn test_lifecycle_on_remove() {
let (lifecycle, _, remove_count, _) = CountingLifecycle::new();
let cache = CacheBuilder::new(4096).lifecycle(lifecycle).build();
let key = TestKey(1);
let value = TestValue {
data: "test".to_string(),
};
cache.insert(key.clone(), value);
assert_eq!(remove_count.load(Ordering::Relaxed), 0);
let removed = cache.remove(&key);
assert!(removed.is_some());
assert_eq!(remove_count.load(Ordering::Relaxed), 1, "Expected 1 remove callback");
}
#[test]
fn test_lifecycle_typed_downcast() {
use crate::TypedLifecycle;
let evicted_keys = Arc::new(std::sync::Mutex::new(Vec::new()));
let keys_clone = evicted_keys.clone();
let lifecycle = TypedLifecycle::<TestKey, _>::new(move |key| {
keys_clone.lock().unwrap().push(key.0);
});
let cache = CacheBuilder::new(500).shards(1).lifecycle(lifecycle).build();
for i in 0..20 {
let key = TestKey(i);
let value = TestValue {
data: "x".repeat(50),
};
cache.insert(key, value);
}
let keys = evicted_keys.lock().unwrap();
assert!(!keys.is_empty(), "Expected some evicted keys to be captured");
}
#[test]
fn test_cache_with_lifecycle_is_send_sync() {
fn assert_send<T: Send>() {}
fn assert_sync<T: Sync>() {}
assert_send::<Cache>();
assert_sync::<Cache>();
assert_send::<Cache<CountingLifecycle>>();
assert_sync::<Cache<CountingLifecycle>>();
}
}