use parking_lot::{RwLock, RwLockReadGuard, RwLockUpgradableReadGuard, RwLockWriteGuard};
use std::collections::HashMap;
use std::hash::Hash;
use std::ops::Deref;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
#[derive(Debug, Default)]
pub struct ReadStateStats {
reads: AtomicU64,
writes: AtomicU64,
upgrades: AtomicU64,
upgrade_failures: AtomicU64,
}
impl ReadStateStats {
pub fn reads(&self) -> u64 {
self.reads.load(Ordering::Relaxed)
}
pub fn writes(&self) -> u64 {
self.writes.load(Ordering::Relaxed)
}
pub fn upgrades(&self) -> u64 {
self.upgrades.load(Ordering::Relaxed)
}
pub fn read_write_ratio(&self) -> f64 {
let writes = self.writes();
if writes == 0 {
f64::INFINITY
} else {
self.reads() as f64 / writes as f64
}
}
fn record_read(&self) {
self.reads.fetch_add(1, Ordering::Relaxed);
}
fn record_write(&self) {
self.writes.fetch_add(1, Ordering::Relaxed);
}
fn record_upgrade(&self) {
self.upgrades.fetch_add(1, Ordering::Relaxed);
}
#[allow(dead_code)]
fn record_upgrade_failure(&self) {
self.upgrade_failures.fetch_add(1, Ordering::Relaxed);
}
}
pub static READ_STATE_STATS: ReadStateStats = ReadStateStats {
reads: AtomicU64::new(0),
writes: AtomicU64::new(0),
upgrades: AtomicU64::new(0),
upgrade_failures: AtomicU64::new(0),
};
pub struct ReadState<T> {
inner: RwLock<T>,
version: AtomicU64,
}
impl<T> ReadState<T> {
#[inline]
pub fn new(value: T) -> Self {
Self {
inner: RwLock::new(value),
version: AtomicU64::new(1),
}
}
#[inline]
pub fn read(&self) -> ReadGuard<'_, T> {
READ_STATE_STATS.record_read();
ReadGuard {
guard: self.inner.read(),
}
}
#[inline]
pub fn try_read(&self) -> Option<ReadGuard<'_, T>> {
self.inner.try_read().map(|guard| {
READ_STATE_STATS.record_read();
ReadGuard { guard }
})
}
#[inline]
pub fn upgradeable_read(&self) -> UpgradeableGuard<'_, T> {
READ_STATE_STATS.record_read();
UpgradeableGuard {
guard: self.inner.upgradable_read(),
version: &self.version,
}
}
#[inline]
pub fn write_guard(&self) -> WriteGuard<'_, T> {
READ_STATE_STATS.record_write();
WriteGuard {
guard: self.inner.write(),
version: &self.version,
}
}
#[inline]
pub fn write<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut T) -> R,
{
READ_STATE_STATS.record_write();
let mut guard = self.inner.write();
let result = f(&mut *guard);
self.version.fetch_add(1, Ordering::Release);
result
}
pub fn read_then_write<P, F>(&self, predicate: P, update: F) -> bool
where
P: FnOnce(&T) -> bool,
F: FnOnce(&mut T),
{
let guard = self.inner.upgradable_read();
if predicate(&*guard) {
READ_STATE_STATS.record_upgrade();
let mut write_guard = RwLockUpgradableReadGuard::upgrade(guard);
update(&mut *write_guard);
self.version.fetch_add(1, Ordering::Release);
true
} else {
false
}
}
#[inline]
pub fn version(&self) -> u64 {
self.version.load(Ordering::Acquire)
}
#[inline]
pub fn changed_since(&self, version: u64) -> bool {
self.version() != version
}
#[inline]
pub fn replace(&self, value: T) -> T {
READ_STATE_STATS.record_write();
let mut guard = self.inner.write();
let old = std::mem::replace(&mut *guard, value);
self.version.fetch_add(1, Ordering::Release);
old
}
}
impl<T: Clone> ReadState<T> {
#[inline]
pub fn cloned(&self) -> T {
self.read().clone()
}
#[inline]
pub fn cloned_versioned(&self) -> (T, u64) {
let guard = self.read();
let value = guard.clone();
let version = self.version();
(value, version)
}
}
impl<T: Default> Default for ReadState<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for ReadState<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.try_read() {
Some(guard) => f.debug_struct("ReadState").field("value", &*guard).finish(),
None => f
.debug_struct("ReadState")
.field("value", &"<locked>")
.finish(),
}
}
}
pub struct ReadGuard<'a, T> {
guard: RwLockReadGuard<'a, T>,
}
impl<T> Deref for ReadGuard<'_, T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target {
&self.guard
}
}
impl<T> AsRef<T> for ReadGuard<'_, T> {
#[inline]
fn as_ref(&self) -> &T {
&self.guard
}
}
pub struct UpgradeableGuard<'a, T> {
guard: RwLockUpgradableReadGuard<'a, T>,
version: &'a AtomicU64,
}
impl<T> Deref for UpgradeableGuard<'_, T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target {
&self.guard
}
}
impl<'a, T> UpgradeableGuard<'a, T> {
pub fn upgrade(self) -> WriteGuard<'a, T> {
READ_STATE_STATS.record_upgrade();
WriteGuard {
guard: RwLockUpgradableReadGuard::upgrade(self.guard),
version: self.version,
}
}
pub fn try_upgrade(self) -> Result<WriteGuard<'a, T>, Self> {
match RwLockUpgradableReadGuard::try_upgrade(self.guard) {
Ok(guard) => {
READ_STATE_STATS.record_upgrade();
Ok(WriteGuard {
guard,
version: self.version,
})
}
Err(guard) => Err(UpgradeableGuard {
guard,
version: self.version,
}),
}
}
}
pub struct WriteGuard<'a, T> {
guard: RwLockWriteGuard<'a, T>,
version: &'a AtomicU64,
}
impl<T> Deref for WriteGuard<'_, T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target {
&self.guard
}
}
impl<T> std::ops::DerefMut for WriteGuard<'_, T> {
#[inline]
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.guard
}
}
impl<T> Drop for WriteGuard<'_, T> {
fn drop(&mut self) {
self.version.fetch_add(1, Ordering::Release);
}
}
pub struct ReadCache<K, V> {
inner: RwLock<HashMap<K, V>>,
}
impl<K, V> ReadCache<K, V>
where
K: Eq + Hash,
{
pub fn new() -> Self {
Self {
inner: RwLock::new(HashMap::new()),
}
}
pub fn with_capacity(capacity: usize) -> Self {
Self {
inner: RwLock::new(HashMap::with_capacity(capacity)),
}
}
#[inline]
pub fn get<Q>(&self, key: &Q) -> Option<V>
where
K: std::borrow::Borrow<Q>,
Q: Hash + Eq + ?Sized,
V: Clone,
{
READ_STATE_STATS.record_read();
self.inner.read().get(key).cloned()
}
#[inline]
pub fn contains_key<Q>(&self, key: &Q) -> bool
where
K: std::borrow::Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
READ_STATE_STATS.record_read();
self.inner.read().contains_key(key)
}
#[inline]
pub fn len(&self) -> usize {
self.inner.read().len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.inner.read().is_empty()
}
#[inline]
pub fn insert(&self, key: K, value: V) -> Option<V> {
READ_STATE_STATS.record_write();
self.inner.write().insert(key, value)
}
#[inline]
pub fn remove<Q>(&self, key: &Q) -> Option<V>
where
K: std::borrow::Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
READ_STATE_STATS.record_write();
self.inner.write().remove(key)
}
#[inline]
pub fn clear(&self) {
READ_STATE_STATS.record_write();
self.inner.write().clear();
}
pub fn get_or_insert<F>(&self, key: K, default: F) -> V
where
F: FnOnce() -> V,
V: Clone,
{
{
let guard = self.inner.upgradable_read();
if let Some(value) = guard.get(&key) {
READ_STATE_STATS.record_read();
return value.clone();
}
READ_STATE_STATS.record_upgrade();
let mut write_guard = RwLockUpgradableReadGuard::upgrade(guard);
if let Some(value) = write_guard.get(&key) {
return value.clone();
}
let value = default();
write_guard.insert(key, value.clone());
value
}
}
pub fn extend<I>(&self, iter: I)
where
I: IntoIterator<Item = (K, V)>,
{
READ_STATE_STATS.record_write();
self.inner.write().extend(iter);
}
pub fn keys(&self) -> Vec<K>
where
K: Clone,
{
READ_STATE_STATS.record_read();
self.inner.read().keys().cloned().collect()
}
pub fn values(&self) -> Vec<V>
where
V: Clone,
{
READ_STATE_STATS.record_read();
self.inner.read().values().cloned().collect()
}
pub fn entries(&self) -> Vec<(K, V)>
where
K: Clone,
V: Clone,
{
READ_STATE_STATS.record_read();
self.inner
.read()
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect()
}
pub fn for_each<F>(&self, mut f: F)
where
F: FnMut(&K, &V),
{
READ_STATE_STATS.record_read();
for (k, v) in self.inner.read().iter() {
f(k, v);
}
}
pub fn update<Q, F>(&self, key: &Q, f: F) -> bool
where
K: std::borrow::Borrow<Q>,
Q: Hash + Eq + ?Sized,
F: FnOnce(&mut V),
{
READ_STATE_STATS.record_write();
let mut guard = self.inner.write();
if let Some(value) = guard.get_mut(key) {
f(value);
true
} else {
false
}
}
}
impl<K, V> Default for ReadCache<K, V>
where
K: Eq + Hash,
{
fn default() -> Self {
Self::new()
}
}
impl<K, V> std::fmt::Debug for ReadCache<K, V>
where
K: std::fmt::Debug + Eq + Hash,
V: std::fmt::Debug,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ReadCache")
.field("len", &self.len())
.finish()
}
}
pub struct ReadConfig<T> {
state: ReadState<Arc<T>>,
}
impl<T> ReadConfig<T> {
pub fn new(config: T) -> Self {
Self {
state: ReadState::new(Arc::new(config)),
}
}
#[inline]
pub fn get(&self) -> Arc<T> {
self.state.read().clone()
}
#[inline]
pub fn get_versioned(&self) -> (Arc<T>, u64) {
let config = self.get();
let version = self.version();
(config, version)
}
#[inline]
pub fn version(&self) -> u64 {
self.state.version()
}
#[inline]
pub fn changed_since(&self, version: u64) -> bool {
self.state.changed_since(version)
}
pub fn set(&self, config: T) {
self.state.replace(Arc::new(config));
}
}
impl<T: Clone> ReadConfig<T> {
pub fn update<F>(&self, f: F)
where
F: FnOnce(&T) -> T,
{
self.state.write(|arc| {
let new_value = f(&**arc);
*arc = Arc::new(new_value);
});
}
pub fn update_if<P, F>(&self, predicate: P, update: F) -> bool
where
P: FnOnce(&T) -> bool,
F: FnOnce(&T) -> T,
{
self.state.read_then_write(
|arc| predicate(&**arc),
|arc| {
let new_value = update(&**arc);
*arc = Arc::new(new_value);
},
)
}
}
impl<T: Default> Default for ReadConfig<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for ReadConfig<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ReadConfig")
.field("config", &self.get())
.field("version", &self.version())
.finish()
}
}
pub struct ArcSwapState<T> {
inner: RwLock<Arc<T>>,
}
impl<T> ArcSwapState<T> {
pub fn new(value: T) -> Self {
Self {
inner: RwLock::new(Arc::new(value)),
}
}
#[inline]
pub fn load(&self) -> Arc<T> {
READ_STATE_STATS.record_read();
self.inner.read().clone()
}
pub fn store(&self, value: T) {
READ_STATE_STATS.record_write();
*self.inner.write() = Arc::new(value);
}
pub fn store_arc(&self, value: Arc<T>) {
READ_STATE_STATS.record_write();
*self.inner.write() = value;
}
}
impl<T: Clone> ArcSwapState<T> {
pub fn update<F>(&self, f: F)
where
F: FnOnce(&T) -> T,
{
READ_STATE_STATS.record_write();
let mut guard = self.inner.write();
let new_value = f(&**guard);
*guard = Arc::new(new_value);
}
pub fn swap(&self, value: T) -> Arc<T> {
READ_STATE_STATS.record_write();
std::mem::replace(&mut *self.inner.write(), Arc::new(value))
}
}
impl<T: Default> Default for ArcSwapState<T> {
fn default() -> Self {
Self::new(T::default())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::thread;
#[test]
fn test_read_state_basic() {
let state = ReadState::new(42i32);
assert_eq!(*state.read(), 42);
state.write(|v| *v = 100);
assert_eq!(*state.read(), 100);
}
#[test]
fn test_read_state_concurrent() {
let state = Arc::new(ReadState::new(0i32));
let mut handles = vec![];
for _ in 0..10 {
let s = Arc::clone(&state);
handles.push(thread::spawn(move || {
for _ in 0..1000 {
let _ = *s.read();
}
}));
}
for _ in 0..2 {
let s = Arc::clone(&state);
handles.push(thread::spawn(move || {
for _ in 0..100 {
s.write(|v| *v += 1);
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(*state.read(), 200);
}
#[test]
fn test_read_state_upgrade() {
let state = ReadState::new(10i32);
let upgraded = state.read_then_write(|v| *v == 10, |v| *v = 20);
assert!(upgraded);
assert_eq!(*state.read(), 20);
let upgraded = state.read_then_write(|v| *v == 10, |v| *v = 30);
assert!(!upgraded);
assert_eq!(*state.read(), 20);
}
#[test]
fn test_read_cache_basic() {
let cache: ReadCache<String, i32> = ReadCache::new();
cache.insert("one".to_string(), 1);
cache.insert("two".to_string(), 2);
assert_eq!(cache.get(&"one".to_string()), Some(1));
assert_eq!(cache.get(&"three".to_string()), None);
assert_eq!(cache.len(), 2);
}
#[test]
fn test_read_cache_get_or_insert() {
let cache: ReadCache<String, i32> = ReadCache::new();
let value = cache.get_or_insert("key".to_string(), || 42);
assert_eq!(value, 42);
let value = cache.get_or_insert("key".to_string(), || 100);
assert_eq!(value, 42);
}
#[test]
fn test_read_config() {
#[derive(Clone, Debug)]
struct Config {
value: i32,
}
let config = ReadConfig::new(Config { value: 10 });
let v1 = config.version();
assert_eq!(config.get().value, 10);
config.update(|c| Config {
value: c.value + 10,
});
assert_eq!(config.get().value, 20);
assert!(config.changed_since(v1));
}
#[test]
fn test_arc_swap_state() {
let state = ArcSwapState::new(42i32);
assert_eq!(*state.load(), 42);
state.update(|v| v + 10);
assert_eq!(*state.load(), 52);
state.store(100);
assert_eq!(*state.load(), 100);
}
#[test]
fn test_version_tracking() {
let state = ReadState::new(0i32);
let v1 = state.version();
state.write(|v| *v = 1);
let v2 = state.version();
assert!(v2 > v1);
assert!(state.changed_since(v1));
assert!(!state.changed_since(v2));
}
}