1use 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
16pub const ENGINE_ID: &str = KnownEngines::REDIS;
18
19pub const PRIMARY: DatabaseFromEngine = Database::from_engine("primary", ENGINE_ID);
21
22#[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 pub fn builder() -> crate::config::RedisBackendBuilder {
76 crate::config::RedisBackendBuilder::new()
77 }
78
79 pub async fn from_env() -> Result<Self> {
85 Self::builder().from_env_defaults().build().await
86 }
87
88 pub async fn connect(url: &str) -> Result<Self> {
94 Self::builder().url(url).build().await
95 }
96
97 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)] 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 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 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 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 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 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 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 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}