1use std::collections::HashMap;
16use std::sync::Arc;
17
18use async_trait::async_trait;
19use dashmap::DashMap;
20
21use crate::error::StorageError;
22use crate::key::StateKey;
23
24#[async_trait]
32pub trait StateStorage: Send + Sync + 'static {
33 async fn get_state(&self, key: StateKey) -> Result<Option<String>, StorageError>;
35
36 async fn set_state(&self, key: StateKey, state: String) -> Result<(), StorageError>;
38
39 async fn clear_state(&self, key: StateKey) -> Result<(), StorageError>;
41
42 async fn get_data(
44 &self,
45 key: StateKey,
46 field: &str,
47 ) -> Result<Option<serde_json::Value>, StorageError>;
48
49 async fn set_data(
51 &self,
52 key: StateKey,
53 field: &str,
54 value: serde_json::Value,
55 ) -> Result<(), StorageError>;
56
57 async fn get_all_data(
59 &self,
60 key: StateKey,
61 ) -> Result<HashMap<String, serde_json::Value>, StorageError>;
62
63 async fn clear_data(&self, key: StateKey) -> Result<(), StorageError>;
65
66 async fn clear_all(&self, key: StateKey) -> Result<(), StorageError>;
68}
69
70#[derive(Clone, Default)]
78pub struct MemoryStorage {
79 entries: Arc<DashMap<StateKey, StorageEntry>>,
80}
81
82#[derive(Clone, Default)]
83struct StorageEntry {
84 state: Option<String>,
85 data: HashMap<String, serde_json::Value>,
86}
87
88impl MemoryStorage {
89 pub fn new() -> Self {
91 Self::default()
92 }
93
94 pub fn len(&self) -> usize {
96 self.entries.len()
97 }
98
99 pub fn is_empty(&self) -> bool {
101 self.entries.is_empty()
102 }
103}
104
105#[async_trait]
106impl StateStorage for MemoryStorage {
107 async fn get_state(&self, key: StateKey) -> Result<Option<String>, StorageError> {
108 Ok(self.entries.get(&key).and_then(|e| e.state.clone()))
109 }
110
111 async fn set_state(&self, key: StateKey, state: String) -> Result<(), StorageError> {
112 self.entries.entry(key).or_default().state = Some(state);
113 Ok(())
114 }
115
116 async fn clear_state(&self, key: StateKey) -> Result<(), StorageError> {
117 if let Some(mut entry) = self.entries.get_mut(&key) {
118 entry.state = None;
119 if entry.data.is_empty() {
120 drop(entry);
121 self.entries.remove(&key);
122 }
123 }
124 Ok(())
125 }
126
127 async fn get_data(
128 &self,
129 key: StateKey,
130 field: &str,
131 ) -> Result<Option<serde_json::Value>, StorageError> {
132 Ok(self
133 .entries
134 .get(&key)
135 .and_then(|e| e.data.get(field).cloned()))
136 }
137
138 async fn set_data(
139 &self,
140 key: StateKey,
141 field: &str,
142 value: serde_json::Value,
143 ) -> Result<(), StorageError> {
144 self.entries
145 .entry(key)
146 .or_default()
147 .data
148 .insert(field.to_string(), value);
149 Ok(())
150 }
151
152 async fn get_all_data(
153 &self,
154 key: StateKey,
155 ) -> Result<HashMap<String, serde_json::Value>, StorageError> {
156 Ok(self
157 .entries
158 .get(&key)
159 .map(|e| e.data.clone())
160 .unwrap_or_default())
161 }
162
163 async fn clear_data(&self, key: StateKey) -> Result<(), StorageError> {
164 if let Some(mut entry) = self.entries.get_mut(&key) {
165 entry.data.clear();
166 if entry.state.is_none() {
167 drop(entry);
168 self.entries.remove(&key);
169 }
170 }
171 Ok(())
172 }
173
174 async fn clear_all(&self, key: StateKey) -> Result<(), StorageError> {
175 self.entries.remove(&key);
176 Ok(())
177 }
178}