use super::Row;
use super::{Client, Param};
use crate::errors::Result;
use crate::{ExecuteResult, Syntax};
use async_trait::async_trait;
use std::sync::Mutex;
#[cfg(feature = "mssql")]
use crate::mssql::transaction::MssqlTransaction;
#[cfg(feature = "sqlite-sync")]
use crate::sqlite_sync::SqliteSyncTransaction;
pub struct Transaction<'t> {
inner: Mutex<Option<TransT<'t>>>,
syntax: crate::Syntax,
}
#[maybe_async::maybe_async]
impl<'t> Transaction<'t> {
pub(crate) fn new(inner: TransT<'t>) -> Self {
let syntax = match &inner {
#[cfg(feature = "sqlite")]
TransT::Sqlite(_) => Syntax::Sqlite,
#[cfg(feature = "sqlite-sync")]
TransT::SqliteSync(_) => Syntax::Sqlite,
#[cfg(feature = "mssql")]
TransT::Mssql(_) => Syntax::Mssql,
#[cfg(feature = "postgres")]
TransT::Postgres(_) => Syntax::Postgres,
#[cfg(feature = "mysql")]
TransT::Mysql(_) => Syntax::Mysql,
};
Self {
syntax,
inner: Mutex::new(Some(inner)),
}
}
pub async fn rollback(self) -> Result<()> {
let inner = self.take_conn();
inner.rollback().await?;
Ok(())
}
pub async fn commit(self) -> Result<()> {
let inner = self.take_conn();
inner.commit().await?;
Ok(())
}
}
impl<'t> Transaction<'t> {
fn take_conn(&self) -> TransT<'t> {
let mut placeholder = None;
let mut m = self.inner.lock().unwrap();
let inner: &mut Option<TransT<'t>> = &mut m;
assert!(inner.is_some(), "Pool was already taken");
std::mem::swap(&mut placeholder, inner);
placeholder.unwrap()
}
fn return_conn(&self, conn: TransT<'t>) {
let mut placeholder = Some(conn);
let mut m = self.inner.lock().unwrap();
let inner: &mut Option<TransT<'t>> = &mut m;
assert!(inner.is_none(), "Overriding existing pool");
std::mem::swap(&mut placeholder, inner);
}
}
pub(crate) enum TransT<'t> {
#[cfg(feature = "sqlite")]
Sqlite(sqlx::Transaction<'t, sqlx::Sqlite>),
#[cfg(feature = "sqlite-sync")]
SqliteSync(SqliteSyncTransaction<'t>),
#[cfg(feature = "postgres")]
Postgres(sqlx::Transaction<'t, sqlx::Postgres>),
#[cfg(feature = "mysql")]
Mysql(sqlx::Transaction<'t, sqlx::MySql>),
#[cfg(feature = "mssql")]
Mssql(MssqlTransaction<'t>),
}
#[maybe_async::maybe_async]
impl TransT<'_> {
async fn rollback(self) -> Result<()> {
match self {
#[cfg(feature = "sqlite")]
TransT::Sqlite(t) => t.rollback().await?,
#[cfg(feature = "sqlite-sync")]
TransT::SqliteSync(t) => t.transaction.rollback().await?,
#[cfg(feature = "mssql")]
TransT::Mssql(t) => t.rollback().await?,
#[cfg(feature = "postgres")]
TransT::Postgres(t) => t.rollback().await?,
#[cfg(feature = "mysql")]
TransT::Mysql(t) => t.rollback().await?,
}
Ok(())
}
async fn commit(self) -> Result<()> {
match self {
#[cfg(feature = "sqlite")]
TransT::Sqlite(t) => t.commit().await?,
#[cfg(feature = "sqlite-sync")]
TransT::SqliteSync(t) => t.transaction.commit()?,
#[cfg(feature = "mssql")]
TransT::Mssql(t) => t.commit().await?,
#[cfg(feature = "postgres")]
TransT::Postgres(t) => t.commit().await?,
#[cfg(feature = "mysql")]
TransT::Mysql(t) => t.commit().await?,
}
Ok(())
}
}
#[cfg(feature = "mysql")]
use super::mysql::MysqlParam;
#[cfg(feature = "postgres")]
use super::postgres::PostgresParam;
#[cfg(feature = "sqlite")]
use super::sqlite::SqliteParam;
#[cfg(feature = "sqlite-sync")]
use super::sqlite_sync::SqliteSyncParam;
#[maybe_async::maybe_async]
#[async_trait]
impl Client for Transaction<'_> {
fn syntax(&self) -> crate::Syntax {
self.syntax
}
async fn execute(&self, sql: &str, params: &[&(dyn Param + Sync)]) -> Result<ExecuteResult> {
if sql.trim().is_empty() {
return Ok(ExecuteResult::new(0));
}
let mut inner = self.take_conn();
let results = execute_inner(&mut inner, sql, params).await;
self.return_conn(inner);
results
}
async fn fetch_rows(&self, sql: &str, params: &[&(dyn Param + Sync)]) -> Result<Vec<Row>> {
let mut inner = self.take_conn();
let results = fetch_rows_inner(&mut inner, sql, params).await;
self.return_conn(inner);
results
}
async fn fetch_many<'s, 'args, 'i>(
&self,
fetches: &[crate::Fetch<'s, 'args, 'i>],
) -> Result<Vec<Vec<Row>>> {
let mut datasets = Vec::default();
let mut inner = self.take_conn();
for fetch in fetches {
let sql = fetch.sql;
let params = fetch.params;
let r = fetch_rows_inner(&mut inner, sql, params).await;
let is_err = r.is_err();
datasets.push(r);
if is_err {
break;
}
}
self.return_conn(inner);
datasets.drain(..).collect()
}
}
#[maybe_async::maybe_async]
async fn execute_inner(
inner: &mut TransT<'_>,
sql: &str,
params: &[&(dyn Param + Sync)],
) -> Result<ExecuteResult> {
match inner {
#[cfg(feature = "sqlite")]
TransT::Sqlite(t) => {
let x: &mut <sqlx::Sqlite as sqlx::Database>::Connection = t;
let sql = sqlx::AssertSqlSafe(sql);
let mut query = sqlx::query::<sqlx::Sqlite>(sql);
for param in params {
query = SqliteParam::add_param(*param, query)
}
let t = query.execute(x).await?;
Ok(ExecuteResult {
rows_affected: t.rows_affected(),
})
}
#[cfg(feature = "sqlite-sync")]
TransT::SqliteSync(t) => {
let mut p = Vec::new();
for param in params {
p.push(SqliteSyncParam::to_sql_dyn(param));
}
let r = t.transaction.execute(sql, &*p)?;
Ok(ExecuteResult::new(r as u64))
}
#[cfg(feature = "postgres")]
TransT::Postgres(t) => {
let x: &mut <sqlx::Postgres as sqlx::Database>::Connection = t;
let sql = sqlx::AssertSqlSafe(sql);
let mut query = sqlx::query::<sqlx::Postgres>(sql);
for param in params {
query = PostgresParam::add_param(*param, query)
}
let t = query.execute(x).await?;
Ok(ExecuteResult {
rows_affected: t.rows_affected(),
})
}
#[cfg(feature = "mysql")]
TransT::Mysql(t) => {
let x: &mut <sqlx::MySql as sqlx::Database>::Connection = t;
let sql = sqlx::AssertSqlSafe(sql);
let mut query = sqlx::query::<sqlx::MySql>(sql);
for param in params {
query = MysqlParam::add_param(*param, query)
}
let t = query.execute(x).await?;
Ok(ExecuteResult {
rows_affected: t.rows_affected(),
})
}
#[cfg(feature = "mssql")]
TransT::Mssql(inner) => {
let result = inner.execute(sql, params).await;
if result.is_err() {
let _ = inner.internal_rollback_check().await;
}
result
}
}
}
#[maybe_async::maybe_async]
async fn fetch_rows_inner(
inner: &mut TransT<'_>,
sql: &str,
params: &[&(dyn Param + Sync)],
) -> Result<Vec<Row>> {
match inner {
#[cfg(feature = "sqlite")]
TransT::Sqlite(t) => {
let x: &mut <sqlx::Sqlite as sqlx::Database>::Connection = t;
let sql = sqlx::AssertSqlSafe(sql);
let mut query = sqlx::query::<sqlx::Sqlite>(sql);
for param in params {
query = SqliteParam::add_param(*param, query)
}
let mut raw_rows = query.fetch_all(x).await?;
let rows: Vec<Row> = raw_rows.drain(..).map(Row::from).collect();
Ok(rows)
}
#[cfg(feature = "sqlite-sync")]
TransT::SqliteSync(t) => {
use super::sqlite_sync::SqliteSyncOwnedRow;
use std::sync::Arc;
let mut p = Vec::new();
for param in params {
p.push(SqliteSyncParam::to_sql_dyn(param));
}
let mut stmt = t.transaction.prepare(sql)?;
let column_names: Vec<String> =
stmt.column_names().iter().map(|s| s.to_string()).collect();
let columns = Arc::new(column_names);
let mut raw_rows = stmt.query(&*p)?;
let mut res = Vec::new();
while let Some(row) = raw_rows.next()? {
let mut data = Vec::new();
for i in 0..row.as_ref().column_count() {
data.push(row.get::<_, rusqlite::types::Value>(i)?);
}
res.push(Row::from(SqliteSyncOwnedRow {
data,
columns: Arc::clone(&columns),
}));
}
Ok(res)
}
#[cfg(feature = "postgres")]
TransT::Postgres(t) => {
let x: &mut <sqlx::Postgres as sqlx::Database>::Connection = t;
let sql = sqlx::AssertSqlSafe(sql);
let mut query = sqlx::query::<sqlx::Postgres>(sql);
for param in params {
query = PostgresParam::add_param(*param, query)
}
let mut raw_rows = query.fetch_all(x).await?;
let rows: Vec<Row> = raw_rows.drain(..).map(Row::from).collect();
Ok(rows)
}
#[cfg(feature = "mysql")]
TransT::Mysql(t) => {
let x: &mut <sqlx::MySql as sqlx::Database>::Connection = t;
let sql = sqlx::AssertSqlSafe(sql);
let mut query = sqlx::query::<sqlx::MySql>(sql);
for param in params {
query = MysqlParam::add_param(*param, query)
}
let mut raw_rows = query.fetch_all(x).await?;
let rows: Vec<Row> = raw_rows.drain(..).map(Row::from).collect();
Ok(rows)
}
#[cfg(feature = "mssql")]
TransT::Mssql(inner) => {
let result = inner.fetch_rows(sql, params).await;
if result.is_err() {
let _ = inner.internal_rollback_check().await;
}
result
}
}
}