use std::collections::HashMap;
use std::hash::{BuildHasher, Hash};
pub(crate) struct GenerationalCache<K, V, S> {
cur: HashMap<K, V, S>,
prev: HashMap<K, V, S>,
per_generation_cap: usize,
}
impl<K, V, S> GenerationalCache<K, V, S>
where
K: Eq + Hash,
S: BuildHasher + Default,
{
pub(crate) fn new(per_generation_cap: usize) -> Self {
Self {
cur: HashMap::default(),
prev: HashMap::default(),
per_generation_cap: per_generation_cap.max(1),
}
}
pub(crate) fn get(&self, key: &K) -> Option<&V> {
self.cur.get(key).or_else(|| self.prev.get(key))
}
pub(crate) fn insert(&mut self, key: K, value: V) {
if self.cur.len() >= self.per_generation_cap {
std::mem::swap(&mut self.cur, &mut self.prev);
self.cur.clear();
}
self.cur.insert(key, value);
}
#[cfg(test)]
pub(crate) fn len(&self) -> usize {
self.cur.len() + self.prev.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use rustc_hash::FxBuildHasher;
type Cache<K, V> = GenerationalCache<K, V, FxBuildHasher>;
#[test]
fn keeps_previous_generation_so_cyclic_access_still_hits() {
let cap = 8;
let mut c: Cache<u32, u32> = Cache::new(cap);
for k in 0..(2 * cap as u32) {
c.insert(k, k);
}
let mut hits = 0;
for k in 0..(2 * cap as u32) {
if c.get(&k).copied() == Some(k) {
hits += 1;
}
}
assert_eq!(hits, 2 * cap, "both live generations must be retained");
assert!(c.len() <= 2 * cap);
}
#[test]
fn evicts_oldest_generation_first() {
let cap = 4;
let mut c: Cache<u32, u32> = Cache::new(cap);
for k in 0..4 {
c.insert(k, k);
}
for k in 4..8 {
c.insert(k, k);
}
c.insert(8, 8);
assert_eq!(c.get(&0), None, "oldest generation should be evicted");
assert_eq!(c.get(&4).copied(), Some(4), "prior generation retained");
assert_eq!(c.get(&8).copied(), Some(8), "newest entry present");
}
#[test]
fn insert_overwrites_within_current_generation() {
let mut c: Cache<u32, u32> = Cache::new(8);
c.insert(1, 10);
c.insert(1, 20);
assert_eq!(c.get(&1).copied(), Some(20));
}
}