Skip to main content

valence_backend_redis/
backend.rs

1//! Redis wire [`DatabaseBackend`] using JSON documents in Redis STRING keys.
2
3use redis::aio::ConnectionManager;
4use redis::AsyncCommands;
5use serde_json::{Map, Value};
6
7use valence_core::{
8    BackendCapabilities, CompiledQuery, Database, DatabaseBackend, DatabaseFromEngine, Error,
9    KnownEngines, RecordId, Result,
10};
11
12use crate::config::RedisConfig;
13use crate::keys::Keyspace;
14
15/// Stable engine slug for router keys (`redis:logical_name`).
16pub const ENGINE_ID: &str = KnownEngines::REDIS;
17
18/// Schema evaluator const for `database:` routing.
19pub const PRIMARY: DatabaseFromEngine = Database::from_engine("primary", ENGINE_ID);
20
21/// Redis-backed [`DatabaseBackend`] storing JSON documents per table/id key.
22///
23/// # Examples
24///
25/// ```ignore
26/// use std::sync::Arc;
27/// use valence::{
28///     valence_schema, Database, DatabaseFromEngine, FieldType, RedisBackend, Valence,
29///     REDIS_ENGINE_ID,
30/// };
31///
32/// const COUNTER_DB: DatabaseFromEngine =
33///     Database::from_engine("default", REDIS_ENGINE_ID);
34///
35/// valence_schema! {
36///     Counter {
37///         table: "counter",
38///         version: "0.1.0",
39///         database: COUNTER_DB,
40///         fields: [
41///             id: { r#type: FieldType::String, primary_key: true, required: true },
42///             value: { r#type: FieldType::Integer, required: true },
43///         ],
44///     }
45/// }
46///
47/// // Reads VALENCE_REDIS_URL and optional VALENCE_REDIS_KEY_PREFIX.
48/// let backend = RedisBackend::from_env().await?;
49/// let valence = Valence::builder()
50///     .add_backend("default", Arc::new(backend))
51///     .build()?;
52/// assert_eq!(
53///     valence.backend_for_table("counter")?.engine_id(),
54///     REDIS_ENGINE_ID
55/// );
56/// # Ok::<(), valence::Error>(())
57/// ```
58#[derive(Clone)]
59pub struct RedisBackend {
60    conn: ConnectionManager,
61    keys: Keyspace,
62}
63
64impl std::fmt::Debug for RedisBackend {
65    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
66        f.debug_struct("RedisBackend")
67            .field("keys", &self.keys)
68            .finish_non_exhaustive()
69    }
70}
71
72impl RedisBackend {
73    /// Start a builder for explicit host wiring.
74    pub fn builder() -> crate::config::RedisBackendBuilder {
75        crate::config::RedisBackendBuilder::new()
76    }
77
78    /// Connect using env defaults via builder (shorthand).
79    pub async fn from_env() -> Result<Self> {
80        Self::builder().from_env_defaults().build().await
81    }
82
83    /// Connect to Redis at `url` with default key prefix.
84    pub async fn connect(url: &str) -> Result<Self> {
85        Self::builder().url(url).build().await
86    }
87
88    /// Connect using explicit config.
89    pub async fn connect_with_config(config: RedisConfig) -> Result<Self> {
90        let client =
91            redis::Client::open(config.url.as_str()).map_err(|e| Error::Database(e.to_string()))?;
92        let conn = ConnectionManager::new(client)
93            .await
94            .map_err(|e| Error::Database(e.to_string()))?;
95        Ok(Self {
96            conn,
97            keys: Keyspace::new(config.key_prefix),
98        })
99    }
100
101    fn map_err(e: redis::RedisError) -> Error {
102        Error::Database(e.to_string())
103    }
104
105    fn assert_safe_table(table: &str) -> Result<()> {
106        if table.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
107            Ok(())
108        } else {
109            Err(Error::Validation(format!("unsafe table name: {table}")))
110        }
111    }
112
113    async fn unique_fields(&self, table: &str) -> Result<Vec<String>> {
114        let key = self.keys.uniq_index(table);
115        let mut conn = self.conn.clone();
116        let fields: Vec<String> = conn.smembers(&key).await.map_err(Self::map_err)?;
117        Ok(fields)
118    }
119
120    async fn claim_unique_fields(
121        &self,
122        table: &str,
123        id: &str,
124        record: &Value,
125        exclude_id: Option<&str>,
126    ) -> Result<()> {
127        for field in self.unique_fields(table).await? {
128            let Some(value) = record.get(&field).and_then(|v| v.as_str()) else {
129                continue;
130            };
131            if let Some(exclude) = exclude_id {
132                if let Ok(Some(row)) = self.get_record(table, exclude).await {
133                    if row.get(&field).and_then(|v| v.as_str()) == Some(value) {
134                        continue;
135                    }
136                }
137            }
138            let key = self.keys.uniq(table, &field, value);
139            let mut conn = self.conn.clone();
140            let set: bool = conn.set_nx(&key, id).await.map_err(Self::map_err)?;
141            if !set {
142                let existing: Option<String> = conn.get(&key).await.map_err(Self::map_err)?;
143                if existing.as_deref() != Some(id) {
144                    return Err(Error::Database(format!(
145                        "duplicate unique index value for {table}.{field}"
146                    )));
147                }
148            }
149        }
150        Ok(())
151    }
152
153    async fn release_unique_fields(&self, table: &str, record: &Value) -> Result<()> {
154        for field in self.unique_fields(table).await? {
155            if let Some(value) = record.get(&field).and_then(|v| v.as_str()) {
156                let key = self.keys.uniq(table, &field, value);
157                let mut conn = self.conn.clone();
158                let _: () = conn.del(&key).await.map_err(Self::map_err)?;
159            }
160        }
161        Ok(())
162    }
163
164    async fn rows_for_table(&self, table: &str, limit: Option<usize>) -> Result<Vec<Value>> {
165        Self::assert_safe_table(table)?;
166        let ids_key = self.keys.table_ids(table);
167        let mut conn = self.conn.clone();
168        let ids: Vec<String> = conn.smembers(&ids_key).await.map_err(Self::map_err)?;
169        let mut rows = Vec::new();
170        for id in ids {
171            if let Some(row) = self.get_record(table, &id).await? {
172                rows.push(row);
173            }
174            if limit.is_some_and(|n| rows.len() >= n) {
175                break;
176            }
177        }
178        Ok(rows)
179    }
180
181    fn execute_redis_descriptor(descriptor: &Value) -> Result<(String, Option<usize>)> {
182        let index = descriptor
183            .get("index")
184            .and_then(|v| v.as_str())
185            .ok_or_else(|| Error::Internal("missing index in redis query".into()))?;
186        let table = index
187            .strip_prefix("idx:")
188            .ok_or_else(|| Error::Internal(format!("invalid redis index: {index}")))?;
189        let limit = descriptor
190            .get("limit")
191            .and_then(|v| v.as_u64())
192            .map(|n| n as usize);
193        Ok((table.to_string(), limit))
194    }
195
196    fn parse_sql_select(q: &str) -> Result<(String, Option<usize>, bool)> {
197        let upper = q.to_uppercase();
198        if !upper.starts_with("SELECT ") {
199            return Err(Error::Internal("not a SELECT query".into()));
200        }
201        let from_idx = upper
202            .find(" FROM ")
203            .ok_or_else(|| Error::Internal("missing FROM in select".into()))?;
204        let table = q[from_idx + 6..]
205            .split_whitespace()
206            .next()
207            .unwrap_or("")
208            .trim()
209            .to_string();
210        let id_only = upper.contains("SELECT ID") && !upper.contains("BODY");
211        let limit = upper
212            .rfind(" LIMIT ")
213            .and_then(|idx| q[idx + 7..].trim().parse::<usize>().ok());
214        Ok((table, limit, id_only))
215    }
216}
217
218#[async_trait::async_trait]
219impl DatabaseBackend for RedisBackend {
220    fn engine_id(&self) -> &'static str {
221        ENGINE_ID
222    }
223
224    fn capabilities(&self) -> BackendCapabilities {
225        BackendCapabilities {
226            supports_merge: true,
227            supports_graph_edges: true,
228            telemetry_label: "redis",
229        }
230    }
231
232    async fn execute_compiled_query(&self, compiled: &CompiledQuery) -> Result<Vec<Value>> {
233        let q = compiled.query_string.trim();
234        if let Ok(descriptor) = serde_json::from_str::<Value>(q) {
235            if descriptor.get("index").is_some() {
236                let (table, _limit) = Self::execute_redis_descriptor(&descriptor)?;
237                let mut rows = self.rows_for_table(&table, None).await?;
238                rows = valence_core::query::apply_equality_where(rows, compiled);
239                rows = valence_core::query::apply_order_limit_offset(rows, &compiled.query_string);
240                return Ok(rows);
241            }
242        }
243
244        let (table, _limit, id_only) = match Self::parse_sql_select(q) {
245            Ok(parsed) => parsed,
246            Err(_) => return Ok(vec![]),
247        };
248        if table.is_empty() {
249            return Ok(vec![]);
250        }
251        // Load all candidates; WHERE / ORDER / LIMIT applied in-process (parity with mem).
252        let mut rows = self.rows_for_table(&table, None).await?;
253        rows = valence_core::query::apply_equality_where(rows, compiled);
254        rows = valence_core::query::apply_order_limit_offset(rows, &compiled.query_string);
255        if id_only {
256            // Match mem: IdOnlyRecord deserializes `{ "id": ... }`, not bare strings.
257            return Ok(rows
258                .iter()
259                .filter_map(|r| {
260                    r.get("id")
261                        .and_then(|id| id.get("id").and_then(|x| x.as_str()))
262                        .or_else(|| r.get("id").and_then(|id| id.as_str()))
263                        .map(|id| serde_json::json!({ "id": id }))
264                })
265                .collect());
266        }
267        Ok(rows)
268    }
269
270    async fn ensure_schemaless_table(&self, table: &str) -> Result<()> {
271        Self::assert_safe_table(table)?;
272        Ok(())
273    }
274
275    async fn get_record(&self, table: &str, id: &str) -> Result<Option<Value>> {
276        Self::assert_safe_table(table)?;
277        let key = self.keys.doc(table, id);
278        let mut conn = self.conn.clone();
279        let raw: Option<String> = conn.get(&key).await.map_err(Self::map_err)?;
280        Ok(raw.map(|text| {
281            let body: Value =
282                serde_json::from_str(&text).unwrap_or_else(|_| Value::Object(Map::new()));
283            row_from_body(table, id, body)
284        }))
285    }
286
287    async fn create_record(&self, table: &str, content: Value) -> Result<Value> {
288        Self::assert_safe_table(table)?;
289        let id = storage_id(&content).unwrap_or_else(uuid_simple);
290        let mut record = content;
291        if let Some(obj) = record.as_object_mut() {
292            let has_string_id = obj.get("id").and_then(|v| v.as_str()).is_some();
293            if !has_string_id {
294                obj.insert("id".into(), record_id_json(table, &id));
295            }
296        }
297        self.claim_unique_fields(table, &id, &record, None).await?;
298        let body = strip_id_field(&record);
299        let body_text =
300            serde_json::to_string(&body).map_err(|e| Error::Serialization(e.to_string()))?;
301        let doc_key = self.keys.doc(table, &id);
302        let ids_key = self.keys.table_ids(table);
303        let mut conn = self.conn.clone();
304        let _: () = conn
305            .set(&doc_key, &body_text)
306            .await
307            .map_err(Self::map_err)?;
308        let _: () = conn.sadd(&ids_key, &id).await.map_err(Self::map_err)?;
309        Ok(record)
310    }
311
312    async fn update_record(&self, table: &str, id: &str, content: Value) -> Result<Value> {
313        let existing = self
314            .get_record(table, id)
315            .await?
316            .ok_or_else(|| Error::NotFound(format!("{table}:{id}")))?;
317        self.release_unique_fields(table, &existing).await?;
318        self.claim_unique_fields(table, id, &content, Some(id))
319            .await?;
320        let mut record = content;
321        if let Some(obj) = record.as_object_mut() {
322            obj.insert("id".into(), record_id_json(table, id));
323        }
324        let body = strip_id_field(&record);
325        let body_text =
326            serde_json::to_string(&body).map_err(|e| Error::Serialization(e.to_string()))?;
327        let doc_key = self.keys.doc(table, id);
328        let mut conn = self.conn.clone();
329        let _: () = conn
330            .set(&doc_key, &body_text)
331            .await
332            .map_err(Self::map_err)?;
333        Ok(record)
334    }
335
336    async fn merge_record(&self, table: &str, id: &str, patch: Value) -> Result<Value> {
337        let existing = self
338            .get_record(table, id)
339            .await?
340            .unwrap_or_else(|| row_from_body(table, id, Value::Object(Map::new())));
341        self.release_unique_fields(table, &existing).await?;
342        let mut merged = existing;
343        if let (Some(base), Some(patch_obj)) = (merged.as_object_mut(), patch.as_object()) {
344            for (k, v) in patch_obj {
345                base.insert(k.clone(), v.clone());
346            }
347        }
348        self.claim_unique_fields(table, id, &merged, Some(id))
349            .await?;
350        let body = strip_id_field(&merged);
351        let body_text =
352            serde_json::to_string(&body).map_err(|e| Error::Serialization(e.to_string()))?;
353        let doc_key = self.keys.doc(table, id);
354        let ids_key = self.keys.table_ids(table);
355        let mut conn = self.conn.clone();
356        let _: () = conn
357            .set(&doc_key, &body_text)
358            .await
359            .map_err(Self::map_err)?;
360        let _: () = conn.sadd(&ids_key, id).await.map_err(Self::map_err)?;
361        Ok(merged)
362    }
363
364    async fn upsert_record(&self, table: &str, id: &str, content: Value) -> Result<Value> {
365        if self.get_record(table, id).await?.is_some() {
366            self.update_record(table, id, content).await
367        } else {
368            let mut record = content;
369            if let Some(obj) = record.as_object_mut() {
370                obj.insert("id".into(), record_id_json(table, id));
371            }
372            self.create_record(table, record).await
373        }
374    }
375
376    async fn delete_record(&self, table: &str, id: &str) -> Result<()> {
377        if let Some(existing) = self.get_record(table, id).await? {
378            self.release_unique_fields(table, &existing).await?;
379        }
380        let doc_key = self.keys.doc(table, id);
381        let ids_key = self.keys.table_ids(table);
382        let mut conn = self.conn.clone();
383        let _: () = conn.del(&doc_key).await.map_err(Self::map_err)?;
384        let _: () = conn.srem(&ids_key, id).await.map_err(Self::map_err)?;
385        Ok(())
386    }
387
388    async fn relate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
389        let key = self.keys.edge(edge_table, from.table(), from.id());
390        let member = format!("{}:{}", to.table(), to.id());
391        let mut conn = self.conn.clone();
392        let _: () = conn.sadd(&key, member).await.map_err(Self::map_err)?;
393        Ok(())
394    }
395
396    async fn unrelate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
397        let key = self.keys.edge(edge_table, from.table(), from.id());
398        let member = format!("{}:{}", to.table(), to.id());
399        let mut conn = self.conn.clone();
400        let _: () = conn.srem(&key, member).await.map_err(Self::map_err)?;
401        Ok(())
402    }
403
404    async fn get_edge_targets(&self, from: &RecordId, edge_table: &str) -> Result<Vec<RecordId>> {
405        let key = self.keys.edge(edge_table, from.table(), from.id());
406        let mut conn = self.conn.clone();
407        let members: Vec<String> = conn.smembers(&key).await.map_err(Self::map_err)?;
408        Ok(members
409            .into_iter()
410            .filter_map(|m| {
411                let (table, id) = m.split_once(':')?;
412                Some(RecordId::new(table.to_string(), id.to_string()))
413            })
414            .collect())
415    }
416
417    async fn define_unique_index(&self, table: &str, field: &str) -> Result<()> {
418        Self::assert_safe_table(table)?;
419        let idx_key = self.keys.uniq_index(table);
420        let mut conn = self.conn.clone();
421        let _: () = conn.sadd(&idx_key, field).await.map_err(Self::map_err)?;
422        for row in self.rows_for_table(table, None).await? {
423            if let Some(value) = row.get(field).and_then(|v| v.as_str()) {
424                let id = row
425                    .get("id")
426                    .and_then(|v| v.get("id").and_then(|x| x.as_str()))
427                    .or_else(|| row.get("id").and_then(|v| v.as_str()))
428                    .unwrap_or("");
429                if !id.is_empty() {
430                    let uniq_key = self.keys.uniq(table, field, value);
431                    let _: bool = conn.set_nx(&uniq_key, id).await.map_err(Self::map_err)?;
432                }
433            }
434        }
435        Ok(())
436    }
437}
438
439fn row_from_body(table: &str, id: &str, body: Value) -> Value {
440    let mut obj = body.as_object().cloned().unwrap_or_default();
441    obj.insert("id".into(), record_id_json(table, id));
442    Value::Object(obj)
443}
444
445fn strip_id_field(record: &Value) -> Map<String, Value> {
446    record
447        .as_object()
448        .cloned()
449        .unwrap_or_default()
450        .into_iter()
451        .filter(|(k, _)| k != "id")
452        .collect()
453}
454
455fn record_id_json(table: &str, id: &str) -> Value {
456    serde_json::json!({
457        "table": table,
458        "id": id,
459    })
460}
461
462fn storage_id(content: &Value) -> Option<String> {
463    content.get("id").and_then(|v| {
464        v.get("id")
465            .and_then(|x| x.as_str())
466            .map(str::to_string)
467            .or_else(|| v.as_str().map(str::to_string))
468    })
469}
470
471fn uuid_simple() -> String {
472    uuid::Uuid::new_v4().to_string()
473}