use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::fmt;
use std::sync::{Arc, RwLock};
#[derive(Clone, Default)]
pub struct SharedState {
inner: Arc<RwLock<HashMap<TypeId, Box<dyn Any + Send + Sync>>>>,
}
impl fmt::Debug for SharedState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let entries = self.inner.read().unwrap_or_else(|e| e.into_inner()).len();
f.debug_struct("SharedState")
.field("entries", &entries)
.finish()
}
}
impl SharedState {
pub fn new() -> Self {
Self::default()
}
pub fn insert<T: Any + Send + Sync>(&self, value: T) {
self.inner
.write()
.unwrap_or_else(|e| e.into_inner())
.insert(TypeId::of::<T>(), Box::new(value));
}
pub fn get<T: Any + Clone>(&self) -> Option<T> {
self.inner
.read()
.unwrap_or_else(|e| e.into_inner())
.get(&TypeId::of::<T>())
.and_then(|v| v.downcast_ref::<T>())
.cloned()
}
pub fn with<T: Any>(&self, f: impl FnOnce(&T)) {
let guard = self.inner.read().unwrap_or_else(|e| e.into_inner());
if let Some(value) = guard
.get(&TypeId::of::<T>())
.and_then(|v| v.downcast_ref::<T>())
{
f(value);
}
}
pub fn with_mut<T: Any>(&self, f: impl FnOnce(&mut T)) {
let mut guard = self.inner.write().unwrap_or_else(|e| e.into_inner());
if let Some(value) = guard
.get_mut(&TypeId::of::<T>())
.and_then(|v| v.downcast_mut::<T>())
{
f(value);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, Clone, PartialEq)]
struct Session {
user: String,
}
#[derive(Debug, Clone, PartialEq)]
struct Counter(usize);
#[test]
fn insert_and_get_by_type() {
let state = SharedState::new();
state.insert(Session {
user: "alice".into(),
});
state.insert(Counter(3));
assert_eq!(
state.get::<Session>(),
Some(Session {
user: "alice".into()
})
);
assert_eq!(state.get::<Counter>(), Some(Counter(3)));
}
#[test]
fn same_type_insert_overwrites_others_untouched() {
let state = SharedState::new();
state.insert(Session {
user: "alice".into(),
});
state.insert(Counter(1));
state.insert(Counter(2));
assert_eq!(state.get::<Counter>(), Some(Counter(2)));
assert_eq!(
state.get::<Session>(),
Some(Session {
user: "alice".into()
})
);
}
#[test]
fn get_missing_type_returns_none() {
let state = SharedState::new();
assert_eq!(state.get::<Session>(), None);
}
#[test]
fn with_and_with_mut_lock_inner() {
let state = SharedState::new();
state.insert(Counter(1));
let mut called = false;
state.with::<Session>(|_| called = true);
assert!(!called);
state.with::<Counter>(|c| {
assert_eq!(c.0, 1);
});
state.with_mut::<Counter>(|c| c.0 += 1);
assert_eq!(state.get::<Counter>(), Some(Counter(2)));
}
#[test]
fn clone_shares_same_instance() {
let state = SharedState::new();
let tool_side = state.clone();
state.insert(Counter(5));
assert_eq!(tool_side.get::<Counter>(), Some(Counter(5)));
tool_side.with_mut::<Counter>(|c| c.0 += 1);
assert_eq!(state.get::<Counter>(), Some(Counter(6)));
}
#[test]
fn poisoned_lock_recovers_and_stays_usable() {
let state = SharedState::new();
state.insert(Counter(1));
let panicked = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
state.with_mut::<Counter>(|_| panic!("user closure panicked"));
}));
assert!(panicked.is_err());
state.with_mut::<Counter>(|c| c.0 += 1);
assert_eq!(state.get::<Counter>(), Some(Counter(2)));
}
}