Skip to main content

valence_backend_redis/
backend.rs

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