Skip to main content

uqa_storage_sqlite/
transaction.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! `SQLite` transactions with automatic rollback on drop.
8
9use crate::{ManagedConnection, SQLiteError};
10use uqa_storage::{TransactionError, TxResult};
11
12/// SQLite-backed transaction. Drops without commit roll back so
13/// panics never leak a half-applied write log.
14pub struct SQLiteTransaction {
15    conn: ManagedConnection,
16    finished: bool,
17}
18
19impl SQLiteTransaction {
20    pub fn begin(conn: ManagedConnection) -> Result<Self, SQLiteError> {
21        conn.begin_transaction()?;
22        Ok(Self {
23            conn,
24            finished: false,
25        })
26    }
27
28    pub fn active(&self) -> bool {
29        !self.finished
30    }
31
32    pub fn commit(&mut self) -> TxResult<()> {
33        if self.finished {
34            return Err(TransactionError::Finished);
35        }
36        let result = self
37            .conn
38            .commit_transaction()
39            .map_err(TransactionError::from);
40        if result.is_ok() || !self.conn.in_transaction() {
41            self.finished = true;
42        }
43        result
44    }
45
46    pub fn rollback(&mut self) -> TxResult<()> {
47        if self.finished {
48            return Err(TransactionError::Finished);
49        }
50        let result = self
51            .conn
52            .rollback_transaction()
53            .map_err(TransactionError::from);
54        if result.is_ok() || !self.conn.in_transaction() {
55            self.finished = true;
56        }
57        result
58    }
59
60    pub fn savepoint(&self, name: &str) -> TxResult<()> {
61        if self.finished {
62            return Err(TransactionError::Finished);
63        }
64        self.conn.savepoint(name)?;
65        Ok(())
66    }
67
68    pub fn release_savepoint(&self, name: &str) -> TxResult<()> {
69        if self.finished {
70            return Err(TransactionError::Finished);
71        }
72        self.conn.release_savepoint(name)?;
73        Ok(())
74    }
75
76    pub fn rollback_to(&self, name: &str) -> TxResult<()> {
77        if self.finished {
78            return Err(TransactionError::Finished);
79        }
80        self.conn.rollback_to_savepoint(name)?;
81        Ok(())
82    }
83}
84
85impl Drop for SQLiteTransaction {
86    fn drop(&mut self) {
87        if !self.finished {
88            self.conn.rollback_transaction_on_drop();
89        }
90    }
91}
92
93#[cfg(test)]
94mod tests {
95    use super::*;
96    #[test]
97    fn sqlite_transaction_commits_writes() {
98        let conn = ManagedConnection::open_in_memory().unwrap();
99        conn.with(|c| {
100            c.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)", [])?;
101            Ok(())
102        })
103        .unwrap();
104        let mut tx = SQLiteTransaction::begin(conn.clone()).unwrap();
105        conn.with(|c| {
106            c.execute("INSERT INTO t (id, v) VALUES (1, 'hi')", [])?;
107            Ok(())
108        })
109        .unwrap();
110        tx.commit().unwrap();
111        let got: i64 = conn
112            .with(|c| Ok(c.query_row("SELECT COUNT(*) FROM t", [], |r| r.get(0))?))
113            .unwrap();
114        assert_eq!(got, 1);
115    }
116
117    #[test]
118    fn sqlite_transaction_rolls_back_on_drop() {
119        let conn = ManagedConnection::open_in_memory().unwrap();
120        conn.with(|c| {
121            c.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, v TEXT)", [])?;
122            Ok(())
123        })
124        .unwrap();
125        {
126            let _tx = SQLiteTransaction::begin(conn.clone()).unwrap();
127            conn.with(|c| {
128                c.execute("INSERT INTO t (id, v) VALUES (1, 'hi')", [])?;
129                Ok(())
130            })
131            .unwrap();
132            // Tx drops without commit -> rollback fires automatically.
133        }
134        let got: i64 = conn
135            .with(|c| Ok(c.query_row("SELECT COUNT(*) FROM t", [], |r| r.get(0))?))
136            .unwrap();
137        assert_eq!(got, 0);
138    }
139}