use crate::{
Turso, TursoConnection,
error::{engine_error, unsupported},
};
use sqlx_core::{
error::Error,
sql_str::{SqlSafeStr, SqlStr},
transaction::TransactionManager,
};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub(crate) enum TransactionState {
#[default]
Clean,
Active,
RollbackPending,
Unknown,
}
#[derive(Debug)]
pub struct TursoTransactionManager;
impl TursoConnection {
pub(crate) async fn ready(&mut self) -> Result<(), Error> {
if self.transaction_state == TransactionState::Unknown {
return Err(unsupported(
"reuse after an uncertain transaction transition",
));
}
self.clear_pending_statement()?;
if matches!(
self.transaction_state,
TransactionState::Active | TransactionState::RollbackPending
) {
let autocommit = self.inner.is_autocommit().map_err(|error| {
self.transaction_state = TransactionState::Unknown;
engine_error(error)
})?;
if autocommit {
self.transaction_state = TransactionState::Unknown;
return Err(unsupported(
"reuse after the SQLx transaction ended unexpectedly",
));
}
}
if self.transaction_state == TransactionState::RollbackPending {
self.transition("ROLLBACK".into_sql_str(), TransactionState::Clean)
.await?;
}
Ok(())
}
async fn transition(&mut self, sql: SqlStr, target: TransactionState) -> Result<(), Error> {
self.clear_pending_statement()?;
self.transaction_state = TransactionState::Unknown;
#[cfg(test)]
let (before, after) = match target {
TransactionState::Active => (tests::Phase::BeforeBegin, tests::Phase::AfterBegin),
_ if sql.as_str() == "COMMIT" => {
(tests::Phase::BeforeCommit, tests::Phase::AfterCommit)
}
_ => (tests::Phase::BeforeRollback, tests::Phase::AfterRollback),
};
#[cfg(test)]
self.transition_barrier(before).await;
#[cfg(test)]
if sql.as_str() == "ROLLBACK" && std::mem::take(&mut self.rollback_fault) {
return Err(engine_error(turso::Error::IoError(
std::io::ErrorKind::Other,
"simulated rollback boundary failure",
)));
}
let mut statement = self
.inner
.prepare(sql.as_str())
.await
.map_err(engine_error)?;
self.pending_statement = Some(statement.clone());
statement.execute(()).await.map_err(engine_error)?;
#[cfg(test)]
self.transition_barrier(after).await;
drop(statement);
self.clear_pending_statement()?;
let autocommit = self.inner.is_autocommit().map_err(engine_error)?;
if autocommit != (target == TransactionState::Clean) {
return Err(unsupported(
"transaction transition did not establish the expected engine state",
));
}
self.transaction_state = target;
Ok(())
}
#[cfg(test)]
async fn transition_barrier(&mut self, phase: tests::Phase) {
if self
.transition_hook
.as_ref()
.is_some_and(|hook| hook.phase == phase)
{
let hook = self.transition_hook.take().expect("matched hook");
let _ = hook.reached.send(());
let _ = hook.resume.await;
}
}
}
impl TransactionManager for TursoTransactionManager {
type Database = Turso;
async fn begin(conn: &mut TursoConnection, statement: Option<SqlStr>) -> Result<(), Error> {
if conn.transaction_state == TransactionState::Active {
conn.rejected_begin_drop = true;
return Err(if statement.is_some() {
Error::InvalidSavePointStatement
} else {
unsupported("nested transactions")
});
}
conn.ready().await?;
let sql = statement.unwrap_or_else(|| "BEGIN".into_sql_str());
validate_begin(sql.as_str())?;
conn.transition(sql, TransactionState::Active).await
}
async fn commit(conn: &mut TursoConnection) -> Result<(), Error> {
conn.ready().await?;
if conn.transaction_state != TransactionState::Active {
return Err(unsupported("commit without an active SQLx transaction"));
}
conn.transition("COMMIT".into_sql_str(), TransactionState::Clean)
.await
}
async fn rollback(conn: &mut TursoConnection) -> Result<(), Error> {
conn.ready().await?;
if conn.transaction_state != TransactionState::Active {
return Err(unsupported("rollback without an active SQLx transaction"));
}
conn.transition("ROLLBACK".into_sql_str(), TransactionState::Clean)
.await
}
fn start_rollback(conn: &mut TursoConnection) {
if std::mem::take(&mut conn.rejected_begin_drop) {
return;
}
if conn.transaction_state == TransactionState::Active {
conn.transaction_state = TransactionState::RollbackPending;
}
}
fn get_transaction_depth(conn: &TursoConnection) -> usize {
usize::from(conn.transaction_state != TransactionState::Clean)
}
}
fn validate_begin(sql: &str) -> Result<(), Error> {
use turso_parser::ast::{Cmd, Stmt};
let parse_error =
|error: turso::core::LimboError| engine_error(turso::Error::Error(error.to_string()));
let (command, end) = turso::core::dialect::sqlite::parse(sql).map_err(parse_error)?;
let (trailing, _) = turso::core::dialect::sqlite::parse(&sql[end..]).map_err(parse_error)?;
if !matches!(command, Some(Cmd::Stmt(Stmt::Begin { .. }))) || trailing.is_some() {
return Err(unsupported("begin_with requires a single BEGIN statement"));
}
Ok(())
}
#[cfg(test)]
pub(crate) mod tests {
use crate::{
Turso, TursoConnectOptions, TursoConnection, TursoTransactionManager, connect_pool,
};
use sqlx::{ConnectOptions, Connection, Executor};
use sqlx_core::transaction::TransactionManager;
use std::{future::Future, io::ErrorKind, time::Duration};
use tokio::sync::oneshot;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum Phase {
BeforeBegin,
AfterBegin,
BeforeCommit,
AfterCommit,
BeforeRollback,
AfterRollback,
}
#[derive(Debug)]
pub(crate) struct Hook {
pub(crate) phase: Phase,
pub(crate) reached: oneshot::Sender<()>,
pub(crate) resume: oneshot::Receiver<()>,
}
async fn bounded<T>(phase: &str, future: impl Future<Output = T>) -> T {
tokio::time::timeout(Duration::from_secs(3), future)
.await
.unwrap_or_else(|_| panic!("transaction phase timed out: {phase}"))
}
type TestResult = Result<(), Box<dyn std::error::Error>>;
async fn cancellation_case(phase: Phase) -> TestResult {
bounded(&format!("{phase:?}"), async {
for memory in [false, true] {
let directory = tempfile::tempdir()?;
let options = if memory { TursoConnectOptions::memory() } else { TursoConnectOptions::file(directory.path().join("cancel.db"))? };
let pool = connect_pool(options.clone(), 1).await?;
let mut connection = pool.acquire().await?;
connection.execute("CREATE TABLE items(value INTEGER)").await?;
connection.execute("INSERT INTO items VALUES (7)").await?;
connection.execute("CREATE TEMP TABLE local_state(value INTEGER)").await?;
let (reached, mut at_barrier) = oneshot::channel();
let (_resume, resume) = oneshot::channel();
connection.transition_hook = Some(Hook { phase, reached, resume });
let mut transition: std::pin::Pin<Box<dyn Future<Output = Result<(), sqlx::Error>> + Send + '_>> = match phase {
Phase::BeforeBegin | Phase::AfterBegin => Box::pin(async {
connection.begin().await?.rollback().await
}),
_ => {
let mut tx = connection.begin().await?;
tx.execute("INSERT INTO items VALUES (8)").await?;
match phase {
Phase::BeforeCommit | Phase::AfterCommit => Box::pin(tx.commit()),
_ => Box::pin(tx.rollback()),
}
}
};
tokio::select! {
result = &mut transition => panic!("{phase:?}: transition completed before barrier: {result:?}"),
reached = &mut at_barrier => reached.expect("barrier sender lost"),
}
drop(transition); assert!(matches!(connection.ping().await, Err(sqlx::Error::Protocol(_))), "{phase:?}: uncertain connection passed ping: {connection:?}");
assert!(connection.execute("INSERT INTO items VALUES (99)").await.is_err(), "{phase:?}: uncertain executor accepted write");
drop(connection);
let next = pool.acquire().await;
if memory {
assert!(matches!(next, Err(sqlx::Error::Configuration(_))), "{phase:?}: memory loss was hidden: {next:?}");
} else {
let mut next = next?;
assert!(next.execute("INSERT INTO local_state VALUES (1)").await.is_err(), "{phase:?}: old physical connection reused");
let expected = if phase == Phase::AfterCommit { 15 } else { 7 };
assert_eq!(sqlx::query_scalar::<Turso, i64>("SELECT sum(value) FROM items").fetch_one(&mut *next).await?, expected, "{phase:?}: transaction leaked");
next.execute("INSERT INTO items VALUES (10)").await?;
assert_eq!(sqlx::query_scalar::<Turso, i64>("SELECT sum(value) FROM items").fetch_one(&mut *next).await?, expected + 10, "{phase:?}: replacement cannot read/write");
drop(next);
}
pool.close().await;
assert_eq!(pool.size(), 0, "{phase:?}: capacity leaked");
}
TestResult::Ok(())
}).await
}
#[tokio::test]
async fn cancelled_transition_never_leaks_transaction_before_begin() -> TestResult {
cancellation_case(Phase::BeforeBegin).await
}
#[tokio::test]
async fn cancelled_transition_never_leaks_transaction_after_begin() -> TestResult {
cancellation_case(Phase::AfterBegin).await
}
#[tokio::test]
async fn cancelled_transition_never_leaks_transaction_before_commit() -> TestResult {
cancellation_case(Phase::BeforeCommit).await
}
#[tokio::test]
async fn cancelled_transition_never_leaks_transaction_after_commit() -> TestResult {
cancellation_case(Phase::AfterCommit).await
}
#[tokio::test]
async fn cancelled_transition_never_leaks_transaction_before_rollback() -> TestResult {
cancellation_case(Phase::BeforeRollback).await
}
#[tokio::test]
async fn cancelled_transition_never_leaks_transaction_after_rollback() -> TestResult {
cancellation_case(Phase::AfterRollback).await
}
#[tokio::test]
async fn pool_owned_cancellation_all_six_boundaries() -> TestResult {
for phase in [
Phase::BeforeBegin,
Phase::AfterBegin,
Phase::BeforeCommit,
Phase::AfterCommit,
Phase::BeforeRollback,
Phase::AfterRollback,
] {
bounded(&format!("pool-owned {phase:?}"), async {
for memory in [false, true] {
let directory = tempfile::tempdir()?;
let options = if memory { TursoConnectOptions::memory() } else { TursoConnectOptions::file(directory.path().join("owned-cancel.db"))? };
let pool = connect_pool(options, 1).await?;
let mut connection = pool.acquire().await?;
connection.execute("CREATE TABLE items(value INTEGER)").await?;
connection.execute("INSERT INTO items VALUES (7)").await?;
connection.execute("CREATE TEMP TABLE local_state(value INTEGER)").await?;
let (reached, mut at_barrier) = oneshot::channel();
let (_resume, resume) = oneshot::channel();
connection.transition_hook = Some(Hook { phase, reached, resume });
drop(connection);
let mut transition: std::pin::Pin<Box<dyn Future<Output = Result<(), sqlx::Error>> + Send + '_>> = match phase {
Phase::BeforeBegin | Phase::AfterBegin => Box::pin(async { pool.begin().await?.rollback().await }),
_ => {
let mut tx = pool.begin().await?;
tx.execute("INSERT INTO items VALUES (8)").await?;
match phase {
Phase::BeforeCommit | Phase::AfterCommit => Box::pin(tx.commit()),
_ => Box::pin(tx.rollback()),
}
}
};
tokio::select! {
result = &mut transition => panic!("pool-owned {phase:?}: transition completed before barrier: {result:?}"),
reached = &mut at_barrier => reached.expect("pool-owned barrier sender lost"),
}
drop(transition);
let next = pool.acquire().await;
if memory {
assert!(matches!(next, Err(sqlx::Error::Configuration(_))), "pool-owned {phase:?}: {next:?}");
} else {
let mut next = next?;
assert!(next.execute("INSERT INTO local_state VALUES (1)").await.is_err(), "pool-owned {phase:?}: old connection reused");
let expected = if phase == Phase::AfterCommit { 15 } else { 7 };
assert_eq!(sqlx::query_scalar::<Turso, i64>("SELECT sum(value) FROM items").fetch_one(&mut *next).await?, expected, "pool-owned {phase:?}");
next.execute("INSERT INTO items VALUES (10)").await?;
assert_eq!(sqlx::query_scalar::<Turso, i64>("SELECT sum(value) FROM items").fetch_one(&mut *next).await?, expected + 10);
drop(next);
}
pool.close().await;
assert_eq!(pool.size(), 0, "pool-owned {phase:?}: capacity leaked");
}
TestResult::Ok(())
}).await?;
}
Ok(())
}
#[tokio::test]
async fn simulated_rollback_failure_quarantines_connection() -> TestResult {
bounded("simulated rollback failure", async {
for memory in [false, true] {
let directory = tempfile::tempdir()?;
let options = if memory {
TursoConnectOptions::memory()
} else {
TursoConnectOptions::file(directory.path().join("rollback-fault.db"))?
};
let pool = connect_pool(options, 1).await?;
let mut connection = pool.acquire().await?;
connection
.execute("CREATE TABLE items(value INTEGER)")
.await?;
connection.execute("INSERT INTO items VALUES (7)").await?;
connection
.execute("CREATE TEMP TABLE local_state(value INTEGER)")
.await?;
let mut tx = connection.begin().await?;
tx.execute("INSERT INTO items VALUES (8)").await?;
tx.rollback_fault = true;
drop(tx);
let error = connection.flush().await.unwrap_err();
let mapped = error
.as_database_error()
.expect("original rollback error must retain Database mapping");
let source =
std::error::Error::source(mapped.downcast_ref::<crate::TursoDatabaseError>())
.unwrap();
assert!(matches!(
source.downcast_ref::<turso::Error>(),
Some(turso::Error::IoError(
ErrorKind::Other,
"simulated rollback boundary failure"
))
));
TursoTransactionManager::start_rollback(&mut connection);
assert!(connection.ping().await.is_err(), "fault quarantine cleared");
drop(connection);
let next = pool.acquire().await;
if memory {
assert!(
matches!(next, Err(sqlx::Error::Configuration(_))),
"{next:?}"
);
} else {
let mut next = next?;
assert!(
next.execute("INSERT INTO local_state VALUES (1)")
.await
.is_err()
);
assert_eq!(
sqlx::query_scalar::<Turso, i64>("SELECT sum(value) FROM items")
.fetch_one(&mut *next)
.await?,
7
);
next.execute("INSERT INTO items VALUES (10)").await?;
drop(next);
}
pool.close().await;
assert_eq!(pool.size(), 0);
}
TestResult::Ok(())
})
.await
}
#[tokio::test]
async fn actual_begin_error_quarantines_connection() -> TestResult {
bounded("actual begin engine error", async {
let mut connection: TursoConnection = TursoConnectOptions::memory().connect().await?;
connection.execute("BEGIN").await?;
assert!(matches!(
connection.begin().await,
Err(sqlx::Error::Database(_))
));
assert!(connection.ping().await.is_err());
connection.close_hard().await?;
TestResult::Ok(())
})
.await
}
}