use std::ops::Deref;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
#[derive(Debug)]
pub struct Snapshot<T> {
inner: Arc<T>,
version: u64,
}
impl<T> Snapshot<T> {
#[inline]
pub fn new(data: T) -> Self {
Self {
inner: Arc::new(data),
version: 0,
}
}
#[inline]
pub(crate) fn from_arc(arc: Arc<T>, version: u64) -> Self {
Self {
inner: arc,
version,
}
}
#[inline]
pub fn version(&self) -> u64 {
self.version
}
#[inline]
pub fn into_arc(self) -> Arc<T> {
self.inner
}
#[inline]
pub fn as_arc(&self) -> &Arc<T> {
&self.inner
}
#[inline]
pub fn same_source(&self, other: &Snapshot<T>) -> bool {
Arc::ptr_eq(&self.inner, &other.inner)
}
#[inline]
pub fn ref_count(&self) -> usize {
Arc::strong_count(&self.inner)
}
}
impl<T> Clone for Snapshot<T> {
#[inline]
fn clone(&self) -> Self {
COW_STATS.record_snapshot_clone();
Self {
inner: Arc::clone(&self.inner),
version: self.version,
}
}
}
impl<T> Deref for Snapshot<T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target {
&self.inner
}
}
impl<T> AsRef<T> for Snapshot<T> {
#[inline]
fn as_ref(&self) -> &T {
&self.inner
}
}
pub struct CowState<T> {
current: RwLock<Arc<T>>,
version: AtomicU64,
}
impl<T> CowState<T> {
pub fn new(value: T) -> Self {
COW_STATS.record_state_created();
Self {
current: RwLock::new(Arc::new(value)),
version: AtomicU64::new(1),
}
}
#[inline]
pub fn snapshot(&self) -> Snapshot<T> {
let arc = self.current.read().unwrap();
let version = self.version.load(Ordering::Acquire);
COW_STATS.record_snapshot_taken();
Snapshot::from_arc(Arc::clone(&arc), version)
}
#[inline]
pub fn version(&self) -> u64 {
self.version.load(Ordering::Acquire)
}
#[inline]
pub fn is_version(&self, version: u64) -> bool {
self.version() == version
}
}
impl<T: Clone> CowState<T> {
pub fn update<F>(&self, f: F)
where
F: FnOnce(&mut T),
{
let mut guard = self.current.write().unwrap();
let mut new_value = (**guard).clone();
f(&mut new_value);
*guard = Arc::new(new_value);
self.version.fetch_add(1, Ordering::Release);
COW_STATS.record_update();
}
pub fn replace(&self, value: T) {
let mut guard = self.current.write().unwrap();
*guard = Arc::new(value);
self.version.fetch_add(1, Ordering::Release);
COW_STATS.record_update();
}
pub fn update_if<P, F>(&self, predicate: P, f: F) -> bool
where
P: FnOnce(&T) -> bool,
F: FnOnce(&mut T),
{
let mut guard = self.current.write().unwrap();
if predicate(&guard) {
let mut new_value = (**guard).clone();
f(&mut new_value);
*guard = Arc::new(new_value);
self.version.fetch_add(1, Ordering::Release);
COW_STATS.record_update();
true
} else {
false
}
}
pub fn compare_and_swap<F>(&self, expected_version: u64, f: F) -> Result<u64, u64>
where
F: FnOnce(&mut T),
{
let mut guard = self.current.write().unwrap();
let current = self.version.load(Ordering::Acquire);
if current != expected_version {
return Err(current);
}
let mut new_value = (**guard).clone();
f(&mut new_value);
*guard = Arc::new(new_value);
let new_version = self.version.fetch_add(1, Ordering::Release) + 1;
COW_STATS.record_update();
Ok(new_version)
}
}
impl<T: Clone + Default> Default for CowState<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for CowState<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let snap = self.current.read().unwrap();
f.debug_struct("CowState")
.field("value", &**snap)
.field("version", &self.version.load(Ordering::Relaxed))
.finish()
}
}
pub struct VersionedState<T> {
inner: CowState<T>,
#[allow(dead_code)]
name: &'static str,
}
impl<T> VersionedState<T> {
pub fn new(value: T) -> Self {
Self {
inner: CowState::new(value),
name: std::any::type_name::<T>(),
}
}
pub fn with_name(value: T, name: &'static str) -> Self {
Self {
inner: CowState::new(value),
name,
}
}
#[inline]
pub fn version(&self) -> u64 {
self.inner.version()
}
#[inline]
pub fn read(&self) -> Snapshot<T> {
self.inner.snapshot()
}
#[inline]
pub fn read_versioned(&self) -> (Snapshot<T>, u64) {
let snap = self.inner.snapshot();
let version = snap.version();
(snap, version)
}
#[inline]
pub fn changed_since(&self, version: u64) -> bool {
self.version() != version
}
}
impl<T: Clone> VersionedState<T> {
pub fn write<F>(&self, f: F)
where
F: FnOnce(&mut T),
{
self.inner.update(f);
}
pub fn set(&self, value: T) {
self.inner.replace(value);
}
pub fn write_if_version<F>(&self, version: u64, f: F) -> Result<u64, u64>
where
F: FnOnce(&mut T),
{
self.inner.compare_and_swap(version, f)
}
}
impl<T: Clone + Default> Default for VersionedState<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for VersionedState<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("VersionedState")
.field("inner", &self.inner)
.field("name", &self.name)
.finish()
}
}
pub struct AtomicState<T> {
current: parking_lot::RwLock<Arc<T>>,
version: AtomicU64,
}
impl<T> AtomicState<T> {
pub fn new(value: Arc<T>) -> Self {
COW_STATS.record_state_created();
Self {
current: parking_lot::RwLock::new(value),
version: AtomicU64::new(1),
}
}
#[inline]
pub fn load(&self) -> Arc<T> {
let arc = Arc::clone(&self.current.read());
COW_STATS.record_snapshot_taken();
arc
}
pub fn store(&self, value: Arc<T>) {
*self.current.write() = value;
self.version.fetch_add(1, Ordering::Release);
COW_STATS.record_update();
}
#[inline]
pub fn version(&self) -> u64 {
self.version.load(Ordering::Acquire)
}
#[inline]
pub fn load_versioned(&self) -> (Arc<T>, u64) {
let value = self.load();
let version = self.version();
(value, version)
}
}
impl<T: Clone> AtomicState<T> {
pub fn update<F>(&self, f: F)
where
F: FnOnce(&T) -> T,
{
let current = self.load();
let new_value = f(¤t);
self.store(Arc::new(new_value));
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for AtomicState<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let value = self.load();
f.debug_struct("AtomicState")
.field("value", &*value)
.field("version", &self.version())
.finish()
}
}
pub struct CachedValue<T> {
value: RwLock<Option<CacheEntry<T>>>,
ttl: std::time::Duration,
}
#[derive(Clone)]
struct CacheEntry<T> {
value: Arc<T>,
created_at: std::time::Instant,
version: u64,
}
impl<T> CachedValue<T> {
pub fn new(ttl: std::time::Duration) -> Self {
Self {
value: RwLock::new(None),
ttl,
}
}
pub fn is_valid(&self) -> bool {
let guard = self.value.read().unwrap();
guard
.as_ref()
.map(|e| e.created_at.elapsed() < self.ttl)
.unwrap_or(false)
}
pub fn get(&self) -> Option<Arc<T>> {
let guard = self.value.read().unwrap();
guard.as_ref().and_then(|e| {
if e.created_at.elapsed() < self.ttl {
COW_STATS.record_cache_hit();
Some(Arc::clone(&e.value))
} else {
COW_STATS.record_cache_miss();
None
}
})
}
pub fn set(&self, value: T) {
let mut guard = self.value.write().unwrap();
let version = guard.as_ref().map(|e| e.version + 1).unwrap_or(1);
*guard = Some(CacheEntry {
value: Arc::new(value),
created_at: std::time::Instant::now(),
version,
});
}
pub fn invalidate(&self) {
let mut guard = self.value.write().unwrap();
*guard = None;
}
pub fn ttl(&self) -> std::time::Duration {
self.ttl
}
pub fn remaining_ttl(&self) -> Option<std::time::Duration> {
let guard = self.value.read().unwrap();
guard.as_ref().and_then(|e| {
let elapsed = e.created_at.elapsed();
if elapsed < self.ttl {
Some(self.ttl - elapsed)
} else {
None
}
})
}
}
impl<T: Clone> CachedValue<T> {
pub fn get_or_compute<F>(&self, f: F) -> Arc<T>
where
F: FnOnce() -> T,
{
{
let guard = self.value.read().unwrap();
if let Some(ref entry) = *guard
&& entry.created_at.elapsed() < self.ttl
{
COW_STATS.record_cache_hit();
return Arc::clone(&entry.value);
}
}
COW_STATS.record_cache_miss();
let mut guard = self.value.write().unwrap();
if let Some(ref entry) = *guard
&& entry.created_at.elapsed() < self.ttl
{
return Arc::clone(&entry.value);
}
let value = Arc::new(f());
let version = guard.as_ref().map(|e| e.version + 1).unwrap_or(1);
*guard = Some(CacheEntry {
value: Arc::clone(&value),
created_at: std::time::Instant::now(),
version,
});
value
}
pub async fn get_or_compute_async<F, Fut>(&self, f: F) -> Arc<T>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = T>,
{
{
let guard = self.value.read().unwrap();
if let Some(ref entry) = *guard
&& entry.created_at.elapsed() < self.ttl
{
COW_STATS.record_cache_hit();
return Arc::clone(&entry.value);
}
}
COW_STATS.record_cache_miss();
let value = Arc::new(f().await);
let mut guard = self.value.write().unwrap();
let version = guard.as_ref().map(|e| e.version + 1).unwrap_or(1);
*guard = Some(CacheEntry {
value: Arc::clone(&value),
created_at: std::time::Instant::now(),
version,
});
value
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for CachedValue<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let is_valid = self.is_valid();
let remaining = self.remaining_ttl();
f.debug_struct("CachedValue")
.field("is_valid", &is_valid)
.field("remaining_ttl", &remaining)
.field("ttl", &self.ttl)
.finish()
}
}
#[derive(Debug, Default)]
pub struct CowStats {
states_created: AtomicU64,
snapshots_taken: AtomicU64,
snapshot_clones: AtomicU64,
updates: AtomicU64,
cache_hits: AtomicU64,
cache_misses: AtomicU64,
}
impl CowStats {
fn record_state_created(&self) {
self.states_created.fetch_add(1, Ordering::Relaxed);
}
fn record_snapshot_taken(&self) {
self.snapshots_taken.fetch_add(1, Ordering::Relaxed);
}
fn record_snapshot_clone(&self) {
self.snapshot_clones.fetch_add(1, Ordering::Relaxed);
}
fn record_update(&self) {
self.updates.fetch_add(1, Ordering::Relaxed);
}
fn record_cache_hit(&self) {
self.cache_hits.fetch_add(1, Ordering::Relaxed);
}
fn record_cache_miss(&self) {
self.cache_misses.fetch_add(1, Ordering::Relaxed);
}
pub fn states_created(&self) -> u64 {
self.states_created.load(Ordering::Relaxed)
}
pub fn snapshots_taken(&self) -> u64 {
self.snapshots_taken.load(Ordering::Relaxed)
}
pub fn snapshot_clones(&self) -> u64 {
self.snapshot_clones.load(Ordering::Relaxed)
}
pub fn updates(&self) -> u64 {
self.updates.load(Ordering::Relaxed)
}
pub fn cache_hits(&self) -> u64 {
self.cache_hits.load(Ordering::Relaxed)
}
pub fn cache_misses(&self) -> u64 {
self.cache_misses.load(Ordering::Relaxed)
}
pub fn cache_hit_ratio(&self) -> f64 {
let hits = self.cache_hits() as f64;
let total = hits + self.cache_misses() as f64;
if total > 0.0 { hits / total } else { 0.0 }
}
pub fn read_write_ratio(&self) -> f64 {
let reads = self.snapshots_taken() as f64;
let writes = self.updates() as f64;
if writes > 0.0 {
reads / writes
} else {
f64::INFINITY
}
}
}
static COW_STATS: CowStats = CowStats {
states_created: AtomicU64::new(0),
snapshots_taken: AtomicU64::new(0),
snapshot_clones: AtomicU64::new(0),
updates: AtomicU64::new(0),
cache_hits: AtomicU64::new(0),
cache_misses: AtomicU64::new(0),
};
pub fn cow_stats() -> &'static CowStats {
&COW_STATS
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_snapshot_basic() {
let snap = Snapshot::new(42);
assert_eq!(*snap, 42);
assert_eq!(snap.version(), 0);
}
#[test]
fn test_snapshot_clone_is_cheap() {
let snap = Snapshot::new(vec![1, 2, 3, 4, 5]);
assert_eq!(snap.ref_count(), 1);
let snap2 = snap.clone();
assert_eq!(snap.ref_count(), 2);
assert_eq!(snap2.ref_count(), 2);
assert!(snap.same_source(&snap2));
}
#[test]
fn test_cow_state_basic() {
let state = CowState::new(10);
assert_eq!(*state.snapshot(), 10);
assert_eq!(state.version(), 1);
}
#[test]
fn test_cow_state_update() {
let state = CowState::new(vec![1, 2, 3]);
let snap1 = state.snapshot();
assert_eq!(state.version(), 1);
state.update(|v| v.push(4));
assert_eq!(state.version(), 2);
let snap2 = state.snapshot();
assert_eq!(*snap1, vec![1, 2, 3]);
assert_eq!(*snap2, vec![1, 2, 3, 4]);
}
#[test]
fn test_cow_state_replace() {
let state = CowState::new("old".to_string());
state.replace("new".to_string());
assert_eq!(*state.snapshot(), "new");
}
#[test]
fn test_cow_state_update_if() {
let state = CowState::new(5);
let updated = state.update_if(|&v| v < 10, |v| *v += 1);
assert!(updated);
assert_eq!(*state.snapshot(), 6);
let updated = state.update_if(|&v| v > 100, |v| *v += 1);
assert!(!updated);
assert_eq!(*state.snapshot(), 6);
}
#[test]
fn test_cow_state_compare_and_swap() {
let state = CowState::new(100);
let version = state.version();
let result = state.compare_and_swap(version, |v| *v += 1);
assert!(result.is_ok());
assert_eq!(*state.snapshot(), 101);
let result = state.compare_and_swap(version, |v| *v += 1);
assert!(result.is_err());
}
#[test]
fn test_versioned_state() {
let state = VersionedState::new(vec!["a", "b"]);
let v1 = state.version();
state.write(|v| v.push("c"));
let v2 = state.version();
assert_ne!(v1, v2);
assert!(state.changed_since(v1));
assert!(!state.changed_since(v2));
}
#[test]
fn test_atomic_state() {
let state = AtomicState::new(Arc::new(42));
assert_eq!(*state.load(), 42);
state.store(Arc::new(100));
assert_eq!(*state.load(), 100);
}
#[test]
fn test_atomic_state_update() {
let state = AtomicState::new(Arc::new(10));
state.update(|&v| v * 2);
assert_eq!(*state.load(), 20);
}
#[test]
fn test_atomic_state_concurrent_load_store() {
let state = Arc::new(AtomicState::new(Arc::new(vec![0u64; 64])));
let readers: Vec<_> = (0..4)
.map(|_| {
let state = Arc::clone(&state);
std::thread::spawn(move || {
for _ in 0..1000 {
let value = state.load();
let first = value[0];
assert!(value.iter().all(|&v| v == first));
}
})
})
.collect();
let writer = {
let state = Arc::clone(&state);
std::thread::spawn(move || {
for i in 1..500u64 {
state.store(Arc::new(vec![i; 64]));
}
})
};
for reader in readers {
reader.join().unwrap();
}
writer.join().unwrap();
assert_eq!(state.load()[0], 499);
}
#[test]
fn test_cached_value() {
let cache: CachedValue<i32> = CachedValue::new(std::time::Duration::from_secs(60));
assert!(!cache.is_valid());
assert!(cache.get().is_none());
cache.set(42);
assert!(cache.is_valid());
assert_eq!(*cache.get().unwrap(), 42);
}
#[test]
fn test_cached_value_compute() {
let cache: CachedValue<i32> = CachedValue::new(std::time::Duration::from_secs(60));
let mut computed_count = 0;
let value = cache.get_or_compute(|| {
computed_count += 1;
42
});
assert_eq!(*value, 42);
let value2 = cache.get_or_compute(|| {
computed_count += 1;
99
});
assert_eq!(*value2, 42); }
#[test]
fn test_cow_stats() {
let stats = cow_stats();
let _ = stats.states_created();
let _ = stats.snapshots_taken();
let _ = stats.updates();
let _ = stats.cache_hit_ratio();
let _ = stats.read_write_ratio();
}
}