1use std::collections::{HashMap, HashSet};
4
5use tokio::sync::{RwLock, RwLockReadGuard};
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 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 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 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 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}