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};
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
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.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
361pub 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}