use std::future::Future;
use std::sync::atomic::Ordering;
use cratestack_core::{CratestackError, TransactionIsolation};
use super::BoundTx;
use crate::descriptor::SqlxRuntime;
use crate::error::cratestack_error_from_sqlx;
use crate::sqlx;
const JOIN_SAVEPOINT: &str = "cratestack_isolated_join";
pub(crate) async fn join_bound<F, Fut, T>(
runtime: &SqlxRuntime,
bound: &BoundTx,
isolation: TransactionIsolation,
body: F,
) -> Result<T, CratestackError>
where
F: FnOnce(SqlxRuntime) -> Fut,
Fut: Future<Output = Result<T, CratestackError>>,
{
if strength(isolation) > strength(bound.isolation) {
return Err(CratestackError::Internal(format!(
"an @isolation(\"{inner}\") procedure was called inside an @isolation(\"{outer}\") \
transaction; a transaction's isolation level cannot be raised after it has begun. \
Declare the calling procedure {inner} or stricter, or call this one outside it",
inner = isolation.as_sql(),
outer = bound.isolation.as_sql(),
)));
}
if bound
.joined
.compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
.is_err()
{
let error = CratestackError::Internal(
"a nested @isolation call started while another one was still running on the same \
transaction (concurrently, or from inside it); joined calls must run one at a time"
.to_owned(),
);
poison(bound, &error);
return Err(error);
}
let mut guard = JoinGuard {
bound,
finished: false,
};
let mark = bound.audit_mark();
let begun = statement(bound, "SAVEPOINT", JOIN_SAVEPOINT).await;
bound.observe(&begun);
if let Err(error) = begun {
guard.finished = true;
return Err(error);
}
let result = body(runtime.clone())
.await
.map_err(CratestackError::propagate_transaction_abort);
let result = match result {
Ok(value) => match statement(bound, "RELEASE SAVEPOINT", JOIN_SAVEPOINT).await {
Ok(()) => Ok(value),
Err(error) => {
poison(bound, &error);
Err(error)
}
},
Err(error) => {
bound.truncate_audit(mark);
let rolled_back = match statement(bound, "ROLLBACK TO SAVEPOINT", JOIN_SAVEPOINT).await
{
Ok(()) => statement(bound, "RELEASE SAVEPOINT", JOIN_SAVEPOINT).await,
Err(rollback_error) => Err(rollback_error),
};
if let Err(rollback_error) = rolled_back {
poison(bound, &rollback_error);
}
Err(error)
}
};
guard.finished = true;
bound.observe(&result);
result
}
struct JoinGuard<'a> {
bound: &'a BoundTx,
finished: bool,
}
impl Drop for JoinGuard<'_> {
fn drop(&mut self) {
if !self.finished {
poison(
self.bound,
&CratestackError::Internal(
"a nested @isolation call was dropped before it finished".to_owned(),
),
);
}
self.bound.joined.store(false, Ordering::SeqCst);
}
}
fn strength(level: TransactionIsolation) -> u8 {
match level {
TransactionIsolation::ReadCommitted => 0,
TransactionIsolation::RepeatableRead => 1,
TransactionIsolation::Serializable => 2,
}
}
async fn statement(bound: &BoundTx, verb: &str, name: &str) -> Result<(), CratestackError> {
let mut guard = bound.lock()?;
let tx = guard.tx()?;
sqlx::query(sqlx::AssertSqlSafe(format!("{verb} {name}")))
.execute(&mut ***tx)
.await
.map(|_| ())
.map_err(cratestack_error_from_sqlx)
}
fn poison(bound: &BoundTx, cause: &CratestackError) {
tracing::warn!(
target: "cratestack",
cratestack_error = cause.code(),
"a nested @isolation call could not close its savepoint; the attempt will not be \
committed",
);
bound.poison(CratestackError::Internal(format!(
"a nested @isolation call could not close its savepoint ({}); the procedure's work was \
rolled back",
cause.code(),
)));
}
impl BoundTx {
fn audit_mark(&self) -> usize {
self.audit_events
.lock()
.map(|events| events.len())
.unwrap_or(0)
}
fn truncate_audit(&self, mark: usize) {
if let Ok(mut events) = self.audit_events.lock() {
events.truncate(mark);
}
}
}