guinea_core/
shared_state.rs1use std::any::{Any, TypeId};
2use std::collections::HashMap;
3use std::sync::{Arc, RwLock};
4
5type Service = Arc<dyn Any + Send + Sync>;
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub struct Poisoned;
10
11impl std::fmt::Display for Poisoned {
12 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
13 f.write_str("shared state lock was poisoned by a panic")
14 }
15}
16
17impl std::error::Error for Poisoned {}
18
19#[derive(Clone, Default)]
20pub struct SharedState {
21 inner: Arc<RwLock<HashMap<TypeId, Service>>>,
22}
23
24impl SharedState {
25 pub fn new() -> Self {
26 Self {
27 inner: Arc::new(RwLock::new(HashMap::new())),
28 }
29 }
30
31 pub fn insert<T>(&self, value: T) -> Option<Arc<T>>
32 where
33 T: Send + Sync + 'static,
34 {
35 let id = TypeId::of::<T>();
36 let mut map = self.inner.write().ok()?;
37 let previous = map.insert(id, Arc::new(value));
38 previous.and_then(|svc| svc.downcast::<T>().ok())
39 }
40
41 pub fn insert_arc<T>(&self, value: Arc<T>) -> Option<Arc<T>>
42 where
43 T: Send + Sync + 'static,
44 {
45 let id = TypeId::of::<T>();
46 let mut map = self.inner.write().ok()?;
47 let previous = map.insert(id, value);
48 previous.and_then(|svc| svc.downcast::<T>().ok())
49 }
50
51 pub fn get<T>(&self) -> Option<Arc<T>>
52 where
53 T: Send + Sync + 'static,
54 {
55 self.try_get().ok().flatten()
56 }
57
58 pub fn try_get<T>(&self) -> Result<Option<Arc<T>>, Poisoned>
61 where
62 T: Send + Sync + 'static,
63 {
64 let map = self.inner.read().map_err(|_| Poisoned)?;
65 Ok(map
66 .get(&TypeId::of::<T>())
67 .and_then(|svc| svc.clone().downcast::<T>().ok()))
68 }
69
70 pub fn remove<T>(&self) -> Option<Arc<T>>
71 where
72 T: Send + Sync + 'static,
73 {
74 let id = TypeId::of::<T>();
75 let mut map = self.inner.write().ok()?;
76 map.remove(&id)?.downcast::<T>().ok()
77 }
78
79 pub fn contains<T>(&self) -> bool
80 where
81 T: Send + Sync + 'static,
82 {
83 let id = TypeId::of::<T>();
84 self.inner
85 .read()
86 .map(|m| m.contains_key(&id))
87 .unwrap_or(false)
88 }
89}
90
91#[cfg(test)]
92mod tests {
93 use super::{Poisoned, SharedState};
94 use std::sync::Arc;
95
96 #[test]
97 fn test_shared_state_clone() {
98 let shared = SharedState::new();
99 let clone = shared.clone();
100
101 shared.insert("hello".to_string());
102
103 assert_eq!(
104 clone.get::<String>().as_ref().map(|s| s.as_str()),
105 Some("hello")
106 );
107 }
108
109 #[test]
110 fn insert_get_remove_roundtrip() {
111 let shared = SharedState::new();
112
113 assert!(!shared.contains::<String>());
114 assert!(shared.insert("hello".to_string()).is_none());
115 assert!(shared.contains::<String>());
116 assert_eq!(
117 shared.get::<String>().as_ref().map(|s| s.as_str()),
118 Some("hello")
119 );
120
121 let removed = shared.remove::<String>();
122 assert_eq!(removed.as_ref().map(|s| s.as_str()), Some("hello"));
123 assert!(shared.get::<String>().is_none());
124 }
125
126 #[test]
127 fn overwrite_returns_previous() {
128 let shared = SharedState::new();
129 shared.insert(7u32);
130 let prev = shared.insert(9u32);
131 assert_eq!(prev.map(|v| *v), Some(7u32));
132 assert_eq!(shared.get::<u32>().map(|v| *v), Some(9u32));
133 }
134
135 #[test]
136 fn insert_arc_keeps_identity() {
137 let shared = SharedState::new();
138 let service = Arc::new(42usize);
139 let ptr = Arc::as_ptr(&service);
140 shared.insert_arc(Arc::clone(&service));
141
142 let fetched = shared.get::<usize>().expect("service should exist");
143 assert_eq!(Arc::as_ptr(&fetched), ptr);
144 }
145
146 #[test]
147 fn try_get_reports_a_missing_service_and_a_poisoned_lock_differently() {
148 let shared = SharedState::new();
149 assert_eq!(shared.try_get::<u32>().map(|v| v.is_none()), Ok(true));
150
151 let poisoner = shared.clone();
152 let _ = std::thread::spawn(move || {
153 let _guard = poisoner.inner.write().unwrap();
154 panic!("poison the lock");
155 })
156 .join();
157
158 assert_eq!(shared.try_get::<u32>().err(), Some(Poisoned));
159 assert_eq!(shared.get::<u32>().map(|v| *v), None);
160 }
161}