use crate::config::CacheConfig;
use crate::deps::*;
use crate::entry::CacheEntry;
use crate::stats::CacheStats;
use tracing::debug;
pub struct Cache<K, V>
where
K: Eq + Hash + Clone,
{
data: RwLock<HashMap<K, CacheEntry<V>>>,
config: CacheConfig,
stats: CacheStats,
}
impl<K, V> Cache<K, V>
where
K: Eq + Hash + Clone,
{
pub fn new(config: CacheConfig) -> Self {
Self {
data: RwLock::new(HashMap::with_capacity(config.max_capacity)),
config,
stats: CacheStats::new(),
}
}
pub fn with_defaults() -> Self {
Self::new(CacheConfig::default())
}
pub fn insert(&self, key: K, value: V) {
self.insert_with_ttl(key, value, self.config.default_ttl)
}
pub fn insert_with_ttl(&self, key: K, value: V, ttl: Option<Duration>) {
let entry = CacheEntry::new(value, ttl);
let mut data = self.data.write();
if data.len() >= self.config.max_capacity && !data.contains_key(&key) {
self.evict_lru(&mut data);
}
data.insert(key, entry);
if self.config.enable_stats {
self.stats.record_insert();
}
}
pub fn get(&self, key: &K) -> Option<V>
where
V: Clone,
{
let mut data = self.data.write();
if let Some(entry) = data.get_mut(key) {
if entry.is_expired() {
data.remove(key);
if self.config.enable_stats {
self.stats.record_expiration();
self.stats.record_miss();
}
return None;
}
entry.touch();
if self.config.enable_stats {
self.stats.record_hit();
}
return Some(entry.value().clone());
}
if self.config.enable_stats {
self.stats.record_miss();
}
None
}
pub fn with_value<R, F>(&self, key: &K, f: F) -> Option<R>
where
F: FnOnce(&V) -> R,
{
let mut data = self.data.write();
if let Some(entry) = data.get_mut(key) {
if entry.is_expired() {
data.remove(key);
if self.config.enable_stats {
self.stats.record_expiration();
self.stats.record_miss();
}
return None;
}
entry.touch();
if self.config.enable_stats {
self.stats.record_hit();
}
return Some(f(entry.value()));
}
if self.config.enable_stats {
self.stats.record_miss();
}
None
}
pub fn contains_key(&self, key: &K) -> bool {
let data = self.data.read();
if let Some(entry) = data.get(key) {
!entry.is_expired()
} else {
false
}
}
pub fn remove(&self, key: &K) -> Option<V> {
self.data.write().remove(key).map(|e| e.into_value())
}
pub fn get_or_insert_with<F>(&self, key: K, f: F) -> V
where
V: Clone,
F: FnOnce() -> V,
{
if let Some(value) = self.get(&key) {
return value;
}
let value = f();
self.insert(key, value.clone());
value
}
pub fn clear(&self) {
self.data.write().clear();
}
pub fn len(&self) -> usize {
self.data.read().len()
}
pub fn is_empty(&self) -> bool {
self.data.read().is_empty()
}
pub fn stats(&self) -> crate::stats::CacheStatsSnapshot {
self.stats.snapshot()
}
pub fn reset_stats(&self) {
self.stats.reset();
}
pub fn cleanup_expired(&self) -> usize {
let mut data = self.data.write();
let before = data.len();
data.retain(|_, entry| {
let expired = entry.is_expired();
if expired && self.config.enable_stats {
self.stats.record_expiration();
}
!expired
});
let removed = before - data.len();
if removed > 0 {
debug!(removed = removed, "清理过期缓存条目");
}
removed
}
fn evict_lru(&self, data: &mut HashMap<K, CacheEntry<V>>) {
let to_evict: Vec<K> = data
.iter()
.filter(|(_, entry)| entry.is_expired())
.map(|(k, _)| k.clone())
.take(self.config.eviction_batch_size)
.collect();
let mut to_evict = to_evict;
if to_evict.len() < self.config.eviction_batch_size {
let needed = self.config.eviction_batch_size - to_evict.len();
let mut entries: Vec<_> = data
.iter()
.filter(|(k, _)| !to_evict.contains(k))
.map(|(k, e)| (k.clone(), e.last_accessed()))
.collect();
entries.sort_by_key(|(_, t)| *t);
to_evict.extend(entries.into_iter().take(needed).map(|(k, _)| k));
}
for key in to_evict {
data.remove(&key);
if self.config.enable_stats {
self.stats.record_eviction();
}
}
}
pub fn keys(&self) -> Vec<K> {
self.data.read().keys().cloned().collect()
}
pub fn insert_many<I>(&self, items: I)
where
I: IntoIterator<Item = (K, V)>,
{
let mut data = self.data.write();
for (key, value) in items {
let entry = CacheEntry::new(value, self.config.default_ttl);
data.insert(key, entry);
if self.config.enable_stats {
self.stats.record_insert();
}
}
}
pub fn update<F>(&self, key: &K, f: F) -> bool
where
F: FnOnce(&mut V),
{
let mut data = self.data.write();
if let Some(entry) = data.get_mut(key) {
if entry.is_expired() {
data.remove(key);
if self.config.enable_stats {
self.stats.record_expiration();
}
return false;
}
f(entry.value_mut());
entry.touch();
true
} else {
false
}
}
}
impl<K, V> Default for Cache<K, V>
where
K: Eq + Hash + Clone,
{
fn default() -> Self {
Self::with_defaults()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_insert_and_get() {
let cache: Cache<&str, i32> = Cache::with_defaults();
cache.insert("key1", 42);
assert_eq!(cache.get(&"key1"), Some(42));
assert_eq!(cache.get(&"key2"), None);
}
#[test]
fn test_ttl_expiration() {
let config = CacheConfig::default().default_ttl(Duration::from_millis(10));
let cache: Cache<&str, i32> = Cache::new(config);
cache.insert("key1", 42);
assert!(cache.contains_key(&"key1"));
std::thread::sleep(Duration::from_millis(20));
assert!(!cache.contains_key(&"key1"));
}
#[test]
fn test_remove() {
let cache: Cache<&str, i32> = Cache::with_defaults();
cache.insert("key1", 42);
assert_eq!(cache.remove(&"key1"), Some(42));
assert!(!cache.contains_key(&"key1"));
}
#[test]
fn test_get_or_insert() {
let cache: Cache<&str, i32> = Cache::with_defaults();
let value = cache.get_or_insert_with("key1", || 42);
assert_eq!(value, 42);
let value = cache.get_or_insert_with("key1", || 100);
assert_eq!(value, 42); }
#[test]
fn test_stats() {
let cache: Cache<&str, i32> = Cache::with_defaults();
cache.insert("key1", 42);
cache.get(&"key1");
cache.get(&"key2");
let stats = cache.stats();
assert_eq!(stats.hits, 1);
assert_eq!(stats.misses, 1);
assert_eq!(stats.inserts, 1);
}
#[test]
fn test_lru_eviction() {
let config = CacheConfig::default()
.max_capacity(3)
.no_ttl()
.eviction_batch_size(1);
let cache: Cache<i32, i32> = Cache::new(config);
cache.insert(1, 10);
std::thread::sleep(Duration::from_millis(1));
cache.insert(2, 20);
std::thread::sleep(Duration::from_millis(1));
cache.insert(3, 30);
std::thread::sleep(Duration::from_millis(1));
cache.get(&1);
std::thread::sleep(Duration::from_millis(1));
cache.insert(4, 40);
assert!(cache.contains_key(&1));
assert!(!cache.contains_key(&2)); assert!(cache.contains_key(&3));
assert!(cache.contains_key(&4));
}
}