use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use async_trait::async_trait;
use turso_sql::Statement;
use crate::connection::{ConnectionTrait, StreamTrait, TransactionTrait};
use crate::database::{PooledConnection, retry_busy};
use crate::error::{Error, Result};
use crate::executor::{self, Conn, ExecResult, Row, RowStream};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
#[non_exhaustive]
pub enum TransactionMode {
#[default]
Deferred,
Immediate,
Exclusive,
Concurrent,
}
impl TransactionMode {
fn sql(self) -> &'static str {
match self {
Self::Deferred => "BEGIN DEFERRED",
Self::Immediate => "BEGIN IMMEDIATE",
Self::Exclusive => "BEGIN EXCLUSIVE",
Self::Concurrent => "BEGIN CONCURRENT",
}
}
}
struct State {
savepoints: Vec<u64>,
last_id: u64,
pending_rollback: Option<usize>,
closed: bool,
}
struct Shared {
conn: PooledConnection,
state: Mutex<State>,
}
impl Shared {
fn state(&self) -> MutexGuard<'_, State> {
self.state.lock().unwrap_or_else(PoisonError::into_inner)
}
async fn settle(&self) -> Result<()> {
let pending = {
let mut state = self.state();
state
.pending_rollback
.take()
.and_then(|index| Some((index, *state.savepoints.get(index)?)))
};
if let Some((index, id)) = pending {
tracing::warn!(
depth = index + 1,
"rolling back nested transaction dropped without commit"
);
rollback_to(&self.conn, id).await?;
self.state().savepoints.truncate(index);
}
Ok(())
}
}
async fn rollback_to(conn: &Conn, id: u64) -> Result<()> {
conn.execute_raw(&format!("ROLLBACK TO SAVEPOINT sp{id}"))
.await?;
conn.execute_raw(&format!("RELEASE SAVEPOINT sp{id}"))
.await?;
Ok(())
}
#[must_use = "a transaction must be committed or rolled back"]
pub struct Transaction {
shared: Arc<Shared>,
savepoint: Option<u64>,
depth: u32,
open: bool,
}
impl fmt::Debug for Transaction {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Transaction")
.field("depth", &self.depth)
.field("open", &self.open)
.finish_non_exhaustive()
}
}
impl Transaction {
pub(crate) async fn begin_top(conn: PooledConnection, mode: TransactionMode) -> Result<Self> {
let budget = conn.options().busy_timeout.unwrap_or_default();
retry_busy(budget, || async { conn.execute_raw(mode.sql()).await }).await?;
Ok(Self {
shared: Arc::new(Shared {
conn,
state: Mutex::new(State {
savepoints: Vec::new(),
last_id: 0,
pending_rollback: None,
closed: false,
}),
}),
savepoint: None,
depth: 0,
open: true,
})
}
async fn begin_nested(&self) -> Result<Self> {
self.prepare().await?;
let id = {
let mut state = self.shared.state();
state.last_id += 1;
state.last_id
};
self.conn()
.execute_raw(&format!("SAVEPOINT sp{id}"))
.await?;
self.shared.state().savepoints.push(id);
Ok(Self {
shared: Arc::clone(&self.shared),
savepoint: Some(id),
depth: self.depth + 1,
open: true,
})
}
pub fn depth(&self) -> u32 {
self.depth
}
fn conn(&self) -> &Conn {
&self.shared.conn
}
async fn prepare(&self) -> Result<()> {
if self.shared.state().closed {
return Err(Error::Misuse("transaction already finished".into()));
}
self.shared.settle().await?;
let state = self.shared.state();
let innermost = match self.savepoint {
None => state.savepoints.is_empty(),
Some(id) => state.savepoints.last() == Some(&id),
};
if innermost {
return Ok(());
}
let alive = self
.savepoint
.is_none_or(|id| state.savepoints.contains(&id));
Err(Error::Misuse(
if alive {
"a nested transaction is still open"
} else {
"nested transaction rolled back with its parent"
}
.into(),
))
}
async fn finish(mut self, rollback: bool) -> Result<()> {
self.prepare().await?;
self.open = false;
match self.savepoint {
None => {
self.shared.state().closed = true;
self.conn()
.execute_raw(if rollback { "ROLLBACK" } else { "COMMIT" })
.await?;
}
Some(id) => {
if rollback {
rollback_to(self.conn(), id).await?;
} else {
self.conn()
.execute_raw(&format!("RELEASE SAVEPOINT sp{id}"))
.await?;
}
self.shared.state().savepoints.pop();
}
}
Ok(())
}
pub async fn commit(self) -> Result<()> {
self.finish(false).await
}
pub async fn rollback(self) -> Result<()> {
self.finish(true).await
}
}
impl Drop for Transaction {
fn drop(&mut self) {
if !self.open {
return;
}
let mut state = self.shared.state();
let Some(id) = self.savepoint else {
state.closed = true;
drop(state);
tracing::warn!("transaction dropped without commit or rollback; discarding connection");
self.shared.conn.discard();
return;
};
if state.closed {
return;
}
if let Some(index) = state.savepoints.iter().position(|&s| s == id) {
state.pending_rollback = Some(state.pending_rollback.map_or(index, |p| p.min(index)));
}
}
}
#[async_trait]
impl ConnectionTrait for Transaction {
async fn execute(&self, statement: Statement) -> Result<ExecResult> {
self.prepare().await?;
executor::execute(self.conn(), &statement).await
}
async fn execute_unprepared(&self, sql: &str) -> Result<ExecResult> {
self.prepare().await?;
executor::execute_unprepared(self.conn(), sql).await
}
async fn query_one(&self, statement: Statement) -> Result<Option<Row>> {
self.prepare().await?;
executor::query_one(self.conn(), &statement).await
}
async fn query_all(&self, statement: Statement) -> Result<Vec<Row>> {
self.prepare().await?;
executor::query_all(self.conn(), &statement).await
}
}
impl StreamTrait for Transaction {
fn stream<'a>(
&'a self,
statement: Statement,
) -> Pin<Box<dyn Future<Output = Result<RowStream<'a>>> + Send + 'a>> {
Box::pin(async move {
self.prepare().await?;
executor::stream(self.conn(), &statement, ()).await
})
}
}
#[async_trait]
impl TransactionTrait for Transaction {
async fn begin(&self) -> Result<Transaction> {
self.begin_nested().await
}
async fn begin_with_mode(&self, _mode: TransactionMode) -> Result<Transaction> {
self.begin_nested().await
}
}