Skip to main content

valence_backend_mem/
backend.rs

1//! Minimal in-memory storage engine for tests and embedded hosts.
2
3use std::collections::{HashMap, HashSet};
4use std::sync::RwLock;
5
6use valence_core::{BackendCapabilities, CompiledQuery, DatabaseBackend, Error, RecordId, Result};
7
8/// Stable engine slug for router keys (`inmemory_mem:logical_name`).
9pub const ENGINE_ID: &str = valence_core::KnownEngines::INMEMORY_MEM;
10
11/// In-memory [`DatabaseBackend`] storing rows and graph edges in process memory.
12///
13/// # Examples
14///
15/// ```ignore
16/// use std::sync::Arc;
17/// use valence::{
18///     valence_schema, Database, DatabaseFromEngine, FieldType, InMemoryBackend, Valence,
19///     MEM_ENGINE_ID,
20/// };
21///
22/// const COUNTER_DB: DatabaseFromEngine =
23///     Database::from_engine("default", MEM_ENGINE_ID);
24///
25/// valence_schema! {
26///     Counter {
27///         table: "counter",
28///         version: "0.1.0",
29///         database: COUNTER_DB,
30///         fields: [
31///             id: { r#type: FieldType::String, primary_key: true, required: true },
32///             value: { r#type: FieldType::Integer, required: true },
33///         ],
34///     }
35/// }
36///
37/// let valence = Valence::builder()
38///     .add_backend("default", Arc::new(InMemoryBackend::new()))
39///     .build()
40///     .expect("build");
41/// assert_eq!(
42///     valence
43///         .backend_for_table("counter")
44///         .expect("counter backend")
45///         .engine_id(),
46///     MEM_ENGINE_ID
47/// );
48/// ```
49#[derive(Debug, Default)]
50pub struct InMemoryBackend {
51    tables: RwLock<HashMap<String, HashMap<String, serde_json::Value>>>,
52    edges: RwLock<HashMap<String, HashSet<(String, String)>>>,
53}
54
55impl InMemoryBackend {
56    /// Create an empty in-memory backend.
57    pub fn new() -> Self {
58        Self::default()
59    }
60
61    fn table_records(
62        &self,
63        _table: &str,
64    ) -> Result<std::sync::RwLockWriteGuard<'_, HashMap<String, HashMap<String, serde_json::Value>>>>
65    {
66        self.tables
67            .write()
68            .map_err(|_| Error::Internal("mem backend lock poisoned".into()))
69    }
70
71    fn table_records_read(
72        &self,
73        _table: &str,
74    ) -> Result<std::sync::RwLockReadGuard<'_, HashMap<String, HashMap<String, serde_json::Value>>>>
75    {
76        self.tables
77            .read()
78            .map_err(|_| Error::Internal("mem backend lock poisoned".into()))
79    }
80}
81
82#[async_trait::async_trait]
83impl DatabaseBackend for InMemoryBackend {
84    fn engine_id(&self) -> &'static str {
85        ENGINE_ID
86    }
87
88    fn capabilities(&self) -> BackendCapabilities {
89        BackendCapabilities::mem()
90    }
91
92    async fn execute_compiled_query(
93        &self,
94        compiled: &CompiledQuery,
95    ) -> Result<Vec<serde_json::Value>> {
96        let q = compiled.query_string.trim();
97        let upper = q.to_uppercase();
98        if upper.starts_with("RETURN ") && upper.contains("OWNERSHIP_STATUS") {
99            let table = compiled_param_str(compiled, "table")
100                .ok_or_else(|| Error::Internal("missing table param".into()))?;
101            let record_id = compiled_param_str(compiled, "record_id")
102                .ok_or_else(|| Error::Internal("missing record_id param".into()))?;
103            let ownership_id = compiled_param_str(compiled, "ownership_id")
104                .ok_or_else(|| Error::Internal("missing ownership_id param".into()))?;
105            let row = self.get_record(&table, &record_id).await?;
106            let ownership_status = self
107                .get_record("valence_data_ownership", &ownership_id)
108                .await?
109                .and_then(|r| r.get("status").cloned());
110            return Ok(vec![serde_json::json!({
111                "row": row,
112                "ownership_status": ownership_status,
113            })]);
114        }
115
116        if upper.starts_with("SELECT ") {
117            if upper.contains("COUNT(") {
118                if let Some(from_idx) = upper.find(" FROM ") {
119                    let table = q[from_idx + 6..]
120                        .split_whitespace()
121                        .next()
122                        .unwrap_or("")
123                        .trim();
124                    if !table.is_empty() {
125                        let tables = self.table_records_read(table)?;
126                        let count = tables.get(table).map(|m| m.len() as i64).unwrap_or(0);
127                        return Ok(vec![serde_json::json!(count)]);
128                    }
129                }
130            }
131
132            if upper.contains("SELECT id") && !upper.contains("body") {
133                if let Some(from_idx) = upper.find(" FROM ") {
134                    let table = q[from_idx + 6..]
135                        .split_whitespace()
136                        .next()
137                        .unwrap_or("")
138                        .trim();
139                    if !table.is_empty() {
140                        let tables = self.table_records_read(table)?;
141                        let rows: Vec<serde_json::Value> = tables
142                            .get(table)
143                            .map(|m| {
144                                m.keys()
145                                    .map(|id| serde_json::Value::String(id.clone()))
146                                    .collect()
147                            })
148                            .unwrap_or_default();
149                        return Ok(rows);
150                    }
151                }
152            }
153
154            if upper.contains("body") {
155                if let Some(from_idx) = upper.find(" FROM ") {
156                    let table = q[from_idx + 6..]
157                        .split_whitespace()
158                        .next()
159                        .unwrap_or("")
160                        .trim();
161                    if !table.is_empty() {
162                        let tables = self.table_records_read(table)?;
163                        let mut rows: Vec<serde_json::Value> = tables
164                            .get(table)
165                            .map(|m| m.values().cloned().collect())
166                            .unwrap_or_default();
167                        if let Some(limit_idx) = upper.rfind(" LIMIT ") {
168                            if let Ok(limit) = q[limit_idx + 7..].trim().parse::<usize>() {
169                                rows.truncate(limit);
170                            }
171                        }
172                        return Ok(rows);
173                    }
174                }
175            }
176
177            if let Some(from_idx) = upper.find(" FROM ") {
178                let table = q[from_idx + 6..]
179                    .split_whitespace()
180                    .next()
181                    .unwrap_or("")
182                    .trim();
183                if !table.is_empty() {
184                    let tables = self.table_records_read(table)?;
185                    let mut rows: Vec<serde_json::Value> = tables
186                        .get(table)
187                        .map(|m| m.values().cloned().collect())
188                        .unwrap_or_default();
189                    rows = crate::query_filter::apply_equality_where(rows, compiled);
190                    rows =
191                        crate::query_filter::apply_order_limit_offset(rows, &compiled.query_string);
192                    if upper.contains("SELECT VALUE") || upper.contains("SELECT id") {
193                        return Ok(rows
194                            .into_iter()
195                            .filter_map(|r| {
196                                r.get("id").cloned().map(|id| serde_json::json!({"id": id}))
197                            })
198                            .collect());
199                    }
200                    return Ok(rows);
201                }
202            }
203        }
204        Ok(vec![])
205    }
206
207    async fn get_record(&self, table: &str, id: &str) -> Result<Option<serde_json::Value>> {
208        let tables = self.table_records_read(table)?;
209        Ok(tables.get(table).and_then(|rows| rows.get(id).cloned()))
210    }
211
212    async fn create_record(
213        &self,
214        table: &str,
215        content: serde_json::Value,
216    ) -> Result<serde_json::Value> {
217        let mut tables = self.table_records(table)?;
218        let rows = tables.entry(table.to_string()).or_default();
219        let id = storage_id_from_content(&content).unwrap_or_else(uuid_simple);
220        let mut record = content;
221        if let Some(obj) = record.as_object_mut() {
222            let has_string_id = obj.get("id").and_then(|v| v.as_str()).is_some();
223            if !has_string_id {
224                obj.insert("id".into(), record_id_json(table, &id));
225            }
226        }
227        rows.insert(id, record.clone());
228        Ok(record)
229    }
230
231    async fn update_record(
232        &self,
233        table: &str,
234        id: &str,
235        content: serde_json::Value,
236    ) -> Result<serde_json::Value> {
237        let mut tables = self.table_records(table)?;
238        let rows = tables
239            .get_mut(table)
240            .ok_or_else(|| Error::NotFound(format!("table {table}")))?;
241        if !rows.contains_key(id) {
242            return Err(Error::NotFound(format!("{table}:{id}")));
243        }
244        rows.insert(id.to_string(), content.clone());
245        Ok(content)
246    }
247
248    async fn merge_record(
249        &self,
250        table: &str,
251        id: &str,
252        patch: serde_json::Value,
253    ) -> Result<serde_json::Value> {
254        let mut tables = self.table_records(table)?;
255        let rows = tables.entry(table.to_string()).or_default();
256        let existing = rows
257            .entry(id.to_string())
258            .or_insert_with(|| serde_json::json!({}));
259        if let (Some(base), Some(patch_obj)) = (existing.as_object_mut(), patch.as_object()) {
260            for (k, v) in patch_obj {
261                base.insert(k.clone(), v.clone());
262            }
263        }
264        Ok(existing.clone())
265    }
266
267    async fn upsert_record(
268        &self,
269        table: &str,
270        id: &str,
271        content: serde_json::Value,
272    ) -> Result<serde_json::Value> {
273        let mut tables = self.table_records(table)?;
274        let rows = tables.entry(table.to_string()).or_default();
275        let mut record = content;
276        if let Some(obj) = record.as_object_mut() {
277            obj.insert("id".into(), record_id_json(table, id));
278        }
279        rows.insert(id.to_string(), record.clone());
280        Ok(record)
281    }
282
283    async fn delete_record(&self, table: &str, id: &str) -> Result<()> {
284        let mut tables = self.table_records(table)?;
285        if let Some(rows) = tables.get_mut(table) {
286            rows.remove(id);
287        }
288        Ok(())
289    }
290
291    async fn relate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
292        let key = edge_key(edge_table, from);
293        let mut edges = self
294            .edges
295            .write()
296            .map_err(|_| Error::Internal("mem edges lock poisoned".into()))?;
297        edges
298            .entry(key)
299            .or_default()
300            .insert((to.table().to_string(), to.id().to_string()));
301        Ok(())
302    }
303
304    async fn unrelate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
305        let key = edge_key(edge_table, from);
306        let mut edges = self
307            .edges
308            .write()
309            .map_err(|_| Error::Internal("mem edges lock poisoned".into()))?;
310        if let Some(set) = edges.get_mut(&key) {
311            set.remove(&(to.table().to_string(), to.id().to_string()));
312        }
313        Ok(())
314    }
315
316    async fn get_edge_targets(&self, from: &RecordId, edge_table: &str) -> Result<Vec<RecordId>> {
317        let key = edge_key(edge_table, from);
318        let edges = self
319            .edges
320            .read()
321            .map_err(|_| Error::Internal("mem edges lock poisoned".into()))?;
322        Ok(edges
323            .get(&key)
324            .map(|set| {
325                set.iter()
326                    .map(|(table, id)| RecordId::new(table.clone(), id.clone()))
327                    .collect()
328            })
329            .unwrap_or_default())
330    }
331}
332
333fn edge_key(edge_table: &str, from: &RecordId) -> String {
334    format!("{edge_table}:{}:{}", from.table(), from.id())
335}
336
337fn compiled_param_str(compiled: &CompiledQuery, key: &str) -> Option<String> {
338    compiled
339        .params
340        .iter()
341        .find(|(k, _)| k == key)
342        .and_then(|(_, v)| v.as_str().map(|s| s.to_string()))
343}
344
345fn record_id_json(table: &str, id: &str) -> serde_json::Value {
346    serde_json::json!({
347        "table": table,
348        "id": id,
349    })
350}
351
352fn storage_id_from_content(content: &serde_json::Value) -> Option<String> {
353    let id_val = content.get("id")?;
354    if let Some(id) = id_val.get("id").and_then(|v| v.as_str()) {
355        return Some(id.to_string());
356    }
357    id_val.as_str().map(|s| s.to_string())
358}
359
360fn uuid_simple() -> String {
361    use std::time::{SystemTime, UNIX_EPOCH};
362    let nanos = SystemTime::now()
363        .duration_since(UNIX_EPOCH)
364        .map_or(0, |d| d.as_nanos());
365    format!("mem-{nanos}")
366}
367
368#[cfg(test)]
369mod tests {
370    use super::*;
371
372    #[tokio::test]
373    async fn crud_round_trip() {
374        let backend = InMemoryBackend::new();
375        let created = backend
376            .create_record("user", serde_json::json!({"name": "Ada"}))
377            .await
378            .unwrap();
379        let id = storage_id_from_content(&created).expect("record id");
380        let fetched = backend.get_record("user", &id).await.unwrap().unwrap();
381        assert_eq!(fetched["name"], "Ada");
382
383        let merged = backend
384            .merge_record("user", &id, serde_json::json!({"name": "Grace"}))
385            .await
386            .unwrap();
387        assert_eq!(merged["name"], "Grace");
388
389        backend.delete_record("user", &id).await.unwrap();
390        assert!(backend.get_record("user", &id).await.unwrap().is_none());
391    }
392}