use crate::kv::*;
use std::array;
use std::fmt;
use std::sync::RwLock;
use std::sync::atomic::{AtomicBool, Ordering};
pub trait CacheKeyHash {
fn shard_index(&self, num_shards: usize) -> usize;
}
macro_rules! impl_cache_hash_int {
($($t:ty),*) => {
$(
impl CacheKeyHash for $t {
#[inline(always)]
fn shard_index(&self, num_shards: usize) -> usize {
if num_shards <= 1 {
return 0;
}
let shift = 64 - num_shards.trailing_zeros();
((*self as u64).wrapping_mul(11400714819323198485) >> shift) as usize
}
}
)*
};
}
impl_cache_hash_int!(usize, isize, u64, i64, u32, i32, u16, i16, u8, i8);
macro_rules! impl_cache_hash_float {
($($t:ty),*) => {
$(impl CacheKeyHash for $t {
#[inline(always)]
fn shard_index(&self, num_shards: usize) -> usize {
if num_shards <= 1 {
return 0;
}
let bits = if *self == 0.0 {
0u64
} else {
self.to_bits() as u64
};
let shift = 64 - num_shards.trailing_zeros();
(bits.wrapping_mul(11400714819323198485) >> shift) as usize
}
})*
};
}
impl_cache_hash_float!(f32, f64);
struct Shard<K, V, const CAP: usize>
where
K: CacheKey,
V: CacheValue,
{
hand: usize,
len: usize,
keys: [K; CAP],
values: [V; CAP],
refs: [AtomicBool; CAP],
}
impl<K, V, const CAP: usize> Shard<K, V, CAP>
where
K: CacheKey,
V: CacheValue,
{
fn new() -> Self {
const {
assert!(CAP.is_power_of_two(), "CAP must be a power of two");
}
Self {
hand: 0,
len: 0,
keys: [K::GUARD; CAP],
values: [V::default(); CAP],
refs: array::from_fn(|_| AtomicBool::new(false)),
}
}
#[inline(always)]
fn get(&self, key: &K) -> Option<V> {
let index = K::find_key(&self.keys, key)?;
if index >= self.len {
return None;
}
if !self.refs[index].load(Ordering::Relaxed) {
self.refs[index].store(true, Ordering::Relaxed);
}
Some(self.values[index])
}
#[inline(always)]
fn insert_unchecked(&mut self, key: &K, value: V) {
loop {
let i = self.hand;
if self.refs[i].load(Ordering::Relaxed) {
self.refs[i].store(false, Ordering::Relaxed);
} else {
self.keys[i] = *key;
self.values[i] = value;
self.hand = (self.hand + 1) & (CAP - 1);
if self.len < CAP {
self.len += 1;
}
return;
}
self.hand = (self.hand + 1) & (CAP - 1);
}
}
}
pub trait ConcurrentCache<K, V>: Send + Sync
where
K: CacheKey,
V: CacheValue,
{
fn get(&self, key: &K) -> V;
}
#[cfg_attr(target_arch = "aarch64", repr(align(128)))]
#[cfg_attr(not(target_arch = "aarch64"), repr(align(64)))]
struct AlignedShard<K, V, const SHARD_CAP: usize>
where
K: CacheKey,
V: CacheValue,
{
lock: RwLock<Shard<K, V, SHARD_CAP>>,
}
pub struct ConcurrentClockCache<K, V, F, const SHARD_CAP: usize, const SHARDS: usize>
where
K: CacheKey + CacheKeyHash + Send + Sync,
V: CacheValue + Send + Sync,
F: Fn(&K) -> V + Send + Sync,
{
fetcher: F,
shards: [AlignedShard<K, V, SHARD_CAP>; SHARDS],
}
impl<K, V, F, const SHARD_CAP: usize, const SHARDS: usize>
ConcurrentClockCache<K, V, F, SHARD_CAP, SHARDS>
where
K: CacheKey + CacheKeyHash + Send + Sync,
V: CacheValue + Send + Sync,
F: Fn(&K) -> V + Send + Sync,
{
pub fn new(fetcher: F) -> Self {
const {
assert!(SHARDS.is_power_of_two(), "SHARDS must be a power of two");
assert!(
SHARD_CAP.is_power_of_two(),
"SHARD_CAP must be a power of two"
);
}
let shards = array::from_fn(|_| AlignedShard {
lock: RwLock::new(Shard::new()),
});
Self { fetcher, shards }
}
#[inline]
pub fn len(&self) -> usize {
self.shards.iter().map(|s| s.lock.read().unwrap().len).sum()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline]
pub const fn capacity(&self) -> usize {
SHARD_CAP * SHARDS
}
#[inline]
pub const fn num_shards(&self) -> usize {
SHARDS
}
#[inline]
pub const fn shard_capacity(&self) -> usize {
SHARD_CAP
}
#[inline]
pub fn get(&self, key: &K) -> V {
let shard_index = key.shard_index(SHARDS);
let shard = &self.shards[shard_index];
{
let store = shard.lock.read().unwrap();
if let Some(value) = store.get(key) {
return value;
}
}
let mut store = shard.lock.write().unwrap();
if let Some(val) = store.get(key) {
return val;
}
let value = (self.fetcher)(key);
store.insert_unchecked(key, value);
value
}
}
impl<K, V, F, const SHARD_CAP: usize, const SHARDS: usize> ConcurrentCache<K, V>
for ConcurrentClockCache<K, V, F, SHARD_CAP, SHARDS>
where
K: CacheKey + CacheKeyHash + Send + Sync,
V: CacheValue + Send + Sync,
F: Fn(&K) -> V + Send + Sync,
{
#[inline]
fn get(&self, key: &K) -> V {
ConcurrentClockCache::get(self, key)
}
}
impl<K, V, F, const SHARD_CAP: usize, const SHARDS: usize> fmt::Debug
for ConcurrentClockCache<K, V, F, SHARD_CAP, SHARDS>
where
K: CacheKey + CacheKeyHash + Send + Sync,
V: CacheValue + Send + Sync,
F: Fn(&K) -> V + Send + Sync,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ConcurrentClockCache")
.field("len", &self.len())
.field("capacity", &self.capacity())
.field("shards", &SHARDS)
.field("shard_cap", &SHARD_CAP)
.finish()
}
}
#[cfg(test)]
mod multi_threaded_tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::thread;
fn mock_fetcher(key: &usize) -> usize {
key.wrapping_mul(2654435761)
}
#[test]
fn concurrent_access() {
const SHARD_CAP: usize = 64;
const SHARDS: usize = 16;
let cache =
Arc::new(ConcurrentClockCache::<usize, usize, _, SHARD_CAP, SHARDS>::new(mock_fetcher));
let mut handles = vec![];
let n_threads = thread::available_parallelism().unwrap().get();
for t in 0..n_threads {
let cache_ref = Arc::clone(&cache);
handles.push(thread::spawn(move || {
for i in 0..100_000 {
let key = (i + t * 10) % 500;
let val = cache_ref.get(&key);
assert_eq!(val, mock_fetcher(&key));
}
}));
}
for handle in handles {
handle.join().unwrap();
}
}
#[test]
fn fetcher_deduplication_and_correctness() {
const SHARD_CAP: usize = 32;
const SHARDS: usize = 4;
let fetch_count = Arc::new(AtomicUsize::new(0));
let fetch_count_clone = Arc::clone(&fetch_count);
let cache = Arc::new(
ConcurrentClockCache::<usize, usize, _, SHARD_CAP, SHARDS>::new(move |key: &usize| {
fetch_count_clone.fetch_add(1, Ordering::Relaxed);
key * 10
}),
);
let mut handles = vec![];
for _ in 0..10 {
let cache_ref = Arc::clone(&cache);
handles.push(thread::spawn(move || {
for _ in 0..1_000 {
for key in 0..10 {
let val = cache_ref.get(&key);
assert_eq!(val, key * 10);
}
}
}));
}
for handle in handles {
handle.join().unwrap();
}
assert_eq!(fetch_count.load(Ordering::Relaxed), 10);
}
#[test]
fn fetcher_deduplication_with_barrier() {
use std::sync::Barrier;
const SHARD_CAP: usize = 32;
const SHARDS: usize = 4;
const NUM_THREADS: usize = 10;
let fetch_count = Arc::new(AtomicUsize::new(0));
let fetch_count_clone = Arc::clone(&fetch_count);
let cache = Arc::new(
ConcurrentClockCache::<usize, usize, _, SHARD_CAP, SHARDS>::new(move |key: &usize| {
fetch_count_clone.fetch_add(1, Ordering::Relaxed);
key * 10
}),
);
let barrier = Arc::new(Barrier::new(NUM_THREADS));
let mut handles = vec![];
for _ in 0..NUM_THREADS {
let cache_ref = Arc::clone(&cache);
let b = Arc::clone(&barrier);
handles.push(thread::spawn(move || {
b.wait();
for key in 0..10 {
let val = cache_ref.get(&key);
assert_eq!(val, key * 10);
}
}));
}
for handle in handles {
handle.join().unwrap();
}
assert_eq!(fetch_count.load(Ordering::Relaxed), 10);
}
#[test]
fn guard_key() {
const SHARD_CAP: usize = 16;
const SHARDS: usize = 4;
let cache = Arc::new(
ConcurrentClockCache::<usize, usize, _, SHARD_CAP, SHARDS>::new(|key: &usize| {
key.wrapping_add(100)
}),
);
let mut handles = vec![];
for _ in 0..8 {
let cache_ref = Arc::clone(&cache);
handles.push(thread::spawn(move || {
for _ in 0..1_000 {
let val = cache_ref.get(&usize::MAX);
assert_eq!(val, usize::MAX.wrapping_add(100));
}
}));
}
for handle in handles {
handle.join().unwrap();
}
}
#[test]
fn shard_distribution() {
const SHARDS: usize = 16;
let mut shard_counts = [0usize; SHARDS];
for i in 0..256 {
let key = i * 16;
let shard = key.shard_index(SHARDS);
assert!(shard < SHARDS);
shard_counts[shard] += 1;
}
let non_empty_shards = shard_counts.iter().filter(|&&c| c > 0).count();
assert!(
non_empty_shards > 1,
"Multiples of 16 must not all collapse into a single shard! Got distribution: {:?}",
shard_counts
);
assert_eq!(non_empty_shards, SHARDS);
}
#[test]
fn float_zeros_mapping() {
const SHARDS: usize = 16;
const SHARD_CAP: usize = 16;
let pos_zero: f32 = 0.0;
let neg_zero: f32 = -0.0;
assert_eq!(pos_zero.shard_index(SHARDS), neg_zero.shard_index(SHARDS));
let fetch_count = Arc::new(AtomicUsize::new(0));
let fetch_count_clone = Arc::clone(&fetch_count);
let cache =
ConcurrentClockCache::<f32, u32, _, SHARD_CAP, SHARDS>::new(move |_key: &f32| {
fetch_count_clone.fetch_add(1, Ordering::Relaxed);
42
});
assert_eq!(cache.get(&pos_zero), 42);
assert_eq!(fetch_count.load(Ordering::Relaxed), 1);
assert_eq!(cache.get(&neg_zero), 42);
assert_eq!(fetch_count.load(Ordering::Relaxed), 1);
}
#[test]
fn cache_inherent_methods() {
const SHARD_CAP: usize = 8;
const SHARDS: usize = 4;
let cache = ConcurrentClockCache::<usize, usize, _, SHARD_CAP, SHARDS>::new(|k| *k);
assert!(cache.is_empty());
assert_eq!(cache.len(), 0);
assert_eq!(cache.capacity(), 32);
assert_eq!(cache.num_shards(), 4);
assert_eq!(cache.shard_capacity(), 8);
cache.get(&1);
assert!(!cache.is_empty());
assert_eq!(cache.len(), 1);
let debug_str = format!("{:?}", cache);
assert!(debug_str.contains("ConcurrentClockCache"));
assert!(debug_str.contains("len: 1"));
}
#[test]
fn signed_negative_keys() {
const SHARD_CAP: usize = 8;
const SHARDS: usize = 4;
let cache =
ConcurrentClockCache::<i32, i32, _, SHARD_CAP, SHARDS>::new(|k| k.wrapping_mul(2));
assert_eq!(cache.get(&-10), -20);
assert_eq!(cache.get(&-10), -20);
assert_eq!(cache.get(&i32::MAX), i32::MAX.wrapping_mul(2));
assert_eq!(cache.get(&i32::MIN), i32::MIN.wrapping_mul(2));
let cache_i8 = ConcurrentClockCache::<i8, i8, _, SHARD_CAP, SHARDS>::new(|k| *k);
assert_eq!(cache_i8.get(&-1), -1);
assert_eq!(cache_i8.get(&i8::MAX), i8::MAX);
assert_eq!(cache_i8.get(&i8::MIN), i8::MIN);
}
}
#[cfg(test)]
mod multi_threaded_load_tests {
use super::*;
use std::hint::black_box;
use std::sync::Arc;
use std::thread;
use std::time::Instant;
fn mock_fetcher(key: &usize) -> usize {
key.wrapping_mul(2654435761)
}
const NUM_THREADS: usize = 8;
#[test]
fn high_hit_rate() {
const SHARDS: usize = 16;
const SHARD_CAP: usize = 16;
const TOTAL_OPERATIONS: usize = 5_000_000;
const OPS_PER_THREAD: usize = TOTAL_OPERATIONS / NUM_THREADS;
let cache =
Arc::new(ConcurrentClockCache::<usize, usize, _, SHARD_CAP, SHARDS>::new(mock_fetcher));
let thread_keys: Vec<Vec<usize>> = (0..NUM_THREADS)
.map(|t| {
let mut state: u64 = 12345 + t as u64;
(0..OPS_PER_THREAD)
.map(|_| {
state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
let rand_val = (state >> 33) as usize;
if rand_val % 100 < 80 {
rand_val % 50 } else {
50 + (rand_val % 1950) }
})
.collect()
})
.collect();
println!(
"Running Multithreaded High Hit Rate Test ({} threads, {} total ops)...",
NUM_THREADS, TOTAL_OPERATIONS
);
let start = Instant::now();
let handles: Vec<_> = thread_keys
.into_iter()
.map(|keys| {
let cache_ref = Arc::clone(&cache);
thread::spawn(move || {
for key in keys {
black_box(cache_ref.get(&key));
}
})
})
.collect();
for handle in handles {
handle.join().unwrap();
}
let elapsed = start.elapsed();
let total_ops = OPS_PER_THREAD * NUM_THREADS;
println!(
"High Hit Rate Elapsed: {:?} ({:.2} ns/op, {:.2} Mops/sec)",
elapsed,
elapsed.as_nanos() as f64 / total_ops as f64,
(total_ops as f64 / 1_000_000.0) / elapsed.as_secs_f64()
);
}
#[test]
fn thrashing() {
const SHARDS: usize = 16;
const SHARD_CAP: usize = 32; const TOTAL_OPERATIONS: usize = 5_000_000;
const OPS_PER_THREAD: usize = TOTAL_OPERATIONS / NUM_THREADS;
let cache =
Arc::new(ConcurrentClockCache::<usize, usize, _, SHARD_CAP, SHARDS>::new(mock_fetcher));
let thread_keys: Vec<Vec<usize>> = (0..NUM_THREADS)
.map(|t| (0..OPS_PER_THREAD).map(|i| (i + t * 100) % 1000).collect())
.collect();
println!(
"Running Multithreaded Cache Thrashing Test ({} threads, {} total ops)...",
NUM_THREADS, TOTAL_OPERATIONS
);
let start = Instant::now();
let handles: Vec<_> = thread_keys
.into_iter()
.map(|keys| {
let cache_ref = Arc::clone(&cache);
thread::spawn(move || {
for key in keys {
black_box(cache_ref.get(&key));
}
})
})
.collect();
for handle in handles {
handle.join().unwrap();
}
let elapsed = start.elapsed();
let total_ops = OPS_PER_THREAD * NUM_THREADS;
println!(
"Thrashing Elapsed: {:?} ({:.2} ns/op, {:.2} Mops/sec)",
elapsed,
elapsed.as_nanos() as f64 / total_ops as f64,
(total_ops as f64 / 1_000_000.0) / elapsed.as_secs_f64()
);
}
#[test]
fn large_linear_scan_overhead() {
const SHARDS: usize = 16;
const SHARD_CAP: usize = 256; const TOTAL_OPERATIONS: usize = 1_000_000;
const OPS_PER_THREAD: usize = TOTAL_OPERATIONS / NUM_THREADS;
let cache =
Arc::new(ConcurrentClockCache::<usize, usize, _, SHARD_CAP, SHARDS>::new(mock_fetcher));
let thread_keys: Vec<Vec<usize>> = (0..NUM_THREADS)
.map(|t| {
let mut state: u64 = 54321 + t as u64;
(0..OPS_PER_THREAD)
.map(|_| {
state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
((state >> 33) as usize) % (SHARD_CAP * SHARDS * 2)
})
.collect()
})
.collect();
println!(
"Running Multithreaded Large Cache Test ({} threads, {} total ops)...",
NUM_THREADS, TOTAL_OPERATIONS
);
let start = Instant::now();
let handles: Vec<_> = thread_keys
.into_iter()
.map(|keys| {
let cache_ref = Arc::clone(&cache);
thread::spawn(move || {
for key in keys {
black_box(cache_ref.get(&key));
}
})
})
.collect();
for handle in handles {
handle.join().unwrap();
}
let elapsed = start.elapsed();
let total_ops = OPS_PER_THREAD * NUM_THREADS;
println!(
"Large Cache Elapsed: {:?} ({:.2} ns/op, {:.2} Mops/sec)",
elapsed,
elapsed.as_nanos() as f64 / total_ops as f64,
(total_ops as f64 / 1_000_000.0) / elapsed.as_secs_f64()
);
}
}