Skip to main content

sz_orm_pool/
transaction.rs

1//! Transaction support
2//!
3//! Provides ACID transaction management
4
5use 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// TransactionState 定义在 `error` 模块以避免 `transaction` ↔ `error` 循环依赖;
12// 通过 `pub use error::*;` 在 crate 根重导出,外部访问路径仍为 `sz_orm_core::TransactionState`。
13
14#[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/// 事务传播行为
44#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
45pub enum PropagationBehavior {
46    /// 如果当前存在事务,则加入;否则新建事务(默认)
47    #[default]
48    Required,
49    /// 必须在事务中执行,否则抛出异常
50    Mandatory,
51    /// 必须在没有事务的环境中执行,否则抛出异常
52    Never,
53    /// 如果当前存在事务,则加入;否则无事务执行
54    Supports,
55    /// 总是新建事务,挂起当前事务
56    RequiresNew,
57    /// 如果当前存在事务,则嵌套执行(保存点);否则新建事务
58    Nested,
59}
60
61pub struct TransactOptions {
62    pub isolation_level: Option<IsolationLevel>,
63    pub read_only: bool,
64    pub timeout: Option<Duration>,
65    /// H-8 修复:嵌套事务(保存点)最大深度限制
66    ///
67    /// 默认 `DEFAULT_MAX_NESTING_DEPTH`(8),防止递归事务导致数据库连接耗尽或
68    /// 保存点栈溢出。设为 0 表示禁用嵌套事务(首次 `savepoint()` 即报错)。
69    pub max_nesting_depth: u32,
70    /// 事务传播行为(默认 Required)
71    pub propagation: PropagationBehavior,
72}
73
74/// H-8 默认最大嵌套深度
75pub 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    /// H-8 修复:设置最大嵌套深度
106    pub fn with_max_nesting_depth(mut self, max_depth: u32) -> Self {
107        self.max_nesting_depth = max_depth;
108        self
109    }
110
111    /// 设置事务传播行为
112    pub fn with_propagation(mut self, propagation: PropagationBehavior) -> Self {
113        self.propagation = propagation;
114        self
115    }
116}
117
118/// 校验保存点名称(防止 SQL 注入)
119///
120/// 保存点名称规则:
121/// - 非空
122/// - 只能包含 ASCII 字母、数字、下划线
123/// - 不能以数字开头
124fn 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
137/// 事务对象,封装一个数据库事务
138///
139/// 内部连接以 `Option<Box<dyn Connection>>` 形式持有:
140/// - 事务执行期间,连接存在
141/// - 调用 `take_connection()` 可在 commit/rollback 后取回连接归还到连接池
142/// - Drop 时若事务仍 Active,会尝试 spawn 后台 rollback 任务
143pub struct Transaction {
144    conn: Arc<Mutex<Option<Box<dyn Connection>>>>,
145    state: TransactionState,
146    options: TransactOptions,
147    savepoint_counter: u32,
148    /// 事务超时截止时间(由 options.timeout 计算,在 commit/check 时检查)
149    deadline: Option<Instant>,
150}
151
152impl Transaction {
153    /// 创建新事务(调用方应先通过 connection.begin_transaction() 启动事务)
154    ///
155    /// L-5 修复:补充示例文档
156    ///
157    /// 通常通过 `Connection::begin_transaction()` 创建,而不是直接调用此方法。
158    ///
159    /// # 示例
160    ///
161    /// ```ignore
162    /// use sz_orm_pool::transaction::{Transaction, TransactOptions};
163    ///
164    /// # async fn example(conn: Box<dyn sz_orm_pool::pool::Connection>) {
165    /// // 通常通过 Connection::begin_transaction 创建
166    /// let tx = Transaction::new(conn, TransactOptions::default());
167    /// assert!(tx.is_active());
168    /// # }
169    /// ```
170    pub fn new(conn: Box<dyn Connection>, options: TransactOptions) -> Self {
171        // 任务1:根据 options.timeout 计算事务截止时间
172        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    /// 获取当前事务状态
183    pub fn state(&self) -> TransactionState {
184        self.state
185    }
186
187    /// 检查事务是否仍然活跃
188    pub fn is_active(&self) -> bool {
189        self.state == TransactionState::Active
190    }
191
192    /// 提交事务
193    ///
194    /// L-5 修复:补充示例文档
195    ///
196    /// 提交后事务状态变为 `Committed`,不可再次 commit/rollback。
197    /// 若未提交就 drop,会自动 rollback。
198    ///
199    /// # 示例
200    ///
201    /// ```ignore
202    /// # use sz_orm_pool::transaction::TransactOptions;
203    /// # async fn example(mut tx: sz_orm_pool::transaction::Transaction) -> Result<(), Box<dyn std::error::Error>> {
204    /// tx.execute("INSERT INTO users (name) VALUES ('Alice')").await?;
205    /// tx.execute("INSERT INTO users (name) VALUES ('Bob')").await?;
206    /// tx.commit().await?; // 提交两条 INSERT
207    /// # Ok(())
208    /// # }
209    /// ```
210    pub async fn commit(&mut self) -> Result<(), TxError> {
211        if self.state != TransactionState::Active {
212            return Err(TxError::NotActive(self.state));
213        }
214        // 任务1:检查事务超时,超时则回滚并返回错误
215        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    /// 回滚事务
231    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    /// 在事务中执行 SQL(在事务未结束时执行)
245    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    /// 在事务中执行查询
259    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    /// 创建保存点(用于嵌套事务)
276    ///
277    /// 返回自动生成的保存点名(格式 `sp_<N>`,N 单调递增)。
278    ///
279    /// H-8 修复:检查嵌套深度,超过 `options.max_nesting_depth` 时返回
280    /// `TxError::MaxNestingDepthExceeded`。
281    ///
282    /// L-5 修复:补充示例文档
283    ///
284    /// # 示例
285    ///
286    /// ```ignore
287    /// # async fn example(mut tx: sz_orm_pool::transaction::Transaction) -> Result<(), Box<dyn std::error::Error>> {
288    /// // 在外层事务中创建保存点
289    /// let sp = tx.savepoint().await?;
290    /// // 执行一些操作
291    /// tx.execute("INSERT INTO orders (id) VALUES (1)").await?;
292    /// // 出错时回滚到保存点(不影响外层事务的其他操作)
293    /// tx.rollback_to_savepoint(&sp).await?;
294    /// // 不再需要保存点时释放
295    /// tx.release_savepoint(&sp).await?;
296    /// # Ok(())
297    /// # }
298    /// ```
299    pub async fn savepoint(&mut self) -> Result<String, TxError> {
300        if self.state != TransactionState::Active {
301            return Err(TxError::NotActive(self.state));
302        }
303        // H-8 修复:嵌套深度检查
304        // savepoint_counter 表示已创建的保存点数;新保存点的深度为 counter + 1
305        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        // 内部生成的名称已通过命名规则(sp_ + 数字),但为防御性编程仍校验
315        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    /// 回滚到保存点
326    ///
327    /// `name` 必须是合法的保存点名称(仅 ASCII 字母/数字/下划线,且不以数字开头)。
328    /// 通常使用 `savepoint()` 返回的名称。
329    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    /// 释放保存点
344    ///
345    /// `name` 必须是合法的保存点名称(仅 ASCII 字母/数字/下划线,且不以数字开头)。
346    /// 通常使用 `savepoint()` 返回的名称。
347    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    /// 取出底层连接(用于归还到连接池)
362    ///
363    /// 仅在事务已 commit/rollback 后才能调用,否则返回 `NotActive` 错误。
364    /// 重复调用返回 `ConnectionTaken` 错误。
365    ///
366    /// 典型用法:
367    /// ```ignore
368    /// tx.commit().await?;
369    /// let conn = tx.take_connection().await?;
370    /// pool.release(conn).await;
371    /// ```
372    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    /// 获取事务选项
381    pub fn options(&self) -> &TransactOptions {
382        &self.options
383    }
384}
385
386/// M-8 修复:检测错误字符串是否表示死锁
387///
388/// 各数据库死锁错误码:
389/// - MySQL: 1213 (ER_LOCK_DEADLOCK)
390/// - PostgreSQL: 40P01 (deadlock_detected)
391/// - SQLite: "database is locked" / "database table is locked"
392/// - Oracle: ORA-00060
393/// - SQL Server: 1205 (Transaction was deadlocked)
394pub fn is_deadlock_error(err_msg: &str) -> bool {
395    let lower = err_msg.to_lowercase();
396    // MySQL: "Deadlock found when trying to get lock" (error 1213)
397    if lower.contains("deadlock found when trying to get lock") {
398        return true;
399    }
400    // MySQL error code 1213
401    if lower.contains("error 1213") || lower.contains("(1213)") {
402        return true;
403    }
404    // PostgreSQL: "deadlock detected" (SQLSTATE 40P01)
405    if lower.contains("deadlock detected") || lower.contains("40p01") {
406        return true;
407    }
408    // SQLite: "database is locked"
409    if lower.contains("database is locked") || lower.contains("database table is locked") {
410        return true;
411    }
412    // Oracle: ORA-00060: deadlock detected while waiting for resource
413    if lower.contains("ora-00060") {
414        return true;
415    }
416    // SQL Server: 1205 (Transaction was deadlocked on lock resources)
417    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
426/// M-8 修复:在死锁时自动重试事务
427///
428/// `operation` 是一个异步闭包,返回 `Result<T, TxError>`。如果返回的错误包含
429/// 死锁信息(通过 `is_deadlock_error` 判断),则等待 `backoff` 后重试。
430///
431/// # 参数
432///
433/// - `max_attempts`: 最大重试次数(含首次执行),默认 3
434/// - `backoff`: 每次重试前的等待时间,默认 50ms
435/// - `operation`: 异步闭包,接受当前尝试次数(从 1 开始),返回 `Result<T, TxError>`
436///
437/// # 返回值
438///
439/// - 成功时返回 `Ok(T)`
440/// - 所有重试都失败时返回最后一次的 `Err(TxError::DeadlockDetected)`
441///
442/// # 示例
443///
444/// ```ignore
445/// let result = retry_on_deadlock(3, Duration::from_millis(50), |attempt| async move {
446///     // 执行事务操作
447///     tx.execute("UPDATE accounts SET balance = balance - 100 WHERE id = 1").await
448/// }).await;
449/// ```
450pub 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                // 检查是否为死锁错误
465                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                // 非死锁错误或已达最大重试次数,直接返回
475                return Err(e);
476            }
477        }
478    }
479    // 理论上不会到达(循环内所有路径都会 return),但为类型安全保留
480    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        // 如果事务未被显式提交或回滚,在 drop 时尝试回滚
489        // 注意:无法在 Drop 中 await,所以这里 spawn 一个后台任务执行 rollback
490        if self.state == TransactionState::Active {
491            let conn = self.conn.clone();
492            // 尝试获取当前 tokio 运行时句柄;若不存在(如非 async 上下文)则跳过
493            if let Ok(handle) = tokio::runtime::Handle::try_current() {
494                // spawn 后台任务:锁连接 → rollback → 连接随 Arc 释放而 Drop
495                // 若任务因 runtime 关闭未执行,连接 Drop 时由驱动/池策略兜底
496                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            // 标记为已回滚(即使后台任务未完成,状态机也需要前进)
504            self.state = TransactionState::RolledBack;
505        }
506    }
507}
508
509/// 事务管理器,管理多个事务
510pub 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    /// 开始新事务
522    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    /// 提交事务
539    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    /// 回滚事务
548    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    /// 获取事务状态
557    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    /// 列出所有事务 ID
563    pub async fn list(&self) -> Vec<String> {
564        let txs = self.transactions.lock().await;
565        txs.keys().cloned().collect()
566    }
567
568    /// 移除已完成的事务
569    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    /// 测试用的模拟连接
588    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        // 再次 commit 应该失败(NotActive)
739        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        // 再次 rollback 应该失败(NotActive)
759        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        // 非法名称:包含单引号(SQL 注入尝试)
822        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        // 非法名称:以数字开头
830        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        // 非法名称:空字符串
838        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        // 合法名称:字母+下划线+数字
846        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        // Active 状态下不能取连接
856        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        // commit 后可以取连接
864        tx.commit().await?;
865        let conn = tx.take_connection().await;
866        assert!(conn.is_ok());
867
868        // 重复取连接应失败
869        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    /// 验证 Drop 时若事务仍 Active,会 spawn 后台 rollback 任务
936    #[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            // _tx 在块结束时 drop,状态为 Active,应触发后台 rollback
1012        }
1013        // 给 spawn 的任务一点时间执行
1014        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    // ==================== H-8 事务嵌套深度限制测试 ====================
1022
1023    #[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        // 默认深度 8,创建 8 个保存点应全部成功
1042        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        // 创建 8 个保存点(达到上限)
1055        for _ in 0..8 {
1056            tx.savepoint().await?;
1057        }
1058
1059        // 第 9 个保存点应失败
1060        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        // 3 个保存点应成功
1081        for i in 1..=3 {
1082            let sp = tx.savepoint().await?;
1083            assert_eq!(sp, format!("sp_{}", i));
1084        }
1085
1086        // 第 4 个应失败
1087        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        // 深度 0 表示禁用嵌套,首次 savepoint 即失败
1108        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        // 即使回滚到保存点,savepoint_counter 不减少(保存点栈可能仍存在),
1125        // 因此深度检查仍以 savepoint_counter 为准
1126        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        // 回滚到 sp1(不重置计数器)
1133        tx.rollback_to_savepoint(&sp1).await?;
1134        tx.release_savepoint(&sp2).await?;
1135
1136        // 第 3 个保存点仍应失败(计数器不回退)
1137        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    // ==================== M-8 死锁检测重试测试 ====================
1165
1166    #[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                        // 前两次返回死锁错误
1247                        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        // 应返回最后一次的死锁错误
1282        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        // 非死锁错误应立即返回,不重试
1304        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}