use std::collections::{HashMap, VecDeque};
#[derive(Debug, Clone, Default)]
pub struct CacheStats {
pub hits: u64,
pub misses: u64,
pub predictions: u64,
pub prediction_hits: u64,
pub evictions: u64,
}
impl CacheStats {
#[must_use]
pub fn hit_rate(&self) -> f64 {
let total = self.hits + self.misses;
if total == 0 {
0.0
} else {
self.hits as f64 / total as f64
}
}
#[must_use]
pub fn prediction_accuracy(&self) -> f64 {
if self.predictions == 0 {
0.0
} else {
self.prediction_hits as f64 / self.predictions as f64
}
}
}
struct Entry<V> {
value: V,
order: u64,
}
pub struct PredictiveCache<V> {
max_size: usize,
prediction_depth: usize,
max_history: usize,
cache: HashMap<String, Entry<V>>,
lru_counter: u64,
access_history: VecDeque<String>,
access_patterns: HashMap<String, HashMap<String, u64>>,
prewarmed: HashMap<String, bool>,
pub stats: CacheStats,
}
impl<V: Clone> PredictiveCache<V> {
#[must_use]
pub fn new(max_size: usize, prediction_depth: usize) -> Self {
Self {
max_size,
prediction_depth,
max_history: 100,
cache: HashMap::new(),
lru_counter: 0,
access_history: VecDeque::with_capacity(100),
access_patterns: HashMap::new(),
prewarmed: HashMap::new(),
stats: CacheStats::default(),
}
}
pub fn get(&mut self, key: &str) -> Option<V> {
if self.prewarmed.remove(key).is_some() {
self.stats.prediction_hits += 1;
}
let found = self.cache.get(key).map(|e| e.value.clone());
if let Some(value) = found {
self.stats.hits += 1;
self.lru_counter += 1;
if let Some(entry) = self.cache.get_mut(key) {
entry.order = self.lru_counter;
}
self.record_access(key);
self.predict_next(key);
Some(value)
} else {
self.stats.misses += 1;
self.record_access(key);
None
}
}
pub fn set(&mut self, key: &str, value: V) {
self.lru_counter += 1;
if let Some(entry) = self.cache.get_mut(key) {
entry.value = value;
entry.order = self.lru_counter;
return;
}
self.cache.insert(
key.to_string(),
Entry {
value,
order: self.lru_counter,
},
);
if self.cache.len() > self.max_size {
self.evict_lru();
}
}
pub fn prewarm<F>(&mut self, loader: F, keys: &[String])
where
F: Fn(&str) -> Option<V>,
{
for key in keys {
if !self.cache.contains_key(key) {
if let Some(value) = loader(key) {
self.set(key, value);
}
}
}
}
pub fn invalidate(&mut self, key: &str) {
self.cache.remove(key);
self.prewarmed.remove(key);
}
pub fn clear(&mut self) {
self.cache.clear();
self.prewarmed.clear();
self.access_history.clear();
self.access_patterns.clear();
self.stats = CacheStats::default();
}
#[must_use]
pub fn len(&self) -> usize {
self.cache.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.cache.is_empty()
}
#[must_use]
pub fn likely_next(&self, current_key: &str, top_n: usize) -> Vec<(String, f64)> {
match self.access_patterns.get(current_key) {
None => Vec::new(),
Some(transitions) => {
let total: u64 = transitions.values().sum();
if total == 0 {
return Vec::new();
}
let mut result: Vec<(String, f64)> = transitions
.iter()
.map(|(k, &count)| (k.clone(), count as f64 / total as f64))
.collect();
result.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
result.truncate(top_n);
result
}
}
}
#[must_use]
pub fn hot_keys(&self, top_n: usize) -> Vec<(String, u64)> {
let mut counts: HashMap<String, u64> = HashMap::new();
for key in &self.access_history {
*counts.entry(key.clone()).or_insert(0) += 1;
}
let mut result: Vec<(String, u64)> = counts.into_iter().collect();
result.sort_by_key(|x| std::cmp::Reverse(x.1));
result.truncate(top_n);
result
}
fn evict_lru(&mut self) {
if self.cache.is_empty() {
return;
}
let min_key = self
.cache
.iter()
.min_by_key(|(_, e)| e.order)
.map(|(k, _)| k.clone());
if let Some(key) = min_key {
self.cache.remove(&key);
self.stats.evictions += 1;
}
}
fn record_access(&mut self, key: &str) {
self.access_history.push_back(key.to_string());
if self.access_history.len() > self.max_history {
self.access_history.pop_front();
}
if self.access_history.len() >= 2 {
let prev_key = self
.access_history
.get(self.access_history.len() - 2)
.cloned();
if let Some(prev) = prev_key {
*self
.access_patterns
.entry(prev)
.or_default()
.entry(key.to_string())
.or_insert(0) += 1;
}
}
}
fn predict_next(&mut self, current_key: &str) {
let transitions = match self.access_patterns.get(current_key) {
None => return,
Some(t) => t.clone(),
};
let total: u64 = transitions.values().sum();
if total == 0 {
return;
}
let mut likely: Vec<(String, f64)> = transitions
.iter()
.map(|(k, &count)| (k.clone(), count as f64 / total as f64))
.collect();
likely.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
likely.truncate(self.prediction_depth);
for (next_key, probability) in likely {
if probability > 0.3 {
self.prewarmed.insert(next_key, true);
self.stats.predictions += 1;
}
}
}
}
impl<V: Clone> Default for PredictiveCache<V> {
fn default() -> Self {
Self::new(1000, 3)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cache_miss_then_hit() {
let mut cache: PredictiveCache<String> = PredictiveCache::new(10, 3);
assert!(cache.get("missing").is_none());
assert_eq!(cache.stats.misses, 1);
cache.set("key", "value".to_string());
assert_eq!(cache.get("key"), Some("value".to_string()));
assert_eq!(cache.stats.hits, 1);
}
#[test]
fn cache_eviction() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(3, 1);
cache.set("a", 1);
cache.set("b", 2);
cache.set("c", 3);
cache.set("d", 4);
assert_eq!(cache.len(), 3);
assert!(cache.get("a").is_none());
assert!(cache.get("d").is_some());
assert_eq!(cache.stats.evictions, 1);
}
#[test]
fn cache_lru_order_updates_on_get() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(3, 1);
cache.set("a", 1);
cache.set("b", 2);
cache.set("c", 3);
let _ = cache.get("a");
cache.set("d", 4);
assert!(cache.get("a").is_some());
assert!(cache.get("b").is_none());
}
#[test]
fn cache_markov_prediction() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(10, 3);
for key in &["a", "b", "c", "a", "b", "c"] {
cache.set(key, 1);
let _ = cache.get(key);
}
let likely = cache.likely_next("a", 5);
assert!(!likely.is_empty());
assert_eq!(likely[0].0, "b");
}
#[test]
fn cache_prediction_hit_tracking() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(10, 3);
cache.set("a", 1);
let _ = cache.get("a");
cache.set("b", 2);
let _ = cache.get("b");
let _ = cache.get("a");
assert!(cache.stats.predictions > 0);
let _ = cache.get("b");
assert!(cache.stats.prediction_hits > 0);
}
#[test]
fn cache_hot_keys() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(10, 1);
for _ in 0..5 {
let _ = cache.get("hot");
}
for _ in 0..2 {
let _ = cache.get("warm");
}
let hot = cache.hot_keys(10);
assert_eq!(hot[0].0, "hot");
assert!(hot[0].1 > hot[1].1);
}
#[test]
fn cache_invalidate() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(10, 1);
cache.set("a", 1);
cache.invalidate("a");
assert!(cache.get("a").is_none());
}
#[test]
fn cache_clear() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(10, 1);
cache.set("a", 1);
cache.set("b", 2);
let _ = cache.get("a");
cache.clear();
assert!(cache.is_empty());
assert_eq!(cache.stats.hits, 0);
}
#[test]
fn cache_prewarm() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(10, 1);
let loader = |key: &str| -> Option<i32> { i32::try_from(key.len()).ok() };
cache.prewarm(loader, &["alpha".to_string(), "beta".to_string()]);
assert_eq!(cache.get("alpha"), Some(5));
assert_eq!(cache.get("beta"), Some(4));
}
#[test]
fn cache_stats_hit_rate() {
let mut cache: PredictiveCache<i32> = PredictiveCache::new(10, 1);
cache.set("a", 1);
let _ = cache.get("a"); let _ = cache.get("a"); let _ = cache.get("b");
assert!((cache.stats.hit_rate() - 2.0 / 3.0).abs() < 0.01);
}
#[test]
fn cache_likely_next_empty_for_unknown() {
let cache: PredictiveCache<i32> = PredictiveCache::new(10, 3);
assert!(cache.likely_next("unknown", 5).is_empty());
}
}