Skip to main content

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