1use crate::error::{TransactionState, TxError};
6use crate::pool::Connection;
7use std::sync::Arc;
8use std::time::{Duration, Instant};
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, crate::value::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, crate::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, crate::value::Value>>,
625 crate::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<(), crate::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<(), crate::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<(), crate::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<(), crate::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() {
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.unwrap();
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 }
748
749 #[tokio::test]
750 async fn test_transaction_rollback() {
751 let conn = Box::new(MockConnection::new());
752 let mut tx = Transaction::new(conn, TransactOptions::default());
753
754 tx.rollback().await.unwrap();
755 assert_eq!(tx.state(), TransactionState::RolledBack);
756
757 let result = tx.rollback().await;
759 assert!(result.is_err());
760 match result {
761 Err(TxError::NotActive(state)) => {
762 assert_eq!(state, TransactionState::RolledBack);
763 }
764 _ => panic!("Expected NotActive error"),
765 }
766 }
767
768 #[tokio::test]
769 async fn test_transaction_execute_after_commit() {
770 let conn = Box::new(MockConnection::new());
771 let mut tx = Transaction::new(conn, TransactOptions::default());
772 tx.commit().await.unwrap();
773
774 let result = tx.execute("SELECT 1").await;
775 assert!(result.is_err());
776 match result {
777 Err(TxError::NotActive(_)) => {}
778 _ => panic!("Expected NotActive error"),
779 }
780 }
781
782 #[tokio::test]
783 async fn test_transaction_query_after_commit_returns_not_active() {
784 let conn = Box::new(MockConnection::new());
785 let mut tx = Transaction::new(conn, TransactOptions::default());
786 tx.commit().await.unwrap();
787
788 let result = tx.query("SELECT 1").await;
789 assert!(result.is_err());
790 match result {
791 Err(TxError::NotActive(_)) => {}
792 _ => panic!("Expected NotActive error"),
793 }
794 }
795
796 #[tokio::test]
797 async fn test_transaction_savepoint() {
798 let conn = Box::new(MockConnection::new());
799 let mut tx = Transaction::new(conn, TransactOptions::default());
800
801 let sp1 = tx.savepoint().await.unwrap();
802 assert_eq!(sp1, "sp_1");
803
804 let sp2 = tx.savepoint().await.unwrap();
805 assert_eq!(sp2, "sp_2");
806
807 tx.rollback_to_savepoint(&sp1).await.unwrap();
808 tx.release_savepoint(&sp2).await.unwrap();
809 }
810
811 #[tokio::test]
812 async fn test_transaction_savepoint_name_validation() {
813 let conn = Box::new(MockConnection::new());
814 let mut tx = Transaction::new(conn, TransactOptions::default());
815
816 let result = tx.rollback_to_savepoint("sp'; DROP TABLE--").await;
818 assert!(result.is_err());
819 match result {
820 Err(TxError::InvalidSavepointName(_)) => {}
821 _ => panic!("Expected InvalidSavepointName error"),
822 }
823
824 let result = tx.release_savepoint("1sp").await;
826 assert!(result.is_err());
827 match result {
828 Err(TxError::InvalidSavepointName(_)) => {}
829 _ => panic!("Expected InvalidSavepointName error"),
830 }
831
832 let result = tx.rollback_to_savepoint("").await;
834 assert!(result.is_err());
835 match result {
836 Err(TxError::InvalidSavepointName(_)) => {}
837 _ => panic!("Expected InvalidSavepointName error"),
838 }
839
840 let result = tx.rollback_to_savepoint("sp_test_1").await;
842 assert!(result.is_ok());
843 }
844
845 #[tokio::test]
846 async fn test_transaction_take_connection() {
847 let conn = Box::new(MockConnection::new());
848 let mut tx = Transaction::new(conn, TransactOptions::default());
849
850 let result = tx.take_connection().await;
852 assert!(result.is_err());
853 match result {
854 Err(TxError::NotActive(_)) => {}
855 _ => panic!("Expected NotActive error"),
856 }
857
858 tx.commit().await.unwrap();
860 let conn = tx.take_connection().await;
861 assert!(conn.is_ok());
862
863 let result = tx.take_connection().await;
865 assert!(result.is_err());
866 match result {
867 Err(TxError::ConnectionTaken) => {}
868 _ => panic!("Expected ConnectionTaken error"),
869 }
870 }
871
872 #[tokio::test]
873 async fn test_transaction_manager() {
874 let mgr = TransactionManager::new();
875 let conn = Box::new(MockConnection::new());
876
877 mgr.begin("tx1".to_string(), conn, TransactOptions::default())
878 .await
879 .unwrap();
880
881 let state = mgr.state("tx1").await;
882 assert_eq!(state, Some(TransactionState::Active));
883
884 mgr.commit("tx1").await.unwrap();
885 let state = mgr.state("tx1").await;
886 assert_eq!(state, Some(TransactionState::Committed));
887
888 let list = mgr.list().await;
889 assert!(list.contains(&"tx1".to_string()));
890 }
891
892 #[tokio::test]
893 async fn test_transaction_manager_rollback() {
894 let mgr = TransactionManager::new();
895 let conn = Box::new(MockConnection::new());
896
897 mgr.begin("tx2".to_string(), conn, TransactOptions::default())
898 .await
899 .unwrap();
900
901 mgr.rollback("tx2").await.unwrap();
902 let state = mgr.state("tx2").await;
903 assert_eq!(state, Some(TransactionState::RolledBack));
904 }
905
906 #[tokio::test]
907 async fn test_transaction_manager_not_found() {
908 let mgr = TransactionManager::new();
909 let result = mgr.commit("nonexistent").await;
910 assert!(result.is_err());
911 }
912
913 #[tokio::test]
914 async fn test_transaction_manager_remove() {
915 let mgr = TransactionManager::new();
916 let conn = Box::new(MockConnection::new());
917
918 mgr.begin("tx3".to_string(), conn, TransactOptions::default())
919 .await
920 .unwrap();
921
922 let removed = mgr.remove("tx3").await;
923 assert!(removed.is_some());
924
925 let state = mgr.state("tx3").await;
926 assert_eq!(state, None);
927 }
928
929 #[tokio::test]
931 async fn test_transaction_drop_rolls_back_when_active() {
932 use std::sync::atomic::{AtomicBool, Ordering};
933 use std::sync::Arc as StdArc;
934
935 struct TrackingConnection {
936 rollback_called: StdArc<AtomicBool>,
937 }
938
939 impl Connection for TrackingConnection {
940 fn execute<'a>(
941 &'a mut self,
942 _sql: &'a str,
943 ) -> Pin<Box<dyn Future<Output = Result<u64, crate::DbError>> + Send + 'a>>
944 {
945 Box::pin(async { Ok(1) })
946 }
947 fn query<'a>(
948 &'a mut self,
949 _sql: &'a str,
950 ) -> Pin<
951 Box<
952 dyn Future<
953 Output = Result<
954 Vec<std::collections::HashMap<String, crate::value::Value>>,
955 crate::DbError,
956 >,
957 > + Send
958 + 'a,
959 >,
960 > {
961 Box::pin(async { Ok(vec![]) })
962 }
963 fn begin_transaction<'a>(
964 &'a mut self,
965 ) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
966 Box::pin(async { Ok(()) })
967 }
968 fn commit<'a>(
969 &'a mut self,
970 ) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
971 Box::pin(async { Ok(()) })
972 }
973 fn rollback<'a>(
974 &'a mut self,
975 ) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
976 let flag = self.rollback_called.clone();
977 Box::pin(async move {
978 flag.store(true, Ordering::SeqCst);
979 Ok(())
980 })
981 }
982 fn is_connected(&self) -> bool {
983 true
984 }
985 fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
986 Box::pin(async { true })
987 }
988 fn close<'a>(
989 &'a mut self,
990 ) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
991 Box::pin(async { Ok(()) })
992 }
993 }
994
995 let rollback_flag = StdArc::new(AtomicBool::new(false));
996 let conn = Box::new(TrackingConnection {
997 rollback_called: rollback_flag.clone(),
998 });
999 {
1000 let _tx = Transaction::new(conn, TransactOptions::default());
1001 }
1003 tokio::time::sleep(Duration::from_millis(50)).await;
1005 assert!(
1006 rollback_flag.load(Ordering::SeqCst),
1007 "Drop should have triggered rollback"
1008 );
1009 }
1010
1011 #[test]
1014 fn test_h8_default_max_nesting_depth_is_8() {
1015 let opts = TransactOptions::default();
1016 assert_eq!(opts.max_nesting_depth, DEFAULT_MAX_NESTING_DEPTH);
1017 assert_eq!(opts.max_nesting_depth, 8);
1018 }
1019
1020 #[test]
1021 fn test_h8_with_max_nesting_depth_builder() {
1022 let opts = TransactOptions::default().with_max_nesting_depth(3);
1023 assert_eq!(opts.max_nesting_depth, 3);
1024 }
1025
1026 #[tokio::test]
1027 async fn test_h8_savepoint_within_default_depth_succeeds() {
1028 let conn = Box::new(MockConnection::new());
1029 let mut tx = Transaction::new(conn, TransactOptions::default());
1030
1031 for i in 1..=8 {
1033 let sp = tx.savepoint().await.unwrap();
1034 assert_eq!(sp, format!("sp_{}", i));
1035 }
1036 }
1037
1038 #[tokio::test]
1039 async fn test_h8_savepoint_exceeding_default_depth_fails() {
1040 let conn = Box::new(MockConnection::new());
1041 let mut tx = Transaction::new(conn, TransactOptions::default());
1042
1043 for _ in 0..8 {
1045 tx.savepoint().await.unwrap();
1046 }
1047
1048 let result = tx.savepoint().await;
1050 assert!(result.is_err());
1051 match result {
1052 Err(TxError::MaxNestingDepthExceeded {
1053 current_depth,
1054 max_depth,
1055 }) => {
1056 assert_eq!(current_depth, 9);
1057 assert_eq!(max_depth, 8);
1058 }
1059 _ => panic!("Expected MaxNestingDepthExceeded error"),
1060 }
1061 }
1062
1063 #[tokio::test]
1064 async fn test_h8_savepoint_with_custom_depth_3() {
1065 let conn = Box::new(MockConnection::new());
1066 let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(3));
1067
1068 for i in 1..=3 {
1070 let sp = tx.savepoint().await.unwrap();
1071 assert_eq!(sp, format!("sp_{}", i));
1072 }
1073
1074 let result = tx.savepoint().await;
1076 assert!(result.is_err());
1077 match result {
1078 Err(TxError::MaxNestingDepthExceeded {
1079 current_depth,
1080 max_depth,
1081 }) => {
1082 assert_eq!(current_depth, 4);
1083 assert_eq!(max_depth, 3);
1084 }
1085 _ => panic!("Expected MaxNestingDepthExceeded error"),
1086 }
1087 }
1088
1089 #[tokio::test]
1090 async fn test_h8_savepoint_depth_zero_disables_nesting() {
1091 let conn = Box::new(MockConnection::new());
1092 let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(0));
1093
1094 let result = tx.savepoint().await;
1096 assert!(result.is_err());
1097 match result {
1098 Err(TxError::MaxNestingDepthExceeded {
1099 current_depth,
1100 max_depth,
1101 }) => {
1102 assert_eq!(current_depth, 1);
1103 assert_eq!(max_depth, 0);
1104 }
1105 _ => panic!("Expected MaxNestingDepthExceeded error"),
1106 }
1107 }
1108
1109 #[tokio::test]
1110 async fn test_h8_savepoint_after_rollback_to_still_respects_depth() {
1111 let conn = Box::new(MockConnection::new());
1114 let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(2));
1115
1116 let sp1 = tx.savepoint().await.unwrap();
1117 let sp2 = tx.savepoint().await.unwrap();
1118
1119 tx.rollback_to_savepoint(&sp1).await.unwrap();
1121 tx.release_savepoint(&sp2).await.unwrap();
1122
1123 let result = tx.savepoint().await;
1125 assert!(result.is_err());
1126 match result {
1127 Err(TxError::MaxNestingDepthExceeded {
1128 current_depth,
1129 max_depth,
1130 }) => {
1131 assert_eq!(current_depth, 3);
1132 assert_eq!(max_depth, 2);
1133 }
1134 _ => panic!("Expected MaxNestingDepthExceeded error"),
1135 }
1136 }
1137
1138 #[tokio::test]
1139 async fn test_h8_max_nesting_depth_error_display() {
1140 let err = TxError::MaxNestingDepthExceeded {
1141 current_depth: 10,
1142 max_depth: 8,
1143 };
1144 let msg = format!("{}", err);
1145 assert!(msg.contains("10"));
1146 assert!(msg.contains("8"));
1147 assert!(msg.contains("exceeds"));
1148 }
1149
1150 #[test]
1153 fn test_m8_is_deadlock_error_mysql() {
1154 assert!(is_deadlock_error(
1155 "Deadlock found when trying to get lock; try restarting transaction"
1156 ));
1157 assert!(is_deadlock_error("Error 1213: Deadlock found"));
1158 assert!(is_deadlock_error("MySQL error (1213)"));
1159 }
1160
1161 #[test]
1162 fn test_m8_is_deadlock_error_postgresql() {
1163 assert!(is_deadlock_error("deadlock detected"));
1164 assert!(is_deadlock_error("ERROR: deadlock detected (40P01)"));
1165 assert!(is_deadlock_error("SQLSTATE 40P01"));
1166 }
1167
1168 #[test]
1169 fn test_m8_is_deadlock_error_sqlite() {
1170 assert!(is_deadlock_error("database is locked"));
1171 assert!(is_deadlock_error("database table is locked"));
1172 }
1173
1174 #[test]
1175 fn test_m8_is_deadlock_error_oracle() {
1176 assert!(is_deadlock_error(
1177 "ORA-00060: deadlock detected while waiting for resource"
1178 ));
1179 }
1180
1181 #[test]
1182 fn test_m8_is_deadlock_error_sql_server() {
1183 assert!(is_deadlock_error(
1184 "Transaction (Process ID 52) was deadlocked on lock resources"
1185 ));
1186 assert!(is_deadlock_error("Error 1205: Transaction was deadlocked"));
1187 }
1188
1189 #[test]
1190 fn test_m8_is_deadlock_error_non_deadlock() {
1191 assert!(!is_deadlock_error("connection refused"));
1192 assert!(!is_deadlock_error("syntax error near SELECT"));
1193 assert!(!is_deadlock_error("permission denied for table users"));
1194 assert!(!is_deadlock_error(""));
1195 }
1196
1197 #[tokio::test]
1198 async fn test_m8_retry_on_deadlock_succeeds_first_attempt() {
1199 use std::sync::atomic::{AtomicU32, Ordering};
1200
1201 let counter = Arc::new(AtomicU32::new(0));
1202 let counter_clone = counter.clone();
1203
1204 let result: Result<u32, TxError> =
1205 retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1206 let c = counter_clone.clone();
1207 async move {
1208 c.fetch_add(1, Ordering::SeqCst);
1209 Ok(42u32)
1210 }
1211 })
1212 .await;
1213
1214 assert_eq!(result.unwrap(), 42);
1215 assert_eq!(counter.load(Ordering::SeqCst), 1);
1216 }
1217
1218 #[tokio::test]
1219 async fn test_m8_retry_on_deadlock_retries_on_deadlock_error() {
1220 use std::sync::atomic::{AtomicU32, Ordering};
1221
1222 let counter = Arc::new(AtomicU32::new(0));
1223 let counter_clone = counter.clone();
1224
1225 let result: Result<u32, TxError> =
1226 retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1227 let c = counter_clone.clone();
1228 async move {
1229 let n = c.fetch_add(1, Ordering::SeqCst);
1230 if n < 2 {
1231 Err(TxError::CommitFailed(
1233 "Deadlock found when trying to get lock".to_string(),
1234 ))
1235 } else {
1236 Ok(42u32)
1237 }
1238 }
1239 })
1240 .await;
1241
1242 assert_eq!(result.unwrap(), 42);
1243 assert_eq!(counter.load(Ordering::SeqCst), 3);
1244 }
1245
1246 #[tokio::test]
1247 async fn test_m8_retry_on_deadlock_returns_error_after_max_attempts() {
1248 use std::sync::atomic::{AtomicU32, Ordering};
1249
1250 let counter = Arc::new(AtomicU32::new(0));
1251 let counter_clone = counter.clone();
1252
1253 let result: Result<u32, TxError> =
1254 retry_on_deadlock(2, Duration::from_millis(1), |_attempt| {
1255 let c = counter_clone.clone();
1256 async move {
1257 c.fetch_add(1, Ordering::SeqCst);
1258 Err(TxError::CommitFailed(
1259 "Deadlock found when trying to get lock".to_string(),
1260 ))
1261 }
1262 })
1263 .await;
1264
1265 assert!(result.is_err());
1267 assert_eq!(counter.load(Ordering::SeqCst), 2);
1268 }
1269
1270 #[tokio::test]
1271 async fn test_m8_retry_on_deadlock_does_not_retry_non_deadlock_errors() {
1272 use std::sync::atomic::{AtomicU32, Ordering};
1273
1274 let counter = Arc::new(AtomicU32::new(0));
1275 let counter_clone = counter.clone();
1276
1277 let result: Result<u32, TxError> =
1278 retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1279 let c = counter_clone.clone();
1280 async move {
1281 c.fetch_add(1, Ordering::SeqCst);
1282 Err(TxError::CommitFailed("syntax error".to_string()))
1283 }
1284 })
1285 .await;
1286
1287 assert!(result.is_err());
1289 assert_eq!(counter.load(Ordering::SeqCst), 1);
1290 }
1291
1292 #[test]
1293 fn test_m8_deadlock_error_display() {
1294 let err = TxError::DeadlockDetected {
1295 attempt: 2,
1296 max_attempts: 3,
1297 };
1298 let msg = format!("{}", err);
1299 assert!(msg.contains("2"));
1300 assert!(msg.contains("3"));
1301 assert!(msg.contains("Deadlock"));
1302 }
1303}