use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use async_trait::async_trait;
use turso_sql::Statement;
use crate::connection::{ConnectionTrait, StreamTrait, TransactionTrait};
use crate::database::{PooledConnection, retry_busy};
use crate::error::Result;
use crate::executor::{self, Conn, ExecResult, Row, RowStream};
const NO_PENDING: u32 = u32::MAX;
#[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 Shared {
conn: PooledConnection,
depth: AtomicU32,
pending_rollback: AtomicU32,
}
impl Shared {
async fn settle(&self) -> Result<()> {
let pending = self.pending_rollback.swap(NO_PENDING, Ordering::AcqRel);
if pending != NO_PENDING {
tracing::warn!(
depth = pending,
"rolling back nested transaction dropped without commit"
);
rollback_to(&self.conn, pending).await?;
self.depth.store(pending - 1, Ordering::Release);
}
Ok(())
}
}
async fn rollback_to(conn: &Conn, depth: u32) -> Result<()> {
conn.execute_raw(&format!("ROLLBACK TO SAVEPOINT sp{depth}"))
.await?;
conn.execute_raw(&format!("RELEASE SAVEPOINT sp{depth}"))
.await?;
Ok(())
}
#[must_use = "a transaction must be committed or rolled back"]
pub struct Transaction {
shared: Arc<Shared>,
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,
depth: AtomicU32::new(0),
pending_rollback: AtomicU32::new(NO_PENDING),
}),
depth: 0,
open: true,
})
}
async fn begin_nested(&self) -> Result<Self> {
self.shared.settle().await?;
let depth = self.shared.depth.load(Ordering::Acquire) + 1;
self.shared
.conn
.execute_raw(&format!("SAVEPOINT sp{depth}"))
.await?;
self.shared.depth.store(depth, Ordering::Release);
Ok(Self {
shared: Arc::clone(&self.shared),
depth,
open: true,
})
}
pub fn depth(&self) -> u32 {
self.depth
}
fn conn(&self) -> &Conn {
&self.shared.conn
}
async fn finish(mut self, rollback: bool) -> Result<()> {
self.open = false;
self.shared.settle().await?;
if self.depth == 0 {
self.conn()
.execute_raw(if rollback { "ROLLBACK" } else { "COMMIT" })
.await?;
} else if rollback {
rollback_to(self.conn(), self.depth).await?;
self.shared.depth.store(self.depth - 1, Ordering::Release);
} else {
self.conn()
.execute_raw(&format!("RELEASE SAVEPOINT sp{}", self.depth))
.await?;
self.shared.depth.store(self.depth - 1, Ordering::Release);
}
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;
}
if self.depth == 0 {
tracing::warn!("transaction dropped without commit or rollback; discarding connection");
self.shared.conn.discard();
} else {
self.shared
.pending_rollback
.fetch_min(self.depth, Ordering::AcqRel);
}
}
}
#[async_trait]
impl ConnectionTrait for Transaction {
async fn execute(&self, statement: Statement) -> Result<ExecResult> {
self.shared.settle().await?;
executor::execute(self.conn(), &statement).await
}
async fn execute_unprepared(&self, sql: &str) -> Result<ExecResult> {
self.shared.settle().await?;
executor::execute_unprepared(self.conn(), sql).await
}
async fn query_one(&self, statement: Statement) -> Result<Option<Row>> {
self.shared.settle().await?;
executor::query_one(self.conn(), &statement).await
}
async fn query_all(&self, statement: Statement) -> Result<Vec<Row>> {
self.shared.settle().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.shared.settle().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
}
}