Skip to main content

doc_db/
lib.rs

1mod memory;
2pub mod mock;
3mod runtime;
4mod trx;
5mod turso;
6
7use anyhow::Result;
8use bytes::Bytes;
9pub use libsql_hrana::proto::Value;
10use memory::{MemoryDatabase, MemoryTransaction};
11use std::future::Future;
12pub use trx::{
13    ConflictDetails, ConflictKey, DocGet, DocHandle, DocKey, Document, Trx, TrxControl, TrxRead,
14    TrxResult,
15};
16use turso::{StoredDoc, TursoDatabase, TursoTransaction};
17
18pub fn text_value(s: impl Into<String>) -> Value {
19    Value::Text {
20        value: s.into().into(),
21    }
22}
23
24pub fn integer_value(i: i64) -> Value {
25    Value::Integer { value: i }
26}
27
28pub enum BatchOp<'a> {
29    Put {
30        pk: &'a str,
31        sk: &'a str,
32        data: &'a [u8],
33    },
34    Delete {
35        pk: &'a str,
36        sk: &'a str,
37    },
38}
39
40#[derive(Clone)]
41pub enum WriteOp {
42    Insert {
43        pk: String,
44        sk: String,
45        data: Vec<u8>,
46    },
47    Update {
48        pk: String,
49        sk: String,
50        expected_version: i64,
51        data: Vec<u8>,
52    },
53    Delete {
54        pk: String,
55        sk: String,
56        expected_version: i64,
57    },
58}
59
60pub struct CommitOutcome {
61    pub affected_counts: Vec<u64>,
62    pub conflict: Option<ConflictInfo>,
63}
64
65pub struct RawStatement {
66    pub sql: String,
67    pub args: Vec<Value>,
68}
69
70pub struct RawStatementResult {
71    pub column_names: Vec<String>,
72    pub rows: Vec<Vec<Value>>,
73    pub affected_row_count: u64,
74    pub rows_read: u64,
75    pub rows_written: u64,
76    pub query_duration_ms: f64,
77}
78
79pub enum RawTransactionOutcome {
80    Committed {
81        statement_results: Vec<RawStatementResult>,
82    },
83    RolledBack {
84        failed_statement_index: usize,
85        error_message: String,
86    },
87}
88
89pub struct ConflictInfo {
90    pub step_index: usize,
91    pub message: String,
92}
93
94pub fn turso() -> Database {
95    let url = std::env::var("TURSO_URL").expect("TURSO_URL must be set");
96    let auth_token = std::env::var("TURSO_AUTH_TOKEN").expect("TURSO_AUTH_TOKEN must be set");
97    turso_with_config(url, auth_token)
98}
99
100pub fn turso_with_config(url: String, auth_token: String) -> Database {
101    Database {
102        inner: DatabaseInner::Turso(TursoDatabase::new(url, auth_token)),
103        mock_state: mock::MockState::default(),
104    }
105}
106
107pub fn memory() -> Database {
108    Database {
109        inner: DatabaseInner::Memory(MemoryDatabase::new()),
110        mock_state: mock::MockState::default(),
111    }
112}
113
114#[derive(Clone)]
115pub struct Database {
116    inner: DatabaseInner,
117    mock_state: mock::MockState,
118}
119
120impl Database {
121    pub async fn get(&self, pk: &str, sk: &str) -> Result<Option<Bytes>> {
122        if let Some(result) = self.mock_state.try_match(mock::MockOp::Get, pk, sk) {
123            return match result {
124                mock::MockResult::OkGet(data) => Ok(data.map(Bytes::from)),
125                mock::MockResult::Err(msg) => Err(anyhow::anyhow!("{}", msg)),
126                _ => unreachable!(),
127            };
128        }
129        match &self.inner {
130            DatabaseInner::Turso(db) => db.get(pk, sk).await,
131            DatabaseInner::Memory(db) => db.get(pk, sk).await,
132        }
133    }
134
135    pub async fn put(&self, pk: &str, sk: &str, data: &[u8]) -> Result<()> {
136        if let Some(result) = self.mock_state.try_match(mock::MockOp::Put, pk, sk) {
137            return match result {
138                mock::MockResult::OkVoid => Ok(()),
139                mock::MockResult::Err(msg) => Err(anyhow::anyhow!("{}", msg)),
140                _ => unreachable!(),
141            };
142        }
143        match &self.inner {
144            DatabaseInner::Turso(db) => db.put(pk, sk, data).await,
145            DatabaseInner::Memory(db) => db.put(pk, sk, data).await,
146        }
147    }
148
149    pub async fn delete(&self, pk: &str, sk: &str) -> Result<()> {
150        if let Some(result) = self.mock_state.try_match(mock::MockOp::Delete, pk, sk) {
151            return match result {
152                mock::MockResult::OkVoid => Ok(()),
153                mock::MockResult::Err(msg) => Err(anyhow::anyhow!("{}", msg)),
154                _ => unreachable!(),
155            };
156        }
157        match &self.inner {
158            DatabaseInner::Turso(db) => db.delete(pk, sk).await,
159            DatabaseInner::Memory(db) => db.delete(pk, sk).await,
160        }
161    }
162
163    // --- Mock API ---
164
165    pub fn mock_get(&self, pk: &str, sk: &str) -> mock::MockGetBuilder<'_> {
166        mock::MockGetBuilder::new(self, pk.to_string(), sk.to_string())
167    }
168
169    pub fn mock_put(&self, pk: &str, sk: &str) -> mock::MockPutBuilder<'_> {
170        mock::MockPutBuilder::new(self, pk.to_string(), sk.to_string())
171    }
172
173    pub fn mock_delete(&self, pk: &str, sk: &str) -> mock::MockDeleteBuilder<'_> {
174        mock::MockDeleteBuilder::new(self, pk.to_string(), sk.to_string())
175    }
176
177    pub fn clear_mocks(&self) {
178        self.mock_state.clear();
179    }
180
181    pub(crate) fn add_mock_rule(&self, rule: mock::MockRule) {
182        self.mock_state.push(rule);
183    }
184
185    #[tracing::instrument(skip_all, fields(pk = %pk.as_ref(), limit = limit))]
186    pub async fn query<S1: AsRef<str>, S2: AsRef<str>>(
187        &self,
188        pk: S1,
189        after_sk: Option<S2>,
190        limit: usize,
191    ) -> Result<Vec<(String, Bytes)>> {
192        match &self.inner {
193            DatabaseInner::Turso(db) => db.query(pk, after_sk, limit).await,
194            DatabaseInner::Memory(db) => db.query(pk, after_sk, limit).await,
195        }
196    }
197
198    #[tracing::instrument(skip_all, fields(limit = limit))]
199    pub async fn scan(
200        &self,
201        after: Option<(&str, &str)>,
202        limit: usize,
203    ) -> Result<Vec<(String, String, Bytes)>> {
204        match &self.inner {
205            DatabaseInner::Turso(db) => db.scan(after, limit).await,
206            DatabaseInner::Memory(db) => db.scan(after, limit).await,
207        }
208    }
209
210    #[tracing::instrument(skip_all, fields(ops = ops.len()))]
211    pub async fn batch(&self, ops: &[BatchOp<'_>]) -> Result<()> {
212        match &self.inner {
213            DatabaseInner::Turso(db) => db.batch(ops).await,
214            DatabaseInner::Memory(db) => db.batch(ops).await,
215        }
216    }
217
218    #[tracing::instrument(skip_all)]
219    pub async fn transaction(&self) -> Result<Transaction> {
220        match &self.inner {
221            DatabaseInner::Turso(db) => Ok(Transaction {
222                inner: TransactionInner::Turso(db.transaction().await?),
223            }),
224            DatabaseInner::Memory(db) => Ok(Transaction {
225                inner: TransactionInner::Memory(db.transaction().await?),
226            }),
227        }
228    }
229
230    #[tracing::instrument(skip_all)]
231    pub async fn trx<F, Fut, Out, Cancel, E>(&self, f: F) -> TrxResult<Out, Cancel, E>
232    where
233        F: FnMut(Trx) -> Fut,
234        Fut: Future<Output = Result<TrxControl<Out, Cancel>, E>>,
235        E: From<anyhow::Error>,
236    {
237        trx::run(self.clone(), f).await
238    }
239
240    #[tracing::instrument(skip_all, fields(sql = %sql))]
241    pub async fn execute_raw(
242        &self,
243        sql: &str,
244        args: Vec<Value>,
245        want_rows: bool,
246    ) -> Result<Vec<Vec<Value>>> {
247        match &self.inner {
248            DatabaseInner::Turso(db) => db.execute_raw(sql, args, want_rows).await,
249            DatabaseInner::Memory(db) => db.execute_raw(sql, args, want_rows).await,
250        }
251    }
252
253    /// Runs the statements as a single transaction (all-or-nothing) and
254    /// returns, per statement, the result rows together with the engine's own
255    /// `rows_read` / `rows_written` counters — the numbers Turso bills by.
256    /// A statement failure rolls the whole transaction back and is reported as
257    /// [`RawTransactionOutcome::RolledBack`], not as `Err`.
258    ///
259    /// Turso backend only; the in-memory test backend rejects it.
260    #[tracing::instrument(skip_all, fields(statements = statements.len()))]
261    pub async fn execute_raw_transactional(
262        &self,
263        statements: &[RawStatement],
264    ) -> Result<RawTransactionOutcome> {
265        match &self.inner {
266            DatabaseInner::Turso(db) => db.execute_raw_transactional(statements).await,
267            DatabaseInner::Memory(_) => {
268                anyhow::bail!("execute_raw_transactional is only supported on the Turso backend")
269            }
270        }
271    }
272
273    #[tracing::instrument(skip_all, fields(ops = ops.len()))]
274    pub async fn execute_ops(&self, ops: Vec<DbOp>) -> Result<Vec<DbResult>> {
275        match &self.inner {
276            DatabaseInner::Turso(db) => db.execute_ops(ops).await,
277            DatabaseInner::Memory(db) => db.execute_ops(ops).await,
278        }
279    }
280
281    #[tracing::instrument(skip_all, fields(pk = %pk, sk = %sk))]
282    pub(crate) async fn get_with_version(&self, pk: &str, sk: &str) -> Result<Option<StoredDoc>> {
283        match &self.inner {
284            DatabaseInner::Turso(db) => db.get_with_version(pk, sk).await,
285            DatabaseInner::Memory(db) => db.get_with_version(pk, sk).await,
286        }
287    }
288
289    #[tracing::instrument(skip_all, fields(reads = keys.len()))]
290    pub(crate) async fn batch_get_with_version(
291        &self,
292        keys: &[(String, String)],
293    ) -> Result<Vec<Option<StoredDoc>>> {
294        if keys.is_empty() {
295            return Ok(vec![]);
296        }
297        let mut out = Vec::with_capacity(keys.len());
298        for (pk, sk) in keys {
299            out.push(self.get_with_version(pk, sk).await?);
300        }
301        Ok(out)
302    }
303
304    #[tracing::instrument(skip_all, fields(reads = keys.len()))]
305    pub(crate) async fn begin_immediate_with_reads(
306        &self,
307        keys: &[(String, String)],
308    ) -> Result<(Transaction, Vec<Option<StoredDoc>>)> {
309        match &self.inner {
310            DatabaseInner::Turso(db) => {
311                let (tx, docs) = db.begin_immediate_with_reads(keys).await?;
312                Ok((
313                    Transaction {
314                        inner: TransactionInner::Turso(tx),
315                    },
316                    docs,
317                ))
318            }
319            DatabaseInner::Memory(db) => {
320                let (tx, docs) = db.begin_immediate_with_reads(keys).await?;
321                Ok((
322                    Transaction {
323                        inner: TransactionInner::Memory(tx),
324                    },
325                    docs,
326                ))
327            }
328        }
329    }
330}
331
332#[derive(Clone)]
333enum DatabaseInner {
334    Turso(TursoDatabase),
335    Memory(MemoryDatabase),
336}
337
338pub struct Transaction {
339    inner: TransactionInner,
340}
341
342enum TransactionInner {
343    Turso(TursoTransaction),
344    Memory(MemoryTransaction),
345}
346
347impl Transaction {
348    #[tracing::instrument(skip_all, fields(pk = %pk, sk = %sk))]
349    pub async fn get(&mut self, pk: &str, sk: &str) -> Result<Option<Bytes>> {
350        match &mut self.inner {
351            TransactionInner::Turso(tx) => tx.get(pk, sk).await,
352            TransactionInner::Memory(tx) => tx.get(pk, sk).await,
353        }
354    }
355
356    #[tracing::instrument(skip_all, fields(pk = %pk, sk = %sk, bytes = data.len()))]
357    pub async fn put(&mut self, pk: &str, sk: &str, data: &[u8]) -> Result<()> {
358        match &mut self.inner {
359            TransactionInner::Turso(tx) => tx.put(pk, sk, data).await,
360            TransactionInner::Memory(tx) => tx.put(pk, sk, data).await,
361        }
362    }
363
364    #[tracing::instrument(skip_all, fields(pk = %pk, sk = %sk))]
365    pub async fn delete(&mut self, pk: &str, sk: &str) -> Result<()> {
366        match &mut self.inner {
367            TransactionInner::Turso(tx) => tx.delete(pk, sk).await,
368            TransactionInner::Memory(tx) => tx.delete(pk, sk).await,
369        }
370    }
371
372    #[tracing::instrument(skip_all)]
373    pub async fn commit(self) -> Result<()> {
374        match self.inner {
375            TransactionInner::Turso(tx) => tx.commit().await,
376            TransactionInner::Memory(tx) => tx.commit().await,
377        }
378    }
379
380    #[tracing::instrument(skip_all)]
381    pub async fn rollback(self) -> Result<()> {
382        match self.inner {
383            TransactionInner::Turso(tx) => tx.rollback().await,
384            TransactionInner::Memory(tx) => tx.rollback().await,
385        }
386    }
387
388    #[tracing::instrument(skip_all, fields(writes = writes.len()))]
389    pub(crate) async fn apply_writes_and_commit(
390        &mut self,
391        writes: &[WriteOp],
392    ) -> Result<CommitOutcome> {
393        match &mut self.inner {
394            TransactionInner::Turso(tx) => tx.apply_writes_and_commit(writes).await,
395            TransactionInner::Memory(tx) => tx.apply_writes_and_commit(writes).await,
396        }
397    }
398
399    #[tracing::instrument(skip_all, fields(reads = keys.len()))]
400    pub(crate) async fn batch_get_with_version(
401        &mut self,
402        keys: &[(String, String)],
403    ) -> Result<Vec<Option<StoredDoc>>> {
404        match &mut self.inner {
405            TransactionInner::Turso(tx) => tx.batch_get_with_version(keys).await,
406            TransactionInner::Memory(tx) => tx.batch_get_with_version(keys).await,
407        }
408    }
409}
410
411pub enum DbOp {
412    Get {
413        pk: String,
414        sk: String,
415    },
416    Query {
417        pk: String,
418        after_sk: Option<String>,
419        limit: Option<usize>,
420    },
421    Put {
422        pk: String,
423        sk: String,
424        data: Vec<u8>,
425    },
426    Delete {
427        pk: String,
428        sk: String,
429    },
430}
431
432pub enum DbResult {
433    Single(Option<Bytes>),
434    Multiple(Vec<(String, Bytes)>),
435    Done,
436}
437
438pub type DbResultParser<O> = Box<dyn FnOnce(&mut std::vec::IntoIter<DbResult>) -> Result<O> + Send>;
439
440pub struct Prepared<O> {
441    pub ops: Vec<DbOp>,
442    pub parse: DbResultParser<O>,
443}
444
445#[allow(async_fn_in_trait)]
446pub trait DbRequest: Sized {
447    type Output;
448    fn prepare(self) -> Prepared<Self::Output>;
449
450    async fn send_with(self, db: &Database) -> Result<Self::Output> {
451        let prepared = self.prepare();
452        let results = db.execute_ops(prepared.ops).await?;
453        let mut iter = results.into_iter();
454        (prepared.parse)(&mut iter)
455    }
456}
457
458macro_rules! impl_db_request_tuple {
459    ($($T:ident),+) => {
460        #[allow(non_snake_case)]
461        impl<$($T: DbRequest),+> DbRequest for ($($T,)+)
462        where $($T::Output: 'static),+
463        {
464            type Output = ($($T::Output,)+);
465            fn prepare(self) -> Prepared<Self::Output> {
466                let ($($T,)+) = self;
467                $(let $T = $T.prepare();)+
468                let mut ops = Vec::new();
469                $(ops.extend($T.ops);)+
470                Prepared {
471                    ops,
472                    parse: Box::new(move |iter| {
473                        Ok(($(($T.parse)(iter)?,)+))
474                    }),
475                }
476            }
477        }
478    };
479}
480
481impl_db_request_tuple!(A);
482impl_db_request_tuple!(A, B);
483impl_db_request_tuple!(A, B, C);
484impl_db_request_tuple!(A, B, C, D);
485impl_db_request_tuple!(A, B, C, D, E);
486impl_db_request_tuple!(A, B, C, D, E, F);
487impl_db_request_tuple!(A, B, C, D, E, F, G);
488impl_db_request_tuple!(A, B, C, D, E, F, G, H);
489impl_db_request_tuple!(A, B, C, D, E, F, G, H, I);
490impl_db_request_tuple!(A, B, C, D, E, F, G, H, I, J);
491impl_db_request_tuple!(A, B, C, D, E, F, G, H, I, J, K);
492impl_db_request_tuple!(A, B, C, D, E, F, G, H, I, J, K, L);
493
494impl<T: DbRequest> DbRequest for Vec<T>
495where
496    T::Output: 'static,
497{
498    type Output = Vec<T::Output>;
499    fn prepare(self) -> Prepared<Self::Output> {
500        let mut all_ops = Vec::new();
501        let mut parsers: Vec<DbResultParser<T::Output>> = Vec::new();
502        for item in self {
503            let p = item.prepare();
504            all_ops.extend(p.ops);
505            parsers.push(p.parse);
506        }
507        Prepared {
508            ops: all_ops,
509            parse: Box::new(move |iter| parsers.into_iter().map(|p| p(iter)).collect()),
510        }
511    }
512}