1use std::collections::{HashMap, HashSet};
4use std::sync::RwLock;
5
6use valence_core::{BackendCapabilities, CompiledQuery, DatabaseBackend, Error, RecordId, Result};
7
8pub const ENGINE_ID: &str = valence_core::KnownEngines::INMEMORY_MEM;
10
11#[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 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}