use std::cell::UnsafeCell;
use std::mem::MaybeUninit;
use std::ops::{Deref, DerefMut};
use std::sync::atomic::{AtomicU64, AtomicUsize, 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: UnsafeCell<Option<T>>,
version: UnsafeCell<u64>,
hits: UnsafeCell<u64>,
misses: UnsafeCell<u64>,
}
impl<T> LocalStateCache<T> {
pub const fn new() -> Self {
Self {
value: UnsafeCell::new(None),
version: UnsafeCell::new(0),
hits: UnsafeCell::new(0),
misses: UnsafeCell::new(0),
}
}
#[inline]
pub fn get_or_refresh<F>(&self, current_version: u64, refresh: F) -> &T
where
F: FnOnce() -> T,
{
unsafe {
let cached_version = *self.version.get();
if cached_version == current_version
&& let Some(ref value) = *self.value.get()
{
*self.hits.get() += 1;
LOCALITY_STATS.record_cache_hit();
return value;
}
*self.misses.get() += 1;
LOCALITY_STATS.record_cache_miss();
let new_value = refresh();
*self.value.get() = Some(new_value);
*self.version.get() = current_version;
(*self.value.get()).as_ref().unwrap()
}
}
#[inline]
pub fn invalidate(&self) {
unsafe {
*self.value.get() = None;
*self.version.get() = 0;
}
}
pub fn stats(&self) -> LocalCacheStats {
unsafe {
LocalCacheStats {
hits: *self.hits.get(),
misses: *self.misses.get(),
}
}
}
}
impl<T> Default for LocalStateCache<T> {
fn default() -> Self {
Self::new()
}
}
unsafe impl<T: Send> Send for LocalStateCache<T> {}
#[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<[MaybeUninit<f64>; CAPACITY]>; FIELDS],
len: AtomicUsize,
}
impl<const FIELDS: usize, const CAPACITY: usize> SoaStorage<FIELDS, CAPACITY> {
pub fn new() -> Self {
Self {
data: std::array::from_fn(|_| {
Box::new(unsafe {
MaybeUninit::<[MaybeUninit<f64>; CAPACITY]>::uninit().assume_init()
})
}),
len: AtomicUsize::new(0),
}
}
#[inline]
pub fn len(&self) -> usize {
self.len.load(Ordering::Relaxed)
}
#[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());
unsafe { self.data[field][index].assume_init() }
}
#[inline]
pub fn set_field(&self, field: usize, index: usize, value: f64) {
assert!(field < FIELDS && index < CAPACITY);
unsafe {
let ptr = self.data[field].as_ptr() as *mut MaybeUninit<f64>;
(*ptr.add(index)).write(value);
}
}
pub fn push(&self) -> Option<usize> {
let index = self.len.fetch_add(1, Ordering::Relaxed);
if index < CAPACITY {
Some(index)
} else {
self.len.fetch_sub(1, Ordering::Relaxed);
None
}
}
pub fn field_slice(&self, field: usize) -> &[f64] {
assert!(field < FIELDS);
let len = self.len();
unsafe { std::slice::from_raw_parts(self.data[field].as_ptr() as *const f64, 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_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 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 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_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();
}
}