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};
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, json_to_surreal_content_value, map_looks_like_surreal_thing_only,
21    record_map_to_json_object, select_record_json, thing_only_key_from_tb_id_map, thing_to_id_only,
22    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.into_iter().next().expect("len checked");
132    if let Some(m) = try_value_as_record_map(&row) {
133        if map_looks_like_surreal_thing_only(&m) {
134            let id = thing_only_key_from_tb_id_map(&m)?;
135            return select_record_json(db, table, &id)
136                .await?
137                .ok_or_else(|| Error::Validation("Failed to read record after create".into()));
138        }
139        return Ok({
140            let mut json = record_map_to_json_object(&m);
141            valence_core::row_json::normalize_record_id_field(table, &mut json);
142            json
143        });
144    }
145
146    Err(Error::Validation(
147        "Failed to decode create response from database".into(),
148    ))
149}
150
151#[async_trait::async_trait]
152impl DatabaseBackend for SurrealEmbeddedBackend {
153    fn engine_id(&self) -> &'static str {
154        ENGINE_ID
155    }
156
157    fn capabilities(&self) -> BackendCapabilities {
158        surreal_capabilities()
159    }
160
161    fn as_any_local(&self) -> Option<&dyn Any> {
162        Some(self as &dyn Any)
163    }
164
165    async fn use_namespace(&self, ns: &str, db_name: &str) -> Result<()> {
166        self.db.use_ns(ns).use_db(db_name).await.map_err(db_err)?;
167        Ok(())
168    }
169
170    async fn execute_compiled_query(
171        &self,
172        compiled: &CompiledQuery,
173    ) -> Result<Vec<serde_json::Value>> {
174        execute_compiled_query_inner(&self.db, &compiled.query_string, &compiled.params).await
175    }
176
177    async fn get_record(&self, table: &str, id: &str) -> Result<Option<serde_json::Value>> {
178        ensure_schemaless_table(&self.db, table).await?;
179        select_record_json(&self.db, table, id).await
180    }
181
182    async fn create_record(
183        &self,
184        table: &str,
185        content: serde_json::Value,
186    ) -> Result<serde_json::Value> {
187        ensure_schemaless_table(&self.db, table).await?;
188        let explicit_id = content
189            .get("id")
190            .and_then(|v| v.as_str())
191            .map(|s| thing_to_id_only(s.to_string()))
192            .filter(|s| !s.is_empty());
193        let json_content = strip_id_from_content(content);
194        let resource = match explicit_id.as_deref() {
195            Some(id) => surrealdb::opt::Resource::from((table, id)),
196            None => surrealdb::opt::Resource::from(table),
197        };
198        let surreal_content = json_to_surreal_content_value(json_content);
199        let raw: SurrealValueType = self
200            .db
201            .create(resource)
202            .content(surreal_content)
203            .await
204            .map_err(db_err)?;
205
206        if let Some(id_for_get) = explicit_id {
207            return select_record_json(&self.db, table, &id_for_get)
208                .await?
209                .ok_or_else(|| Error::Validation("Failed to read record after create".into()));
210        }
211
212        row_json_after_create(&self.db, table, raw).await
213    }
214
215    async fn update_record(
216        &self,
217        table: &str,
218        id: &str,
219        content: serde_json::Value,
220    ) -> Result<serde_json::Value> {
221        ensure_schemaless_table(&self.db, table).await?;
222        let resource = surrealdb::opt::Resource::from((table, id));
223        let content = json_to_surreal_content_value(strip_id_from_content(content));
224        let _: SurrealValueType = self
225            .db
226            .update(resource)
227            .content(content)
228            .await
229            .map_err(db_err)?;
230        select_record_json(&self.db, table, id)
231            .await?
232            .ok_or_else(|| Error::Validation("Failed to read record after update".into()))
233    }
234
235    async fn merge_record(
236        &self,
237        table: &str,
238        id: &str,
239        patch: serde_json::Value,
240    ) -> Result<serde_json::Value> {
241        ensure_schemaless_table(&self.db, table).await?;
242        let resource = surrealdb::opt::Resource::from((table, id));
243        let patch = json_to_surreal_content_value(strip_id_from_content(patch));
244        let _: SurrealValueType = self
245            .db
246            .update(resource)
247            .merge(patch)
248            .await
249            .map_err(db_err)?;
250        select_record_json(&self.db, table, id)
251            .await?
252            .ok_or_else(|| Error::Validation("Failed to read record after merge".into()))
253    }
254
255    async fn upsert_record(
256        &self,
257        table: &str,
258        id: &str,
259        content: serde_json::Value,
260    ) -> Result<serde_json::Value> {
261        ensure_schemaless_table(&self.db, table).await?;
262        let resource = surrealdb::opt::Resource::from((table, id));
263        let content = json_to_surreal_content_value(strip_id_from_content(content));
264        let _: SurrealValueType = self
265            .db
266            .upsert(resource)
267            .content(content)
268            .await
269            .map_err(db_err)?;
270        select_record_json(&self.db, table, id)
271            .await?
272            .ok_or_else(|| Error::Validation("Failed to read record after upsert".into()))
273    }
274
275    async fn delete_record(&self, table: &str, id: &str) -> Result<()> {
276        let resource = surrealdb::opt::Resource::from((table, id));
277        let _: SurrealValueType = self.db.delete(resource).await.map_err(db_err)?;
278        Ok(())
279    }
280
281    async fn relate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
282        let from_t = surreal_from_valence(from);
283        let to_t = surreal_from_valence(to);
284        let q = format!("RELATE $from->{edge_table}->$to RETURN NONE");
285        ensure_schemaless_table(&self.db, edge_table).await?;
286        self.db
287            .query(&q)
288            .bind(("from", from_t))
289            .bind(("to", to_t))
290            .await
291            .map_err(db_err)?;
292        Ok(())
293    }
294
295    async fn unrelate_edge(&self, from: &RecordId, edge_table: &str, to: &RecordId) -> Result<()> {
296        let from_t = surreal_from_valence(from);
297        let to_t = surreal_from_valence(to);
298        let q = format!("DELETE $from->{edge_table} WHERE `out` = $to RETURN NONE");
299        ensure_schemaless_table(&self.db, edge_table).await?;
300        self.db
301            .query(&q)
302            .bind(("from", from_t))
303            .bind(("to", to_t))
304            .await
305            .map_err(db_err)?;
306        Ok(())
307    }
308
309    async fn get_edge_targets(&self, from: &RecordId, edge_table: &str) -> Result<Vec<RecordId>> {
310        use crate::query_exec::query_err_is_missing_table;
311
312        let from_t = surreal_from_valence(from);
313        let q = format!("SELECT VALUE `out` FROM {edge_table} WHERE `in` = $from");
314        let mut response = match self.db.query(&q).bind(("from", from_t)).await {
315            Ok(r) => r,
316            Err(e) if query_err_is_missing_table(&e.to_string()) => {
317                return Ok(vec![]);
318            }
319            Err(e) => return Err(db_err(e)),
320        };
321        let outs: Vec<surrealdb::types::RecordId> = match response.take(0) {
322            Ok(r) => r,
323            Err(e) if query_err_is_missing_table(&e.to_string()) => {
324                return Ok(vec![]);
325            }
326            Err(e) => return Err(db_err(e)),
327        };
328        Ok(outs.into_iter().map(valence_from_surreal).collect())
329    }
330
331    async fn ensure_schemaless_table(&self, table: &str) -> Result<()> {
332        ensure_schemaless_table(&self.db, table).await
333    }
334
335    async fn define_unique_index(&self, table: &str, field: &str) -> Result<()> {
336        ensure_schemaless_table(&self.db, table).await?;
337        let index_name = format!("idx_{table}_{field}_unique");
338        let query = format!("DEFINE INDEX {index_name} ON TABLE {table} COLUMNS {field} UNIQUE");
339        match self.db.query(&query).await {
340            Ok(_) => Ok(()),
341            Err(e) => {
342                let message = e.to_string().to_lowercase();
343                if message.contains("already") && message.contains("index") {
344                    Ok(())
345                } else {
346                    Err(db_err(e))
347                }
348            }
349        }
350    }
351
352    fn ttl_capability(&self) -> BackendTtlCapability {
353        BackendTtlCapability::Deferred
354    }
355
356    async fn apply_ttl_policy(&self, _table: &str, _policy: &SchemaTtlPolicy) -> Result<()> {
357        Ok(())
358    }
359}
360
361/// Alias for [`SurrealEmbeddedBackend`] (historical template name).
362pub type SurrealMemBackend = SurrealEmbeddedBackend;
363
364#[cfg(test)]
365mod tests {
366    use super::*;
367    use surrealdb::engine::local::Mem;
368
369    async fn mem_backend() -> SurrealEmbeddedBackend {
370        let db = SDb::init();
371        db.connect::<Mem>(()).await.unwrap();
372        db.use_ns("test").use_db("test").await.unwrap();
373        SurrealEmbeddedBackend::new(db)
374    }
375
376    #[tokio::test]
377    async fn create_record_id_field_is_table_object() {
378        let b = mem_backend().await;
379        let row = b
380            .create_record("widget", serde_json::json!({"name": "alpha"}))
381            .await
382            .expect("create");
383        let id = row.get("id").expect("id field");
384        assert!(
385            id.get("table").is_some() && id.get("id").is_some(),
386            "expected RecordId object, got {id:?}"
387        );
388    }
389
390    #[tokio::test]
391    async fn define_unique_index_idempotent() {
392        let b = mem_backend().await;
393        b.define_unique_index("uniq_tbl", "email")
394            .await
395            .expect("first define");
396        b.define_unique_index("uniq_tbl", "email")
397            .await
398            .expect("second define idempotent");
399    }
400}