use std::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: bool,
}
impl<'a> Transaction<'a> {
async fn new(conn: &'a Connection) -> Result<Self, FrankenError> {
conn.begin_transaction().await?;
Ok(Self {
conn,
finalized: false,
})
}
pub async fn commit(&mut self) -> Result<(), FrankenError> {
self.conn.commit_transaction().await?;
self.finalized = true;
Ok(())
}
pub async fn rollback(&mut self) -> Result<(), FrankenError> {
self.conn.rollback_transaction().await?;
self.finalized = true;
Ok(())
}
pub async fn execute(&self, sql: &str) -> Result<usize, FrankenError> {
self.conn.execute(sql).await
}
pub async fn execute_with_params(
&self,
sql: &str,
params: &[SqliteValue],
) -> Result<usize, FrankenError> {
self.conn.execute_with_params(sql, params).await
}
pub async fn execute_with_params_skip_statement_savepoint(
&self,
sql: &str,
params: &[SqliteValue],
) -> Result<usize, FrankenError> {
self.conn
.execute_with_params_skip_statement_savepoint_in_explicit_txn(sql, params)
.await
}
pub async fn execute_compat(
&self,
sql: &str,
params: &[ParamValue],
) -> Result<usize, FrankenError> {
let values: Vec<SqliteValue> = params.iter().map(|p| p.0.clone()).collect();
self.conn.execute_with_params(sql, &values).await
}
pub async fn query(&self, sql: &str) -> Result<Vec<Row>, FrankenError> {
self.conn.query(sql).await
}
pub async fn query_with_params(
&self,
sql: &str,
params: &[SqliteValue],
) -> Result<Vec<Row>, FrankenError> {
self.conn.query_with_params(sql, params).await
}
pub async fn query_params(
&self,
sql: &str,
params: &[ParamValue],
) -> Result<Vec<Row>, FrankenError> {
let values: Vec<SqliteValue> = params.iter().map(|p| p.0.clone()).collect();
self.conn.query_with_params(sql, &values).await
}
pub async fn query_row(&self, sql: &str) -> Result<Row, FrankenError> {
self.conn.query_row(sql).await
}
pub async fn query_row_with_params(
&self,
sql: &str,
params: &[SqliteValue],
) -> Result<Row, FrankenError> {
self.conn.query_row_with_params(sql, params).await
}
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>,
{
let values: Vec<SqliteValue> = params.iter().map(|p| p.0.clone()).collect();
let row = self.conn.query_row_with_params(sql, &values).await?;
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>,
{
let values: Vec<SqliteValue> = params.iter().map(|p| p.0.clone()).collect();
let mut mapped = Vec::new();
self.conn
.query_with_params_for_each(sql, &values, |row| {
mapped.push(f(row)?);
Ok(())
})
.await?;
Ok(mapped)
}
pub async fn execute_batch(&self, sql: &str) -> Result<(), FrankenError> {
Connection::execute_batch(self.conn, sql).await
}
pub fn last_insert_rowid(&self) -> Result<i64, FrankenError> {
Ok(self.conn.last_insert_rowid())
}
}
impl Drop for Transaction<'_> {
fn drop(&mut self) {
if !self.finalized {
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 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());
});
}
}