use std::fmt::Display;
pub use crate::{MemoizedMemorySize, MemorySize};
#[derive(Debug, Clone, PartialEq)]
pub struct MemoizedValue<T: Display> {
maybe_value: Option<T>,
is_dirty: bool,
}
impl<T: Display> Default for MemoizedValue<T> {
fn default() -> Self {
Self {
maybe_value: None,
is_dirty: true,
}
}
}
impl<T: Display> MemoizedValue<T> {
#[must_use]
pub fn new() -> Self { Self::default() }
pub fn invalidate(&mut self) { self.is_dirty = true; }
pub fn upsert<F>(&mut self, calculate_fn: F)
where
F: FnOnce() -> T,
{
if self.is_dirty || self.maybe_value.is_none() {
self.maybe_value = Some(calculate_fn());
self.is_dirty = false;
}
}
#[must_use]
pub fn get_cached(&self) -> Option<&T> {
if self.is_dirty {
None
} else {
self.maybe_value.as_ref()
}
}
pub fn get_or_insert_with<F>(&mut self, calculate_fn: F) -> &T
where
F: FnOnce() -> T,
{
self.upsert(calculate_fn);
debug_assert!(
self.maybe_value.is_some(),
"Cached value should be set after upsert"
);
self.maybe_value
.as_ref()
.expect("Value should be cached after upsert")
}
#[must_use]
pub fn is_dirty(&self) -> bool { self.is_dirty || self.maybe_value.is_none() }
}
#[cfg(test)]
mod tests {
use std::{cell::RefCell, rc::Rc};
use super::*;
#[derive(Debug, Clone, PartialEq)]
struct TestValue {
value: String,
calculation_count: Rc<RefCell<usize>>,
}
impl TestValue {
fn with_shared_counter(value: &str, counter: Rc<RefCell<usize>>) -> Self {
Self {
value: value.to_string(),
calculation_count: counter,
}
}
}
impl Display for TestValue {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
*self.calculation_count.borrow_mut() += 1;
write!(f, "{}", self.value)
}
}
#[test]
fn test_new_creates_empty_dirty_cache() {
let cache: MemoizedValue<TestValue> = MemoizedValue::new();
assert!(cache.is_dirty());
assert!(cache.get_cached().is_none());
}
#[test]
fn test_default_creates_empty_dirty_cache() {
let cache: MemoizedValue<TestValue> = MemoizedValue::default();
assert!(cache.is_dirty());
assert!(cache.get_cached().is_none());
}
#[test]
fn test_upsert_calculates_when_dirty() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
let counter_clone = counter.clone();
cache.upsert(|| TestValue::with_shared_counter("test", counter_clone));
assert!(!cache.is_dirty());
assert_eq!(*counter.borrow(), 0);
let cached = cache.get_cached().unwrap();
assert_eq!(cached.value, "test");
}
#[test]
fn test_upsert_does_not_recalculate_when_clean() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
let counter_clone = counter.clone();
cache.upsert(|| TestValue::with_shared_counter("test1", counter_clone.clone()));
cache.upsert(|| TestValue::with_shared_counter("test2", counter_clone));
let cached = cache.get_cached().unwrap();
assert_eq!(cached.value, "test1"); }
#[test]
fn test_invalidate_marks_cache_dirty() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
cache.upsert(|| TestValue::with_shared_counter("test", counter));
assert!(!cache.is_dirty());
cache.invalidate();
assert!(cache.is_dirty());
}
#[test]
fn test_get_cached_returns_none_when_dirty() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
cache.upsert(|| TestValue::with_shared_counter("test", counter));
assert!(cache.get_cached().is_some());
cache.invalidate();
assert!(cache.get_cached().is_none());
}
#[test]
fn test_get_or_insert_with_calculates_when_needed() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
let counter_clone = counter.clone();
let result = cache
.get_or_insert_with(|| TestValue::with_shared_counter("test", counter_clone));
assert_eq!(result.value, "test");
assert!(!cache.is_dirty());
}
#[test]
fn test_get_or_insert_with_returns_cached_when_clean() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
let counter_clone = counter.clone();
let result1_value = cache
.get_or_insert_with(|| {
TestValue::with_shared_counter("test1", counter_clone.clone())
})
.value
.clone();
let result2_value = cache
.get_or_insert_with(|| TestValue::with_shared_counter("test2", counter_clone))
.value
.clone();
assert_eq!(result1_value, "test1");
assert_eq!(result2_value, "test1"); }
#[test]
fn test_is_dirty_tracks_state_correctly() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
assert!(cache.is_dirty());
cache.upsert(|| TestValue::with_shared_counter("test", counter));
assert!(!cache.is_dirty());
cache.invalidate();
assert!(cache.is_dirty());
}
#[test]
fn test_clone_preserves_state() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
cache.upsert(|| TestValue::with_shared_counter("test", counter));
let cloned = cache.clone();
assert_eq!(cache.is_dirty(), cloned.is_dirty());
assert_eq!(
cache.get_cached().map(|v| &v.value),
cloned.get_cached().map(|v| &v.value)
);
}
#[test]
fn test_partial_eq_works_correctly() {
let mut cache1 = MemoizedValue::new();
let mut cache2 = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
assert_eq!(cache1, cache2);
cache1.upsert(|| TestValue::with_shared_counter("test", counter.clone()));
cache2.upsert(|| TestValue::with_shared_counter("test", counter));
assert_eq!(cache1, cache2);
cache1.invalidate();
assert_ne!(cache1, cache2);
}
#[test]
fn test_debug_format_shows_internal_state() {
let mut cache = MemoizedValue::new();
let counter = Rc::new(RefCell::new(0));
let debug_empty = format!("{cache:?}");
assert!(debug_empty.contains("maybe_value: None"));
assert!(debug_empty.contains("is_dirty: true"));
cache.upsert(|| TestValue::with_shared_counter("test", counter));
let debug_filled = format!("{cache:?}");
assert!(debug_filled.contains("maybe_value: Some"));
assert!(debug_filled.contains("is_dirty: false"));
}
#[test]
fn test_with_string_type() {
let mut cache: MemoizedValue<String> = MemoizedValue::new();
let call_count = Rc::new(RefCell::new(0));
let call_count_clone = call_count.clone();
let result = cache.get_or_insert_with(|| {
*call_count_clone.borrow_mut() += 1;
"Hello, World!".to_string()
});
assert_eq!(result, "Hello, World!");
assert_eq!(*call_count.borrow(), 1);
let result2 = cache.get_or_insert_with(|| {
*call_count_clone.borrow_mut() += 1;
"Should not be called".to_string()
});
assert_eq!(result2, "Hello, World!");
assert_eq!(*call_count.borrow(), 1); }
#[test]
fn test_with_numeric_type() {
let mut cache: MemoizedValue<i32> = MemoizedValue::new();
let calculation_calls = Rc::new(RefCell::new(0));
let calculation_calls_clone = calculation_calls.clone();
cache.upsert(|| {
*calculation_calls_clone.borrow_mut() += 1;
42
});
assert_eq!(cache.get_cached(), Some(&42));
assert_eq!(*calculation_calls.borrow(), 1);
cache.upsert(|| {
*calculation_calls_clone.borrow_mut() += 1;
99
});
assert_eq!(cache.get_cached(), Some(&42)); assert_eq!(*calculation_calls.borrow(), 1); }
}