use dashmap::DashMap;
use ntex::time::interval;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use crate::error::{AuthError, AuthResult};
#[derive(Debug, Clone)]
pub struct CacheEntry {
pub value: bool,
pub expires_at: u64,
pub created_at: u64,
pub access_count: u64,
pub last_accessed: u64,
}
impl CacheEntry {
pub fn new(value: bool, ttl_seconds: u64) -> Self {
let now = current_timestamp();
Self {
value,
expires_at: now + ttl_seconds,
created_at: now,
access_count: 1,
last_accessed: now,
}
}
pub fn is_expired(&self) -> bool {
current_timestamp() > self.expires_at
}
pub fn age_seconds(&self) -> u64 {
current_timestamp().saturating_sub(self.created_at)
}
pub fn time_since_last_access(&self) -> u64 {
current_timestamp().saturating_sub(self.last_accessed)
}
pub fn mark_accessed(&mut self) {
self.access_count += 1;
self.last_accessed = current_timestamp();
}
pub fn hotness_score(&self) -> f64 {
let age = self.age_seconds() as f64;
let access_rate = self.access_count as f64 / age.max(1.0);
let recency = 1.0 / (self.time_since_last_access() as f64 + 1.0);
access_rate * recency
}
}
#[derive(Debug, Clone)]
pub struct CacheConfig {
pub max_size: usize,
pub ttl_seconds: u64,
pub cleanup_interval_seconds: u64,
pub auto_cleanup: bool,
pub soft_limit_ratio: f64,
pub cleanup_batch_size: usize,
}
impl Default for CacheConfig {
fn default() -> Self {
Self {
max_size: 1000,
ttl_seconds: 300, cleanup_interval_seconds: 60, auto_cleanup: true,
soft_limit_ratio: 0.8, cleanup_batch_size: 100,
}
}
}
impl CacheConfig {
pub fn new() -> Self {
Self::default()
}
pub fn max_size(mut self, size: usize) -> Self {
self.max_size = size;
self
}
pub fn ttl_seconds(mut self, seconds: u64) -> Self {
self.ttl_seconds = seconds;
self
}
pub fn ttl_minutes(self, minutes: u64) -> Self {
self.ttl_seconds(minutes * 60)
}
pub fn ttl_hours(self, hours: u64) -> Self {
self.ttl_seconds(hours * 3600)
}
pub fn cleanup_interval_seconds(mut self, seconds: u64) -> Self {
self.cleanup_interval_seconds = seconds;
self
}
pub fn disable_auto_cleanup(mut self) -> Self {
self.auto_cleanup = false;
self
}
pub fn soft_limit_ratio(mut self, ratio: f64) -> Self {
self.soft_limit_ratio = ratio;
self
}
pub fn cleanup_batch_size(mut self, size: usize) -> Self {
self.cleanup_batch_size = size;
self
}
pub fn validate(&self) -> AuthResult<()> {
if self.max_size == 0 {
return Err(AuthError::ConfigError(
"max_size must be greater than 0".to_string(),
));
}
if self.ttl_seconds == 0 {
return Err(AuthError::ConfigError(
"ttl_seconds must be greater than 0".to_string(),
));
}
if self.cleanup_interval_seconds == 0 {
return Err(AuthError::ConfigError(
"cleanup_interval_seconds must be greater than 0".to_string(),
));
}
if !(0.1..=0.95).contains(&self.soft_limit_ratio) {
return Err(AuthError::ConfigError(
"soft_limit_ratio must be between 0.1 and 0.95".to_string(),
));
}
Ok(())
}
}
pub struct AuthCache {
cache: Arc<DashMap<[u8; 32], CacheEntry>>,
config: CacheConfig,
stats: CacheStatistics,
_cleanup_handle: Mutex<Option<ntex::rt::JoinHandle<()>>>,
}
impl AuthCache {
pub fn new(config: CacheConfig) -> AuthResult<Self> {
config.validate()?;
let cache = Arc::new(DashMap::new());
let stats = CacheStatistics::new();
let cleanup_handle = if config.auto_cleanup {
Some(Self::start_cleanup_task(
Arc::clone(&cache),
config.clone(),
stats.clone(),
))
} else {
None
};
Ok(Self {
cache,
config,
stats,
_cleanup_handle: Mutex::new(cleanup_handle),
})
}
fn start_cleanup_task(
cache: Arc<DashMap<[u8; 32], CacheEntry>>,
config: CacheConfig,
stats: CacheStatistics,
) -> ntex::rt::JoinHandle<()> {
ntex::rt::spawn(async move {
let interval = interval(Duration::from_secs(config.cleanup_interval_seconds));
loop {
interval.tick().await;
let expired_count = Self::cleanup_expired(&cache);
stats.add_expired_cleaned(expired_count);
let soft_limit = (config.max_size as f64 * config.soft_limit_ratio) as usize;
if cache.len() > soft_limit {
let cleaned = Self::cleanup_by_hotness(&cache, config.cleanup_batch_size);
stats.add_size_cleaned(cleaned);
}
}
})
}
pub fn get(&self, key: &[u8; 32]) -> Option<bool> {
self.stats.add_access();
if let Some(mut entry) = self.cache.get_mut(key) {
if entry.is_expired() {
drop(entry); self.cache.remove(key);
self.stats.add_miss();
None
} else {
entry.mark_accessed();
let value = entry.value;
self.stats.add_hit();
Some(value)
}
} else {
self.stats.add_miss();
None
}
}
pub fn insert(&self, key: [u8; 32], value: bool) -> AuthResult<()> {
if self.cache.len() >= self.config.max_size {
self.force_cleanup();
}
let entry = CacheEntry::new(value, self.config.ttl_seconds);
self.cache.insert(key, entry);
self.stats.add_insertion();
Ok(())
}
pub fn remove(&self, key: &[u8; 32]) -> Option<bool> {
self.cache.remove(key).map(|(_, entry)| {
self.stats.add_removal();
entry.value
})
}
pub fn force_cleanup(&self) {
let expired_count = Self::cleanup_expired(&self.cache);
self.stats.add_expired_cleaned(expired_count);
if self.cache.len() > self.config.max_size {
let cleaned = Self::cleanup_by_hotness(&self.cache, self.config.cleanup_batch_size);
self.stats.add_size_cleaned(cleaned);
}
}
pub fn clear(&self) {
let count = self.cache.len();
self.cache.clear();
self.stats.add_cleared(count);
}
pub fn stats(&self) -> CacheStats {
let total_entries = self.cache.len() as u64;
let expired_count = self.cache.iter().filter(|entry| entry.is_expired()).count() as u64;
let (total_age, min_age, max_age) = if total_entries > 0 {
let ages: Vec<u64> = self.cache.iter().map(|entry| entry.age_seconds()).collect();
let total: u64 = ages.iter().sum();
let min = *ages.iter().min().unwrap_or(&0);
let max = *ages.iter().max().unwrap_or(&0);
(total, min, max)
} else {
(0, 0, 0)
};
CacheStats {
total_entries,
expired_entries: expired_count,
valid_entries: total_entries - expired_count,
average_age_seconds: total_age.checked_div(total_entries).unwrap_or(0),
min_age_seconds: min_age,
max_age_seconds: max_age,
memory_usage_estimate: total_entries as usize * std::mem::size_of::<CacheEntry>(),
hit_count: self.stats.hit_count.load(Ordering::Relaxed),
miss_count: self.stats.miss_count.load(Ordering::Relaxed),
total_accesses: self.stats.total_accesses.load(Ordering::Relaxed),
insertions: self.stats.insertions.load(Ordering::Relaxed),
removals: self.stats.removals.load(Ordering::Relaxed),
expired_cleaned: self.stats.expired_cleaned.load(Ordering::Relaxed),
size_cleaned: self.stats.size_cleaned.load(Ordering::Relaxed),
}
}
pub fn contains_key(&self, key: &[u8; 32]) -> bool {
self.cache.contains_key(key)
}
pub fn len(&self) -> usize {
self.cache.len()
}
pub fn is_empty(&self) -> bool {
self.cache.is_empty()
}
pub fn config(&self) -> &CacheConfig {
&self.config
}
fn cleanup_expired(cache: &DashMap<[u8; 32], CacheEntry>) -> u64 {
let initial_len = cache.len();
cache.retain(|_, entry| !entry.is_expired());
(initial_len - cache.len()) as u64
}
fn cleanup_by_hotness(cache: &DashMap<[u8; 32], CacheEntry>, max_remove: usize) -> u64 {
if cache.is_empty() {
return 0;
}
let mut entries: Vec<([u8; 32], f64)> = cache
.iter()
.map(|item| (*item.key(), item.value().hotness_score()))
.collect();
entries.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
let remove_count = max_remove.min(entries.len());
let mut removed = 0;
for (key, _) in entries.into_iter().take(remove_count) {
if cache.remove(&key).is_some() {
removed += 1;
}
}
removed
}
}
#[derive(Debug, Clone)]
pub struct CacheStats {
pub total_entries: u64,
pub expired_entries: u64,
pub valid_entries: u64,
pub average_age_seconds: u64,
pub min_age_seconds: u64,
pub max_age_seconds: u64,
pub memory_usage_estimate: usize,
pub hit_count: u64,
pub miss_count: u64,
pub total_accesses: u64,
pub insertions: u64,
pub removals: u64,
pub expired_cleaned: u64,
pub size_cleaned: u64,
}
impl CacheStats {
pub fn hit_ratio(&self) -> f64 {
if self.total_accesses == 0 {
0.0
} else {
self.hit_count as f64 / self.total_accesses as f64
}
}
pub fn miss_ratio(&self) -> f64 {
1.0 - self.hit_ratio()
}
pub fn is_healthy(&self) -> bool {
self.hit_ratio() > 0.8 && self.expired_entries < self.total_entries / 4
}
pub fn efficiency_score(&self) -> f64 {
let hit_ratio = self.hit_ratio();
let expired_ratio = if self.total_entries > 0 {
self.expired_entries as f64 / self.total_entries as f64
} else {
0.0
};
hit_ratio * (1.0 - expired_ratio)
}
}
#[derive(Debug, Clone)]
struct CacheStatistics {
hit_count: Arc<AtomicU64>,
miss_count: Arc<AtomicU64>,
total_accesses: Arc<AtomicU64>,
insertions: Arc<AtomicU64>,
removals: Arc<AtomicU64>,
expired_cleaned: Arc<AtomicU64>,
size_cleaned: Arc<AtomicU64>,
}
impl CacheStatistics {
fn new() -> Self {
Self {
hit_count: Arc::new(AtomicU64::new(0)),
miss_count: Arc::new(AtomicU64::new(0)),
total_accesses: Arc::new(AtomicU64::new(0)),
insertions: Arc::new(AtomicU64::new(0)),
removals: Arc::new(AtomicU64::new(0)),
expired_cleaned: Arc::new(AtomicU64::new(0)),
size_cleaned: Arc::new(AtomicU64::new(0)),
}
}
fn add_hit(&self) {
self.hit_count.fetch_add(1, Ordering::Relaxed);
}
fn add_miss(&self) {
self.miss_count.fetch_add(1, Ordering::Relaxed);
}
fn add_access(&self) {
self.total_accesses.fetch_add(1, Ordering::Relaxed);
}
fn add_insertion(&self) {
self.insertions.fetch_add(1, Ordering::Relaxed);
}
fn add_removal(&self) {
self.removals.fetch_add(1, Ordering::Relaxed);
}
fn add_expired_cleaned(&self, count: u64) {
self.expired_cleaned.fetch_add(count, Ordering::Relaxed);
}
fn add_size_cleaned(&self, count: u64) {
self.size_cleaned.fetch_add(count, Ordering::Relaxed);
}
fn add_cleared(&self, count: usize) {
self.removals.fetch_add(count as u64, Ordering::Relaxed);
}
}
fn current_timestamp() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::time::{Duration, sleep};
#[test]
fn test_cache_entry() {
let entry = CacheEntry::new(true, 60);
assert!(entry.value);
assert!(!entry.is_expired());
assert_eq!(entry.access_count, 1);
}
#[ntex::test]
async fn test_cache_basic_operations() {
let config = CacheConfig::new().max_size(100).ttl_seconds(60);
let cache = AuthCache::new(config).unwrap();
let key1 = [1u8; 32];
let key2 = [2u8; 32];
cache.insert(key1, true).unwrap();
assert_eq!(cache.get(&key1), Some(true));
assert_eq!(cache.get(&key2), None);
assert_eq!(cache.remove(&key1), Some(true));
assert_eq!(cache.get(&key1), None);
}
#[tokio::test]
async fn test_cache_expiration() {
let config = CacheConfig::new()
.max_size(100)
.ttl_seconds(1)
.disable_auto_cleanup();
let cache = AuthCache::new(config).unwrap();
let key1 = [1u8; 32];
cache.insert(key1, true).unwrap();
assert_eq!(cache.get(&key1), Some(true));
sleep(Duration::from_secs(2)).await;
assert_eq!(cache.get(&key1), None);
}
#[test]
fn test_cache_stats() {
let config = CacheConfig::new().disable_auto_cleanup();
let cache = AuthCache::new(config).unwrap();
let key1 = [1u8; 32];
let key2 = [2u8; 32];
let key3 = [3u8; 32];
cache.insert(key1, true).unwrap();
cache.insert(key2, false).unwrap();
cache.get(&key1);
cache.get(&key2);
cache.get(&key3);
let stats = cache.stats();
assert_eq!(stats.total_entries, 2);
assert!(stats.hit_ratio() > 0.0);
assert!(stats.efficiency_score() > 0.0);
}
#[test]
fn test_hotness_score() {
let entry1 = CacheEntry::new(true, 60);
let mut entry2 = CacheEntry::new(true, 60);
for _ in 0..10 {
entry2.mark_accessed();
}
assert!(entry2.hotness_score() > entry1.hotness_score());
}
#[test]
fn test_config_validation() {
assert!(CacheConfig::new().max_size(0).validate().is_err());
assert!(CacheConfig::new().ttl_seconds(0).validate().is_err());
assert!(CacheConfig::new().soft_limit_ratio(1.5).validate().is_err());
assert!(CacheConfig::new().validate().is_ok());
}
}