Skip to main content

guinea_core/
shared_state.rs

1use std::any::{Any, TypeId};
2use std::collections::HashMap;
3use std::sync::{Arc, RwLock};
4
5type Service = Arc<dyn Any + Send + Sync>;
6
7/// A lock poisoned by a panic elsewhere - distinct from "no such service".
8#[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    /// Like [`get`](Self::get), but tells a poisoned lock apart from a missing
59    /// service.
60    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}