1use crate::pool::Connection;
6use std::sync::Arc;
7use std::time::{Duration, Instant};
8use sz_orm_model::error::{TransactionState, TxError};
9use tokio::sync::Mutex;
10
11#[derive(Debug, Clone, PartialEq, Default)]
15pub enum IsolationLevel {
16 ReadUncommitted,
17 ReadCommitted,
18 #[default]
19 RepeatableRead,
20 Serializable,
21 Snapshot,
22}
23
24impl std::fmt::Display for IsolationLevel {
25 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26 match self {
27 IsolationLevel::ReadUncommitted => write!(f, "READ UNCOMMITTED"),
28 IsolationLevel::ReadCommitted => write!(f, "READ COMMITTED"),
29 IsolationLevel::RepeatableRead => write!(f, "REPEATABLE READ"),
30 IsolationLevel::Serializable => write!(f, "SERIALIZABLE"),
31 IsolationLevel::Snapshot => write!(f, "SNAPSHOT"),
32 }
33 }
34}
35
36#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
37pub enum AutoCommit {
38 #[default]
39 On,
40 Off,
41}
42
43#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
45pub enum PropagationBehavior {
46 #[default]
48 Required,
49 Mandatory,
51 Never,
53 Supports,
55 RequiresNew,
57 Nested,
59}
60
61pub struct TransactOptions {
62 pub isolation_level: Option<IsolationLevel>,
63 pub read_only: bool,
64 pub timeout: Option<Duration>,
65 pub max_nesting_depth: u32,
70 pub propagation: PropagationBehavior,
72}
73
74pub const DEFAULT_MAX_NESTING_DEPTH: u32 = 8;
76
77impl Default for TransactOptions {
78 fn default() -> Self {
79 Self {
80 isolation_level: None,
81 read_only: false,
82 timeout: None,
83 max_nesting_depth: DEFAULT_MAX_NESTING_DEPTH,
84 propagation: PropagationBehavior::default(),
85 }
86 }
87}
88
89impl TransactOptions {
90 pub fn with_isolation(mut self, level: IsolationLevel) -> Self {
91 self.isolation_level = Some(level);
92 self
93 }
94
95 pub fn read_only(mut self) -> Self {
96 self.read_only = true;
97 self
98 }
99
100 pub fn with_timeout(mut self, timeout: Duration) -> Self {
101 self.timeout = Some(timeout);
102 self
103 }
104
105 pub fn with_max_nesting_depth(mut self, max_depth: u32) -> Self {
107 self.max_nesting_depth = max_depth;
108 self
109 }
110
111 pub fn with_propagation(mut self, propagation: PropagationBehavior) -> Self {
113 self.propagation = propagation;
114 self
115 }
116}
117
118fn validate_savepoint_name(name: &str) -> Result<(), TxError> {
125 if name.is_empty() {
126 return Err(TxError::InvalidSavepointName(name.to_string()));
127 }
128 if !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
129 return Err(TxError::InvalidSavepointName(name.to_string()));
130 }
131 if name.starts_with(|c: char| c.is_ascii_digit()) {
132 return Err(TxError::InvalidSavepointName(name.to_string()));
133 }
134 Ok(())
135}
136
137pub struct Transaction {
144 conn: Arc<Mutex<Option<Box<dyn Connection>>>>,
145 state: TransactionState,
146 options: TransactOptions,
147 savepoint_counter: u32,
148 deadline: Option<Instant>,
150}
151
152impl Transaction {
153 pub fn new(conn: Box<dyn Connection>, options: TransactOptions) -> Self {
171 let deadline = options.timeout.map(|t| Instant::now() + t);
173 Self {
174 conn: Arc::new(Mutex::new(Some(conn))),
175 state: TransactionState::Active,
176 options,
177 savepoint_counter: 0,
178 deadline,
179 }
180 }
181
182 pub fn state(&self) -> TransactionState {
184 self.state
185 }
186
187 pub fn is_active(&self) -> bool {
189 self.state == TransactionState::Active
190 }
191
192 pub async fn commit(&mut self) -> Result<(), TxError> {
211 if self.state != TransactionState::Active {
212 return Err(TxError::NotActive(self.state));
213 }
214 if let Some(deadline) = self.deadline {
216 if Instant::now() > deadline {
217 self.rollback().await.ok();
218 return Err(TxError::CommitFailed("Transaction timeout".to_string()));
219 }
220 }
221 let mut conn_guard = self.conn.lock().await;
222 let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
223 conn.commit()
224 .await
225 .map_err(|e| TxError::CommitFailed(e.to_string()))?;
226 self.state = TransactionState::Committed;
227 Ok(())
228 }
229
230 pub async fn rollback(&mut self) -> Result<(), TxError> {
232 if self.state != TransactionState::Active {
233 return Err(TxError::NotActive(self.state));
234 }
235 let mut conn_guard = self.conn.lock().await;
236 let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
237 conn.rollback()
238 .await
239 .map_err(|e| TxError::RollbackFailed(e.to_string()))?;
240 self.state = TransactionState::RolledBack;
241 Ok(())
242 }
243
244 pub async fn execute(&mut self, sql: &str) -> Result<u64, TxError> {
246 if self.state != TransactionState::Active {
247 return Err(TxError::NotActive(self.state));
248 }
249 let mut conn_guard = self.conn.lock().await;
250 let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
251 let result = conn
252 .execute(sql)
253 .await
254 .map_err(|e| TxError::CommitFailed(e.to_string()))?;
255 Ok(result)
256 }
257
258 pub async fn query(
260 &mut self,
261 sql: &str,
262 ) -> Result<Vec<std::collections::HashMap<String, sz_orm_model::Value>>, TxError> {
263 if self.state != TransactionState::Active {
264 return Err(TxError::NotActive(self.state));
265 }
266 let mut conn_guard = self.conn.lock().await;
267 let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
268 let result = conn
269 .query(sql)
270 .await
271 .map_err(|e| TxError::CommitFailed(e.to_string()))?;
272 Ok(result)
273 }
274
275 pub async fn savepoint(&mut self) -> Result<String, TxError> {
300 if self.state != TransactionState::Active {
301 return Err(TxError::NotActive(self.state));
302 }
303 let next_depth = self.savepoint_counter + 1;
306 if next_depth > self.options.max_nesting_depth {
307 return Err(TxError::MaxNestingDepthExceeded {
308 current_depth: next_depth,
309 max_depth: self.options.max_nesting_depth,
310 });
311 }
312 self.savepoint_counter += 1;
313 let name = format!("sp_{}", self.savepoint_counter);
314 validate_savepoint_name(&name)?;
316 let sql = format!("SAVEPOINT {}", name);
317 let mut conn_guard = self.conn.lock().await;
318 let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
319 conn.execute(&sql)
320 .await
321 .map_err(|e| TxError::SavepointError(e.to_string()))?;
322 Ok(name)
323 }
324
325 pub async fn rollback_to_savepoint(&mut self, name: &str) -> Result<(), TxError> {
330 if self.state != TransactionState::Active {
331 return Err(TxError::NotActive(self.state));
332 }
333 validate_savepoint_name(name)?;
334 let sql = format!("ROLLBACK TO SAVEPOINT {}", name);
335 let mut conn_guard = self.conn.lock().await;
336 let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
337 conn.execute(&sql)
338 .await
339 .map_err(|e| TxError::SavepointError(e.to_string()))?;
340 Ok(())
341 }
342
343 pub async fn release_savepoint(&mut self, name: &str) -> Result<(), TxError> {
348 if self.state != TransactionState::Active {
349 return Err(TxError::NotActive(self.state));
350 }
351 validate_savepoint_name(name)?;
352 let sql = format!("RELEASE SAVEPOINT {}", name);
353 let mut conn_guard = self.conn.lock().await;
354 let conn = conn_guard.as_mut().ok_or(TxError::ConnectionTaken)?;
355 conn.execute(&sql)
356 .await
357 .map_err(|e| TxError::SavepointError(e.to_string()))?;
358 Ok(())
359 }
360
361 pub async fn take_connection(&mut self) -> Result<Box<dyn Connection>, TxError> {
373 if self.state == TransactionState::Active {
374 return Err(TxError::NotActive(self.state));
375 }
376 let mut conn_guard = self.conn.lock().await;
377 conn_guard.take().ok_or(TxError::ConnectionTaken)
378 }
379
380 pub fn options(&self) -> &TransactOptions {
382 &self.options
383 }
384}
385
386pub fn is_deadlock_error(err_msg: &str) -> bool {
395 let lower = err_msg.to_lowercase();
396 if lower.contains("deadlock found when trying to get lock") {
398 return true;
399 }
400 if lower.contains("error 1213") || lower.contains("(1213)") {
402 return true;
403 }
404 if lower.contains("deadlock detected") || lower.contains("40p01") {
406 return true;
407 }
408 if lower.contains("database is locked") || lower.contains("database table is locked") {
410 return true;
411 }
412 if lower.contains("ora-00060") {
414 return true;
415 }
416 if lower.contains("transaction (process id") && lower.contains("was deadlocked") {
418 return true;
419 }
420 if lower.contains("error 1205") || lower.contains("(1205)") {
421 return true;
422 }
423 false
424}
425
426pub async fn retry_on_deadlock<F, Fut, T>(
451 max_attempts: u32,
452 backoff: Duration,
453 operation: F,
454) -> Result<T, TxError>
455where
456 F: Fn(u32) -> Fut,
457 Fut: std::future::Future<Output = Result<T, TxError>>,
458{
459 let mut last_err: Option<TxError> = None;
460 for attempt in 1..=max_attempts {
461 match operation(attempt).await {
462 Ok(v) => return Ok(v),
463 Err(e) => {
464 let err_msg = format!("{}", e);
466 if is_deadlock_error(&err_msg) && attempt < max_attempts {
467 tokio::time::sleep(backoff).await;
468 last_err = Some(TxError::DeadlockDetected {
469 attempt,
470 max_attempts,
471 });
472 continue;
473 }
474 return Err(e);
476 }
477 }
478 }
479 Err(last_err.unwrap_or(TxError::DeadlockDetected {
481 attempt: max_attempts,
482 max_attempts,
483 }))
484}
485
486impl Drop for Transaction {
487 fn drop(&mut self) {
488 if self.state == TransactionState::Active {
491 let conn = self.conn.clone();
492 if let Ok(handle) = tokio::runtime::Handle::try_current() {
494 handle.spawn(async move {
497 let mut conn_guard = conn.lock().await;
498 if let Some(ref mut conn) = *conn_guard {
499 let _ = conn.rollback().await;
500 }
501 });
502 }
503 self.state = TransactionState::RolledBack;
505 }
506 }
507}
508
509pub struct TransactionManager {
511 transactions: Arc<Mutex<std::collections::HashMap<String, Transaction>>>,
512}
513
514impl TransactionManager {
515 pub fn new() -> Self {
516 Self {
517 transactions: Arc::new(Mutex::new(std::collections::HashMap::new())),
518 }
519 }
520
521 pub async fn begin(
523 &self,
524 id: String,
525 conn: Box<dyn Connection>,
526 options: TransactOptions,
527 ) -> Result<(), TxError> {
528 let mut conn = conn;
529 conn.begin_transaction()
530 .await
531 .map_err(|e| TxError::CommitFailed(e.to_string()))?;
532 let tx = Transaction::new(conn, options);
533 let mut txs = self.transactions.lock().await;
534 txs.insert(id, tx);
535 Ok(())
536 }
537
538 pub async fn commit(&self, id: &str) -> Result<(), TxError> {
540 let mut txs = self.transactions.lock().await;
541 let tx = txs
542 .get_mut(id)
543 .ok_or_else(|| TxError::SavepointError(format!("Transaction {} not found", id)))?;
544 tx.commit().await
545 }
546
547 pub async fn rollback(&self, id: &str) -> Result<(), TxError> {
549 let mut txs = self.transactions.lock().await;
550 let tx = txs
551 .get_mut(id)
552 .ok_or_else(|| TxError::SavepointError(format!("Transaction {} not found", id)))?;
553 tx.rollback().await
554 }
555
556 pub async fn state(&self, id: &str) -> Option<TransactionState> {
558 let txs = self.transactions.lock().await;
559 txs.get(id).map(|tx| tx.state())
560 }
561
562 pub async fn list(&self) -> Vec<String> {
564 let txs = self.transactions.lock().await;
565 txs.keys().cloned().collect()
566 }
567
568 pub async fn remove(&self, id: &str) -> Option<Transaction> {
570 let mut txs = self.transactions.lock().await;
571 txs.remove(id)
572 }
573}
574
575impl Default for TransactionManager {
576 fn default() -> Self {
577 Self::new()
578 }
579}
580
581#[cfg(test)]
582mod tests {
583 use super::*;
584 use std::future::Future;
585 use std::pin::Pin;
586
587 struct MockConnection {
589 begin_called: bool,
590 commit_called: bool,
591 rollback_called: bool,
592 executed_sql: Vec<String>,
593 }
594
595 impl MockConnection {
596 fn new() -> Self {
597 Self {
598 begin_called: false,
599 commit_called: false,
600 rollback_called: false,
601 executed_sql: Vec::new(),
602 }
603 }
604 }
605
606 impl Connection for MockConnection {
607 fn execute<'a>(
608 &'a mut self,
609 sql: &'a str,
610 ) -> Pin<Box<dyn Future<Output = Result<u64, sz_orm_model::DbError>> + Send + 'a>> {
611 Box::pin(async move {
612 self.executed_sql.push(sql.to_string());
613 Ok(1)
614 })
615 }
616
617 fn query<'a>(
618 &'a mut self,
619 _sql: &'a str,
620 ) -> Pin<
621 Box<
622 dyn Future<
623 Output = Result<
624 Vec<std::collections::HashMap<String, sz_orm_model::Value>>,
625 sz_orm_model::DbError,
626 >,
627 > + Send
628 + 'a,
629 >,
630 > {
631 Box::pin(async move { Ok(vec![]) })
632 }
633
634 fn begin_transaction<'a>(
635 &'a mut self,
636 ) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>> {
637 Box::pin(async move {
638 self.begin_called = true;
639 Ok(())
640 })
641 }
642
643 fn commit<'a>(
644 &'a mut self,
645 ) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>> {
646 Box::pin(async move {
647 self.commit_called = true;
648 Ok(())
649 })
650 }
651
652 fn rollback<'a>(
653 &'a mut self,
654 ) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>> {
655 Box::pin(async move {
656 self.rollback_called = true;
657 Ok(())
658 })
659 }
660
661 fn is_connected(&self) -> bool {
662 true
663 }
664
665 fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
666 Box::pin(async move { true })
667 }
668
669 fn close<'a>(
670 &'a mut self,
671 ) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>> {
672 Box::pin(async move { Ok(()) })
673 }
674 }
675
676 #[test]
677 fn test_isolation_level_display() {
678 assert_eq!(IsolationLevel::ReadCommitted.to_string(), "READ COMMITTED");
679 assert_eq!(IsolationLevel::Serializable.to_string(), "SERIALIZABLE");
680 }
681
682 #[test]
683 fn test_transaction_state_default() {
684 let opts = TransactOptions::default();
685 assert!(opts.isolation_level.is_none());
686 assert!(!opts.read_only);
687 }
688
689 #[test]
690 fn test_transact_options_builder() {
691 let opts = TransactOptions {
692 isolation_level: Some(IsolationLevel::Serializable),
693 read_only: true,
694 timeout: Some(Duration::from_secs(30)),
695 max_nesting_depth: DEFAULT_MAX_NESTING_DEPTH,
696 propagation: PropagationBehavior::default(),
697 };
698
699 assert_eq!(opts.isolation_level, Some(IsolationLevel::Serializable));
700 assert!(opts.read_only);
701 assert_eq!(opts.timeout, Some(Duration::from_secs(30)));
702 }
703
704 #[test]
705 fn test_auto_commit_default() {
706 assert_eq!(AutoCommit::default(), AutoCommit::On);
707 }
708
709 #[test]
710 fn test_transaction_state() {
711 assert_eq!(TransactionState::Active, TransactionState::Active);
712 assert_ne!(TransactionState::Active, TransactionState::Committed);
713 }
714
715 #[test]
716 fn test_transact_options_chaining() {
717 let opts = TransactOptions::default()
718 .with_isolation(IsolationLevel::Serializable)
719 .read_only()
720 .with_timeout(Duration::from_secs(60));
721 assert_eq!(opts.isolation_level, Some(IsolationLevel::Serializable));
722 assert!(opts.read_only);
723 assert_eq!(opts.timeout, Some(Duration::from_secs(60)));
724 }
725
726 #[tokio::test]
727 async fn test_transaction_commit() -> Result<(), TxError> {
728 let conn = Box::new(MockConnection::new());
729 let mut tx = Transaction::new(conn, TransactOptions::default());
730 assert!(tx.is_active());
731
732 let result = tx.execute("INSERT INTO users VALUES (1)").await;
733 assert!(result.is_ok());
734
735 tx.commit().await?;
736 assert_eq!(tx.state(), TransactionState::Committed);
737
738 let result = tx.commit().await;
740 assert!(result.is_err());
741 match result {
742 Err(TxError::NotActive(state)) => {
743 assert_eq!(state, TransactionState::Committed);
744 }
745 _ => panic!("Expected NotActive error"),
746 }
747 Ok(())
748 }
749
750 #[tokio::test]
751 async fn test_transaction_rollback() -> Result<(), TxError> {
752 let conn = Box::new(MockConnection::new());
753 let mut tx = Transaction::new(conn, TransactOptions::default());
754
755 tx.rollback().await?;
756 assert_eq!(tx.state(), TransactionState::RolledBack);
757
758 let result = tx.rollback().await;
760 assert!(result.is_err());
761 match result {
762 Err(TxError::NotActive(state)) => {
763 assert_eq!(state, TransactionState::RolledBack);
764 }
765 _ => panic!("Expected NotActive error"),
766 }
767 Ok(())
768 }
769
770 #[tokio::test]
771 async fn test_transaction_execute_after_commit() -> Result<(), TxError> {
772 let conn = Box::new(MockConnection::new());
773 let mut tx = Transaction::new(conn, TransactOptions::default());
774 tx.commit().await?;
775
776 let result = tx.execute("SELECT 1").await;
777 assert!(result.is_err());
778 match result {
779 Err(TxError::NotActive(_)) => {}
780 _ => panic!("Expected NotActive error"),
781 }
782 Ok(())
783 }
784
785 #[tokio::test]
786 async fn test_transaction_query_after_commit_returns_not_active() -> Result<(), TxError> {
787 let conn = Box::new(MockConnection::new());
788 let mut tx = Transaction::new(conn, TransactOptions::default());
789 tx.commit().await?;
790
791 let result = tx.query("SELECT 1").await;
792 assert!(result.is_err());
793 match result {
794 Err(TxError::NotActive(_)) => {}
795 _ => panic!("Expected NotActive error"),
796 }
797 Ok(())
798 }
799
800 #[tokio::test]
801 async fn test_transaction_savepoint() -> Result<(), TxError> {
802 let conn = Box::new(MockConnection::new());
803 let mut tx = Transaction::new(conn, TransactOptions::default());
804
805 let sp1 = tx.savepoint().await?;
806 assert_eq!(sp1, "sp_1");
807
808 let sp2 = tx.savepoint().await?;
809 assert_eq!(sp2, "sp_2");
810
811 tx.rollback_to_savepoint(&sp1).await?;
812 tx.release_savepoint(&sp2).await?;
813 Ok(())
814 }
815
816 #[tokio::test]
817 async fn test_transaction_savepoint_name_validation() {
818 let conn = Box::new(MockConnection::new());
819 let mut tx = Transaction::new(conn, TransactOptions::default());
820
821 let result = tx.rollback_to_savepoint("sp'; DROP TABLE--").await;
823 assert!(result.is_err());
824 match result {
825 Err(TxError::InvalidSavepointName(_)) => {}
826 _ => panic!("Expected InvalidSavepointName error"),
827 }
828
829 let result = tx.release_savepoint("1sp").await;
831 assert!(result.is_err());
832 match result {
833 Err(TxError::InvalidSavepointName(_)) => {}
834 _ => panic!("Expected InvalidSavepointName error"),
835 }
836
837 let result = tx.rollback_to_savepoint("").await;
839 assert!(result.is_err());
840 match result {
841 Err(TxError::InvalidSavepointName(_)) => {}
842 _ => panic!("Expected InvalidSavepointName error"),
843 }
844
845 let result = tx.rollback_to_savepoint("sp_test_1").await;
847 assert!(result.is_ok());
848 }
849
850 #[tokio::test]
851 async fn test_transaction_take_connection() -> Result<(), TxError> {
852 let conn = Box::new(MockConnection::new());
853 let mut tx = Transaction::new(conn, TransactOptions::default());
854
855 let result = tx.take_connection().await;
857 assert!(result.is_err());
858 match result {
859 Err(TxError::NotActive(_)) => {}
860 _ => panic!("Expected NotActive error"),
861 }
862
863 tx.commit().await?;
865 let conn = tx.take_connection().await;
866 assert!(conn.is_ok());
867
868 let result = tx.take_connection().await;
870 assert!(result.is_err());
871 match result {
872 Err(TxError::ConnectionTaken) => {}
873 _ => panic!("Expected ConnectionTaken error"),
874 }
875 Ok(())
876 }
877
878 #[tokio::test]
879 async fn test_transaction_manager() -> Result<(), TxError> {
880 let mgr = TransactionManager::new();
881 let conn = Box::new(MockConnection::new());
882
883 mgr.begin("tx1".to_string(), conn, TransactOptions::default())
884 .await?;
885
886 let state = mgr.state("tx1").await;
887 assert_eq!(state, Some(TransactionState::Active));
888
889 mgr.commit("tx1").await?;
890 let state = mgr.state("tx1").await;
891 assert_eq!(state, Some(TransactionState::Committed));
892
893 let list = mgr.list().await;
894 assert!(list.contains(&"tx1".to_string()));
895 Ok(())
896 }
897
898 #[tokio::test]
899 async fn test_transaction_manager_rollback() -> Result<(), TxError> {
900 let mgr = TransactionManager::new();
901 let conn = Box::new(MockConnection::new());
902
903 mgr.begin("tx2".to_string(), conn, TransactOptions::default())
904 .await?;
905
906 mgr.rollback("tx2").await?;
907 let state = mgr.state("tx2").await;
908 assert_eq!(state, Some(TransactionState::RolledBack));
909 Ok(())
910 }
911
912 #[tokio::test]
913 async fn test_transaction_manager_not_found() {
914 let mgr = TransactionManager::new();
915 let result = mgr.commit("nonexistent").await;
916 assert!(result.is_err());
917 }
918
919 #[tokio::test]
920 async fn test_transaction_manager_remove() -> Result<(), TxError> {
921 let mgr = TransactionManager::new();
922 let conn = Box::new(MockConnection::new());
923
924 mgr.begin("tx3".to_string(), conn, TransactOptions::default())
925 .await?;
926
927 let removed = mgr.remove("tx3").await;
928 assert!(removed.is_some());
929
930 let state = mgr.state("tx3").await;
931 assert_eq!(state, None);
932 Ok(())
933 }
934
935 #[tokio::test]
937 async fn test_transaction_drop_rolls_back_when_active() {
938 use std::sync::atomic::{AtomicBool, Ordering};
939 use std::sync::Arc as StdArc;
940
941 struct TrackingConnection {
942 rollback_called: StdArc<AtomicBool>,
943 }
944
945 impl Connection for TrackingConnection {
946 fn execute<'a>(
947 &'a mut self,
948 _sql: &'a str,
949 ) -> Pin<Box<dyn Future<Output = Result<u64, sz_orm_model::DbError>> + Send + 'a>>
950 {
951 Box::pin(async { Ok(1) })
952 }
953 fn query<'a>(
954 &'a mut self,
955 _sql: &'a str,
956 ) -> Pin<
957 Box<
958 dyn Future<
959 Output = Result<
960 Vec<std::collections::HashMap<String, sz_orm_model::Value>>,
961 sz_orm_model::DbError,
962 >,
963 > + Send
964 + 'a,
965 >,
966 > {
967 Box::pin(async { Ok(vec![]) })
968 }
969 fn begin_transaction<'a>(
970 &'a mut self,
971 ) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>>
972 {
973 Box::pin(async { Ok(()) })
974 }
975 fn commit<'a>(
976 &'a mut self,
977 ) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>>
978 {
979 Box::pin(async { Ok(()) })
980 }
981 fn rollback<'a>(
982 &'a mut self,
983 ) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>>
984 {
985 let flag = self.rollback_called.clone();
986 Box::pin(async move {
987 flag.store(true, Ordering::SeqCst);
988 Ok(())
989 })
990 }
991 fn is_connected(&self) -> bool {
992 true
993 }
994 fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
995 Box::pin(async { true })
996 }
997 fn close<'a>(
998 &'a mut self,
999 ) -> Pin<Box<dyn Future<Output = Result<(), sz_orm_model::DbError>> + Send + 'a>>
1000 {
1001 Box::pin(async { Ok(()) })
1002 }
1003 }
1004
1005 let rollback_flag = StdArc::new(AtomicBool::new(false));
1006 let conn = Box::new(TrackingConnection {
1007 rollback_called: rollback_flag.clone(),
1008 });
1009 {
1010 let _tx = Transaction::new(conn, TransactOptions::default());
1011 }
1013 tokio::time::sleep(Duration::from_millis(50)).await;
1015 assert!(
1016 rollback_flag.load(Ordering::SeqCst),
1017 "Drop should have triggered rollback"
1018 );
1019 }
1020
1021 #[test]
1024 fn test_h8_default_max_nesting_depth_is_8() {
1025 let opts = TransactOptions::default();
1026 assert_eq!(opts.max_nesting_depth, DEFAULT_MAX_NESTING_DEPTH);
1027 assert_eq!(opts.max_nesting_depth, 8);
1028 }
1029
1030 #[test]
1031 fn test_h8_with_max_nesting_depth_builder() {
1032 let opts = TransactOptions::default().with_max_nesting_depth(3);
1033 assert_eq!(opts.max_nesting_depth, 3);
1034 }
1035
1036 #[tokio::test]
1037 async fn test_h8_savepoint_within_default_depth_succeeds() -> Result<(), TxError> {
1038 let conn = Box::new(MockConnection::new());
1039 let mut tx = Transaction::new(conn, TransactOptions::default());
1040
1041 for i in 1..=8 {
1043 let sp = tx.savepoint().await?;
1044 assert_eq!(sp, format!("sp_{}", i));
1045 }
1046 Ok(())
1047 }
1048
1049 #[tokio::test]
1050 async fn test_h8_savepoint_exceeding_default_depth_fails() -> Result<(), TxError> {
1051 let conn = Box::new(MockConnection::new());
1052 let mut tx = Transaction::new(conn, TransactOptions::default());
1053
1054 for _ in 0..8 {
1056 tx.savepoint().await?;
1057 }
1058
1059 let result = tx.savepoint().await;
1061 assert!(result.is_err());
1062 match result {
1063 Err(TxError::MaxNestingDepthExceeded {
1064 current_depth,
1065 max_depth,
1066 }) => {
1067 assert_eq!(current_depth, 9);
1068 assert_eq!(max_depth, 8);
1069 }
1070 _ => panic!("Expected MaxNestingDepthExceeded error"),
1071 }
1072 Ok(())
1073 }
1074
1075 #[tokio::test]
1076 async fn test_h8_savepoint_with_custom_depth_3() -> Result<(), TxError> {
1077 let conn = Box::new(MockConnection::new());
1078 let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(3));
1079
1080 for i in 1..=3 {
1082 let sp = tx.savepoint().await?;
1083 assert_eq!(sp, format!("sp_{}", i));
1084 }
1085
1086 let result = tx.savepoint().await;
1088 assert!(result.is_err());
1089 match result {
1090 Err(TxError::MaxNestingDepthExceeded {
1091 current_depth,
1092 max_depth,
1093 }) => {
1094 assert_eq!(current_depth, 4);
1095 assert_eq!(max_depth, 3);
1096 }
1097 _ => panic!("Expected MaxNestingDepthExceeded error"),
1098 }
1099 Ok(())
1100 }
1101
1102 #[tokio::test]
1103 async fn test_h8_savepoint_depth_zero_disables_nesting() {
1104 let conn = Box::new(MockConnection::new());
1105 let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(0));
1106
1107 let result = tx.savepoint().await;
1109 assert!(result.is_err());
1110 match result {
1111 Err(TxError::MaxNestingDepthExceeded {
1112 current_depth,
1113 max_depth,
1114 }) => {
1115 assert_eq!(current_depth, 1);
1116 assert_eq!(max_depth, 0);
1117 }
1118 _ => panic!("Expected MaxNestingDepthExceeded error"),
1119 }
1120 }
1121
1122 #[tokio::test]
1123 async fn test_h8_savepoint_after_rollback_to_still_respects_depth() -> Result<(), TxError> {
1124 let conn = Box::new(MockConnection::new());
1127 let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(2));
1128
1129 let sp1 = tx.savepoint().await?;
1130 let sp2 = tx.savepoint().await?;
1131
1132 tx.rollback_to_savepoint(&sp1).await?;
1134 tx.release_savepoint(&sp2).await?;
1135
1136 let result = tx.savepoint().await;
1138 assert!(result.is_err());
1139 match result {
1140 Err(TxError::MaxNestingDepthExceeded {
1141 current_depth,
1142 max_depth,
1143 }) => {
1144 assert_eq!(current_depth, 3);
1145 assert_eq!(max_depth, 2);
1146 }
1147 _ => panic!("Expected MaxNestingDepthExceeded error"),
1148 }
1149 Ok(())
1150 }
1151
1152 #[tokio::test]
1153 async fn test_h8_max_nesting_depth_error_display() {
1154 let err = TxError::MaxNestingDepthExceeded {
1155 current_depth: 10,
1156 max_depth: 8,
1157 };
1158 let msg = format!("{}", err);
1159 assert!(msg.contains("10"));
1160 assert!(msg.contains("8"));
1161 assert!(msg.contains("exceeds"));
1162 }
1163
1164 #[test]
1167 fn test_m8_is_deadlock_error_mysql() {
1168 assert!(is_deadlock_error(
1169 "Deadlock found when trying to get lock; try restarting transaction"
1170 ));
1171 assert!(is_deadlock_error("Error 1213: Deadlock found"));
1172 assert!(is_deadlock_error("MySQL error (1213)"));
1173 }
1174
1175 #[test]
1176 fn test_m8_is_deadlock_error_postgresql() {
1177 assert!(is_deadlock_error("deadlock detected"));
1178 assert!(is_deadlock_error("ERROR: deadlock detected (40P01)"));
1179 assert!(is_deadlock_error("SQLSTATE 40P01"));
1180 }
1181
1182 #[test]
1183 fn test_m8_is_deadlock_error_sqlite() {
1184 assert!(is_deadlock_error("database is locked"));
1185 assert!(is_deadlock_error("database table is locked"));
1186 }
1187
1188 #[test]
1189 fn test_m8_is_deadlock_error_oracle() {
1190 assert!(is_deadlock_error(
1191 "ORA-00060: deadlock detected while waiting for resource"
1192 ));
1193 }
1194
1195 #[test]
1196 fn test_m8_is_deadlock_error_sql_server() {
1197 assert!(is_deadlock_error(
1198 "Transaction (Process ID 52) was deadlocked on lock resources"
1199 ));
1200 assert!(is_deadlock_error("Error 1205: Transaction was deadlocked"));
1201 }
1202
1203 #[test]
1204 fn test_m8_is_deadlock_error_non_deadlock() {
1205 assert!(!is_deadlock_error("connection refused"));
1206 assert!(!is_deadlock_error("syntax error near SELECT"));
1207 assert!(!is_deadlock_error("permission denied for table users"));
1208 assert!(!is_deadlock_error(""));
1209 }
1210
1211 #[tokio::test]
1212 async fn test_m8_retry_on_deadlock_succeeds_first_attempt() -> Result<(), TxError> {
1213 use std::sync::atomic::{AtomicU32, Ordering};
1214
1215 let counter = Arc::new(AtomicU32::new(0));
1216 let counter_clone = counter.clone();
1217
1218 let result: Result<u32, TxError> =
1219 retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1220 let c = counter_clone.clone();
1221 async move {
1222 c.fetch_add(1, Ordering::SeqCst);
1223 Ok(42u32)
1224 }
1225 })
1226 .await;
1227
1228 assert_eq!(result?, 42);
1229 assert_eq!(counter.load(Ordering::SeqCst), 1);
1230 Ok(())
1231 }
1232
1233 #[tokio::test]
1234 async fn test_m8_retry_on_deadlock_retries_on_deadlock_error() -> Result<(), TxError> {
1235 use std::sync::atomic::{AtomicU32, Ordering};
1236
1237 let counter = Arc::new(AtomicU32::new(0));
1238 let counter_clone = counter.clone();
1239
1240 let result: Result<u32, TxError> =
1241 retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1242 let c = counter_clone.clone();
1243 async move {
1244 let n = c.fetch_add(1, Ordering::SeqCst);
1245 if n < 2 {
1246 Err(TxError::CommitFailed(
1248 "Deadlock found when trying to get lock".to_string(),
1249 ))
1250 } else {
1251 Ok(42u32)
1252 }
1253 }
1254 })
1255 .await;
1256
1257 assert_eq!(result?, 42);
1258 assert_eq!(counter.load(Ordering::SeqCst), 3);
1259 Ok(())
1260 }
1261
1262 #[tokio::test]
1263 async fn test_m8_retry_on_deadlock_returns_error_after_max_attempts() {
1264 use std::sync::atomic::{AtomicU32, Ordering};
1265
1266 let counter = Arc::new(AtomicU32::new(0));
1267 let counter_clone = counter.clone();
1268
1269 let result: Result<u32, TxError> =
1270 retry_on_deadlock(2, Duration::from_millis(1), |_attempt| {
1271 let c = counter_clone.clone();
1272 async move {
1273 c.fetch_add(1, Ordering::SeqCst);
1274 Err(TxError::CommitFailed(
1275 "Deadlock found when trying to get lock".to_string(),
1276 ))
1277 }
1278 })
1279 .await;
1280
1281 assert!(result.is_err());
1283 assert_eq!(counter.load(Ordering::SeqCst), 2);
1284 }
1285
1286 #[tokio::test]
1287 async fn test_m8_retry_on_deadlock_does_not_retry_non_deadlock_errors() {
1288 use std::sync::atomic::{AtomicU32, Ordering};
1289
1290 let counter = Arc::new(AtomicU32::new(0));
1291 let counter_clone = counter.clone();
1292
1293 let result: Result<u32, TxError> =
1294 retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1295 let c = counter_clone.clone();
1296 async move {
1297 c.fetch_add(1, Ordering::SeqCst);
1298 Err(TxError::CommitFailed("syntax error".to_string()))
1299 }
1300 })
1301 .await;
1302
1303 assert!(result.is_err());
1305 assert_eq!(counter.load(Ordering::SeqCst), 1);
1306 }
1307
1308 #[test]
1309 fn test_m8_deadlock_error_display() {
1310 let err = TxError::DeadlockDetected {
1311 attempt: 2,
1312 max_attempts: 3,
1313 };
1314 let msg = format!("{}", err);
1315 assert!(msg.contains("2"));
1316 assert!(msg.contains("3"));
1317 assert!(msg.contains("Deadlock"));
1318 }
1319}