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 ConflictInfo {
66    pub step_index: usize,
67    pub message: String,
68}
69
70pub fn turso() -> Database {
71    let url = std::env::var("TURSO_URL").expect("TURSO_URL must be set");
72    let auth_token = std::env::var("TURSO_AUTH_TOKEN").expect("TURSO_AUTH_TOKEN must be set");
73    turso_with_config(url, auth_token)
74}
75
76pub fn turso_with_config(url: String, auth_token: String) -> Database {
77    Database {
78        inner: DatabaseInner::Turso(TursoDatabase::new(url, auth_token)),
79        mock_state: mock::MockState::default(),
80    }
81}
82
83pub fn memory() -> Database {
84    Database {
85        inner: DatabaseInner::Memory(MemoryDatabase::new()),
86        mock_state: mock::MockState::default(),
87    }
88}
89
90#[derive(Clone)]
91pub struct Database {
92    inner: DatabaseInner,
93    mock_state: mock::MockState,
94}
95
96impl Database {
97    pub async fn get(&self, pk: &str, sk: &str) -> Result<Option<Bytes>> {
98        if let Some(result) = self.mock_state.try_match(mock::MockOp::Get, pk, sk) {
99            return match result {
100                mock::MockResult::OkGet(data) => Ok(data.map(Bytes::from)),
101                mock::MockResult::Err(msg) => Err(anyhow::anyhow!("{}", msg)),
102                _ => unreachable!(),
103            };
104        }
105        match &self.inner {
106            DatabaseInner::Turso(db) => db.get(pk, sk).await,
107            DatabaseInner::Memory(db) => db.get(pk, sk).await,
108        }
109    }
110
111    pub async fn put(&self, pk: &str, sk: &str, data: &[u8]) -> Result<()> {
112        if let Some(result) = self.mock_state.try_match(mock::MockOp::Put, pk, sk) {
113            return match result {
114                mock::MockResult::OkVoid => Ok(()),
115                mock::MockResult::Err(msg) => Err(anyhow::anyhow!("{}", msg)),
116                _ => unreachable!(),
117            };
118        }
119        match &self.inner {
120            DatabaseInner::Turso(db) => db.put(pk, sk, data).await,
121            DatabaseInner::Memory(db) => db.put(pk, sk, data).await,
122        }
123    }
124
125    pub async fn delete(&self, pk: &str, sk: &str) -> Result<()> {
126        if let Some(result) = self.mock_state.try_match(mock::MockOp::Delete, pk, sk) {
127            return match result {
128                mock::MockResult::OkVoid => Ok(()),
129                mock::MockResult::Err(msg) => Err(anyhow::anyhow!("{}", msg)),
130                _ => unreachable!(),
131            };
132        }
133        match &self.inner {
134            DatabaseInner::Turso(db) => db.delete(pk, sk).await,
135            DatabaseInner::Memory(db) => db.delete(pk, sk).await,
136        }
137    }
138
139    // --- Mock API ---
140
141    pub fn mock_get(&self, pk: &str, sk: &str) -> mock::MockGetBuilder<'_> {
142        mock::MockGetBuilder::new(self, pk.to_string(), sk.to_string())
143    }
144
145    pub fn mock_put(&self, pk: &str, sk: &str) -> mock::MockPutBuilder<'_> {
146        mock::MockPutBuilder::new(self, pk.to_string(), sk.to_string())
147    }
148
149    pub fn mock_delete(&self, pk: &str, sk: &str) -> mock::MockDeleteBuilder<'_> {
150        mock::MockDeleteBuilder::new(self, pk.to_string(), sk.to_string())
151    }
152
153    pub fn clear_mocks(&self) {
154        self.mock_state.clear();
155    }
156
157    pub(crate) fn add_mock_rule(&self, rule: mock::MockRule) {
158        self.mock_state.push(rule);
159    }
160
161    #[tracing::instrument(skip_all, fields(pk = %pk.as_ref(), limit = limit))]
162    pub async fn query<S1: AsRef<str>, S2: AsRef<str>>(
163        &self,
164        pk: S1,
165        after_sk: Option<S2>,
166        limit: usize,
167    ) -> Result<Vec<(String, Bytes)>> {
168        match &self.inner {
169            DatabaseInner::Turso(db) => db.query(pk, after_sk, limit).await,
170            DatabaseInner::Memory(db) => db.query(pk, after_sk, limit).await,
171        }
172    }
173
174    #[tracing::instrument(skip_all, fields(limit = limit))]
175    pub async fn scan(
176        &self,
177        after: Option<(&str, &str)>,
178        limit: usize,
179    ) -> Result<Vec<(String, String, Bytes)>> {
180        match &self.inner {
181            DatabaseInner::Turso(db) => db.scan(after, limit).await,
182            DatabaseInner::Memory(db) => db.scan(after, limit).await,
183        }
184    }
185
186    #[tracing::instrument(skip_all, fields(ops = ops.len()))]
187    pub async fn batch(&self, ops: &[BatchOp<'_>]) -> Result<()> {
188        match &self.inner {
189            DatabaseInner::Turso(db) => db.batch(ops).await,
190            DatabaseInner::Memory(db) => db.batch(ops).await,
191        }
192    }
193
194    #[tracing::instrument(skip_all)]
195    pub async fn transaction(&self) -> Result<Transaction> {
196        match &self.inner {
197            DatabaseInner::Turso(db) => Ok(Transaction {
198                inner: TransactionInner::Turso(db.transaction().await?),
199            }),
200            DatabaseInner::Memory(db) => Ok(Transaction {
201                inner: TransactionInner::Memory(db.transaction().await?),
202            }),
203        }
204    }
205
206    #[tracing::instrument(skip_all)]
207    pub async fn trx<F, Fut, Out, Cancel, E>(&self, f: F) -> TrxResult<Out, Cancel, E>
208    where
209        F: FnMut(Trx) -> Fut,
210        Fut: Future<Output = Result<TrxControl<Out, Cancel>, E>>,
211        E: From<anyhow::Error>,
212    {
213        trx::run(self.clone(), f).await
214    }
215
216    #[tracing::instrument(skip_all, fields(sql = %sql))]
217    pub async fn execute_raw(
218        &self,
219        sql: &str,
220        args: Vec<Value>,
221        want_rows: bool,
222    ) -> Result<Vec<Vec<Value>>> {
223        match &self.inner {
224            DatabaseInner::Turso(db) => db.execute_raw(sql, args, want_rows).await,
225            DatabaseInner::Memory(db) => db.execute_raw(sql, args, want_rows).await,
226        }
227    }
228
229    #[tracing::instrument(skip_all, fields(ops = ops.len()))]
230    pub async fn execute_ops(&self, ops: Vec<DbOp>) -> Result<Vec<DbResult>> {
231        match &self.inner {
232            DatabaseInner::Turso(db) => db.execute_ops(ops).await,
233            DatabaseInner::Memory(db) => db.execute_ops(ops).await,
234        }
235    }
236
237    #[tracing::instrument(skip_all, fields(pk = %pk, sk = %sk))]
238    pub(crate) async fn get_with_version(&self, pk: &str, sk: &str) -> Result<Option<StoredDoc>> {
239        match &self.inner {
240            DatabaseInner::Turso(db) => db.get_with_version(pk, sk).await,
241            DatabaseInner::Memory(db) => db.get_with_version(pk, sk).await,
242        }
243    }
244
245    #[tracing::instrument(skip_all, fields(reads = keys.len()))]
246    pub(crate) async fn batch_get_with_version(
247        &self,
248        keys: &[(String, String)],
249    ) -> Result<Vec<Option<StoredDoc>>> {
250        if keys.is_empty() {
251            return Ok(vec![]);
252        }
253        let mut out = Vec::with_capacity(keys.len());
254        for (pk, sk) in keys {
255            out.push(self.get_with_version(pk, sk).await?);
256        }
257        Ok(out)
258    }
259
260    #[tracing::instrument(skip_all, fields(reads = keys.len()))]
261    pub(crate) async fn begin_immediate_with_reads(
262        &self,
263        keys: &[(String, String)],
264    ) -> Result<(Transaction, Vec<Option<StoredDoc>>)> {
265        match &self.inner {
266            DatabaseInner::Turso(db) => {
267                let (tx, docs) = db.begin_immediate_with_reads(keys).await?;
268                Ok((
269                    Transaction {
270                        inner: TransactionInner::Turso(tx),
271                    },
272                    docs,
273                ))
274            }
275            DatabaseInner::Memory(db) => {
276                let (tx, docs) = db.begin_immediate_with_reads(keys).await?;
277                Ok((
278                    Transaction {
279                        inner: TransactionInner::Memory(tx),
280                    },
281                    docs,
282                ))
283            }
284        }
285    }
286}
287
288#[derive(Clone)]
289enum DatabaseInner {
290    Turso(TursoDatabase),
291    Memory(MemoryDatabase),
292}
293
294pub struct Transaction {
295    inner: TransactionInner,
296}
297
298enum TransactionInner {
299    Turso(TursoTransaction),
300    Memory(MemoryTransaction),
301}
302
303impl Transaction {
304    #[tracing::instrument(skip_all, fields(pk = %pk, sk = %sk))]
305    pub async fn get(&mut self, pk: &str, sk: &str) -> Result<Option<Bytes>> {
306        match &mut self.inner {
307            TransactionInner::Turso(tx) => tx.get(pk, sk).await,
308            TransactionInner::Memory(tx) => tx.get(pk, sk).await,
309        }
310    }
311
312    #[tracing::instrument(skip_all, fields(pk = %pk, sk = %sk, bytes = data.len()))]
313    pub async fn put(&mut self, pk: &str, sk: &str, data: &[u8]) -> Result<()> {
314        match &mut self.inner {
315            TransactionInner::Turso(tx) => tx.put(pk, sk, data).await,
316            TransactionInner::Memory(tx) => tx.put(pk, sk, data).await,
317        }
318    }
319
320    #[tracing::instrument(skip_all, fields(pk = %pk, sk = %sk))]
321    pub async fn delete(&mut self, pk: &str, sk: &str) -> Result<()> {
322        match &mut self.inner {
323            TransactionInner::Turso(tx) => tx.delete(pk, sk).await,
324            TransactionInner::Memory(tx) => tx.delete(pk, sk).await,
325        }
326    }
327
328    #[tracing::instrument(skip_all)]
329    pub async fn commit(self) -> Result<()> {
330        match self.inner {
331            TransactionInner::Turso(tx) => tx.commit().await,
332            TransactionInner::Memory(tx) => tx.commit().await,
333        }
334    }
335
336    #[tracing::instrument(skip_all)]
337    pub async fn rollback(self) -> Result<()> {
338        match self.inner {
339            TransactionInner::Turso(tx) => tx.rollback().await,
340            TransactionInner::Memory(tx) => tx.rollback().await,
341        }
342    }
343
344    #[tracing::instrument(skip_all, fields(writes = writes.len()))]
345    pub(crate) async fn apply_writes_and_commit(
346        &mut self,
347        writes: &[WriteOp],
348    ) -> Result<CommitOutcome> {
349        match &mut self.inner {
350            TransactionInner::Turso(tx) => tx.apply_writes_and_commit(writes).await,
351            TransactionInner::Memory(tx) => tx.apply_writes_and_commit(writes).await,
352        }
353    }
354
355    #[tracing::instrument(skip_all, fields(reads = keys.len()))]
356    pub(crate) async fn batch_get_with_version(
357        &mut self,
358        keys: &[(String, String)],
359    ) -> Result<Vec<Option<StoredDoc>>> {
360        match &mut self.inner {
361            TransactionInner::Turso(tx) => tx.batch_get_with_version(keys).await,
362            TransactionInner::Memory(tx) => tx.batch_get_with_version(keys).await,
363        }
364    }
365}
366
367pub enum DbOp {
368    Get {
369        pk: String,
370        sk: String,
371    },
372    Query {
373        pk: String,
374        after_sk: Option<String>,
375        limit: Option<usize>,
376    },
377    Put {
378        pk: String,
379        sk: String,
380        data: Vec<u8>,
381    },
382    Delete {
383        pk: String,
384        sk: String,
385    },
386}
387
388pub enum DbResult {
389    Single(Option<Bytes>),
390    Multiple(Vec<(String, Bytes)>),
391    Done,
392}
393
394pub type DbResultParser<O> = Box<dyn FnOnce(&mut std::vec::IntoIter<DbResult>) -> Result<O> + Send>;
395
396pub struct Prepared<O> {
397    pub ops: Vec<DbOp>,
398    pub parse: DbResultParser<O>,
399}
400
401#[allow(async_fn_in_trait)]
402pub trait DbRequest: Sized {
403    type Output;
404    fn prepare(self) -> Prepared<Self::Output>;
405
406    async fn send_with(self, db: &Database) -> Result<Self::Output> {
407        let prepared = self.prepare();
408        let results = db.execute_ops(prepared.ops).await?;
409        let mut iter = results.into_iter();
410        (prepared.parse)(&mut iter)
411    }
412}
413
414macro_rules! impl_db_request_tuple {
415    ($($T:ident),+) => {
416        #[allow(non_snake_case)]
417        impl<$($T: DbRequest),+> DbRequest for ($($T,)+)
418        where $($T::Output: 'static),+
419        {
420            type Output = ($($T::Output,)+);
421            fn prepare(self) -> Prepared<Self::Output> {
422                let ($($T,)+) = self;
423                $(let $T = $T.prepare();)+
424                let mut ops = Vec::new();
425                $(ops.extend($T.ops);)+
426                Prepared {
427                    ops,
428                    parse: Box::new(move |iter| {
429                        Ok(($(($T.parse)(iter)?,)+))
430                    }),
431                }
432            }
433        }
434    };
435}
436
437impl_db_request_tuple!(A);
438impl_db_request_tuple!(A, B);
439impl_db_request_tuple!(A, B, C);
440impl_db_request_tuple!(A, B, C, D);
441impl_db_request_tuple!(A, B, C, D, E);
442impl_db_request_tuple!(A, B, C, D, E, F);
443impl_db_request_tuple!(A, B, C, D, E, F, G);
444impl_db_request_tuple!(A, B, C, D, E, F, G, H);
445impl_db_request_tuple!(A, B, C, D, E, F, G, H, I);
446impl_db_request_tuple!(A, B, C, D, E, F, G, H, I, J);
447impl_db_request_tuple!(A, B, C, D, E, F, G, H, I, J, K);
448impl_db_request_tuple!(A, B, C, D, E, F, G, H, I, J, K, L);
449
450impl<T: DbRequest> DbRequest for Vec<T>
451where
452    T::Output: 'static,
453{
454    type Output = Vec<T::Output>;
455    fn prepare(self) -> Prepared<Self::Output> {
456        let mut all_ops = Vec::new();
457        let mut parsers: Vec<DbResultParser<T::Output>> = Vec::new();
458        for item in self {
459            let p = item.prepare();
460            all_ops.extend(p.ops);
461            parsers.push(p.parse);
462        }
463        Prepared {
464            ops: all_ops,
465            parse: Box::new(move |iter| parsers.into_iter().map(|p| p(iter)).collect()),
466        }
467    }
468}