use std::{future::Future, pin::Pin};
use futures_util::TryStreamExt;
use saddle_observability::{ActiveCall, CallKind, Observer};
use sqlx::MySql;
use tokio::sync::mpsc;
use crate::{
CallContext, DbRow, MAX_QUERY_ROWS, MAX_RESULT_BYTES, Result, Statement, WriteResult,
cleanup::RollbackJob,
error::{
map_operation_error, result_limit_exceeded, transaction_commit_failed,
transaction_rollback_failed,
},
row::row_payload_bytes,
};
pub type TransactionFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T>> + Send + 'a>>;
pub struct Transaction {
inner: Option<sqlx::Transaction<'static, MySql>>,
observer: Observer,
context: CallContext,
boundary: Option<ActiveCall>,
cleanup: Option<mpsc::UnboundedSender<RollbackJob>>,
}
impl Transaction {
pub(crate) fn new(
inner: sqlx::Transaction<'static, MySql>,
observer: Observer,
boundary: ActiveCall,
cleanup: mpsc::UnboundedSender<RollbackJob>,
) -> Self {
let context = boundary.context().clone();
Self {
inner: Some(inner),
observer,
context,
boundary: Some(boundary),
cleanup: Some(cleanup),
}
}
fn inner(&mut self) -> &mut sqlx::Transaction<'static, MySql> {
self.inner
.as_mut()
.expect("active transaction has an inner connection")
}
pub async fn query_all(&mut self, statement: Statement) -> Result<Vec<DbRow>> {
statement.validate()?;
let operation = statement.operation().to_owned();
let call = self.observer.start_child_call(
&self.context,
CallKind::Database,
"database",
"database",
operation,
);
let result = async {
let mut stream = statement.query().fetch(&mut **self.inner());
let mut rows = Vec::new();
let mut result_bytes = 0_usize;
while let Some(row) = stream.try_next().await.map_err(map_operation_error)? {
if rows.len() == MAX_QUERY_ROWS {
return Err(result_limit_exceeded());
}
result_bytes = result_bytes.saturating_add(row_payload_bytes(&row)?);
if result_bytes > MAX_RESULT_BYTES {
return Err(result_limit_exceeded());
}
rows.push(DbRow(row));
}
Ok(rows)
}
.await;
finish_call(call, &result);
result
}
pub async fn query_optional(&mut self, statement: Statement) -> Result<Option<DbRow>> {
statement.validate()?;
let operation = statement.operation().to_owned();
let call = self.observer.start_child_call(
&self.context,
CallKind::Database,
"database",
"database",
operation,
);
let result = statement
.query()
.fetch_optional(&mut **self.inner())
.await
.map_err(map_operation_error)
.and_then(|row| {
row.map(|row| {
row_payload_bytes(&row)?;
Ok(DbRow(row))
})
.transpose()
});
finish_call(call, &result);
result
}
pub async fn write(&mut self, statement: Statement) -> Result<WriteResult> {
statement.validate()?;
let operation = statement.operation().to_owned();
let call = self.observer.start_child_call(
&self.context,
CallKind::Database,
"database",
"database",
operation,
);
let result = statement
.query()
.execute(&mut **self.inner())
.await
.map(|result| WriteResult::new(result.rows_affected(), result.last_insert_id()))
.map_err(map_operation_error);
finish_call(call, &result);
result
}
pub(crate) async fn commit(mut self) -> Result<()> {
let inner = self
.inner
.take()
.expect("active transaction has an inner connection");
let boundary = self
.boundary
.take()
.expect("active transaction has a boundary");
let _cleanup = self.cleanup.take();
let commit = self.observer.start_child_call(
boundary.context(),
CallKind::Transaction,
"database",
"database",
"commit",
);
match inner.commit().await {
Ok(()) => {
commit.succeed();
boundary.succeed();
Ok(())
}
Err(_) => {
let error = transaction_commit_failed();
commit.fail(&error);
boundary.fail(&error);
Err(error)
}
}
}
pub(crate) async fn rollback(mut self, work_error: crate::SaddleError) -> crate::SaddleError {
let inner = self
.inner
.take()
.expect("active transaction has an inner connection");
let boundary = self
.boundary
.take()
.expect("active transaction has a boundary");
let _cleanup = self.cleanup.take();
let rollback = self.observer.start_child_call(
boundary.context(),
CallKind::Transaction,
"database",
"database",
"rollback",
);
match inner.rollback().await {
Ok(()) => {
rollback.succeed();
boundary.fail(&work_error);
work_error
}
Err(_) => {
let error = transaction_rollback_failed();
rollback.fail(&error);
boundary.fail(&error);
error
}
}
}
}
impl Drop for Transaction {
fn drop(&mut self) {
let (Some(inner), Some(boundary), Some(cleanup)) =
(self.inner.take(), self.boundary.take(), self.cleanup.take())
else {
return;
};
let rollback = self.observer.start_child_call(
boundary.context(),
CallKind::Transaction,
"database",
"database",
"rollback",
);
let job = RollbackJob {
inner,
boundary,
rollback,
};
if let Err(error) = cleanup.send(job) {
error.0.fail_without_worker();
}
}
}
fn finish_call<T>(call: saddle_observability::ActiveCall, result: &Result<T>) {
match result {
Ok(_) => call.succeed(),
Err(error) => call.fail(error),
}
}