use anyhow::Result;
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use super::backend::{Cache, CacheConfig};
use super::entry::CacheEntry;
use super::key::{CacheCategory, CacheKey};
use super::stats::{AtomicCacheStats, CacheStats, CategoryStats};
pub struct MemoryCache {
entries: Arc<RwLock<HashMap<CacheKey, CacheEntry>>>,
config: CacheConfig,
stats: AtomicCacheStats,
}
impl MemoryCache {
pub fn new(config: CacheConfig) -> Self {
Self {
entries: Arc::new(RwLock::new(HashMap::new())),
config,
stats: AtomicCacheStats::new(),
}
}
pub fn with_entries(config: CacheConfig, entries: HashMap<CacheKey, CacheEntry>) -> Self {
let stats = AtomicCacheStats::new();
for entry in entries.values() {
stats.record_add(entry.size_bytes);
}
Self {
entries: Arc::new(RwLock::new(entries)),
config,
stats,
}
}
async fn maybe_evict(&self) -> Result<()> {
if let Some(max_size) = self.config.max_size_bytes {
let entries = self.entries.read().await;
let current_size: u64 = entries.values().map(|e| e.size_bytes).sum();
if current_size > max_size {
drop(entries);
self.evict_lru().await?;
}
}
Ok(())
}
async fn evict_lru(&self) -> Result<u64> {
let max_size = match self.config.max_size_bytes {
Some(size) => size,
None => return Ok(0),
};
let mut entries = self.entries.write().await;
let mut evicted = 0u64;
while !entries.is_empty() {
let current_size: u64 = entries.values().map(|e| e.size_bytes).sum();
if current_size <= max_size {
break;
}
let lru_key = entries
.iter()
.min_by_key(|(_, e)| e.last_accessed)
.map(|(k, _)| k.clone());
if let Some(key) = lru_key {
if let Some(entry) = entries.remove(&key) {
self.stats.record_remove(entry.size_bytes);
evicted += 1;
}
} else {
break;
}
}
Ok(evicted)
}
async fn remove_expired(&self) -> u64 {
let mut entries = self.entries.write().await;
let expired_keys: Vec<CacheKey> = entries
.iter()
.filter(|(_, e)| e.is_expired())
.map(|(k, _)| k.clone())
.collect();
let mut removed = 0u64;
for key in expired_keys {
if let Some(entry) = entries.remove(&key) {
self.stats.record_remove(entry.size_bytes);
removed += 1;
}
}
removed
}
}
#[async_trait]
impl Cache for MemoryCache {
async fn get(&self, key: &CacheKey) -> Result<Option<CacheEntry>> {
let mut entries = self.entries.write().await;
if let Some(entry) = entries.get_mut(key) {
if entry.is_expired() {
let size = entry.size_bytes;
entries.remove(key);
self.stats.record_remove(size);
self.stats.record_miss();
return Ok(None);
}
entry.touch();
self.stats.record_hit();
Ok(Some(entry.clone()))
} else {
self.stats.record_miss();
Ok(None)
}
}
async fn set(&self, key: &CacheKey, entry: CacheEntry) -> Result<()> {
self.maybe_evict().await?;
let size = entry.size_bytes;
let mut entries = self.entries.write().await;
if let Some(old) = entries.remove(key) {
self.stats.record_remove(old.size_bytes);
}
entries.insert(key.clone(), entry);
self.stats.record_add(size);
Ok(())
}
async fn delete(&self, key: &CacheKey) -> Result<()> {
let mut entries = self.entries.write().await;
if let Some(entry) = entries.remove(key) {
self.stats.record_remove(entry.size_bytes);
}
Ok(())
}
async fn exists(&self, key: &CacheKey) -> Result<bool> {
let entries = self.entries.read().await;
if let Some(entry) = entries.get(key) {
Ok(!entry.is_expired())
} else {
Ok(false)
}
}
async fn clear(&self, category: Option<CacheCategory>) -> Result<u64> {
let mut entries = self.entries.write().await;
let mut cleared = 0u64;
match category {
Some(cat) => {
let keys_to_remove: Vec<CacheKey> = entries
.keys()
.filter(|k| k.category == cat)
.cloned()
.collect();
for key in keys_to_remove {
if let Some(entry) = entries.remove(&key) {
self.stats.record_remove(entry.size_bytes);
cleared += 1;
}
}
}
None => {
cleared = entries.len() as u64;
entries.clear();
self.stats.reset();
}
}
Ok(cleared)
}
async fn stats(&self) -> Result<CacheStats> {
let entries = self.entries.read().await;
let mut stats = self.stats.to_stats();
let mut by_category: HashMap<String, CategoryStats> = HashMap::new();
for (key, entry) in entries.iter() {
let cat_name = key.category.to_string();
let cat_stats = by_category.entry(cat_name).or_default();
cat_stats.entries += 1;
cat_stats.size_bytes += entry.size_bytes;
if entry.is_expired() {
cat_stats.expired += 1;
}
}
stats.by_category = by_category;
if !entries.is_empty() {
stats.oldest_entry = entries.values().map(|e| e.created_at).min();
stats.newest_entry = entries.values().map(|e| e.created_at).max();
}
stats.total_entries = entries.len() as u64;
stats.total_size_bytes = entries.values().map(|e| e.size_bytes).sum();
stats.expired_entries = entries.values().filter(|e| e.is_expired()).count() as u64;
Ok(stats)
}
async fn prune_expired(&self) -> Result<u64> {
Ok(self.remove_expired().await)
}
async fn keys(&self, category: Option<CacheCategory>) -> Result<Vec<CacheKey>> {
let entries = self.entries.read().await;
let keys: Vec<CacheKey> = match category {
Some(cat) => entries
.keys()
.filter(|k| k.category == cat)
.cloned()
.collect(),
None => entries.keys().cloned().collect(),
};
Ok(keys)
}
}
impl Default for MemoryCache {
fn default() -> Self {
Self::new(CacheConfig::memory())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[tokio::test]
async fn memory_cache_basic_operations() {
let cache = MemoryCache::new(CacheConfig::memory());
let key = CacheKey::schema("test-server");
let entry = CacheEntry::new(b"test data".to_vec(), Duration::from_secs(3600));
cache.set(&key, entry.clone()).await.unwrap();
assert!(cache.exists(&key).await.unwrap());
let retrieved = cache.get(&key).await.unwrap();
assert!(retrieved.is_some());
assert_eq!(retrieved.unwrap().data, b"test data".to_vec());
cache.delete(&key).await.unwrap();
assert!(!cache.exists(&key).await.unwrap());
}
#[tokio::test]
async fn memory_cache_expiration() {
let cache = MemoryCache::new(CacheConfig::memory());
let key = CacheKey::schema("test-server");
let entry = CacheEntry::new(b"test".to_vec(), Duration::from_secs(0));
cache.set(&key, entry).await.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
let retrieved = cache.get(&key).await.unwrap();
assert!(retrieved.is_none());
}
#[tokio::test]
async fn memory_cache_lru_eviction() {
let config = CacheConfig::memory().with_max_size(250);
let cache = MemoryCache::new(config);
for i in 0..5 {
let key = CacheKey::schema(&format!("server-{}", i));
let entry = CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600));
cache.set(&key, entry).await.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
}
let stats = cache.stats().await.unwrap();
assert!(
stats.total_entries < 5,
"Expected fewer than 5 entries after eviction, got {}",
stats.total_entries
);
}
#[tokio::test]
async fn memory_cache_clear_category() {
let cache = MemoryCache::new(CacheConfig::memory());
let key1 = CacheKey::schema("server1");
let key2 = CacheKey::validation("server2", "1.0");
cache
.set(&key1, CacheEntry::new(vec![], Duration::from_secs(3600)))
.await
.unwrap();
cache
.set(&key2, CacheEntry::new(vec![], Duration::from_secs(3600)))
.await
.unwrap();
let cleared = cache.clear(Some(CacheCategory::Schema)).await.unwrap();
assert_eq!(cleared, 1);
assert!(cache.exists(&key2).await.unwrap());
}
#[tokio::test]
async fn memory_cache_stats() {
let cache = MemoryCache::new(CacheConfig::memory());
for i in 0..5 {
let key = CacheKey::schema(&format!("server-{}", i));
let entry = CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600));
cache.set(&key, entry).await.unwrap();
}
cache.get(&CacheKey::schema("server-0")).await.unwrap(); cache.get(&CacheKey::schema("nonexistent")).await.unwrap();
let stats = cache.stats().await.unwrap();
assert_eq!(stats.total_entries, 5);
assert_eq!(stats.hits, 1);
assert_eq!(stats.misses, 1);
}
#[tokio::test]
async fn memory_cache_miss_on_nonexistent() {
let cache = MemoryCache::new(CacheConfig::memory());
let key = CacheKey::schema("nonexistent");
let result = cache.get(&key).await.unwrap();
assert!(result.is_none());
let stats = cache.stats().await.unwrap();
assert_eq!(stats.misses, 1);
assert_eq!(stats.hits, 0);
}
#[tokio::test]
async fn memory_cache_replace_existing() {
let cache = MemoryCache::new(CacheConfig::memory());
let key = CacheKey::schema("test-server");
let entry1 = CacheEntry::new(b"data1".to_vec(), Duration::from_secs(3600));
cache.set(&key, entry1).await.unwrap();
let entry2 = CacheEntry::new(b"data2".to_vec(), Duration::from_secs(3600));
cache.set(&key, entry2).await.unwrap();
let retrieved = cache.get(&key).await.unwrap().unwrap();
assert_eq!(retrieved.data, b"data2".to_vec());
}
#[tokio::test]
async fn memory_cache_clear_all() {
let cache = MemoryCache::new(CacheConfig::memory());
let key1 = CacheKey::schema("server1");
let key2 = CacheKey::validation("server2", "1.0");
let key3 = CacheKey::scan_result("server3", "ruleset1");
cache
.set(
&key1,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
cache
.set(
&key2,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
cache
.set(
&key3,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
let cleared = cache.clear(None).await.unwrap();
assert_eq!(cleared, 3);
assert!(!cache.exists(&key1).await.unwrap());
assert!(!cache.exists(&key2).await.unwrap());
assert!(!cache.exists(&key3).await.unwrap());
let stats = cache.stats().await.unwrap();
assert_eq!(stats.total_entries, 0);
}
#[tokio::test]
async fn memory_cache_prune_expired() {
let cache = MemoryCache::new(CacheConfig::memory());
let key1 = CacheKey::schema("expired1");
let key2 = CacheKey::schema("expired2");
let key3 = CacheKey::schema("valid");
cache
.set(
&key1,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(0)),
)
.await
.unwrap();
cache
.set(
&key2,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(0)),
)
.await
.unwrap();
cache
.set(
&key3,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
let pruned = cache.prune_expired().await.unwrap();
assert_eq!(pruned, 2);
assert!(cache.exists(&key3).await.unwrap());
assert!(!cache.exists(&key1).await.unwrap());
assert!(!cache.exists(&key2).await.unwrap());
}
#[tokio::test]
async fn memory_cache_keys_by_category() {
let cache = MemoryCache::new(CacheConfig::memory());
let key1 = CacheKey::schema("server1");
let key2 = CacheKey::schema("server2");
let key3 = CacheKey::validation("server3", "1.0");
cache
.set(&key1, CacheEntry::new(vec![], Duration::from_secs(3600)))
.await
.unwrap();
cache
.set(&key2, CacheEntry::new(vec![], Duration::from_secs(3600)))
.await
.unwrap();
cache
.set(&key3, CacheEntry::new(vec![], Duration::from_secs(3600)))
.await
.unwrap();
let schema_keys = cache.keys(Some(CacheCategory::Schema)).await.unwrap();
assert_eq!(schema_keys.len(), 2);
let all_keys = cache.keys(None).await.unwrap();
assert_eq!(all_keys.len(), 3);
}
#[tokio::test]
async fn memory_cache_exists_with_expired() {
let cache = MemoryCache::new(CacheConfig::memory());
let key = CacheKey::schema("test");
cache
.set(&key, CacheEntry::new(vec![], Duration::from_secs(0)))
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
assert!(!cache.exists(&key).await.unwrap());
}
#[tokio::test]
async fn memory_cache_with_entries_constructor() {
let mut entries = HashMap::new();
let key = CacheKey::schema("preloaded");
let entry = CacheEntry::new(b"data".to_vec(), Duration::from_secs(3600));
entries.insert(key.clone(), entry);
let cache = MemoryCache::with_entries(CacheConfig::memory(), entries);
let retrieved = cache.get(&key).await.unwrap();
assert!(retrieved.is_some());
assert_eq!(retrieved.unwrap().data, b"data".to_vec());
let stats = cache.stats().await.unwrap();
assert_eq!(stats.total_entries, 1);
}
#[tokio::test]
async fn memory_cache_default_constructor() {
let cache = MemoryCache::default();
let key = CacheKey::schema("test");
let entry = CacheEntry::new(b"data".to_vec(), Duration::from_secs(3600));
cache.set(&key, entry).await.unwrap();
let retrieved = cache.get(&key).await.unwrap();
assert!(retrieved.is_some());
}
#[tokio::test]
async fn memory_cache_empty_cache_stats() {
let cache = MemoryCache::new(CacheConfig::memory());
let stats = cache.stats().await.unwrap();
assert_eq!(stats.total_entries, 0);
assert_eq!(stats.total_size_bytes, 0);
assert_eq!(stats.expired_entries, 0);
assert!(stats.oldest_entry.is_none());
assert!(stats.newest_entry.is_none());
}
#[tokio::test]
async fn memory_cache_stats_with_expired() {
let cache = MemoryCache::new(CacheConfig::memory());
let key1 = CacheKey::schema("expired");
let key2 = CacheKey::schema("valid");
cache
.set(
&key1,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(0)),
)
.await
.unwrap();
cache
.set(
&key2,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
let stats = cache.stats().await.unwrap();
assert_eq!(stats.total_entries, 2);
assert_eq!(stats.expired_entries, 1);
}
#[tokio::test]
async fn memory_cache_stats_by_category() {
let cache = MemoryCache::new(CacheConfig::memory());
cache
.set(
&CacheKey::schema("server1"),
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
cache
.set(
&CacheKey::schema("server2"),
CacheEntry::new(vec![0u8; 150], Duration::from_secs(3600)),
)
.await
.unwrap();
cache
.set(
&CacheKey::validation("server3", "1.0"),
CacheEntry::new(vec![0u8; 200], Duration::from_secs(3600)),
)
.await
.unwrap();
let stats = cache.stats().await.unwrap();
assert!(stats.by_category.contains_key("schemas"));
assert!(stats.by_category.contains_key("validation"));
let schema_stats = &stats.by_category["schemas"];
assert_eq!(schema_stats.entries, 2);
assert_eq!(schema_stats.size_bytes, 250);
}
#[tokio::test]
async fn memory_cache_delete_nonexistent() {
let cache = MemoryCache::new(CacheConfig::memory());
let key = CacheKey::schema("nonexistent");
cache.delete(&key).await.unwrap();
let stats = cache.stats().await.unwrap();
assert_eq!(stats.total_entries, 0);
}
#[tokio::test]
async fn memory_cache_no_eviction_when_under_capacity() {
let config = CacheConfig::memory().with_max_size(1000);
let cache = MemoryCache::new(config);
let key = CacheKey::schema("server1");
cache
.set(
&key,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
assert!(cache.exists(&key).await.unwrap());
let stats = cache.stats().await.unwrap();
assert_eq!(stats.total_entries, 1);
}
#[tokio::test]
async fn memory_cache_lru_eviction_multiple_rounds() {
let config = CacheConfig::memory().with_max_size(150);
let cache = MemoryCache::new(config);
let key1 = CacheKey::schema("server1");
cache
.set(
&key1,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
let key2 = CacheKey::schema("server2");
cache
.set(
&key2,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
let key3 = CacheKey::schema("server3");
cache
.set(
&key3,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
assert!(!cache.exists(&key1).await.unwrap());
assert!(cache.exists(&key2).await.unwrap());
assert!(cache.exists(&key3).await.unwrap());
}
#[tokio::test]
async fn memory_cache_touch_updates_last_accessed() {
let cache = MemoryCache::new(CacheConfig::memory());
let key = CacheKey::schema("test");
cache
.set(
&key,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(3600)),
)
.await
.unwrap();
let entry1 = cache.get(&key).await.unwrap().unwrap();
let accessed1 = entry1.last_accessed;
tokio::time::sleep(Duration::from_millis(10)).await;
let entry2 = cache.get(&key).await.unwrap().unwrap();
let accessed2 = entry2.last_accessed;
assert!(accessed2 > accessed1);
}
#[tokio::test]
async fn memory_cache_expired_on_get_removes_entry() {
let cache = MemoryCache::new(CacheConfig::memory());
let key = CacheKey::schema("test");
cache
.set(
&key,
CacheEntry::new(vec![0u8; 100], Duration::from_secs(0)),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(10)).await;
let result = cache.get(&key).await.unwrap();
assert!(result.is_none());
let keys = cache.keys(None).await.unwrap();
assert_eq!(keys.len(), 0);
}
#[tokio::test]
async fn memory_cache_prune_expired_empty_cache() {
let cache = MemoryCache::new(CacheConfig::memory());
let pruned = cache.prune_expired().await.unwrap();
assert_eq!(pruned, 0);
}
#[tokio::test]
async fn memory_cache_clear_category_empty() {
let cache = MemoryCache::new(CacheConfig::memory());
let cleared = cache.clear(Some(CacheCategory::Schema)).await.unwrap();
assert_eq!(cleared, 0);
}
}