1use crate::error::Result;
15use crate::value::Value;
16use std::collections::HashMap;
17use std::sync::{Arc, Mutex};
18
19pub trait StateStore: Send + Sync {
25 fn get(&self, node_id: &str) -> Result<Option<Arc<Value>>>;
27
28 fn set(&self, node_id: &str, state: Value) -> Result<()>;
30
31 fn remove(&self, node_id: &str) -> Result<()>;
33
34 fn clear(&self) -> Result<()>;
36
37 fn keys(&self) -> Result<Vec<String>>;
39}
40
41#[derive(Default)]
46pub struct MemoryStateStore {
47 inner: Mutex<HashMap<String, Arc<Value>>>,
48}
49
50impl MemoryStateStore {
51 pub fn new() -> Self {
53 Self::default()
54 }
55
56 fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<String, Arc<Value>>> {
64 self.inner.lock().unwrap_or_else(|e| e.into_inner())
65 }
66}
67
68impl StateStore for MemoryStateStore {
69 fn get(&self, node_id: &str) -> Result<Option<Arc<Value>>> {
70 Ok(self.lock().get(node_id).cloned())
71 }
72
73 fn set(&self, node_id: &str, state: Value) -> Result<()> {
74 self.lock().insert(node_id.to_string(), Arc::new(state));
75 Ok(())
76 }
77
78 fn remove(&self, node_id: &str) -> Result<()> {
79 self.lock().remove(node_id);
80 Ok(())
81 }
82
83 fn clear(&self) -> Result<()> {
84 self.lock().clear();
85 Ok(())
86 }
87
88 fn keys(&self) -> Result<Vec<String>> {
89 Ok(self.lock().keys().cloned().collect())
90 }
91}
92
93#[cfg(test)]
94mod tests {
95 use super::*;
96
97 #[test]
98 fn memory_store_roundtrip() {
99 let store = MemoryStateStore::new();
100 assert!(store.get("a").unwrap().is_none());
101
102 store
103 .set("a", Value::json(serde_json::json!({"mean": 5.0})))
104 .unwrap();
105 let state = store.get("a").unwrap().unwrap();
106 assert_eq!(state.as_json().unwrap()["mean"], 5.0);
107
108 let s1 = store.get("a").unwrap().unwrap();
110 let s2 = store.get("a").unwrap().unwrap();
111 assert!(Arc::ptr_eq(&s1, &s2));
112 }
113
114 #[test]
115 fn memory_store_remove_and_clear() {
116 let store = MemoryStateStore::new();
117 store.set("a", Value::Empty).unwrap();
118 store.set("b", Value::Empty).unwrap();
119 assert_eq!(store.keys().unwrap().len(), 2);
120
121 store.remove("a").unwrap();
122 assert!(store.get("a").unwrap().is_none());
123 assert!(store.get("b").unwrap().is_some());
124
125 store.clear().unwrap();
126 assert!(store.keys().unwrap().is_empty());
127 }
128
129 #[test]
130 fn memory_store_overwrites() {
131 let store = MemoryStateStore::new();
132 store
133 .set("a", Value::json(serde_json::json!({"v": 1})))
134 .unwrap();
135 store
136 .set("a", Value::json(serde_json::json!({"v": 2})))
137 .unwrap();
138 let state = store.get("a").unwrap().unwrap();
139 assert_eq!(state.as_json().unwrap()["v"], 2);
140 }
141}