use std::cell::{Cell, RefCell};
use std::ops::{Deref, DerefMut};
use std::rc::Rc;
use std::sync::atomic::{AtomicU64, Ordering};
pub const CACHE_LINE_SIZE: usize = 64;
pub const L1_CACHE_SIZE: usize = 32 * 1024;
pub const L2_CACHE_SIZE: usize = 256 * 1024;
pub const L3_CACHE_SIZE_PER_CORE: usize = 2 * 1024 * 1024;
#[repr(C, align(64))]
#[derive(Debug)]
pub struct CacheAligned<T> {
value: T,
}
impl<T> CacheAligned<T> {
#[inline]
pub const fn new(value: T) -> Self {
Self { value }
}
#[inline]
pub fn into_inner(self) -> T {
self.value
}
#[inline]
pub fn get(&self) -> &T {
&self.value
}
#[inline]
pub fn get_mut(&mut self) -> &mut T {
&mut self.value
}
}
impl<T: Default> Default for CacheAligned<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T: Clone> Clone for CacheAligned<T> {
fn clone(&self) -> Self {
Self::new(self.value.clone())
}
}
impl<T> Deref for CacheAligned<T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target {
&self.value
}
}
impl<T> DerefMut for CacheAligned<T> {
#[inline]
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.value
}
}
#[repr(C, align(64))]
#[derive(Debug)]
pub struct Padded<T> {
value: T,
}
impl<T> Padded<T> {
#[inline]
pub const fn new(value: T) -> Self {
Self { value }
}
#[inline]
pub fn into_inner(self) -> T {
self.value
}
}
impl<T: Default> Default for Padded<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T: Clone> Clone for Padded<T> {
fn clone(&self) -> Self {
Self::new(self.value.clone())
}
}
impl<T> Deref for Padded<T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target {
&self.value
}
}
impl<T> DerefMut for Padded<T> {
#[inline]
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.value
}
}
#[derive(Debug)]
pub struct HotCold<H, C> {
pub hot: H,
cold: Box<C>,
}
impl<H, C> HotCold<H, C> {
pub fn new(hot: H, cold: C) -> Self {
Self {
hot,
cold: Box::new(cold),
}
}
#[inline(always)]
pub fn hot(&self) -> &H {
&self.hot
}
#[inline(always)]
pub fn hot_mut(&mut self) -> &mut H {
&mut self.hot
}
#[inline]
pub fn cold(&self) -> &C {
&self.cold
}
#[inline]
pub fn cold_mut(&mut self) -> &mut C {
&mut self.cold
}
pub fn into_parts(self) -> (H, C) {
(self.hot, *self.cold)
}
}
impl<H: Clone, C: Clone> Clone for HotCold<H, C> {
fn clone(&self) -> Self {
Self::new(self.hot.clone(), (*self.cold).clone())
}
}
pub struct LocalStateCache<T> {
value: RefCell<Option<Rc<T>>>,
version: Cell<u64>,
hits: Cell<u64>,
misses: Cell<u64>,
}
impl<T> LocalStateCache<T> {
pub const fn new() -> Self {
Self {
value: RefCell::new(None),
version: Cell::new(0),
hits: Cell::new(0),
misses: Cell::new(0),
}
}
#[inline]
pub fn get_or_refresh<F>(&self, current_version: u64, refresh: F) -> Rc<T>
where
F: FnOnce() -> T,
{
if self.version.get() == current_version
&& let Some(ref value) = *self.value.borrow()
{
self.hits.set(self.hits.get() + 1);
LOCALITY_STATS.record_cache_hit();
return Rc::clone(value);
}
self.misses.set(self.misses.get() + 1);
LOCALITY_STATS.record_cache_miss();
let new_value = Rc::new(refresh());
*self.value.borrow_mut() = Some(Rc::clone(&new_value));
self.version.set(current_version);
new_value
}
#[inline]
pub fn invalidate(&self) {
*self.value.borrow_mut() = None;
self.version.set(0);
}
pub fn stats(&self) -> LocalCacheStats {
LocalCacheStats {
hits: self.hits.get(),
misses: self.misses.get(),
}
}
}
impl<T> Default for LocalStateCache<T> {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Copy)]
pub struct LocalCacheStats {
pub hits: u64,
pub misses: u64,
}
impl LocalCacheStats {
pub fn hit_ratio(&self) -> f64 {
let total = self.hits + self.misses;
if total > 0 {
self.hits as f64 / total as f64
} else {
0.0
}
}
}
#[repr(C, align(64))]
pub struct CompactState<const N: usize> {
values: [AtomicU64; N],
}
impl<const N: usize> CompactState<N> {
pub fn new() -> Self {
const {
assert!(
std::mem::size_of::<[AtomicU64; N]>() <= CACHE_LINE_SIZE,
"CompactState exceeds cache line size"
);
}
Self {
values: std::array::from_fn(|_| AtomicU64::new(0)),
}
}
#[inline]
pub fn get(&self, index: usize) -> u64 {
self.values[index].load(Ordering::Relaxed)
}
#[inline]
pub fn set(&self, index: usize, value: u64) {
self.values[index].store(value, Ordering::Relaxed);
}
#[inline]
pub fn increment(&self, index: usize) -> u64 {
self.values[index].fetch_add(1, Ordering::Relaxed)
}
#[inline]
pub fn add(&self, index: usize, delta: u64) -> u64 {
self.values[index].fetch_add(delta, Ordering::Relaxed)
}
pub fn all(&self) -> [u64; N] {
std::array::from_fn(|i| self.get(i))
}
}
impl<const N: usize> Default for CompactState<N> {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PrefetchLevel {
L1,
L2,
L3,
NonTemporal,
}
#[inline]
pub fn prefetch<T>(ptr: &T, level: PrefetchLevel) {
let addr = ptr as *const T as *const u8;
prefetch_ptr(addr, level);
}
#[inline]
pub fn prefetch_ptr(ptr: *const u8, level: PrefetchLevel) {
#[cfg(target_arch = "x86_64")]
{
use std::arch::x86_64::*;
unsafe {
match level {
PrefetchLevel::L1 => _mm_prefetch(ptr as *const i8, _MM_HINT_T0),
PrefetchLevel::L2 => _mm_prefetch(ptr as *const i8, _MM_HINT_T1),
PrefetchLevel::L3 => _mm_prefetch(ptr as *const i8, _MM_HINT_T2),
PrefetchLevel::NonTemporal => _mm_prefetch(ptr as *const i8, _MM_HINT_NTA),
}
}
}
#[cfg(target_arch = "aarch64")]
{
unsafe {
match level {
PrefetchLevel::L1 => {
std::arch::asm!("prfm pldl1keep, [{x}]", x = in(reg) ptr, options(readonly, nostack));
}
PrefetchLevel::L2 => {
std::arch::asm!("prfm pldl2keep, [{x}]", x = in(reg) ptr, options(readonly, nostack));
}
PrefetchLevel::L3 => {
std::arch::asm!("prfm pldl3keep, [{x}]", x = in(reg) ptr, options(readonly, nostack));
}
PrefetchLevel::NonTemporal => {
std::arch::asm!("prfm pldl1strm, [{x}]", x = in(reg) ptr, options(readonly, nostack));
}
}
}
}
#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
{
let _ = (ptr, level);
}
}
#[inline]
pub fn prefetch_range<T>(slice: &[T], level: PrefetchLevel) {
let ptr = slice.as_ptr() as *const u8;
let len = std::mem::size_of_val(slice);
for offset in (0..len).step_by(CACHE_LINE_SIZE) {
prefetch_ptr(unsafe { ptr.add(offset) }, level);
}
}
pub struct SoaStorage<const FIELDS: usize, const CAPACITY: usize> {
data: [Box<[f64; CAPACITY]>; FIELDS],
len: usize,
}
impl<const FIELDS: usize, const CAPACITY: usize> SoaStorage<FIELDS, CAPACITY> {
pub fn new() -> Self {
Self {
data: std::array::from_fn(|_| {
vec![0.0f64; CAPACITY]
.into_boxed_slice()
.try_into()
.expect("boxed slice length matches CAPACITY")
}),
len: 0,
}
}
#[inline]
pub fn len(&self) -> usize {
self.len
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len == 0
}
#[inline]
pub const fn capacity(&self) -> usize {
CAPACITY
}
#[inline]
pub fn get_field(&self, field: usize, index: usize) -> f64 {
assert!(field < FIELDS && index < self.len);
self.data[field][index]
}
#[inline]
pub fn set_field(&mut self, field: usize, index: usize, value: f64) {
assert!(field < FIELDS && index < self.len);
self.data[field][index] = value;
}
pub fn push(&mut self) -> Option<usize> {
if self.len < CAPACITY {
let index = self.len;
self.len += 1;
Some(index)
} else {
None
}
}
pub fn field_slice(&self, field: usize) -> &[f64] {
assert!(field < FIELDS);
&self.data[field][..self.len]
}
#[inline]
pub fn prefetch_field(&self, field: usize, level: PrefetchLevel) {
assert!(field < FIELDS);
prefetch_range(self.field_slice(field), level);
}
}
impl<const FIELDS: usize, const CAPACITY: usize> Default for SoaStorage<FIELDS, CAPACITY> {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Default)]
pub struct LocalityStats {
cache_hits: AtomicU64,
cache_misses: AtomicU64,
prefetches: AtomicU64,
}
impl LocalityStats {
fn record_cache_hit(&self) {
self.cache_hits.fetch_add(1, Ordering::Relaxed);
}
fn record_cache_miss(&self) {
self.cache_misses.fetch_add(1, Ordering::Relaxed);
}
#[allow(dead_code)]
fn record_prefetch(&self) {
self.prefetches.fetch_add(1, Ordering::Relaxed);
}
pub fn cache_hits(&self) -> u64 {
self.cache_hits.load(Ordering::Relaxed)
}
pub fn cache_misses(&self) -> u64 {
self.cache_misses.load(Ordering::Relaxed)
}
pub fn prefetches(&self) -> u64 {
self.prefetches.load(Ordering::Relaxed)
}
pub fn hit_ratio(&self) -> f64 {
let hits = self.cache_hits() as f64;
let total = hits + self.cache_misses() as f64;
if total > 0.0 { hits / total } else { 0.0 }
}
}
static LOCALITY_STATS: LocalityStats = LocalityStats {
cache_hits: AtomicU64::new(0),
cache_misses: AtomicU64::new(0),
prefetches: AtomicU64::new(0),
};
pub fn locality_stats() -> &'static LocalityStats {
&LOCALITY_STATS
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cache_aligned_size() {
assert_eq!(std::mem::align_of::<CacheAligned<u64>>(), CACHE_LINE_SIZE);
assert!(std::mem::size_of::<CacheAligned<u64>>() >= CACHE_LINE_SIZE);
}
#[test]
fn test_cache_aligned_basic() {
let aligned = CacheAligned::new(42u64);
assert_eq!(*aligned, 42);
}
#[test]
fn test_cache_aligned_deref() {
let mut aligned = CacheAligned::new(vec![1, 2, 3]);
aligned.push(4);
assert_eq!(&*aligned, &vec![1, 2, 3, 4]);
}
#[test]
fn test_padded_size() {
assert!(std::mem::size_of::<Padded<u64>>() >= CACHE_LINE_SIZE);
}
#[test]
fn test_padded_basic() {
let padded = Padded::new(100u64);
assert_eq!(*padded, 100);
}
#[test]
fn test_hot_cold_separation() {
let data: HotCold<u64, Vec<String>> = HotCold::new(42, vec!["cold".into()]);
assert_eq!(*data.hot(), 42);
assert_eq!(data.cold().len(), 1);
}
#[test]
fn test_hot_cold_mutation() {
let mut data: HotCold<u64, Vec<u64>> = HotCold::new(0, vec![]);
*data.hot_mut() = 100;
data.cold_mut().push(1);
assert_eq!(*data.hot(), 100);
assert_eq!(data.cold().len(), 1);
}
#[test]
fn test_local_state_cache() {
let cache: LocalStateCache<u64> = LocalStateCache::new();
let val1 = cache.get_or_refresh(1, || 42);
assert_eq!(*val1, 42);
let val2 = cache.get_or_refresh(1, || 100);
assert_eq!(*val2, 42);
let val3 = cache.get_or_refresh(2, || 100);
assert_eq!(*val3, 100);
}
#[test]
fn test_local_state_cache_value_survives_refresh() {
let cache: LocalStateCache<String> = LocalStateCache::new();
let val1 = cache.get_or_refresh(1, || "first".to_string());
let val2 = cache.get_or_refresh(2, || "second".to_string());
cache.invalidate();
assert_eq!(*val1, "first");
assert_eq!(*val2, "second");
}
#[test]
fn test_local_cache_stats() {
let cache: LocalStateCache<u64> = LocalStateCache::new();
cache.get_or_refresh(1, || 1); cache.get_or_refresh(1, || 2); cache.get_or_refresh(1, || 3);
let stats = cache.stats();
assert_eq!(stats.hits, 2);
assert_eq!(stats.misses, 1);
assert!((stats.hit_ratio() - 0.666).abs() < 0.01);
}
#[test]
fn test_compact_state() {
let state = CompactState::<4>::new();
state.set(0, 100);
state.set(1, 200);
assert_eq!(state.get(0), 100);
assert_eq!(state.get(1), 200);
state.increment(0);
assert_eq!(state.get(0), 101);
state.add(1, 50);
assert_eq!(state.get(1), 250);
}
#[test]
fn test_compact_state_all() {
let state = CompactState::<4>::new();
state.set(0, 1);
state.set(1, 2);
state.set(2, 3);
state.set(3, 4);
let all = state.all();
assert_eq!(all, [1, 2, 3, 4]);
}
#[test]
fn test_soa_storage() {
let mut storage = SoaStorage::<3, 100>::new();
let idx = storage.push().unwrap();
storage.set_field(0, idx, 1.0);
storage.set_field(1, idx, 2.0);
storage.set_field(2, idx, 3.0);
assert_eq!(storage.get_field(0, idx), 1.0);
assert_eq!(storage.get_field(1, idx), 2.0);
assert_eq!(storage.get_field(2, idx), 3.0);
}
#[test]
fn test_soa_field_slice() {
let mut storage = SoaStorage::<2, 100>::new();
for _ in 0..10 {
let idx = storage.push().unwrap();
storage.set_field(0, idx, idx as f64);
}
let slice = storage.field_slice(0);
assert_eq!(slice.len(), 10);
}
#[test]
fn test_soa_unset_fields_are_zeroed() {
let mut storage = SoaStorage::<2, 8>::new();
let idx = storage.push().unwrap();
storage.set_field(0, idx, 7.0);
assert_eq!(storage.get_field(0, idx), 7.0);
assert_eq!(storage.get_field(1, idx), 0.0);
assert_eq!(storage.field_slice(1), &[0.0]);
}
#[test]
fn test_soa_push_clamped_to_capacity() {
let mut storage = SoaStorage::<1, 4>::new();
for _ in 0..4 {
assert!(storage.push().is_some());
}
assert!(storage.push().is_none());
assert!(storage.push().is_none());
assert_eq!(storage.len(), 4);
assert_eq!(storage.field_slice(0).len(), 4);
}
#[test]
fn test_prefetch_levels() {
let data = vec![0u64; 100];
prefetch(&data[0], PrefetchLevel::L1);
prefetch(&data[50], PrefetchLevel::L2);
prefetch(&data[99], PrefetchLevel::L3);
}
#[test]
fn test_locality_stats() {
let stats = locality_stats();
let _ = stats.cache_hits();
let _ = stats.cache_misses();
let _ = stats.hit_ratio();
}
}