1use 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
25pub const ENGINE_ID: &str = KnownEngines::SURREALDB;
27
28pub 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#[derive(Debug, Clone)]
78pub struct SurrealEmbeddedBackend {
79 db: SDb,
80}
81
82impl SurrealEmbeddedBackend {
83 pub fn new(db: SDb) -> Self {
85 Self { db }
86 }
87
88 pub fn inner(&self) -> &SDb {
90 &self.db
91 }
92
93 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
424pub 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}