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};
4
5use tokio::sync::{RwLock, RwLockReadGuard};
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    async fn table_records_read(
62        &self,
63        _table: &str,
64    ) -> RwLockReadGuard<'_, HashMap<String, HashMap<String, serde_json::Value>>> {
65        self.tables.read().await
66    }
67}
68
69#[async_trait::async_trait]
70impl DatabaseBackend for InMemoryBackend {
71    fn engine_id(&self) -> &'static str {
72        ENGINE_ID
73    }
74
75    fn capabilities(&self) -> BackendCapabilities {
76        BackendCapabilities::mem()
77    }
78
79    async fn execute_compiled_query(
80        &self,
81        compiled: &CompiledQuery,
82    ) -> Result<Vec<serde_json::Value>> {
83        let q = compiled.query_string.trim();
84        let upper = q.to_uppercase();
85        if upper.starts_with("RETURN ") && upper.contains("OWNERSHIP_STATUS") {
86            let table = compiled_param_str(compiled, "table")
87                .ok_or_else(|| Error::Internal("missing table param".into()))?;
88            let record_id = compiled_param_str(compiled, "record_id")
89                .ok_or_else(|| Error::Internal("missing record_id param".into()))?;
90            let ownership_id = compiled_param_str(compiled, "ownership_id")
91                .ok_or_else(|| Error::Internal("missing ownership_id param".into()))?;
92            let row = self.get_record(&table, &record_id).await?;
93            let ownership_status = self
94                .get_record("valence_data_ownership", &ownership_id)
95                .await?
96                .and_then(|r| r.get("status").cloned());
97            return Ok(vec![serde_json::json!({
98                "row": row,
99                "ownership_status": ownership_status,
100            })]);
101        }
102
103        if upper.starts_with("SELECT ") {
104            if upper.contains("COUNT(") {
105                if let Some(from_idx) = upper.find(" FROM ") {
106                    let table = q[from_idx + 6..]
107                        .split_whitespace()
108                        .next()
109                        .unwrap_or("")
110                        .trim();
111                    if !table.is_empty() {
112                        let mut rows = {
113                            let tables = self.table_records_read(table).await;
114                            tables
115                                .get(table)
116                                .map(|m| m.values().cloned().collect::<Vec<_>>())
117                                .unwrap_or_default()
118                        };
119                        rows = crate::query_filter::apply_equality_where(rows, compiled);
120                        let count = i64::try_from(rows.len()).unwrap_or(i64::MAX);
121                        return Ok(vec![serde_json::json!(count)]);
122                    }
123                }
124            }
125
126            if upper.contains("SELECT id") && !upper.contains("body") {
127                if let Some(from_idx) = upper.find(" FROM ") {
128                    let table = q[from_idx + 6..]
129                        .split_whitespace()
130                        .next()
131                        .unwrap_or("")
132                        .trim();
133                    if !table.is_empty() {
134                        let rows = {
135                            let tables = self.table_records_read(table).await;
136                            tables
137                                .get(table)
138                                .map(|m| {
139                                    m.keys()
140                                        .map(|id| serde_json::Value::String(id.clone()))
141                                        .collect::<Vec<_>>()
142                                })
143                                .unwrap_or_default()
144                        };
145                        return Ok(rows);
146                    }
147                }
148            }
149
150            if upper.contains("body") {
151                if let Some(from_idx) = upper.find(" FROM ") {
152                    let table = q[from_idx + 6..]
153                        .split_whitespace()
154                        .next()
155                        .unwrap_or("")
156                        .trim();
157                    if !table.is_empty() {
158                        let mut rows = {
159                            let tables = self.table_records_read(table).await;
160                            tables
161                                .get(table)
162                                .map(|m| m.values().cloned().collect::<Vec<_>>())
163                                .unwrap_or_default()
164                        };
165                        if let Some(limit_idx) = upper.rfind(" LIMIT ") {
166                            if let Ok(limit) = q[limit_idx + 7..].trim().parse::<usize>() {
167                                rows.truncate(limit);
168                            }
169                        }
170                        return Ok(rows);
171                    }
172                }
173            }
174
175            if let Some(from_idx) = upper.find(" FROM ") {
176                let table = q[from_idx + 6..]
177                    .split_whitespace()
178                    .next()
179                    .unwrap_or("")
180                    .trim();
181                if !table.is_empty() {
182                    let mut rows = {
183                        let tables = self.table_records_read(table).await;
184                        tables
185                            .get(table)
186                            .map(|m| m.values().cloned().collect::<Vec<_>>())
187                            .unwrap_or_default()
188                    };
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                        // Emit bare id cells (string or RecordId object) — not
194                        // `{"id": …}` wrappers — so `extract_id_from_select_value` works.
195                        return Ok(rows
196                            .into_iter()
197                            .filter_map(|r| r.get("id").cloned())
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).await;
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        if let Ok(layout) = valence_core::storage_layout::StorageLayout::from_registry_table(table)
218        {
219            valence_core::storage_layout::validate_write_types(&layout, &content)?;
220        }
221        let mut content = content;
222        valence_core::ttl::prepare_create_content(table, self, &mut content)?;
223        let id = storage_id_from_content(&content).unwrap_or_else(uuid_simple);
224        let mut record = content;
225        if let Some(obj) = record.as_object_mut() {
226            let has_string_id = obj.get("id").and_then(|v| v.as_str()).is_some();
227            if !has_string_id {
228                obj.insert("id".into(), record_id_json(table, &id));
229            }
230        }
231        self.tables
232            .write()
233            .await
234            .entry(table.to_string())
235            .or_default()
236            .insert(id, record.clone());
237        Ok(record)
238    }
239
240    async fn update_record(
241        &self,
242        table: &str,
243        id: &str,
244        content: serde_json::Value,
245    ) -> Result<serde_json::Value> {
246        let mut tables = self.tables.write().await;
247        let rows = tables
248            .get_mut(table)
249            .ok_or_else(|| Error::NotFound(format!("table {table}")))?;
250        if !rows.contains_key(id) {
251            return Err(Error::NotFound(format!("{table}:{id}")));
252        }
253        rows.insert(id.to_string(), content.clone());
254        drop(tables);
255        Ok(content)
256    }
257
258    async fn merge_record(
259        &self,
260        table: &str,
261        id: &str,
262        patch: serde_json::Value,
263    ) -> Result<serde_json::Value> {
264        let mut tables = self.tables.write().await;
265        let rows = tables.entry(table.to_string()).or_default();
266        let existing = rows
267            .entry(id.to_string())
268            .or_insert_with(|| serde_json::json!({}));
269        if let (Some(base), Some(patch_obj)) = (existing.as_object_mut(), patch.as_object()) {
270            for (k, v) in patch_obj {
271                base.insert(k.clone(), v.clone());
272            }
273        }
274        let merged = existing.clone();
275        drop(tables);
276        Ok(merged)
277    }
278
279    async fn upsert_record(
280        &self,
281        table: &str,
282        id: &str,
283        content: serde_json::Value,
284    ) -> Result<serde_json::Value> {
285        let mut record = content;
286        if let Some(obj) = record.as_object_mut() {
287            obj.insert("id".into(), record_id_json(table, id));
288        }
289        self.tables
290            .write()
291            .await
292            .entry(table.to_string())
293            .or_default()
294            .insert(id.to_string(), record.clone());
295        Ok(record)
296    }
297
298    async fn delete_record(&self, table: &str, id: &str) -> Result<()> {
299        if let Some(rows) = self.tables.write().await.get_mut(table) {
300            rows.remove(id);
301        }
302        Ok(())
303    }
304
305    async fn relate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
306        let key = edge_key(edge_table, from);
307        self.edges
308            .write()
309            .await
310            .entry(key)
311            .or_default()
312            .insert((to.table().to_string(), to.id().to_string()));
313        Ok(())
314    }
315
316    async fn unrelate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
317        let key = edge_key(edge_table, from);
318        if let Some(set) = self.edges.write().await.get_mut(&key) {
319            set.remove(&(to.table().to_string(), to.id().to_string()));
320        }
321        Ok(())
322    }
323
324    async fn get_edge_targets(&self, from: &RecordId, edge_table: &str) -> Result<Vec<RecordId>> {
325        let key = edge_key(edge_table, from);
326        let edges = self.edges.read().await;
327        Ok(edges
328            .get(&key)
329            .map(|set| {
330                set.iter()
331                    .map(|(table, id)| RecordId::new(table.clone(), id.clone()))
332                    .collect()
333            })
334            .unwrap_or_default())
335    }
336
337    async fn get_edge_sources(&self, to: &RecordId, edge_table: &str) -> Result<Vec<RecordId>> {
338        let edges = self.edges.read().await;
339        let prefix = format!("{edge_table}:");
340        let mut sources = Vec::new();
341        let to_key = (to.table().to_string(), to.id().to_string());
342        for (key, set) in edges.iter() {
343            if !key.starts_with(&prefix) || !set.contains(&to_key) {
344                continue;
345            }
346            // key = "{edge_table}:{from_table}:{from_id}"
347            let rest = &key[prefix.len()..];
348            if let Some((ft, fid)) = rest.split_once(':') {
349                if !ft.is_empty() && !fid.is_empty() {
350                    sources.push(RecordId::new(ft, fid));
351                }
352            }
353        }
354        Ok(sources)
355    }
356
357    /// No DDL on mem — uniqueness is enforced by the `SELECT VALUE id` probe in codegen.
358    async fn define_unique_index(&self, _table: &str, _field: &str) -> Result<()> {
359        Err(valence_core::Error::Internal(
360            "unique indexes not supported on in-memory backend".into(),
361        ))
362    }
363
364    fn ttl_capability(&self) -> valence_core::ttl::BackendTtlCapability {
365        valence_core::ttl::BackendTtlCapability::Deferred
366    }
367}
368
369fn edge_key(edge_table: &str, from: &RecordId) -> String {
370    format!("{edge_table}:{}:{}", from.table(), from.id())
371}
372
373fn compiled_param_str(compiled: &CompiledQuery, key: &str) -> Option<String> {
374    compiled
375        .params
376        .iter()
377        .find(|(k, _)| k == key)
378        .and_then(|(_, v)| v.as_str().map(|s| s.to_string()))
379}
380
381fn record_id_json(table: &str, id: &str) -> serde_json::Value {
382    serde_json::json!({
383        "table": table,
384        "id": id,
385    })
386}
387
388fn storage_id_from_content(content: &serde_json::Value) -> Option<String> {
389    let id_val = content.get("id")?;
390    if let Some(id) = id_val.get("id").and_then(|v| v.as_str()) {
391        return Some(id.to_string());
392    }
393    id_val.as_str().map(|s| s.to_string())
394}
395
396fn uuid_simple() -> String {
397    use std::time::{SystemTime, UNIX_EPOCH};
398    let nanos = SystemTime::now()
399        .duration_since(UNIX_EPOCH)
400        .map_or(0, |d| d.as_nanos());
401    format!("mem-{nanos}")
402}
403
404#[cfg(test)]
405mod tests {
406    #![allow(clippy::expect_used, clippy::unwrap_used)]
407
408    use super::*;
409
410    #[tokio::test]
411    async fn crud_round_trip() {
412        let backend = InMemoryBackend::new();
413        let created = backend
414            .create_record("user", serde_json::json!({"name": "Ada"}))
415            .await
416            .unwrap();
417        let id = storage_id_from_content(&created).expect("record id");
418        let fetched = backend.get_record("user", &id).await.unwrap().unwrap();
419        assert_eq!(fetched["name"], "Ada");
420
421        let merged = backend
422            .merge_record("user", &id, serde_json::json!({"name": "Grace"}))
423            .await
424            .unwrap();
425        assert_eq!(merged["name"], "Grace");
426
427        backend.delete_record("user", &id).await.unwrap();
428        assert!(backend.get_record("user", &id).await.unwrap().is_none());
429    }
430
431    #[test]
432    fn ttl_capability_is_deferred() {
433        let backend = InMemoryBackend::new();
434        assert_eq!(
435            backend.ttl_capability(),
436            valence_core::ttl::BackendTtlCapability::Deferred
437        );
438    }
439}