use crate::{NodeId, Region, TileBuf};
use std::collections::{BTreeMap, HashMap};
use std::sync::Mutex;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct TileKey {
pub node: NodeId,
pub region: Region,
}
impl TileKey {
#[must_use]
pub const fn new(node: NodeId, region: Region) -> Self {
Self { node, region }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub struct CacheStats {
pub hits: u64,
pub misses: u64,
pub evictions: u64,
pub insertions: u64,
pub rejections: u64,
}
impl CacheStats {
#[must_use]
pub fn hit_rate(&self) -> Option<f64> {
let total = self.hits + self.misses;
(total > 0).then(|| self.hits as f64 / total as f64)
}
}
#[derive(Debug)]
struct Entry {
tile: std::sync::Arc<TileBuf>,
bytes: usize,
tick: u64,
}
#[derive(Debug)]
struct Inner {
entries: HashMap<TileKey, Entry>,
recency: BTreeMap<u64, TileKey>,
bytes: usize,
next_tick: u64,
stats: CacheStats,
}
#[derive(Debug)]
pub struct TileCache {
inner: Mutex<Inner>,
budget: usize,
}
impl TileCache {
pub const DEFAULT_BUDGET: usize = 64 * 1024 * 1024;
#[must_use]
pub fn new(budget: usize) -> Self {
Self {
inner: Mutex::new(Inner {
entries: HashMap::new(),
recency: BTreeMap::new(),
bytes: 0,
next_tick: 0,
stats: CacheStats::default(),
}),
budget,
}
}
#[must_use]
pub const fn budget(&self) -> usize {
self.budget
}
#[must_use]
pub fn get(&self, key: &TileKey) -> Option<std::sync::Arc<TileBuf>> {
let mut inner = self.lock();
let tick = inner.next_tick;
let Some(entry) = inner.entries.get_mut(key) else {
inner.stats.misses += 1;
return None;
};
let previous = entry.tick;
entry.tick = tick;
let tile = std::sync::Arc::clone(&entry.tile);
inner.next_tick += 1;
inner.recency.remove(&previous);
inner.recency.insert(tick, *key);
inner.stats.hits += 1;
Some(tile)
}
pub fn insert(&self, key: TileKey, tile: std::sync::Arc<TileBuf>) -> std::sync::Arc<TileBuf> {
let bytes = tile.bytes().len();
let mut inner = self.lock();
inner.stats.insertions += 1;
if bytes > self.budget {
inner.stats.rejections += 1;
return tile;
}
if let Some(previous) = inner.entries.remove(&key) {
inner.bytes -= previous.bytes;
inner.recency.remove(&previous.tick);
}
let tick = inner.next_tick;
inner.next_tick += 1;
inner.bytes += bytes;
inner.entries.insert(
key,
Entry {
tile: std::sync::Arc::clone(&tile),
bytes,
tick,
},
);
inner.recency.insert(tick, key);
inner.evict_to_fit(self.budget);
tile
}
#[must_use]
pub fn bytes_used(&self) -> usize {
self.lock().bytes
}
#[must_use]
pub fn len(&self) -> usize {
self.lock().entries.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[must_use]
pub fn stats(&self) -> CacheStats {
self.lock().stats
}
pub fn clear(&self) {
let mut inner = self.lock();
inner.entries.clear();
inner.recency.clear();
inner.bytes = 0;
}
fn lock(&self) -> std::sync::MutexGuard<'_, Inner> {
self.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
}
impl Default for TileCache {
fn default() -> Self {
Self::new(Self::DEFAULT_BUDGET)
}
}
impl Inner {
fn evict_to_fit(&mut self, budget: usize) {
while self.bytes > budget {
let Some((&tick, &key)) = self.recency.iter().next() else {
break;
};
self.recency.remove(&tick);
if let Some(entry) = self.entries.remove(&key) {
self.bytes -= entry.bytes;
self.stats.evictions += 1;
}
}
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::indexing_slicing,
clippy::panic,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
use crate::PixelFormat;
use std::sync::Arc;
fn tile(size: u32) -> Arc<TileBuf> {
Arc::new(TileBuf::zeroed(Region::from_size(size, 1), PixelFormat::Gray8).unwrap())
}
fn key(region: Region) -> TileKey {
TileKey::new(fresh_node_id(), region)
}
fn fresh_node_id() -> NodeId {
use crate::testing::CountingProducer;
use crate::{Format, Image, ImageDescriptor};
let descriptor = ImageDescriptor::new(1, 1, PixelFormat::Gray8).unwrap();
let image = Image::from_producer(Arc::new(CountingProducer::new(descriptor)), Format::Raw);
image.node().id()
}
#[test]
fn a_tile_round_trips_through_the_cache() {
let cache = TileCache::new(1024);
let k = key(Region::from_size(4, 1));
assert!(cache.get(&k).is_none(), "empty cache misses");
cache.insert(k, tile(4));
let found = cache.get(&k).unwrap();
assert_eq!(found.bytes().len(), 4);
assert_eq!(cache.len(), 1);
assert_eq!(cache.bytes_used(), 4);
}
#[test]
fn the_budget_is_never_exceeded_by_retained_bytes() {
let cache = TileCache::new(100);
for _ in 0..50 {
cache.insert(key(Region::from_size(10, 1)), tile(10));
assert!(
cache.bytes_used() <= 100,
"retained {} bytes over a 100 byte budget",
cache.bytes_used()
);
}
assert!(cache.stats().evictions > 0, "nothing was ever evicted");
}
#[test]
fn eviction_removes_the_least_recently_used_entry() {
let cache = TileCache::new(30);
let (a, b, c) = (
key(Region::from_size(1, 1)),
key(Region::from_size(2, 1)),
key(Region::from_size(3, 1)),
);
cache.insert(a, tile(10));
cache.insert(b, tile(10));
assert!(cache.get(&a).is_some());
cache.insert(c, tile(10));
assert_eq!(cache.len(), 3);
let d = key(Region::from_size(4, 1));
cache.insert(d, tile(10));
assert!(cache.get(&b).is_none(), "LRU entry survived");
assert!(cache.get(&a).is_some(), "recently used entry was evicted");
assert!(cache.get(&c).is_some());
assert!(cache.get(&d).is_some());
}
#[test]
fn an_evicted_tile_stays_valid_for_whoever_holds_it() {
let cache = TileCache::new(10);
let k = key(Region::from_size(10, 1));
let held = cache.insert(k, tile(10));
cache.insert(key(Region::from_size(9, 1)), tile(10));
assert!(cache.get(&k).is_none(), "expected eviction");
assert_eq!(held.bytes().len(), 10);
assert!(held.as_tile().is_ok());
}
#[test]
fn a_tile_larger_than_the_budget_is_returned_but_not_retained() {
let cache = TileCache::new(10);
let k = key(Region::from_size(50, 1));
let returned = cache.insert(k, tile(50));
assert_eq!(returned.bytes().len(), 50, "the tile is still usable");
assert!(cache.get(&k).is_none());
assert_eq!(cache.len(), 0);
assert_eq!(cache.stats().rejections, 1);
assert_eq!(cache.stats().evictions, 0);
}
#[test]
fn a_zero_budget_disables_retention() {
let cache = TileCache::new(0);
let k = key(Region::from_size(4, 1));
let returned = cache.insert(k, tile(4));
assert_eq!(returned.bytes().len(), 4, "the tile is still returned");
assert!(cache.get(&k).is_none());
assert_eq!(cache.bytes_used(), 0);
}
#[test]
fn reinserting_a_key_replaces_it_without_double_counting() {
let cache = TileCache::new(1000);
let k = key(Region::from_size(4, 1));
cache.insert(k, tile(10));
cache.insert(k, tile(20));
assert_eq!(cache.len(), 1);
assert_eq!(cache.bytes_used(), 20, "old bytes were not released");
assert_eq!(cache.get(&k).unwrap().bytes().len(), 20);
}
#[test]
fn keys_distinguish_node_and_region() {
let cache = TileCache::new(1000);
let node = fresh_node_id();
let a = TileKey::new(node, Region::from_size(4, 1));
let b = TileKey::new(node, Region::new(4, 0, 4, 1));
cache.insert(a, tile(4));
assert!(
cache.get(&b).is_none(),
"different regions must not collide"
);
let other = TileKey::new(fresh_node_id(), Region::from_size(4, 1));
assert!(
cache.get(&other).is_none(),
"different nodes must not collide"
);
}
#[test]
fn statistics_track_lookups_and_evictions() {
let cache = TileCache::new(10);
let k = key(Region::from_size(4, 1));
assert!(cache.get(&k).is_none());
cache.insert(k, tile(4));
assert!(cache.get(&k).is_some());
let stats = cache.stats();
assert_eq!(stats.hits, 1);
assert_eq!(stats.misses, 1);
assert_eq!(stats.insertions, 1);
assert_eq!(stats.hit_rate(), Some(0.5));
assert_eq!(TileCache::new(1).stats().hit_rate(), None, "no lookups yet");
}
#[test]
fn clear_drops_entries_but_keeps_statistics() {
let cache = TileCache::new(1000);
cache.insert(key(Region::from_size(4, 1)), tile(4));
assert!(!cache.is_empty());
cache.clear();
assert!(cache.is_empty());
assert_eq!(cache.bytes_used(), 0);
assert_eq!(cache.stats().insertions, 1, "statistics are cumulative");
}
#[test]
fn the_cache_is_usable_from_many_threads() {
let cache = Arc::new(TileCache::new(4096));
let keys: Vec<TileKey> = (0..8).map(|i| key(Region::from_size(i + 1, 1))).collect();
std::thread::scope(|scope| {
for _ in 0..8 {
let cache = Arc::clone(&cache);
let keys = keys.clone();
scope.spawn(move || {
for _ in 0..200 {
for (i, k) in keys.iter().enumerate() {
cache.insert(*k, tile(i as u32 + 1));
let _ = cache.get(k);
}
}
});
}
});
let inner = cache.lock();
let actual: usize = inner.entries.values().map(|e| e.bytes).sum();
assert_eq!(inner.bytes, actual, "byte accounting drifted");
assert_eq!(
inner.recency.len(),
inner.entries.len(),
"recency index drifted"
);
assert!(inner.bytes <= 4096);
}
#[test]
fn a_poisoned_lock_does_not_disable_the_cache() {
let cache = Arc::new(TileCache::new(1000));
let k = key(Region::from_size(4, 1));
cache.insert(k, tile(4));
let poisoner = Arc::clone(&cache);
let handle = std::thread::spawn(move || {
let _guard = poisoner.lock();
panic!("poison the mutex");
});
assert!(handle.join().is_err(), "the thread was supposed to panic");
assert!(cache.get(&k).is_some());
cache.insert(key(Region::from_size(8, 1)), tile(8));
}
}