Skip to main content

valence_backend_surreal/
embedded.rs

1//! In-process SurrealDB (`Surreal<Db>`) backend.
2
3use std::any::Any;
4
5use surrealdb::engine::local::Db;
6use surrealdb::types::Value as SurrealValueType;
7use surrealdb::{Connection, Surreal};
8
9use valence_core::backend::{BackendCapabilities, DatabaseBackend};
10use valence_core::compiled_query::CompiledQuery;
11use valence_core::error::{Error, Result};
12use valence_core::record_id::RecordId;
13use valence_core::ttl::{BackendTtlCapability, SchemaTtlPolicy, EXPIRE_AT_FIELD};
14use valence_core::KnownEngines;
15
16use crate::error::db_err;
17use crate::query_exec::execute_compiled_query_inner;
18use crate::record_id::{surreal_from_valence, valence_from_surreal};
19use crate::row_json::{
20    ensure_schemaless_table, ensure_typed_table, json_to_surreal_content_value,
21    map_looks_like_surreal_thing_only, record_map_to_json_object, select_record_json,
22    sync_typed_table, thing_only_key_from_tb_id_map, thing_to_id_only, try_value_as_record_map,
23};
24
25/// Stable engine slug for router keys (`surrealdb:logical_name`).
26pub const ENGINE_ID: &str = KnownEngines::SURREALDB;
27
28/// Embedded Surreal client type.
29pub type SDb = Surreal<Db>;
30
31pub const fn surreal_capabilities() -> BackendCapabilities {
32    BackendCapabilities {
33        supports_merge: true,
34        supports_graph_edges: true,
35        telemetry_label: "surrealdb",
36    }
37}
38
39/// Wraps a local Surreal handle (memory, RocksDB, etc.).
40///
41/// # Examples
42///
43/// ```ignore
44/// use std::sync::Arc;
45/// use valence::{
46///     valence_schema, Database, DatabaseFromEngine, FieldType, SDb, SurrealEmbeddedBackend,
47///     Valence, SURREAL_ENGINE_ID,
48/// };
49///
50/// const COUNTER_DB: DatabaseFromEngine =
51///     Database::from_engine("default", SURREAL_ENGINE_ID);
52///
53/// valence_schema! {
54///     Counter {
55///         table: "counter",
56///         version: "0.1.0",
57///         database: COUNTER_DB,
58///         fields: [
59///             id: { r#type: FieldType::String, primary_key: true, required: true },
60///             value: { r#type: FieldType::Integer, required: true },
61///         ],
62///     }
63/// }
64///
65/// let db = SDb::init();
66/// db.connect::<surrealdb::engine::local::Mem>(()).await?;
67/// db.use_ns("demo").use_db("demo").await?;
68/// let valence = Valence::builder()
69///     .add_backend("default", Arc::new(SurrealEmbeddedBackend::new(db)))
70///     .build()?;
71/// assert_eq!(
72///     valence.backend_for_table("counter")?.engine_id(),
73///     SURREAL_ENGINE_ID
74/// );
75/// # Ok::<(), valence::Error>(())
76/// ```
77#[derive(Debug, Clone)]
78pub struct SurrealEmbeddedBackend {
79    db: SDb,
80}
81
82impl SurrealEmbeddedBackend {
83    /// Wrap an existing embedded Surreal client.
84    pub fn new(db: SDb) -> Self {
85        Self { db }
86    }
87
88    /// Borrow the underlying Surreal client.
89    pub fn inner(&self) -> &SDb {
90        &self.db
91    }
92
93    /// Consume the adapter and return the Surreal client.
94    pub fn into_inner(self) -> SDb {
95        self.db
96    }
97}
98
99pub fn strip_id_from_content(mut content: serde_json::Value) -> serde_json::Value {
100    if let serde_json::Value::Object(ref mut map) = content {
101        map.remove("id");
102    }
103    content
104}
105
106pub async fn row_json_after_create<C>(
107    db: &Surreal<C>,
108    table: &str,
109    raw: SurrealValueType,
110) -> Result<serde_json::Value>
111where
112    C: Connection,
113{
114    let rows: Vec<SurrealValueType> = match raw {
115        SurrealValueType::Array(arr) => arr.into_inner(),
116        other => vec![other],
117    };
118    match rows.len() {
119        0 => {
120            return Err(Error::Validation(
121                "Failed to read record after create (empty response)".into(),
122            ));
123        }
124        1 => {}
125        _ => {
126            return Err(Error::Validation(
127                "Unexpected multi-row create response".into(),
128            ));
129        }
130    }
131    let row = rows
132        .into_iter()
133        .next()
134        .ok_or_else(|| Error::Internal("create response length invariant violated".into()))?;
135    if let Some(m) = try_value_as_record_map(&row) {
136        if map_looks_like_surreal_thing_only(&m) {
137            let id = thing_only_key_from_tb_id_map(&m)?;
138            return select_record_json(db, table, &id)
139                .await?
140                .ok_or_else(|| Error::Validation("Failed to read record after create".into()));
141        }
142        return Ok({
143            let mut json = record_map_to_json_object(&m);
144            valence_core::row_json::normalize_record_id_field(table, &mut json);
145            json
146        });
147    }
148
149    Err(Error::Validation(
150        "Failed to decode create response from database".into(),
151    ))
152}
153
154#[async_trait::async_trait]
155impl DatabaseBackend for SurrealEmbeddedBackend {
156    fn engine_id(&self) -> &'static str {
157        ENGINE_ID
158    }
159
160    fn capabilities(&self) -> BackendCapabilities {
161        surreal_capabilities()
162    }
163
164    fn as_any_local(&self) -> Option<&dyn Any> {
165        Some(self as &dyn Any)
166    }
167
168    async fn use_namespace(&self, ns: &str, db_name: &str) -> Result<()> {
169        self.db.use_ns(ns).use_db(db_name).await.map_err(db_err)?;
170        Ok(())
171    }
172
173    async fn execute_compiled_query(
174        &self,
175        compiled: &CompiledQuery,
176    ) -> Result<Vec<serde_json::Value>> {
177        execute_compiled_query_inner(&self.db, &compiled.query_string, &compiled.params).await
178    }
179
180    async fn get_record(&self, table: &str, id: &str) -> Result<Option<serde_json::Value>> {
181        ensure_schemaless_table(&self.db, table).await?;
182        select_record_json(&self.db, table, id).await
183    }
184
185    async fn create_record(
186        &self,
187        table: &str,
188        content: serde_json::Value,
189    ) -> Result<serde_json::Value> {
190        ensure_schemaless_table(&self.db, table).await?;
191        let mut content = content;
192        valence_core::ttl::prepare_create_content(table, self, &mut content)?;
193        let explicit_id = content
194            .get("id")
195            .and_then(|v| {
196                v.as_str().map(str::to_string).or_else(|| {
197                    v.as_object()
198                        .and_then(|o| o.get("id"))
199                        .and_then(|x| x.as_str())
200                        .map(str::to_string)
201                })
202            })
203            .map(thing_to_id_only)
204            .filter(|s| !s.is_empty());
205        let json_content = strip_id_from_content(content);
206        let resource = match explicit_id.as_deref() {
207            Some(id) => surrealdb::opt::Resource::from((table, id)),
208            None => surrealdb::opt::Resource::from(table),
209        };
210        let surreal_content = json_to_surreal_content_value(json_content);
211        let raw: SurrealValueType = self
212            .db
213            .create(resource)
214            .content(surreal_content)
215            .await
216            .map_err(db_err)?;
217
218        if let Some(id_for_get) = explicit_id {
219            return select_record_json(&self.db, table, &id_for_get)
220                .await?
221                .ok_or_else(|| Error::Validation("Failed to read record after create".into()));
222        }
223
224        row_json_after_create(&self.db, table, raw).await
225    }
226
227    async fn update_record(
228        &self,
229        table: &str,
230        id: &str,
231        content: serde_json::Value,
232    ) -> Result<serde_json::Value> {
233        ensure_schemaless_table(&self.db, table).await?;
234        let resource = surrealdb::opt::Resource::from((table, id));
235        let content = json_to_surreal_content_value(strip_id_from_content(content));
236        let _: SurrealValueType = self
237            .db
238            .update(resource)
239            .content(content)
240            .await
241            .map_err(db_err)?;
242        select_record_json(&self.db, table, id)
243            .await?
244            .ok_or_else(|| Error::Validation("Failed to read record after update".into()))
245    }
246
247    async fn merge_record(
248        &self,
249        table: &str,
250        id: &str,
251        patch: serde_json::Value,
252    ) -> Result<serde_json::Value> {
253        ensure_schemaless_table(&self.db, table).await?;
254        let resource = surrealdb::opt::Resource::from((table, id));
255        let patch = json_to_surreal_content_value(strip_id_from_content(patch));
256        let _: SurrealValueType = self
257            .db
258            .update(resource)
259            .merge(patch)
260            .await
261            .map_err(db_err)?;
262        select_record_json(&self.db, table, id)
263            .await?
264            .ok_or_else(|| Error::Validation("Failed to read record after merge".into()))
265    }
266
267    async fn upsert_record(
268        &self,
269        table: &str,
270        id: &str,
271        content: serde_json::Value,
272    ) -> Result<serde_json::Value> {
273        ensure_schemaless_table(&self.db, table).await?;
274        let resource = surrealdb::opt::Resource::from((table, id));
275        let content = json_to_surreal_content_value(strip_id_from_content(content));
276        let _: SurrealValueType = self
277            .db
278            .upsert(resource)
279            .content(content)
280            .await
281            .map_err(db_err)?;
282        select_record_json(&self.db, table, id)
283            .await?
284            .ok_or_else(|| Error::Validation("Failed to read record after upsert".into()))
285    }
286
287    async fn delete_record(&self, table: &str, id: &str) -> Result<()> {
288        let resource = surrealdb::opt::Resource::from((table, id));
289        let _: SurrealValueType = self.db.delete(resource).await.map_err(db_err)?;
290        Ok(())
291    }
292
293    async fn relate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
294        let from_t = surreal_from_valence(from);
295        let to_t = surreal_from_valence(to);
296        let q = format!("RELATE $from->{edge_table}->$to RETURN NONE");
297        ensure_schemaless_table(&self.db, edge_table).await?;
298        self.db
299            .query(&q)
300            .bind(("from", from_t))
301            .bind(("to", to_t))
302            .await
303            .map_err(db_err)?;
304        Ok(())
305    }
306
307    async fn unrelate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
308        let from_t = surreal_from_valence(from);
309        let to_t = surreal_from_valence(to);
310        let q = format!("DELETE $from->{edge_table} WHERE `out` = $to RETURN NONE");
311        ensure_schemaless_table(&self.db, edge_table).await?;
312        self.db
313            .query(&q)
314            .bind(("from", from_t))
315            .bind(("to", to_t))
316            .await
317            .map_err(db_err)?;
318        Ok(())
319    }
320
321    async fn get_edge_targets(&self, from: &RecordId, edge_table: &str) -> Result<Vec<RecordId>> {
322        use crate::query_exec::query_err_is_missing_table;
323
324        let from_t = surreal_from_valence(from);
325        let q = format!("SELECT VALUE `out` FROM {edge_table} WHERE `in` = $from");
326        let mut response = match self.db.query(&q).bind(("from", from_t)).await {
327            Ok(r) => r,
328            Err(e) if query_err_is_missing_table(&e.to_string()) => {
329                return Ok(vec![]);
330            }
331            Err(e) => return Err(db_err(e)),
332        };
333        let outs: Vec<surrealdb::types::RecordId> = match response.take(0) {
334            Ok(r) => r,
335            Err(e) if query_err_is_missing_table(&e.to_string()) => {
336                return Ok(vec![]);
337            }
338            Err(e) => return Err(db_err(e)),
339        };
340        Ok(outs.into_iter().map(valence_from_surreal).collect())
341    }
342
343    async fn get_edge_sources(&self, to: &RecordId, edge_table: &str) -> Result<Vec<RecordId>> {
344        use crate::query_exec::query_err_is_missing_table;
345
346        let to_t = surreal_from_valence(to);
347        let q = format!("SELECT VALUE `in` FROM {edge_table} WHERE `out` = $to");
348        let mut response = match self.db.query(&q).bind(("to", to_t)).await {
349            Ok(r) => r,
350            Err(e) if query_err_is_missing_table(&e.to_string()) => {
351                return Ok(vec![]);
352            }
353            Err(e) => return Err(db_err(e)),
354        };
355        let ins: Vec<surrealdb::types::RecordId> = match response.take(0) {
356            Ok(r) => r,
357            Err(e) if query_err_is_missing_table(&e.to_string()) => {
358                return Ok(vec![]);
359            }
360            Err(e) => return Err(db_err(e)),
361        };
362        Ok(ins.into_iter().map(valence_from_surreal).collect())
363    }
364
365    async fn ensure_schemaless_table(&self, table: &str) -> Result<()> {
366        ensure_schemaless_table(&self.db, table).await
367    }
368
369    async fn ensure_typed_table(
370        &self,
371        layout: &valence_core::storage_layout::StorageLayout,
372    ) -> Result<()> {
373        ensure_typed_table(&self.db, layout).await
374    }
375
376    async fn sync_typed_table(
377        &self,
378        layout: &valence_core::storage_layout::StorageLayout,
379    ) -> Result<()> {
380        sync_typed_table(&self.db, layout).await
381    }
382
383    async fn define_unique_index(&self, table: &str, field: &str) -> Result<()> {
384        ensure_schemaless_table(&self.db, table).await?;
385        let index_name = format!("idx_{table}_{field}_unique");
386        let query = format!("DEFINE INDEX {index_name} ON TABLE {table} COLUMNS {field} UNIQUE");
387        match self.db.query(&query).await {
388            Ok(_) => Ok(()),
389            Err(e) => {
390                let message = e.to_string().to_lowercase();
391                if message.contains("already") && message.contains("index") {
392                    Ok(())
393                } else {
394                    Err(db_err(e))
395                }
396            }
397        }
398    }
399
400    fn ttl_capability(&self) -> BackendTtlCapability {
401        BackendTtlCapability::Deferred
402    }
403
404    async fn apply_ttl_policy(&self, table: &str, _policy: &SchemaTtlPolicy) -> Result<()> {
405        ensure_schemaless_table(&self.db, table).await?;
406        let field = EXPIRE_AT_FIELD;
407        valence_core::safe_ident::assert_safe_ident(field)?;
408        let index_name = format!("valence_ttl_expire_at_{table}");
409        let query = format!("DEFINE INDEX {index_name} ON TABLE {table} COLUMNS {field}");
410        match self.db.query(&query).await {
411            Ok(_) => Ok(()),
412            Err(e) => {
413                let message = e.to_string().to_lowercase();
414                if message.contains("already") && message.contains("index") {
415                    Ok(())
416                } else {
417                    Err(db_err(e))
418                }
419            }
420        }
421    }
422}
423
424/// Alias for [`SurrealEmbeddedBackend`] (historical template name).
425pub type SurrealMemBackend = SurrealEmbeddedBackend;
426
427#[cfg(test)]
428mod tests {
429    #![allow(
430        clippy::unwrap_used,
431        clippy::expect_used,
432        clippy::print_stdout,
433        clippy::print_stderr
434    )]
435
436    use super::*;
437    use surrealdb::engine::local::Mem;
438
439    async fn mem_backend() -> SurrealEmbeddedBackend {
440        let db = SDb::init();
441        db.connect::<Mem>(()).await.unwrap();
442        db.use_ns("test").use_db("test").await.unwrap();
443        SurrealEmbeddedBackend::new(db)
444    }
445
446    #[tokio::test]
447    async fn create_record_id_field_is_table_object() {
448        let b = mem_backend().await;
449        let row = b
450            .create_record("widget", serde_json::json!({"name": "alpha"}))
451            .await
452            .expect("create");
453        let id = row.get("id").expect("id field");
454        assert!(
455            id.get("table").is_some() && id.get("id").is_some(),
456            "expected RecordId object, got {id:?}"
457        );
458    }
459
460    #[tokio::test]
461    async fn define_unique_index_idempotent() {
462        let b = mem_backend().await;
463        b.define_unique_index("uniq_tbl", "email")
464            .await
465            .expect("first define");
466        b.define_unique_index("uniq_tbl", "email")
467            .await
468            .expect("second define idempotent");
469    }
470}