use crate::pool::Connection;
use std::sync::Arc;
use std::time::{Duration, Instant};
use sz_orm_model::error::{TransactionState, TxError};
use tokio::sync::Mutex;
#[derive(Debug, Clone, PartialEq, Default)]
pub enum IsolationLevel {
ReadUncommitted,
ReadCommitted,
#[default]
RepeatableRead,
Serializable,
Snapshot,
}
impl std::fmt::Display for IsolationLevel {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
IsolationLevel::ReadUncommitted => write!(f, "READ UNCOMMITTED"),
IsolationLevel::ReadCommitted => write!(f, "READ COMMITTED"),
IsolationLevel::RepeatableRead => write!(f, "REPEATABLE READ"),
IsolationLevel::Serializable => write!(f, "SERIALIZABLE"),
IsolationLevel::Snapshot => write!(f, "SNAPSHOT"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum AutoCommit {
#[default]
On,
Off,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum PropagationBehavior {
#[default]
Required,
Mandatory,
Never,
Supports,
RequiresNew,
Nested,
}
pub struct TransactOptions {
pub isolation_level: Option<IsolationLevel>,
pub read_only: bool,
pub timeout: Option<Duration>,
pub max_nesting_depth: u32,
pub propagation: PropagationBehavior,
}
pub const DEFAULT_MAX_NESTING_DEPTH: u32 = 8;
impl Default for TransactOptions {
fn default() -> Self {
Self {
isolation_level: None,
read_only: false,
timeout: None,
max_nesting_depth: DEFAULT_MAX_NESTING_DEPTH,
propagation: PropagationBehavior::default(),
}
}
}
impl TransactOptions {
pub fn with_isolation(mut self, level: IsolationLevel) -> Self {
self.isolation_level = Some(level);
self
}
pub fn read_only(mut self) -> Self {
self.read_only = true;
self
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout);
self
}
pub fn with_max_nesting_depth(mut self, max_depth: u32) -> Self {
self.max_nesting_depth = max_depth;
self
}
pub fn with_propagation(mut self, propagation: PropagationBehavior) -> Self {
self.propagation = propagation;
self
}
}
fn validate_savepoint_name(name: &str) -> Result<(), TxError> {
if name.is_empty() {
return Err(TxError::InvalidSavepointName(name.to_string()));
}
if !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
return Err(TxError::InvalidSavepointName(name.to_string()));
}
if name.starts_with(|c: char| c.is_ascii_digit()) {
return Err(TxError::InvalidSavepointName(name.to_string()));
}
Ok(())
}
pub struct Transaction {
conn: Arc<Mutex<Option<Box<dyn Connection>>>>,
state: TransactionState,
options: TransactOptions,
savepoint_counter: u32,
deadline: Option<Instant>,
}
impl Transaction {
pub fn new(conn: Box<dyn Connection>, options: TransactOptions) -> Self {
let deadline = options.timeout.map(|t| Instant::now() + t);
Self {
conn: Arc::new(Mutex::new(Some(conn))),
state: TransactionState::Active,
options,
savepoint_counter: 0,
deadline,
}
}
pub fn state(&self) -> TransactionState {
self.state
}
pub fn is_active(&self) -> bool {
self.state == TransactionState::Active
}
pub async fn commit(&mut self) -> Result<(), TxError> {
if self.state != TransactionState::Active {
return Err(TxError::NotActive(self.state));
}
if let Some(deadline) = self.deadline {
if Instant::now() > deadline {
self.rollback().await.ok();
return Err(TxError::CommitFailed("Transaction timeout".to_string()));
}
}
let mut conn_guard = self.conn.lock().await;
let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
conn.commit()
.await
.map_err(|e| TxError::CommitFailed(e.to_string()))?;
self.state = TransactionState::Committed;
Ok(())
}
pub async fn rollback(&mut self) -> Result<(), TxError> {
if self.state != TransactionState::Active {
return Err(TxError::NotActive(self.state));
}
let mut conn_guard = self.conn.lock().await;
let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
conn.rollback()
.await
.map_err(|e| TxError::RollbackFailed(e.to_string()))?;
self.state = TransactionState::RolledBack;
Ok(())
}
pub async fn execute(&mut self, sql: &str) -> Result<u64, TxError> {
if self.state != TransactionState::Active {
return Err(TxError::NotActive(self.state));
}
let mut conn_guard = self.conn.lock().await;
let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
let result = conn
.execute(sql)
.await
.map_err(|e| TxError::CommitFailed(e.to_string()))?;
Ok(result)
}
pub async fn query(
&mut self,
sql: &str,
) -> Result<Vec<std::collections::HashMap<String, sz_orm_model::Value>>, TxError> {
if self.state != TransactionState::Active {
return Err(TxError::NotActive(self.state));
}
let mut conn_guard = self.conn.lock().await;
let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
let result = conn
.query(sql)
.await
.map_err(|e| TxError::CommitFailed(e.to_string()))?;
Ok(result)
}
pub async fn savepoint(&mut self) -> Result<String, TxError> {
if self.state != TransactionState::Active {
return Err(TxError::NotActive(self.state));
}
let next_depth = self.savepoint_counter + 1;
if next_depth > self.options.max_nesting_depth {
return Err(TxError::MaxNestingDepthExceeded {
current_depth: next_depth,
max_depth: self.options.max_nesting_depth,
});
}
self.savepoint_counter += 1;
let name = format!("sp_{}", self.savepoint_counter);
validate_savepoint_name(&name)?;
let sql = format!("SAVEPOINT {}", name);
let mut conn_guard = self.conn.lock().await;
let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
conn.execute(&sql)
.await
.map_err(|e| TxError::SavepointError(e.to_string()))?;
Ok(name)
}
pub async fn rollback_to_savepoint(&mut self, name: &str) -> Result<(), TxError> {
if self.state != TransactionState::Active {
return Err(TxError::NotActive(self.state));
}
validate_savepoint_name(name)?;
let sql = format!("ROLLBACK TO SAVEPOINT {}", name);
let mut conn_guard = self.conn.lock().await;
let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
conn.execute(&sql)
.await
.map_err(|e| TxError::SavepointError(e.to_string()))?;
Ok(())
}
pub async fn release_savepoint(&mut self, name: &str) -> Result<(), TxError> {
if self.state != TransactionState::Active {
return Err(TxError::NotActive(self.state));
}
validate_savepoint_name(name)?;
let sql = format!("RELEASE SAVEPOINT {}", name);
let mut conn_guard = self.conn.lock().await;
let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
conn.execute(&sql)
.await
.map_err(|e| TxError::SavepointError(e.to_string()))?;
Ok(())
}
pub async fn take_connection(&mut self) -> Result<Box<dyn Connection>, TxError> {
if self.state == TransactionState::Active {
return Err(TxError::NotActive(self.state));
}
let mut conn_guard = self.conn.lock().await;
conn_guard.take().ok_or(TxError::ConnectionTaken)
}
pub fn options(&self) -> &TransactOptions {
&self.options
}
}
pub fn is_deadlock_error(err_msg: &str) -> bool {
let lower = err_msg.to_lowercase();
if lower.contains("deadlock found when trying to get lock") {
return true;
}
if lower.contains("error 1213") || lower.contains("(1213)") {
return true;
}
if lower.contains("deadlock detected") || lower.contains("40p01") {
return true;
}
if lower.contains("database is locked") || lower.contains("database table is locked") {
return true;
}
if lower.contains("ora-00060") {
return true;
}
if lower.contains("transaction (process id") && lower.contains("was deadlocked") {
return true;
}
if lower.contains("error 1205") || lower.contains("(1205)") {
return true;
}
false
}
pub async fn retry_on_deadlock<F, Fut, T>(
max_attempts: u32,
backoff: Duration,
operation: F,
) -> Result<T, TxError>
where
F: Fn(u32) -> Fut,
Fut: std::future::Future<Output = Result<T, TxError>>,
{
let mut last_err: Option<TxError> = None;
for attempt in 1..=max_attempts {
match operation(attempt).await {
Ok(v) => return Ok(v),
Err(e) => {
let err_msg = format!("{}", e);
if is_deadlock_error(&err_msg) && attempt < max_attempts {
tokio::time::sleep(backoff).await;
last_err = Some(TxError::DeadlockDetected {
attempt,
max_attempts,
});
continue;
}
return Err(e);
}
}
}
Err(last_err.unwrap_or(TxError::DeadlockDetected {
attempt: max_attempts,
max_attempts,
}))
}
impl Drop for Transaction {
fn drop(&mut self) {
if self.state == TransactionState::Active {
let conn = self.conn.clone();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
let mut conn_guard = conn.lock().await;
if let Some(ref mut conn) = *conn_guard {
let _ = conn.rollback().await;
}
});
}
self.state = TransactionState::RolledBack;
}
}
}
pub struct TransactionManager {
transactions: Arc<Mutex<std::collections::HashMap<String, Transaction>>>,
}
impl TransactionManager {
pub fn new() -> Self {
Self {
transactions: Arc::new(Mutex::new(std::collections::HashMap::new())),
}
}
pub async fn begin(
&self,
id: String,
conn: Box<dyn Connection>,
options: TransactOptions,
) -> Result<(), TxError> {
let mut conn = conn;
conn.begin_transaction()
.await
.map_err(|e| TxError::CommitFailed(e.to_string()))?;
let tx = Transaction::new(conn, options);
let mut txs = self.transactions.lock().await;
txs.insert(id, tx);
Ok(())
}
pub async fn commit(&self, id: &str) -> Result<(), TxError> {
let mut txs = self.transactions.lock().await;
let tx = txs
.get_mut(id)
.ok_or_else(|| TxError::SavepointError(format!("Transaction {} not found", id)))?;
tx.commit().await
}
pub async fn rollback(&self, id: &str) -> Result<(), TxError> {
let mut txs = self.transactions.lock().await;
let tx = txs
.get_mut(id)
.ok_or_else(|| TxError::SavepointError(format!("Transaction {} not found", id)))?;
tx.rollback().await
}
pub async fn state(&self, id: &str) -> Option<TransactionState> {
let txs = self.transactions.lock().await;
txs.get(id).map(|tx| tx.state())
}
pub async fn list(&self) -> Vec<String> {
let txs = self.transactions.lock().await;
txs.keys().cloned().collect()
}
pub async fn remove(&self, id: &str) -> Option<Transaction> {
let mut txs = self.transactions.lock().await;
txs.remove(id)
}
}
impl Default for TransactionManager {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::future::Future;
use std::pin::Pin;
struct MockConnection {
begin_called: bool,
commit_called: bool,
rollback_called: bool,
executed_sql: Vec<String>,
}
impl MockConnection {
fn new() -> Self {
Self {
begin_called: false,
commit_called: false,
rollback_called: false,
executed_sql: Vec::new(),
}
}
}
impl Connection for MockConnection {
fn execute<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, sz_orm_model::DbError>> + Send + 'a>> {
Box::pin(async move {
self.executed_sql.push(sql.to_string());
Ok(1)
})
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<
Box<
dyn Future<
Output = Result<
Vec<std::collections::HashMap<String, sz_orm_model::Value>>,
sz_orm_model::DbError,
>,
> + Send
+ 'a,
>,
> {
Box::pin(async move { Ok(vec![]) })
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>> {
Box::pin(async move {
self.begin_called = true;
Ok(())
})
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>> {
Box::pin(async move {
self.commit_called = true;
Ok(())
})
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>> {
Box::pin(async move {
self.rollback_called = true;
Ok(())
})
}
fn is_connected(&self) -> bool {
true
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move { true })
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>> {
Box::pin(async move { Ok(()) })
}
}
#[test]
fn test_isolation_level_display() {
assert_eq!(IsolationLevel::ReadCommitted.to_string(), "READ COMMITTED");
assert_eq!(IsolationLevel::Serializable.to_string(), "SERIALIZABLE");
}
#[test]
fn test_transaction_state_default() {
let opts = TransactOptions::default();
assert!(opts.isolation_level.is_none());
assert!(!opts.read_only);
}
#[test]
fn test_transact_options_builder() {
let opts = TransactOptions {
isolation_level: Some(IsolationLevel::Serializable),
read_only: true,
timeout: Some(Duration::from_secs(30)),
max_nesting_depth: DEFAULT_MAX_NESTING_DEPTH,
propagation: PropagationBehavior::default(),
};
assert_eq!(opts.isolation_level, Some(IsolationLevel::Serializable));
assert!(opts.read_only);
assert_eq!(opts.timeout, Some(Duration::from_secs(30)));
}
#[test]
fn test_auto_commit_default() {
assert_eq!(AutoCommit::default(), AutoCommit::On);
}
#[test]
fn test_transaction_state() {
assert_eq!(TransactionState::Active, TransactionState::Active);
assert_ne!(TransactionState::Active, TransactionState::Committed);
}
#[test]
fn test_transact_options_chaining() {
let opts = TransactOptions::default()
.with_isolation(IsolationLevel::Serializable)
.read_only()
.with_timeout(Duration::from_secs(60));
assert_eq!(opts.isolation_level, Some(IsolationLevel::Serializable));
assert!(opts.read_only);
assert_eq!(opts.timeout, Some(Duration::from_secs(60)));
}
#[tokio::test]
async fn test_transaction_commit() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
assert!(tx.is_active());
let result = tx.execute("INSERT INTO users VALUES (1)").await;
assert!(result.is_ok());
tx.commit().await?;
assert_eq!(tx.state(), TransactionState::Committed);
let result = tx.commit().await;
assert!(result.is_err());
match result {
Err(TxError::NotActive(state)) => {
assert_eq!(state, TransactionState::Committed);
}
_ => panic!("Expected NotActive error"),
}
Ok(())
}
#[tokio::test]
async fn test_transaction_rollback() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
tx.rollback().await?;
assert_eq!(tx.state(), TransactionState::RolledBack);
let result = tx.rollback().await;
assert!(result.is_err());
match result {
Err(TxError::NotActive(state)) => {
assert_eq!(state, TransactionState::RolledBack);
}
_ => panic!("Expected NotActive error"),
}
Ok(())
}
#[tokio::test]
async fn test_transaction_execute_after_commit() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
tx.commit().await?;
let result = tx.execute("SELECT 1").await;
assert!(result.is_err());
match result {
Err(TxError::NotActive(_)) => {}
_ => panic!("Expected NotActive error"),
}
Ok(())
}
#[tokio::test]
async fn test_transaction_query_after_commit_returns_not_active() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
tx.commit().await?;
let result = tx.query("SELECT 1").await;
assert!(result.is_err());
match result {
Err(TxError::NotActive(_)) => {}
_ => panic!("Expected NotActive error"),
}
Ok(())
}
#[tokio::test]
async fn test_transaction_savepoint() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
let sp1 = tx.savepoint().await?;
assert_eq!(sp1, "sp_1");
let sp2 = tx.savepoint().await?;
assert_eq!(sp2, "sp_2");
tx.rollback_to_savepoint(&sp1).await?;
tx.release_savepoint(&sp2).await?;
Ok(())
}
#[tokio::test]
async fn test_transaction_savepoint_name_validation() {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
let result = tx.rollback_to_savepoint("sp'; DROP TABLE--").await;
assert!(result.is_err());
match result {
Err(TxError::InvalidSavepointName(_)) => {}
_ => panic!("Expected InvalidSavepointName error"),
}
let result = tx.release_savepoint("1sp").await;
assert!(result.is_err());
match result {
Err(TxError::InvalidSavepointName(_)) => {}
_ => panic!("Expected InvalidSavepointName error"),
}
let result = tx.rollback_to_savepoint("").await;
assert!(result.is_err());
match result {
Err(TxError::InvalidSavepointName(_)) => {}
_ => panic!("Expected InvalidSavepointName error"),
}
let result = tx.rollback_to_savepoint("sp_test_1").await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_transaction_take_connection() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
let result = tx.take_connection().await;
assert!(result.is_err());
match result {
Err(TxError::NotActive(_)) => {}
_ => panic!("Expected NotActive error"),
}
tx.commit().await?;
let conn = tx.take_connection().await;
assert!(conn.is_ok());
let result = tx.take_connection().await;
assert!(result.is_err());
match result {
Err(TxError::ConnectionTaken) => {}
_ => panic!("Expected ConnectionTaken error"),
}
Ok(())
}
#[tokio::test]
async fn test_transaction_manager() -> Result<(), TxError> {
let mgr = TransactionManager::new();
let conn = Box::new(MockConnection::new());
mgr.begin("tx1".to_string(), conn, TransactOptions::default())
.await?;
let state = mgr.state("tx1").await;
assert_eq!(state, Some(TransactionState::Active));
mgr.commit("tx1").await?;
let state = mgr.state("tx1").await;
assert_eq!(state, Some(TransactionState::Committed));
let list = mgr.list().await;
assert!(list.contains(&"tx1".to_string()));
Ok(())
}
#[tokio::test]
async fn test_transaction_manager_rollback() -> Result<(), TxError> {
let mgr = TransactionManager::new();
let conn = Box::new(MockConnection::new());
mgr.begin("tx2".to_string(), conn, TransactOptions::default())
.await?;
mgr.rollback("tx2").await?;
let state = mgr.state("tx2").await;
assert_eq!(state, Some(TransactionState::RolledBack));
Ok(())
}
#[tokio::test]
async fn test_transaction_manager_not_found() {
let mgr = TransactionManager::new();
let result = mgr.commit("nonexistent").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_transaction_manager_remove() -> Result<(), TxError> {
let mgr = TransactionManager::new();
let conn = Box::new(MockConnection::new());
mgr.begin("tx3".to_string(), conn, TransactOptions::default())
.await?;
let removed = mgr.remove("tx3").await;
assert!(removed.is_some());
let state = mgr.state("tx3").await;
assert_eq!(state, None);
Ok(())
}
#[tokio::test]
async fn test_transaction_drop_rolls_back_when_active() {
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc as StdArc;
struct TrackingConnection {
rollback_called: StdArc<AtomicBool>,
}
impl Connection for TrackingConnection {
fn execute<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, sz_orm_model::DbError>> + Send + 'a>>
{
Box::pin(async { Ok(1) })
}
fn query<'a>(
&'a mut self,
_sql: &'a str,
) -> Pin<
Box<
dyn Future<
Output = Result<
Vec<std::collections::HashMap<String, sz_orm_model::Value>>,
sz_orm_model::DbError,
>,
> + Send
+ 'a,
>,
> {
Box::pin(async { Ok(vec![]) })
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>>
{
Box::pin(async { Ok(()) })
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>>
{
Box::pin(async { Ok(()) })
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>>
{
let flag = self.rollback_called.clone();
Box::pin(async move {
flag.store(true, Ordering::SeqCst);
Ok(())
})
}
fn is_connected(&self) -> bool {
true
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async { true })
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>>
{
Box::pin(async { Ok(()) })
}
}
let rollback_flag = StdArc::new(AtomicBool::new(false));
let conn = Box::new(TrackingConnection {
rollback_called: rollback_flag.clone(),
});
{
let _tx = Transaction::new(conn, TransactOptions::default());
}
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(
rollback_flag.load(Ordering::SeqCst),
"Drop should have triggered rollback"
);
}
#[test]
fn test_h8_default_max_nesting_depth_is_8() {
let opts = TransactOptions::default();
assert_eq!(opts.max_nesting_depth, DEFAULT_MAX_NESTING_DEPTH);
assert_eq!(opts.max_nesting_depth, 8);
}
#[test]
fn test_h8_with_max_nesting_depth_builder() {
let opts = TransactOptions::default().with_max_nesting_depth(3);
assert_eq!(opts.max_nesting_depth, 3);
}
#[tokio::test]
async fn test_h8_savepoint_within_default_depth_succeeds() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
for i in 1..=8 {
let sp = tx.savepoint().await?;
assert_eq!(sp, format!("sp_{}", i));
}
Ok(())
}
#[tokio::test]
async fn test_h8_savepoint_exceeding_default_depth_fails() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default());
for _ in 0..8 {
tx.savepoint().await?;
}
let result = tx.savepoint().await;
assert!(result.is_err());
match result {
Err(TxError::MaxNestingDepthExceeded {
current_depth,
max_depth,
}) => {
assert_eq!(current_depth, 9);
assert_eq!(max_depth, 8);
}
_ => panic!("Expected MaxNestingDepthExceeded error"),
}
Ok(())
}
#[tokio::test]
async fn test_h8_savepoint_with_custom_depth_3() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(3));
for i in 1..=3 {
let sp = tx.savepoint().await?;
assert_eq!(sp, format!("sp_{}", i));
}
let result = tx.savepoint().await;
assert!(result.is_err());
match result {
Err(TxError::MaxNestingDepthExceeded {
current_depth,
max_depth,
}) => {
assert_eq!(current_depth, 4);
assert_eq!(max_depth, 3);
}
_ => panic!("Expected MaxNestingDepthExceeded error"),
}
Ok(())
}
#[tokio::test]
async fn test_h8_savepoint_depth_zero_disables_nesting() {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(0));
let result = tx.savepoint().await;
assert!(result.is_err());
match result {
Err(TxError::MaxNestingDepthExceeded {
current_depth,
max_depth,
}) => {
assert_eq!(current_depth, 1);
assert_eq!(max_depth, 0);
}
_ => panic!("Expected MaxNestingDepthExceeded error"),
}
}
#[tokio::test]
async fn test_h8_savepoint_after_rollback_to_still_respects_depth() -> Result<(), TxError> {
let conn = Box::new(MockConnection::new());
let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(2));
let sp1 = tx.savepoint().await?;
let sp2 = tx.savepoint().await?;
tx.rollback_to_savepoint(&sp1).await?;
tx.release_savepoint(&sp2).await?;
let result = tx.savepoint().await;
assert!(result.is_err());
match result {
Err(TxError::MaxNestingDepthExceeded {
current_depth,
max_depth,
}) => {
assert_eq!(current_depth, 3);
assert_eq!(max_depth, 2);
}
_ => panic!("Expected MaxNestingDepthExceeded error"),
}
Ok(())
}
#[tokio::test]
async fn test_h8_max_nesting_depth_error_display() {
let err = TxError::MaxNestingDepthExceeded {
current_depth: 10,
max_depth: 8,
};
let msg = format!("{}", err);
assert!(msg.contains("10"));
assert!(msg.contains("8"));
assert!(msg.contains("exceeds"));
}
#[test]
fn test_m8_is_deadlock_error_mysql() {
assert!(is_deadlock_error(
"Deadlock found when trying to get lock; try restarting transaction"
));
assert!(is_deadlock_error("Error 1213: Deadlock found"));
assert!(is_deadlock_error("MySQL error (1213)"));
}
#[test]
fn test_m8_is_deadlock_error_postgresql() {
assert!(is_deadlock_error("deadlock detected"));
assert!(is_deadlock_error("ERROR: deadlock detected (40P01)"));
assert!(is_deadlock_error("SQLSTATE 40P01"));
}
#[test]
fn test_m8_is_deadlock_error_sqlite() {
assert!(is_deadlock_error("database is locked"));
assert!(is_deadlock_error("database table is locked"));
}
#[test]
fn test_m8_is_deadlock_error_oracle() {
assert!(is_deadlock_error(
"ORA-00060: deadlock detected while waiting for resource"
));
}
#[test]
fn test_m8_is_deadlock_error_sql_server() {
assert!(is_deadlock_error(
"Transaction (Process ID 52) was deadlocked on lock resources"
));
assert!(is_deadlock_error("Error 1205: Transaction was deadlocked"));
}
#[test]
fn test_m8_is_deadlock_error_non_deadlock() {
assert!(!is_deadlock_error("connection refused"));
assert!(!is_deadlock_error("syntax error near SELECT"));
assert!(!is_deadlock_error("permission denied for table users"));
assert!(!is_deadlock_error(""));
}
#[tokio::test]
async fn test_m8_retry_on_deadlock_succeeds_first_attempt() -> Result<(), TxError> {
use std::sync::atomic::{AtomicU32, Ordering};
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let result: Result<u32, TxError> =
retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
let c = counter_clone.clone();
async move {
c.fetch_add(1, Ordering::SeqCst);
Ok(42u32)
}
})
.await;
assert_eq!(result?, 42);
assert_eq!(counter.load(Ordering::SeqCst), 1);
Ok(())
}
#[tokio::test]
async fn test_m8_retry_on_deadlock_retries_on_deadlock_error() -> Result<(), TxError> {
use std::sync::atomic::{AtomicU32, Ordering};
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let result: Result<u32, TxError> =
retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
let c = counter_clone.clone();
async move {
let n = c.fetch_add(1, Ordering::SeqCst);
if n < 2 {
Err(TxError::CommitFailed(
"Deadlock found when trying to get lock".to_string(),
))
} else {
Ok(42u32)
}
}
})
.await;
assert_eq!(result?, 42);
assert_eq!(counter.load(Ordering::SeqCst), 3);
Ok(())
}
#[tokio::test]
async fn test_m8_retry_on_deadlock_returns_error_after_max_attempts() {
use std::sync::atomic::{AtomicU32, Ordering};
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let result: Result<u32, TxError> =
retry_on_deadlock(2, Duration::from_millis(1), |_attempt| {
let c = counter_clone.clone();
async move {
c.fetch_add(1, Ordering::SeqCst);
Err(TxError::CommitFailed(
"Deadlock found when trying to get lock".to_string(),
))
}
})
.await;
assert!(result.is_err());
assert_eq!(counter.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_m8_retry_on_deadlock_does_not_retry_non_deadlock_errors() {
use std::sync::atomic::{AtomicU32, Ordering};
let counter = Arc::new(AtomicU32::new(0));
let counter_clone = counter.clone();
let result: Result<u32, TxError> =
retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
let c = counter_clone.clone();
async move {
c.fetch_add(1, Ordering::SeqCst);
Err(TxError::CommitFailed("syntax error".to_string()))
}
})
.await;
assert!(result.is_err());
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[test]
fn test_m8_deadlock_error_display() {
let err = TxError::DeadlockDetected {
attempt: 2,
max_attempts: 3,
};
let msg = format!("{}", err);
assert!(msg.contains("2"));
assert!(msg.contains("3"));
assert!(msg.contains("Deadlock"));
}
}