oxigdal_streaming/state/
backend.rs1use crate::error::{Result, StreamingError};
4use async_trait::async_trait;
5use std::collections::HashMap;
6#[cfg(feature = "kv-store")]
7use std::path::PathBuf;
8use std::sync::Arc;
9use tokio::sync::RwLock;
10
11#[cfg(feature = "kv-store")]
12use oxistore_core::{KvStore, StoreError};
13#[cfg(feature = "kv-store")]
14use oxistore_kv_fjall::FjallStore;
15
16#[async_trait]
18pub trait StateBackend: Send + Sync {
19 async fn get(&self, key: &[u8]) -> Result<Option<Vec<u8>>>;
21
22 async fn put(&self, key: &[u8], value: &[u8]) -> Result<()>;
24
25 async fn delete(&self, key: &[u8]) -> Result<()>;
27
28 async fn contains(&self, key: &[u8]) -> Result<bool>;
30
31 async fn clear(&self) -> Result<()>;
33
34 async fn snapshot(&self) -> Result<Vec<u8>>;
36
37 async fn restore(&self, snapshot: &[u8]) -> Result<()>;
39
40 async fn keys(&self) -> Result<Vec<Vec<u8>>>;
42
43 fn name(&self) -> &str;
45}
46
47pub struct MemoryStateBackend {
49 state: Arc<RwLock<HashMap<Vec<u8>, Vec<u8>>>>,
50}
51
52impl MemoryStateBackend {
53 pub fn new() -> Self {
55 Self {
56 state: Arc::new(RwLock::new(HashMap::new())),
57 }
58 }
59}
60
61impl Default for MemoryStateBackend {
62 fn default() -> Self {
63 Self::new()
64 }
65}
66
67#[async_trait]
68impl StateBackend for MemoryStateBackend {
69 async fn get(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
70 Ok(self.state.read().await.get(key).cloned())
71 }
72
73 async fn put(&self, key: &[u8], value: &[u8]) -> Result<()> {
74 self.state
75 .write()
76 .await
77 .insert(key.to_vec(), value.to_vec());
78 Ok(())
79 }
80
81 async fn delete(&self, key: &[u8]) -> Result<()> {
82 self.state.write().await.remove(key);
83 Ok(())
84 }
85
86 async fn contains(&self, key: &[u8]) -> Result<bool> {
87 Ok(self.state.read().await.contains_key(key))
88 }
89
90 async fn clear(&self) -> Result<()> {
91 self.state.write().await.clear();
92 Ok(())
93 }
94
95 async fn snapshot(&self) -> Result<Vec<u8>> {
96 let state = self.state.read().await;
97 oxicode::encode_to_vec(&*state)
99 .map_err(|e| StreamingError::SerializationError(e.to_string()))
100 }
101
102 async fn restore(&self, snapshot: &[u8]) -> Result<()> {
103 let (restored, _): (HashMap<Vec<u8>, Vec<u8>>, _) = oxicode::decode_from_slice(snapshot)
104 .map_err(|e| StreamingError::SerializationError(e.to_string()))?;
105 *self.state.write().await = restored;
106 Ok(())
107 }
108
109 async fn keys(&self) -> Result<Vec<Vec<u8>>> {
110 Ok(self.state.read().await.keys().cloned().collect())
111 }
112
113 fn name(&self) -> &str {
114 "MemoryStateBackend"
115 }
116}
117
118#[cfg(feature = "kv-store")]
124pub struct KvStateBackend {
125 store: Arc<FjallStore>,
126 path: PathBuf,
127}
128
129#[cfg(feature = "kv-store")]
130impl KvStateBackend {
131 pub fn new(path: PathBuf) -> Result<Self> {
135 let store = FjallStore::open(&path).map_err(StoreError::from)?;
136
137 Ok(Self {
138 store: Arc::new(store),
139 path,
140 })
141 }
142
143 pub fn path(&self) -> &PathBuf {
145 &self.path
146 }
147}
148
149#[cfg(feature = "kv-store")]
150#[async_trait]
151impl StateBackend for KvStateBackend {
152 async fn get(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
153 Ok(self.store.get(key)?)
154 }
155
156 async fn put(&self, key: &[u8], value: &[u8]) -> Result<()> {
157 self.store.put(key, value)?;
158 Ok(())
159 }
160
161 async fn delete(&self, key: &[u8]) -> Result<()> {
162 self.store.delete(key)?;
163 Ok(())
164 }
165
166 async fn contains(&self, key: &[u8]) -> Result<bool> {
167 Ok(self.store.contains(key)?)
168 }
169
170 async fn clear(&self) -> Result<()> {
171 let keys: Vec<Vec<u8>> = self
172 .store
173 .keys()?
174 .collect::<std::result::Result<Vec<_>, StoreError>>()?;
175
176 for key in keys {
177 self.store.delete(&key)?;
178 }
179
180 Ok(())
181 }
182
183 async fn snapshot(&self) -> Result<Vec<u8>> {
184 let mut data = Vec::new();
189
190 for item in self.store.iter()? {
191 let (key, value) = item?;
192 let entry = (key, value);
193 let serialized = oxicode::encode_to_vec(&entry)
195 .map_err(|e| StreamingError::SerializationError(e.to_string()))?;
196 data.extend_from_slice(&(serialized.len() as u32).to_le_bytes());
197 data.extend_from_slice(&serialized);
198 }
199
200 Ok(data)
201 }
202
203 async fn restore(&self, snapshot: &[u8]) -> Result<()> {
204 self.clear().await?;
205
206 let mut offset = 0;
207 while offset < snapshot.len() {
208 if offset + 4 > snapshot.len() {
209 break;
210 }
211
212 let len = u32::from_le_bytes([
213 snapshot[offset],
214 snapshot[offset + 1],
215 snapshot[offset + 2],
216 snapshot[offset + 3],
217 ]) as usize;
218
219 offset += 4;
220
221 if offset + len > snapshot.len() {
222 break;
223 }
224
225 let entry_data = &snapshot[offset..offset + len];
226 let ((key, value), _): ((Vec<u8>, Vec<u8>), _) = oxicode::decode_from_slice(entry_data)
227 .map_err(|e| StreamingError::SerializationError(e.to_string()))?;
228 self.store.put(&key, &value)?;
229
230 offset += len;
231 }
232
233 Ok(())
234 }
235
236 async fn keys(&self) -> Result<Vec<Vec<u8>>> {
237 let keys: Vec<Vec<u8>> = self
238 .store
239 .keys()?
240 .collect::<std::result::Result<Vec<_>, StoreError>>()?;
241
242 Ok(keys)
243 }
244
245 fn name(&self) -> &str {
246 "KvStateBackend"
247 }
248}
249
250#[cfg(test)]
251mod tests {
252 use super::*;
253
254 #[tokio::test]
255 async fn test_memory_backend() -> Result<()> {
256 let backend = MemoryStateBackend::new();
257
258 backend.put(b"key1", b"value1").await?;
259 let value = backend.get(b"key1").await?;
260 assert_eq!(value, Some(b"value1".to_vec()));
261
262 assert!(backend.contains(b"key1").await?);
263 assert!(!backend.contains(b"key2").await?);
264
265 backend.delete(b"key1").await?;
266 assert!(!backend.contains(b"key1").await?);
267
268 Ok(())
269 }
270
271 #[tokio::test]
272 async fn test_memory_backend_snapshot() -> Result<()> {
273 let backend = MemoryStateBackend::new();
274
275 backend.put(b"key1", b"value1").await?;
276 backend.put(b"key2", b"value2").await?;
277
278 let snapshot = backend.snapshot().await?;
279
280 let backend2 = MemoryStateBackend::new();
281 backend2.restore(&snapshot).await?;
282
283 assert_eq!(backend2.get(b"key1").await?, Some(b"value1".to_vec()));
284 assert_eq!(backend2.get(b"key2").await?, Some(b"value2".to_vec()));
285
286 Ok(())
287 }
288
289 #[cfg(feature = "kv-store")]
290 #[tokio::test]
291 async fn test_kv_backend() -> Result<()> {
292 let temp_dir = tempfile::tempdir()
293 .map_err(|e| StreamingError::StateError(format!("Failed to create temp dir: {}", e)))?;
294 let backend = KvStateBackend::new(temp_dir.path().to_path_buf())?;
295
296 backend.put(b"key1", b"value1").await?;
297 let value = backend.get(b"key1").await?;
298 assert_eq!(value, Some(b"value1".to_vec()));
299
300 assert!(backend.contains(b"key1").await?);
301
302 backend.delete(b"key1").await?;
303 assert!(!backend.contains(b"key1").await?);
304
305 Ok(())
306 }
307
308 #[cfg(feature = "kv-store")]
309 #[tokio::test]
310 async fn test_kv_backend_snapshot_restore() -> Result<()> {
311 let temp_dir = tempfile::tempdir()
312 .map_err(|e| StreamingError::StateError(format!("Failed to create temp dir: {}", e)))?;
313 let backend = KvStateBackend::new(temp_dir.path().to_path_buf())?;
314
315 backend.put(b"key1", b"value1").await?;
316 backend.put(b"key2", b"value2").await?;
317
318 let snapshot = backend.snapshot().await?;
319 let mut keys = backend.keys().await?;
320 keys.sort();
321 assert_eq!(keys, vec![b"key1".to_vec(), b"key2".to_vec()]);
322
323 backend.clear().await?;
324 assert!(!backend.contains(b"key1").await?);
325
326 backend.restore(&snapshot).await?;
327 assert_eq!(backend.get(b"key1").await?, Some(b"value1".to_vec()));
328 assert_eq!(backend.get(b"key2").await?, Some(b"value2".to_vec()));
329
330 Ok(())
331 }
332}