use std::sync::Arc;
use ahash::AHashMap;
use parking_lot::RwLock;
use crate::error::Result;
#[derive(Debug)]
pub struct SegmentedReaderCache<R> {
inner: RwLock<AHashMap<String, Arc<R>>>,
}
impl<R> Default for SegmentedReaderCache<R> {
fn default() -> Self {
Self {
inner: RwLock::new(AHashMap::default()),
}
}
}
impl<R> SegmentedReaderCache<R> {
pub fn new() -> Self {
Self::default()
}
pub fn get_or_load<F>(&self, segment_id: &str, loader: F) -> Result<Arc<R>>
where
F: FnOnce() -> Result<R>,
{
if let Some(reader) = self.inner.read().get(segment_id) {
return Ok(Arc::clone(reader));
}
let mut guard = self.inner.write();
if let Some(reader) = guard.get(segment_id) {
return Ok(Arc::clone(reader));
}
let reader = Arc::new(loader()?);
guard.insert(segment_id.to_string(), Arc::clone(&reader));
Ok(reader)
}
pub fn invalidate(&self, segment_id: &str) {
self.inner.write().remove(segment_id);
}
pub fn clear(&self) {
self.inner.write().clear();
}
pub fn len(&self) -> usize {
self.inner.read().len()
}
pub fn is_empty(&self) -> bool {
self.inner.read().is_empty()
}
pub fn contains(&self, segment_id: &str) -> bool {
self.inner.read().contains_key(segment_id)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::LaurusError;
use std::sync::atomic::{AtomicUsize, Ordering};
#[test]
fn get_or_load_invokes_loader_on_miss_only() {
let cache: SegmentedReaderCache<()> = SegmentedReaderCache::new();
let calls = AtomicUsize::new(0);
let r1 = cache.get_or_load("seg-a", || {
calls.fetch_add(1, Ordering::SeqCst);
Err(LaurusError::other("intentional miss"))
});
assert!(r1.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert_eq!(cache.len(), 0, "no entry inserted on loader error");
let r2 = cache.get_or_load("seg-a", || {
calls.fetch_add(1, Ordering::SeqCst);
Err(LaurusError::other("still missing"))
});
assert!(r2.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[test]
fn invalidate_is_noop_for_unknown_segment() {
let cache: SegmentedReaderCache<()> = SegmentedReaderCache::new();
cache.invalidate("never-cached");
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
}
#[test]
fn clear_empties_the_cache() {
let cache: SegmentedReaderCache<()> = SegmentedReaderCache::new();
assert!(cache.is_empty());
cache.clear();
assert!(cache.is_empty());
}
#[test]
fn contains_reports_membership() {
let cache: SegmentedReaderCache<()> = SegmentedReaderCache::new();
assert!(!cache.contains("seg-a"));
}
#[test]
fn get_or_load_caches_successful_loads() {
let cache: SegmentedReaderCache<u32> = SegmentedReaderCache::new();
let calls = AtomicUsize::new(0);
let load = || {
calls.fetch_add(1, Ordering::SeqCst);
Ok(42)
};
let r1 = cache.get_or_load("seg-a", load).unwrap();
assert_eq!(*r1, 42);
assert_eq!(cache.len(), 1);
let r2 = cache.get_or_load("seg-a", load).unwrap();
assert_eq!(*r2, 42);
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"second call must hit the cache, not reinvoke the loader"
);
}
}