1use 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
15pub const ENGINE_ID: &str = KnownEngines::REDIS;
17
18pub const PRIMARY: DatabaseFromEngine = Database::from_engine("primary", ENGINE_ID);
20
21#[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 pub fn builder() -> crate::config::RedisBackendBuilder {
75 crate::config::RedisBackendBuilder::new()
76 }
77
78 pub async fn from_env() -> Result<Self> {
80 Self::builder().from_env_defaults().build().await
81 }
82
83 pub async fn connect(url: &str) -> Result<Self> {
85 Self::builder().url(url).build().await
86 }
87
88 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 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 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}