use std::{cell::Cell, future::Future};
use fsqlite_error::FrankenError;
use fsqlite_types::value::SqliteValue;
use crate::{Connection, Row};
use super::params::ParamValue;
pub struct Transaction<'a> {
conn: &'a Connection,
finalized: Cell<bool>,
}
impl<'a> Transaction<'a> {
async fn new(conn: &'a Connection) -> Result<Self, FrankenError> {
conn.begin_transaction().await?;
Ok(Self {
conn,
finalized: Cell::new(false),
})
}
fn ensure_active(&self) -> Result<(), FrankenError> {
if self.finalized.get() {
return Err(FrankenError::NoActiveTransaction);
}
if !self.conn.in_transaction() {
self.finalized.set(true);
return Err(FrankenError::NoActiveTransaction);
}
Ok(())
}
fn observe_transaction_state<T>(
&self,
result: Result<T, FrankenError>,
) -> Result<T, FrankenError> {
if !self.conn.in_transaction() {
self.finalized.set(true);
}
result
}
pub async fn commit(&mut self) -> Result<(), FrankenError> {
self.ensure_active()?;
let result = self.conn.commit_transaction().await;
self.observe_transaction_state(result)
}
pub async fn rollback(&mut self) -> Result<(), FrankenError> {
self.ensure_active()?;
let result = self.conn.rollback_transaction().await;
self.observe_transaction_state(result)
}
pub async fn execute(&self, sql: &str) -> Result<usize, FrankenError> {
self.ensure_active()?;
let result = self.conn.execute(sql).await;
self.observe_transaction_state(result)
}
pub async fn execute_with_params(
&self,
sql: &str,
params: &[SqliteValue],
) -> Result<usize, FrankenError> {
self.ensure_active()?;
let result = self.conn.execute_with_params(sql, params).await;
self.observe_transaction_state(result)
}
pub async fn execute_with_params_skip_statement_savepoint(
&self,
sql: &str,
params: &[SqliteValue],
) -> Result<usize, FrankenError> {
self.ensure_active()?;
let result = self
.conn
.execute_with_params_skip_statement_savepoint_in_explicit_txn(sql, params)
.await;
self.observe_transaction_state(result)
}
pub async fn execute_compat(
&self,
sql: &str,
params: &[ParamValue],
) -> Result<usize, FrankenError> {
self.ensure_active()?;
let values: Vec<SqliteValue> = params.iter().map(|p| p.0.clone()).collect();
let result = self.conn.execute_with_params(sql, &values).await;
self.observe_transaction_state(result)
}
pub async fn query(&self, sql: &str) -> Result<Vec<Row>, FrankenError> {
self.ensure_active()?;
let result = self.conn.query(sql).await;
self.observe_transaction_state(result)
}
pub async fn query_with_params(
&self,
sql: &str,
params: &[SqliteValue],
) -> Result<Vec<Row>, FrankenError> {
self.ensure_active()?;
let result = self.conn.query_with_params(sql, params).await;
self.observe_transaction_state(result)
}
pub async fn query_params(
&self,
sql: &str,
params: &[ParamValue],
) -> Result<Vec<Row>, FrankenError> {
self.ensure_active()?;
let values: Vec<SqliteValue> = params.iter().map(|p| p.0.clone()).collect();
let result = self.conn.query_with_params(sql, &values).await;
self.observe_transaction_state(result)
}
pub async fn query_row(&self, sql: &str) -> Result<Row, FrankenError> {
self.ensure_active()?;
let result = self.conn.query_row(sql).await;
self.observe_transaction_state(result)
}
pub async fn query_row_with_params(
&self,
sql: &str,
params: &[SqliteValue],
) -> Result<Row, FrankenError> {
self.ensure_active()?;
let result = self.conn.query_row_with_params(sql, params).await;
self.observe_transaction_state(result)
}
pub async fn query_row_map<T, F>(
&self,
sql: &str,
params: &[ParamValue],
f: F,
) -> Result<T, FrankenError>
where
F: FnOnce(&Row) -> Result<T, FrankenError>,
{
self.ensure_active()?;
let values: Vec<SqliteValue> = params.iter().map(|p| p.0.clone()).collect();
let result = self.conn.query_row_with_params(sql, &values).await;
let row = self.observe_transaction_state(result)?;
f(&row)
}
pub async fn query_map_collect<T, F>(
&self,
sql: &str,
params: &[ParamValue],
mut f: F,
) -> Result<Vec<T>, FrankenError>
where
F: FnMut(&Row) -> Result<T, FrankenError>,
{
self.ensure_active()?;
let values: Vec<SqliteValue> = params.iter().map(|p| p.0.clone()).collect();
let mut mapped = Vec::new();
let result = self
.conn
.query_with_params_for_each(sql, &values, |row| {
mapped.push(f(row)?);
Ok(())
})
.await;
self.observe_transaction_state(result)?;
Ok(mapped)
}
pub async fn execute_batch(&self, sql: &str) -> Result<(), FrankenError> {
self.ensure_active()?;
let result = Connection::execute_batch(self.conn, sql).await;
self.observe_transaction_state(result)
}
pub fn last_insert_rowid(&self) -> Result<i64, FrankenError> {
self.ensure_active()?;
Ok(self.conn.last_insert_rowid())
}
}
impl Drop for Transaction<'_> {
fn drop(&mut self) {
if !self.finalized.get() {
self.conn.mark_transaction_cleanup_required();
tracing::debug!(
target: "fsqlite::compat",
event = "transaction_drop_without_finalize",
msg = "Transaction dropped without an awaited commit()/rollback(); \
it will be rolled back before the next statement runs"
);
}
}
}
pub trait TransactionExt {
fn transaction(&self) -> impl Future<Output = Result<Transaction<'_>, FrankenError>>;
}
impl TransactionExt for Connection {
async fn transaction(&self) -> Result<Transaction<'_>, FrankenError> {
Transaction::new(self).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::compat::RowExt;
#[test]
fn transaction_commit() {
asupersync::test_utils::run_test(|| async {
let conn = Connection::open(":memory:").await.unwrap();
conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)")
.await
.unwrap();
let mut tx = conn.transaction().await.unwrap();
tx.execute("INSERT INTO t (val) VALUES ('committed')")
.await
.unwrap();
tx.commit().await.unwrap();
let rows = conn.query("SELECT val FROM t").await.unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].get_typed::<String>(0).unwrap(), "committed");
});
}
#[test]
fn finalized_transaction_rejects_later_operations() {
asupersync::test_utils::run_test(|| async {
let conn = Connection::open(":memory:").await.unwrap();
conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)")
.await
.unwrap();
let mut tx = conn.transaction().await.unwrap();
tx.execute("INSERT INTO t (val) VALUES ('committed')")
.await
.unwrap();
tx.commit().await.unwrap();
let error = tx
.execute("INSERT INTO t (val) VALUES ('must_not_autocommit')")
.await
.expect_err("a finalized transaction wrapper must reject later statements");
assert!(matches!(error, FrankenError::NoActiveTransaction));
let rows = conn.query("SELECT val FROM t ORDER BY id").await.unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].get_typed::<String>(0).unwrap(), "committed");
});
}
#[test]
fn transaction_rejects_operations_after_sql_ends_underlying_scope() {
asupersync::test_utils::run_test(|| async {
let conn = Connection::open(":memory:").await.unwrap();
conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)")
.await
.unwrap();
let tx = conn.transaction().await.unwrap();
tx.execute("COMMIT").await.unwrap();
let mut replacement = conn.transaction().await.unwrap();
let error = tx
.execute("INSERT INTO t (val) VALUES ('must_not_autocommit')")
.await
.expect_err("a wrapper must reject statements after SQL ends its transaction");
assert!(matches!(error, FrankenError::NoActiveTransaction));
drop(tx);
replacement
.execute("INSERT INTO t (val) VALUES ('replacement_transaction')")
.await
.unwrap();
replacement.commit().await.unwrap();
let rows = conn.query("SELECT val FROM t ORDER BY id").await.unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(
rows[0].get_typed::<String>(0).unwrap(),
"replacement_transaction"
);
});
}
#[test]
fn transaction_drop_rolls_back_before_next_statement() {
asupersync::test_utils::run_test(|| async {
let conn = Connection::open(":memory:").await.unwrap();
conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)")
.await
.unwrap();
{
let tx = conn.transaction().await.unwrap();
tx.execute("INSERT INTO t (val) VALUES ('not_rolled_back')")
.await
.unwrap();
}
let rows = conn.query("SELECT val FROM t").await.unwrap();
assert!(
rows.is_empty(),
"the next statement must roll back an abandoned transaction before it reads"
);
assert!(
!conn.in_transaction(),
"settling the deferred rollback must leave the connection idle"
);
});
}
#[test]
fn transaction_explicit_rollback() {
asupersync::test_utils::run_test(|| async {
let conn = Connection::open(":memory:").await.unwrap();
conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, val TEXT)")
.await
.unwrap();
let mut tx = conn.transaction().await.unwrap();
tx.execute("INSERT INTO t (val) VALUES ('rolled_back')")
.await
.unwrap();
tx.rollback().await.unwrap();
let rows = conn.query("SELECT val FROM t").await.unwrap();
assert!(rows.is_empty());
});
}
}