Skip to main content

fn0_doc_db/
trx.rs

1use crate::{Database, Transaction};
2use anyhow::{Result, anyhow, bail};
3use serde::{Serialize, de::DeserializeOwned};
4use std::{
5    any::type_name,
6    cell::UnsafeCell,
7    collections::HashMap,
8    marker::PhantomData,
9    ops::{Deref, DerefMut},
10    sync::{
11        Arc, Mutex, Weak,
12        atomic::{AtomicBool, Ordering},
13    },
14};
15use tracing::Instrument;
16
17/// Wrapper around UnsafeCell that asserts Send/Sync when T: Send.
18/// DocHandle holds the only outstanding reference within a trx, so concurrent
19/// access cannot occur in practice; the trx future itself is the unit of work.
20#[repr(transparent)]
21pub(crate) struct SyncUnsafeCell<T>(UnsafeCell<T>);
22
23unsafe impl<T: Send> Send for SyncUnsafeCell<T> {}
24unsafe impl<T: Send> Sync for SyncUnsafeCell<T> {}
25
26impl<T> SyncUnsafeCell<T> {
27    fn new(value: T) -> Self {
28        Self(UnsafeCell::new(value))
29    }
30    fn get(&self) -> *mut T {
31        self.0.get()
32    }
33}
34
35#[derive(Clone, Debug, Eq, Hash, PartialEq)]
36pub struct DocKey {
37    pub pk: String,
38    pub sk: String,
39}
40
41impl DocKey {
42    pub fn new(pk: impl Into<String>, sk: impl Into<String>) -> Self {
43        Self {
44            pk: pk.into(),
45            sk: sk.into(),
46        }
47    }
48}
49
50pub trait Document: Serialize + DeserializeOwned + Send + Sync + 'static {
51    fn key(&self) -> DocKey;
52}
53
54pub trait DocGet {
55    type Doc: Document;
56    fn key(&self) -> DocKey;
57}
58
59#[allow(async_fn_in_trait)]
60pub trait TrxRead: Sized {
61    type Output;
62    fn collect_keys(&self, keys: &mut Vec<DocKey>);
63    async fn finalize(
64        self,
65        tx: &Trx,
66        results: &mut std::vec::IntoIter<Option<crate::turso::StoredDoc>>,
67    ) -> Result<Self::Output>;
68}
69
70impl<R> TrxRead for R
71where
72    R: DocGet,
73{
74    type Output = Option<DocHandle<R::Doc>>;
75
76    fn collect_keys(&self, keys: &mut Vec<DocKey>) {
77        keys.push(self.key());
78    }
79
80    async fn finalize(
81        self,
82        tx: &Trx,
83        results: &mut std::vec::IntoIter<Option<crate::turso::StoredDoc>>,
84    ) -> Result<Self::Output> {
85        let stored = results
86            .next()
87            .ok_or_else(|| anyhow!("trx batch result missing for read"))?;
88        let key = self.key();
89        tx.inner
90            .lock()
91            .unwrap()
92            .register_loaded::<R::Doc>(key, stored)
93    }
94}
95
96macro_rules! impl_trx_read_tuple {
97    ($($T:ident),+) => {
98        #[allow(non_snake_case)]
99        impl<$($T: TrxRead),+> TrxRead for ($($T,)+) {
100            type Output = ($($T::Output,)+);
101
102            fn collect_keys(&self, keys: &mut Vec<DocKey>) {
103                let ($($T,)+) = self;
104                $($T.collect_keys(keys);)+
105            }
106
107            async fn finalize(
108                self,
109                tx: &Trx,
110                results: &mut std::vec::IntoIter<Option<crate::turso::StoredDoc>>,
111            ) -> Result<Self::Output> {
112                let ($($T,)+) = self;
113                Ok(($($T.finalize(tx, results).await?,)+))
114            }
115        }
116    };
117}
118
119impl_trx_read_tuple!(A);
120impl_trx_read_tuple!(A, B);
121impl_trx_read_tuple!(A, B, C);
122impl_trx_read_tuple!(A, B, C, D);
123impl_trx_read_tuple!(A, B, C, D, E);
124impl_trx_read_tuple!(A, B, C, D, E, F);
125impl_trx_read_tuple!(A, B, C, D, E, F, G);
126impl_trx_read_tuple!(A, B, C, D, E, F, G, H);
127impl_trx_read_tuple!(A, B, C, D, E, F, G, H, I);
128impl_trx_read_tuple!(A, B, C, D, E, F, G, H, I, J);
129impl_trx_read_tuple!(A, B, C, D, E, F, G, H, I, J, K);
130impl_trx_read_tuple!(A, B, C, D, E, F, G, H, I, J, K, L);
131
132pub struct DocHandle<T> {
133    data: Arc<SyncUnsafeCell<T>>,
134    dirty: Arc<AtomicBool>,
135    deleted: Arc<AtomicBool>,
136    _alive: Arc<()>,
137    _marker: PhantomData<Arc<T>>,
138}
139
140impl<T> DocHandle<T> {
141    pub fn delete(&self) {
142        self.deleted.store(true, Ordering::Release);
143    }
144}
145
146impl<T> Deref for DocHandle<T> {
147    type Target = T;
148
149    fn deref(&self) -> &Self::Target {
150        unsafe { &*self.data.get() }
151    }
152}
153
154impl<T> DerefMut for DocHandle<T> {
155    fn deref_mut(&mut self) -> &mut Self::Target {
156        self.dirty.store(true, Ordering::Release);
157        unsafe { &mut *self.data.get() }
158    }
159}
160
161pub struct Trx {
162    inner: Arc<Mutex<TrxState>>,
163}
164
165impl Trx {
166    #[tracing::instrument(skip_all)]
167    pub async fn get<R>(&self, request: R) -> Result<R::Output>
168    where
169        R: TrxRead,
170    {
171        let mut keys = Vec::new();
172        request.collect_keys(&mut keys);
173
174        {
175            let state = self.inner.lock().unwrap();
176            for key in &keys {
177                if state.index.contains_key(key) {
178                    bail!("duplicate trx key access: {}/{}", key.pk, key.sk);
179                }
180            }
181        }
182
183        let key_pairs: Vec<(String, String)> =
184            keys.iter().map(|k| (k.pk.clone(), k.sk.clone())).collect();
185        let stored = self.batch_load(&key_pairs).await?;
186
187        let mut iter = stored.into_iter();
188        request.finalize(self, &mut iter).await
189    }
190
191    pub fn create<T>(&self, doc: T) -> Result<DocHandle<T>>
192    where
193        T: Document,
194    {
195        self.inner.lock().unwrap().create(doc)
196    }
197
198    pub fn commit<Out, Cancel>(self, out: Out) -> Result<TrxControl<Out, Cancel>> {
199        Ok(TrxControl {
200            inner: TrxControlInner::Commit(out),
201        })
202    }
203
204    pub fn cancel<Out, Cancel>(self, reason: Cancel) -> Result<TrxControl<Out, Cancel>> {
205        Ok(TrxControl {
206            inner: TrxControlInner::Cancel(reason),
207        })
208    }
209
210    async fn batch_load(
211        &self,
212        keys: &[(String, String)],
213    ) -> Result<Vec<Option<crate::turso::StoredDoc>>> {
214        if keys.is_empty() {
215            return Ok(vec![]);
216        }
217        let mut tx_opt = self.inner.lock().unwrap().tx.take();
218        let result = match &mut tx_opt {
219            Some(tx) => tx.batch_get_with_version(keys).await,
220            None => {
221                let db = self.inner.lock().unwrap().db.clone();
222                match db.begin_immediate_with_reads(keys).await {
223                    Ok((tx, docs)) => {
224                        tx_opt = Some(tx);
225                        Ok(docs)
226                    }
227                    Err(e) => Err(e),
228                }
229            }
230        };
231        self.inner.lock().unwrap().tx = tx_opt;
232        result
233    }
234}
235
236pub struct TrxControl<Out, Cancel> {
237    inner: TrxControlInner<Out, Cancel>,
238}
239
240enum TrxControlInner<Out, Cancel> {
241    Commit(Out),
242    Cancel(Cancel),
243}
244
245#[derive(Debug)]
246pub enum TrxResult<Out, Cancel, Err> {
247    Committed(Out),
248    Cancelled(Cancel),
249    Conflict(ConflictDetails),
250    Err(Err),
251}
252
253#[derive(Clone, Debug, Default)]
254pub struct ConflictDetails {
255    pub keys: Vec<ConflictKey>,
256}
257
258#[derive(Clone, Debug)]
259pub struct ConflictKey {
260    pub key: DocKey,
261    pub expected_version: Option<i64>,
262    pub actual_version: Option<i64>,
263}
264
265const MAX_ATTEMPTS: u32 = 5;
266const BACKOFF_BASE_MS: u64 = 50;
267const BACKOFF_CAP_MS: u64 = 1000;
268
269pub(crate) async fn run<F, Fut, Out, Cancel, E>(db: Database, mut f: F) -> TrxResult<Out, Cancel, E>
270where
271    F: FnMut(Trx) -> Fut,
272    Fut: std::future::Future<Output = Result<TrxControl<Out, Cancel>, E>>,
273    E: From<anyhow::Error>,
274{
275    let mut attempt: u32 = 0;
276    loop {
277        let attempt_span = tracing::info_span!("trx_attempt", attempt = attempt);
278        let result = run_attempt(&db, &mut f).instrument(attempt_span).await;
279        match result {
280            AttemptOutcome::Done(r) => return r,
281            AttemptOutcome::Conflict(details) => {
282                if attempt + 1 >= MAX_ATTEMPTS {
283                    return TrxResult::Conflict(details);
284                }
285                let backoff_span = tracing::info_span!("trx_backoff", attempt = attempt);
286                async {
287                    let backoff = conflict_backoff(attempt).await;
288                    crate::runtime::sleep(backoff).await;
289                }
290                .instrument(backoff_span)
291                .await;
292                attempt += 1;
293            }
294        }
295    }
296}
297
298enum AttemptOutcome<Out, Cancel, E> {
299    Done(TrxResult<Out, Cancel, E>),
300    Conflict(ConflictDetails),
301}
302
303async fn run_attempt<F, Fut, Out, Cancel, E>(
304    db: &Database,
305    f: &mut F,
306) -> AttemptOutcome<Out, Cancel, E>
307where
308    F: FnMut(Trx) -> Fut,
309    Fut: std::future::Future<Output = Result<TrxControl<Out, Cancel>, E>>,
310    E: From<anyhow::Error>,
311{
312    let state = Arc::new(Mutex::new(TrxState::new(db.clone())));
313    let tx = Trx {
314        inner: state.clone(),
315    };
316
317    let user_span = tracing::info_span!("trx_user_closure");
318    let control = match f(tx).instrument(user_span).await {
319        Ok(control) => control,
320        Err(err) => return AttemptOutcome::Done(TrxResult::Err(err)),
321    };
322
323    match control.inner {
324        TrxControlInner::Commit(out) => {
325            let (commit_db, tx, entries_result) = take_entries_and_tx(state);
326            let entries = match entries_result {
327                Ok(e) => e,
328                Err(err) => {
329                    if let Some(tx) = tx {
330                        let _ = tx.rollback().await;
331                    }
332                    return AttemptOutcome::Done(TrxResult::Err(E::from(err)));
333                }
334            };
335
336            match commit_entries(commit_db, tx, entries).await {
337                Ok(()) => AttemptOutcome::Done(TrxResult::Committed(out)),
338                Err(CommitFailure::Conflict(details)) => AttemptOutcome::Conflict(details),
339                Err(CommitFailure::Err(err)) => AttemptOutcome::Done(TrxResult::Err(E::from(err))),
340            }
341        }
342        TrxControlInner::Cancel(reason) => {
343            let tx_to_rollback = state.lock().unwrap().tx.take();
344            if let Some(tx) = tx_to_rollback {
345                let _ = tx.rollback().await;
346            }
347            AttemptOutcome::Done(TrxResult::Cancelled(reason))
348        }
349    }
350}
351
352async fn conflict_backoff(attempt: u32) -> std::time::Duration {
353    let ceiling = BACKOFF_BASE_MS
354        .checked_shl(attempt)
355        .unwrap_or(BACKOFF_CAP_MS)
356        .min(BACKOFF_CAP_MS);
357    let mut buf = [0u8; 8];
358    crate::runtime::random_bytes(&mut buf).await;
359    let raw = u64::from_le_bytes(buf);
360    let delay_ms = raw % (ceiling + 1);
361    std::time::Duration::from_millis(delay_ms)
362}
363
364struct TrxState {
365    db: Database,
366    entries: Vec<TrackedEntry>,
367    index: HashMap<DocKey, usize>,
368    tx: Option<Transaction>,
369}
370
371impl TrxState {
372    fn new(db: Database) -> Self {
373        Self {
374            db,
375            entries: Vec::new(),
376            index: HashMap::new(),
377            tx: None,
378        }
379    }
380
381    fn register_loaded<T>(
382        &mut self,
383        key: DocKey,
384        stored: Option<crate::turso::StoredDoc>,
385    ) -> Result<Option<DocHandle<T>>>
386    where
387        T: Document,
388    {
389        if self.index.contains_key(&key) {
390            bail!("duplicate trx key access: {}/{}", key.pk, key.sk);
391        }
392
393        let idx = self.entries.len();
394        self.index.insert(key.clone(), idx);
395
396        match stored {
397            Some(stored) => {
398                let doc = serde_json::from_slice::<T>(&stored.data).map_err(|err| {
399                    anyhow!(
400                        "failed to deserialize {} at {}/{}: {}",
401                        type_name::<T>(),
402                        key.pk,
403                        key.sk,
404                        err
405                    )
406                })?;
407                let (shared, handle) = new_shared_doc(doc);
408                self.entries.push(TrackedEntry {
409                    key,
410                    expected_version: Some(stored.version),
411                    state: TrackedState::Managed {
412                        shared,
413                        created: false,
414                    },
415                });
416                Ok(Some(handle))
417            }
418            None => {
419                self.entries.push(TrackedEntry {
420                    key,
421                    expected_version: None,
422                    state: TrackedState::Missing,
423                });
424                Ok(None)
425            }
426        }
427    }
428
429    fn create<T>(&mut self, doc: T) -> Result<DocHandle<T>>
430    where
431        T: Document,
432    {
433        let key = doc.key();
434        let (shared, handle) = new_shared_doc(doc);
435
436        match self.index.get(&key).copied() {
437            None => {
438                let idx = self.entries.len();
439                self.index.insert(key.clone(), idx);
440                self.entries.push(TrackedEntry {
441                    key,
442                    expected_version: None,
443                    state: TrackedState::Managed {
444                        shared,
445                        created: true,
446                    },
447                });
448                Ok(handle)
449            }
450            Some(idx) => match self.entries.get_mut(idx) {
451                Some(TrackedEntry {
452                    expected_version: None,
453                    state: TrackedState::Missing,
454                    ..
455                }) => {
456                    self.entries[idx].state = TrackedState::Managed {
457                        shared,
458                        created: true,
459                    };
460                    Ok(handle)
461                }
462                _ => bail!("duplicate trx key access: {}/{}", key.pk, key.sk),
463            },
464        }
465    }
466
467    fn take_entries_and_tx(
468        &mut self,
469    ) -> (Database, Option<Transaction>, Result<Vec<TrackedEntry>>) {
470        let tx = self.tx.take();
471        let db = self.db.clone();
472        for entry in &self.entries {
473            if let TrackedState::Managed { shared, .. } = &entry.state
474                && shared.handle_alive.upgrade().is_some()
475            {
476                let err = anyhow!(
477                    "live doc handle escaped trx for key {}/{}; commit outputs must not contain DocHandle values",
478                    entry.key.pk,
479                    entry.key.sk
480                );
481                self.entries.clear();
482                return (db, tx, Err(err));
483            }
484        }
485        (db, tx, Ok(std::mem::take(&mut self.entries)))
486    }
487}
488
489struct TrackedEntry {
490    key: DocKey,
491    expected_version: Option<i64>,
492    state: TrackedState,
493}
494
495impl TrackedEntry {
496    fn write(&self) -> Result<PendingWrite> {
497        match &self.state {
498            TrackedState::Missing => Ok(PendingWrite::None),
499            TrackedState::Managed { shared, created } => {
500                if shared.deleted.load(Ordering::Acquire) {
501                    if *created {
502                        return Ok(PendingWrite::None);
503                    }
504                    let expected_version = self.expected_version.ok_or_else(|| {
505                        anyhow!("existing tracked doc missing expected version for delete")
506                    })?;
507                    return Ok(PendingWrite::Delete(expected_version));
508                }
509
510                if *created {
511                    return Ok(PendingWrite::Insert((shared.serialize)()?));
512                }
513
514                if shared.dirty.load(Ordering::Acquire) {
515                    let expected_version = self.expected_version.ok_or_else(|| {
516                        anyhow!("existing tracked doc missing expected version for update")
517                    })?;
518                    return Ok(PendingWrite::Update {
519                        expected_version,
520                        data: (shared.serialize)()?,
521                    });
522                }
523
524                Ok(PendingWrite::None)
525            }
526        }
527    }
528}
529
530enum TrackedState {
531    Missing,
532    Managed { shared: SharedDoc, created: bool },
533}
534
535struct SharedDoc {
536    dirty: Arc<AtomicBool>,
537    deleted: Arc<AtomicBool>,
538    handle_alive: Weak<()>,
539    serialize: Box<dyn Fn() -> Result<Vec<u8>> + Send + Sync>,
540}
541
542fn new_shared_doc<T>(doc: T) -> (SharedDoc, DocHandle<T>)
543where
544    T: Document + Send + Sync,
545{
546    let data = Arc::new(SyncUnsafeCell::new(doc));
547    let dirty = Arc::new(AtomicBool::new(false));
548    let deleted = Arc::new(AtomicBool::new(false));
549    let alive = Arc::new(());
550
551    let serialize_data = data.clone();
552    let shared = SharedDoc {
553        dirty: dirty.clone(),
554        deleted: deleted.clone(),
555        handle_alive: Arc::downgrade(&alive),
556        serialize: Box::new(move || {
557            let doc_ref = unsafe { &*serialize_data.get() };
558            serde_json::to_vec(doc_ref).map_err(Into::into)
559        }),
560    };
561
562    let handle = DocHandle {
563        data,
564        dirty,
565        deleted,
566        _alive: alive,
567        _marker: PhantomData,
568    };
569
570    (shared, handle)
571}
572
573fn take_entries_and_tx(
574    state: Arc<Mutex<TrxState>>,
575) -> (Database, Option<Transaction>, Result<Vec<TrackedEntry>>) {
576    state.lock().unwrap().take_entries_and_tx()
577}
578
579enum PendingWrite {
580    None,
581    Insert(Vec<u8>),
582    Update {
583        expected_version: i64,
584        data: Vec<u8>,
585    },
586    Delete(i64),
587}
588
589enum CommitFailure {
590    Conflict(ConflictDetails),
591    Err(anyhow::Error),
592}
593
594#[tracing::instrument(skip_all, fields(entries = entries.len(), reused_tx = tx.is_some()))]
595async fn commit_entries(
596    db: Database,
597    tx: Option<Transaction>,
598    entries: Vec<TrackedEntry>,
599) -> std::result::Result<(), CommitFailure> {
600    let mut writes: Vec<crate::WriteOp> = Vec::new();
601    for entry in &entries {
602        match entry.write().map_err(CommitFailure::Err)? {
603            PendingWrite::None => {}
604            PendingWrite::Insert(data) => writes.push(crate::WriteOp::Insert {
605                pk: entry.key.pk.clone(),
606                sk: entry.key.sk.clone(),
607                data,
608            }),
609            PendingWrite::Update {
610                expected_version,
611                data,
612            } => writes.push(crate::WriteOp::Update {
613                pk: entry.key.pk.clone(),
614                sk: entry.key.sk.clone(),
615                expected_version,
616                data,
617            }),
618            PendingWrite::Delete(expected_version) => writes.push(crate::WriteOp::Delete {
619                pk: entry.key.pk.clone(),
620                sk: entry.key.sk.clone(),
621                expected_version,
622            }),
623        }
624    }
625
626    if writes.is_empty() && tx.is_none() {
627        return Ok(());
628    }
629
630    let mut tx = match tx {
631        Some(t) => t,
632        None => {
633            let begin_span = tracing::info_span!("commit_begin_tx");
634            async { db.transaction().await.map_err(CommitFailure::Err) }
635                .instrument(begin_span)
636                .await?
637        }
638    };
639
640    let outcome = tx
641        .apply_writes_and_commit(&writes)
642        .await
643        .map_err(CommitFailure::Err)?;
644
645    let mut conflicts = Vec::new();
646
647    if let Some(info) = outcome.conflict
648        && let Some(op) = writes.get(info.step_index)
649    {
650        let (pk, sk, expected) = write_key_and_expected(op);
651        conflicts.push(ConflictKey {
652            key: DocKey { pk, sk },
653            expected_version: expected,
654            actual_version: None,
655        });
656    }
657
658    for (i, count) in outcome.affected_counts.iter().enumerate() {
659        if *count == 1 {
660            continue;
661        }
662        let Some(op) = writes.get(i) else { continue };
663        let (pk, sk, expected) = write_key_and_expected(op);
664        let key = DocKey { pk, sk };
665        if conflicts.iter().any(|c| c.key == key) {
666            continue;
667        }
668        conflicts.push(ConflictKey {
669            key,
670            expected_version: expected,
671            actual_version: None,
672        });
673    }
674
675    if !conflicts.is_empty() {
676        let key_pairs: Vec<(String, String)> = conflicts
677            .iter()
678            .map(|c| (c.key.pk.clone(), c.key.sk.clone()))
679            .collect();
680        if let Ok(stored) = db.batch_get_with_version(&key_pairs).await {
681            for (c, slot) in conflicts.iter_mut().zip(stored.into_iter()) {
682                c.actual_version = slot.map(|d| d.version);
683            }
684        }
685        return Err(CommitFailure::Conflict(ConflictDetails { keys: conflicts }));
686    }
687
688    Ok(())
689}
690
691fn write_key_and_expected(op: &crate::WriteOp) -> (String, String, Option<i64>) {
692    match op {
693        crate::WriteOp::Insert { pk, sk, .. } => (pk.clone(), sk.clone(), None),
694        crate::WriteOp::Update {
695            pk,
696            sk,
697            expected_version,
698            ..
699        }
700        | crate::WriteOp::Delete {
701            pk,
702            sk,
703            expected_version,
704        } => (pk.clone(), sk.clone(), Some(*expected_version)),
705    }
706}
707
708#[cfg(test)]
709mod tests {
710    use super::*;
711    use crate::turso_with_config;
712
713    #[derive(Clone, serde::Serialize, serde::Deserialize)]
714    struct TestDoc {
715        id: String,
716        value: i32,
717    }
718
719    impl Document for TestDoc {
720        fn key(&self) -> DocKey {
721            DocKey::new("TestDoc", format!("id={}", self.id))
722        }
723    }
724
725    struct TestDocGet {
726        id: String,
727    }
728
729    impl DocGet for TestDocGet {
730        type Doc = TestDoc;
731
732        fn key(&self) -> DocKey {
733            DocKey::new("TestDoc", format!("id={}", self.id))
734        }
735    }
736
737    fn test_state() -> TrxState {
738        TrxState::new(turso_with_config(
739            "http://127.0.0.1:0".to_string(),
740            String::new(),
741        ))
742    }
743
744    #[test]
745    fn create_produces_insert_even_without_mutation() {
746        let mut state = test_state();
747        let _handle = state
748            .create(TestDoc {
749                id: "a".into(),
750                value: 1,
751            })
752            .expect("create should succeed");
753
754        let write = state.entries[0].write().expect("write plan");
755        match write {
756            PendingWrite::Insert(data) => {
757                let doc: TestDoc = serde_json::from_slice(&data).expect("deserialize insert");
758                assert_eq!(doc.id, "a");
759                assert_eq!(doc.value, 1);
760            }
761            _ => panic!("expected insert"),
762        }
763    }
764
765    #[test]
766    fn loaded_doc_marks_dirty_on_deref_mut() {
767        let mut state = test_state();
768        let doc = TestDoc {
769            id: "a".into(),
770            value: 1,
771        };
772        let key = TestDocGet { id: "a".into() }.key();
773        let mut handle = state
774            .register_loaded::<TestDoc>(
775                key,
776                Some(crate::turso::StoredDoc {
777                    data: serde_json::to_vec(&doc).expect("serialize").into(),
778                    version: 7,
779                }),
780            )
781            .expect("load should succeed")
782            .expect("doc should exist");
783
784        handle.value = 5;
785
786        match state.entries[0].write().expect("write plan") {
787            PendingWrite::Update {
788                expected_version,
789                data,
790            } => {
791                assert_eq!(expected_version, 7);
792                let doc: TestDoc = serde_json::from_slice(&data).expect("deserialize update");
793                assert_eq!(doc.value, 5);
794            }
795            _ => panic!("expected update"),
796        }
797    }
798
799    #[test]
800    fn loaded_doc_delete_produces_delete_write() {
801        let mut state = test_state();
802        let doc = TestDoc {
803            id: "a".into(),
804            value: 1,
805        };
806        let key = TestDocGet { id: "a".into() }.key();
807        let handle = state
808            .register_loaded::<TestDoc>(
809                key,
810                Some(crate::turso::StoredDoc {
811                    data: serde_json::to_vec(&doc).expect("serialize").into(),
812                    version: 7,
813                }),
814            )
815            .expect("load should succeed")
816            .expect("doc should exist");
817
818        handle.delete();
819
820        match state.entries[0].write().expect("write plan") {
821            PendingWrite::Delete(expected_version) => assert_eq!(expected_version, 7),
822            _ => panic!("expected delete"),
823        }
824    }
825
826    #[test]
827    fn missing_read_can_be_promoted_to_create() {
828        let mut state = test_state();
829        let key = TestDocGet { id: "a".into() }.key();
830        let loaded = state
831            .register_loaded::<TestDoc>(key, None)
832            .expect("register missing should succeed");
833        assert!(loaded.is_none());
834
835        let handle = state
836            .create(TestDoc {
837                id: "a".into(),
838                value: 3,
839            })
840            .expect("create after missing get should succeed");
841        assert_eq!(handle.value, 3);
842
843        assert!(matches!(
844            state.entries[0].write().expect("write plan"),
845            PendingWrite::Insert(_)
846        ));
847    }
848
849    #[test]
850    fn duplicate_key_access_is_rejected() {
851        let mut state = test_state();
852        let first = state.register_loaded::<TestDoc>(TestDocGet { id: "a".into() }.key(), None);
853        assert!(first.is_ok());
854
855        let second = state.register_loaded::<TestDoc>(TestDocGet { id: "a".into() }.key(), None);
856        assert!(second.is_err());
857    }
858
859    #[test]
860    fn live_handle_cannot_escape_commit_boundary() {
861        let mut state = test_state();
862        let _handle = state
863            .create(TestDoc {
864                id: "a".into(),
865                value: 1,
866            })
867            .expect("create should succeed");
868
869        let (_, _, result) = state.take_entries_and_tx();
870        match result {
871            Ok(_) => panic!("live handle should fail"),
872            Err(err) => assert!(err.to_string().contains("live doc handle escaped trx")),
873        }
874    }
875}