Skip to main content

sz_orm_core/
transaction.rs

1//! Transaction support
2//!
3//! Provides ACID transaction management
4
5use crate::error::{TransactionState, TxError};
6use crate::pool::Connection;
7use std::sync::Arc;
8use std::time::{Duration, Instant};
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_core::transaction::{Transaction, TransactOptions};
163    ///
164    /// # async fn example(conn: Box<dyn sz_orm_core::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_core::transaction::TransactOptions;
203    /// # async fn example(mut tx: sz_orm_core::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, 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    /// 创建保存点(用于嵌套事务)
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_core::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, 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() -> 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, crate::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, crate::value::Value>>,
961                                crate::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<(), crate::DbError>> + Send + 'a>> {
972                Box::pin(async { Ok(()) })
973            }
974            fn commit<'a>(
975                &'a mut self,
976            ) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
977                Box::pin(async { Ok(()) })
978            }
979            fn rollback<'a>(
980                &'a mut self,
981            ) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
982                let flag = self.rollback_called.clone();
983                Box::pin(async move {
984                    flag.store(true, Ordering::SeqCst);
985                    Ok(())
986                })
987            }
988            fn is_connected(&self) -> bool {
989                true
990            }
991            fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
992                Box::pin(async { true })
993            }
994            fn close<'a>(
995                &'a mut self,
996            ) -> Pin<Box<dyn Future<Output = Result<(), crate::DbError>> + Send + 'a>> {
997                Box::pin(async { Ok(()) })
998            }
999        }
1000
1001        let rollback_flag = StdArc::new(AtomicBool::new(false));
1002        let conn = Box::new(TrackingConnection {
1003            rollback_called: rollback_flag.clone(),
1004        });
1005        {
1006            let _tx = Transaction::new(conn, TransactOptions::default());
1007            // _tx 在块结束时 drop,状态为 Active,应触发后台 rollback
1008        }
1009        // 给 spawn 的任务一点时间执行
1010        tokio::time::sleep(Duration::from_millis(50)).await;
1011        assert!(
1012            rollback_flag.load(Ordering::SeqCst),
1013            "Drop should have triggered rollback"
1014        );
1015    }
1016
1017    // ==================== H-8 事务嵌套深度限制测试 ====================
1018
1019    #[test]
1020    fn test_h8_default_max_nesting_depth_is_8() {
1021        let opts = TransactOptions::default();
1022        assert_eq!(opts.max_nesting_depth, DEFAULT_MAX_NESTING_DEPTH);
1023        assert_eq!(opts.max_nesting_depth, 8);
1024    }
1025
1026    #[test]
1027    fn test_h8_with_max_nesting_depth_builder() {
1028        let opts = TransactOptions::default().with_max_nesting_depth(3);
1029        assert_eq!(opts.max_nesting_depth, 3);
1030    }
1031
1032    #[tokio::test]
1033    async fn test_h8_savepoint_within_default_depth_succeeds() -> Result<(), TxError> {
1034        let conn = Box::new(MockConnection::new());
1035        let mut tx = Transaction::new(conn, TransactOptions::default());
1036
1037        // 默认深度 8,创建 8 个保存点应全部成功
1038        for i in 1..=8 {
1039            let sp = tx.savepoint().await?;
1040            assert_eq!(sp, format!("sp_{}", i));
1041        }
1042        Ok(())
1043    }
1044
1045    #[tokio::test]
1046    async fn test_h8_savepoint_exceeding_default_depth_fails() -> Result<(), TxError> {
1047        let conn = Box::new(MockConnection::new());
1048        let mut tx = Transaction::new(conn, TransactOptions::default());
1049
1050        // 创建 8 个保存点(达到上限)
1051        for _ in 0..8 {
1052            tx.savepoint().await?;
1053        }
1054
1055        // 第 9 个保存点应失败
1056        let result = tx.savepoint().await;
1057        assert!(result.is_err());
1058        match result {
1059            Err(TxError::MaxNestingDepthExceeded {
1060                current_depth,
1061                max_depth,
1062            }) => {
1063                assert_eq!(current_depth, 9);
1064                assert_eq!(max_depth, 8);
1065            }
1066            _ => panic!("Expected MaxNestingDepthExceeded error"),
1067        }
1068        Ok(())
1069    }
1070
1071    #[tokio::test]
1072    async fn test_h8_savepoint_with_custom_depth_3() -> Result<(), TxError> {
1073        let conn = Box::new(MockConnection::new());
1074        let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(3));
1075
1076        // 3 个保存点应成功
1077        for i in 1..=3 {
1078            let sp = tx.savepoint().await?;
1079            assert_eq!(sp, format!("sp_{}", i));
1080        }
1081
1082        // 第 4 个应失败
1083        let result = tx.savepoint().await;
1084        assert!(result.is_err());
1085        match result {
1086            Err(TxError::MaxNestingDepthExceeded {
1087                current_depth,
1088                max_depth,
1089            }) => {
1090                assert_eq!(current_depth, 4);
1091                assert_eq!(max_depth, 3);
1092            }
1093            _ => panic!("Expected MaxNestingDepthExceeded error"),
1094        }
1095        Ok(())
1096    }
1097
1098    #[tokio::test]
1099    async fn test_h8_savepoint_depth_zero_disables_nesting() {
1100        let conn = Box::new(MockConnection::new());
1101        let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(0));
1102
1103        // 深度 0 表示禁用嵌套,首次 savepoint 即失败
1104        let result = tx.savepoint().await;
1105        assert!(result.is_err());
1106        match result {
1107            Err(TxError::MaxNestingDepthExceeded {
1108                current_depth,
1109                max_depth,
1110            }) => {
1111                assert_eq!(current_depth, 1);
1112                assert_eq!(max_depth, 0);
1113            }
1114            _ => panic!("Expected MaxNestingDepthExceeded error"),
1115        }
1116    }
1117
1118    #[tokio::test]
1119    async fn test_h8_savepoint_after_rollback_to_still_respects_depth() -> Result<(), TxError> {
1120        // 即使回滚到保存点,savepoint_counter 不减少(保存点栈可能仍存在),
1121        // 因此深度检查仍以 savepoint_counter 为准
1122        let conn = Box::new(MockConnection::new());
1123        let mut tx = Transaction::new(conn, TransactOptions::default().with_max_nesting_depth(2));
1124
1125        let sp1 = tx.savepoint().await?;
1126        let sp2 = tx.savepoint().await?;
1127
1128        // 回滚到 sp1(不重置计数器)
1129        tx.rollback_to_savepoint(&sp1).await?;
1130        tx.release_savepoint(&sp2).await?;
1131
1132        // 第 3 个保存点仍应失败(计数器不回退)
1133        let result = tx.savepoint().await;
1134        assert!(result.is_err());
1135        match result {
1136            Err(TxError::MaxNestingDepthExceeded {
1137                current_depth,
1138                max_depth,
1139            }) => {
1140                assert_eq!(current_depth, 3);
1141                assert_eq!(max_depth, 2);
1142            }
1143            _ => panic!("Expected MaxNestingDepthExceeded error"),
1144        }
1145        Ok(())
1146    }
1147
1148    #[tokio::test]
1149    async fn test_h8_max_nesting_depth_error_display() {
1150        let err = TxError::MaxNestingDepthExceeded {
1151            current_depth: 10,
1152            max_depth: 8,
1153        };
1154        let msg = format!("{}", err);
1155        assert!(msg.contains("10"));
1156        assert!(msg.contains("8"));
1157        assert!(msg.contains("exceeds"));
1158    }
1159
1160    // ==================== M-8 死锁检测重试测试 ====================
1161
1162    #[test]
1163    fn test_m8_is_deadlock_error_mysql() {
1164        assert!(is_deadlock_error(
1165            "Deadlock found when trying to get lock; try restarting transaction"
1166        ));
1167        assert!(is_deadlock_error("Error 1213: Deadlock found"));
1168        assert!(is_deadlock_error("MySQL error (1213)"));
1169    }
1170
1171    #[test]
1172    fn test_m8_is_deadlock_error_postgresql() {
1173        assert!(is_deadlock_error("deadlock detected"));
1174        assert!(is_deadlock_error("ERROR: deadlock detected (40P01)"));
1175        assert!(is_deadlock_error("SQLSTATE 40P01"));
1176    }
1177
1178    #[test]
1179    fn test_m8_is_deadlock_error_sqlite() {
1180        assert!(is_deadlock_error("database is locked"));
1181        assert!(is_deadlock_error("database table is locked"));
1182    }
1183
1184    #[test]
1185    fn test_m8_is_deadlock_error_oracle() {
1186        assert!(is_deadlock_error(
1187            "ORA-00060: deadlock detected while waiting for resource"
1188        ));
1189    }
1190
1191    #[test]
1192    fn test_m8_is_deadlock_error_sql_server() {
1193        assert!(is_deadlock_error(
1194            "Transaction (Process ID 52) was deadlocked on lock resources"
1195        ));
1196        assert!(is_deadlock_error("Error 1205: Transaction was deadlocked"));
1197    }
1198
1199    #[test]
1200    fn test_m8_is_deadlock_error_non_deadlock() {
1201        assert!(!is_deadlock_error("connection refused"));
1202        assert!(!is_deadlock_error("syntax error near SELECT"));
1203        assert!(!is_deadlock_error("permission denied for table users"));
1204        assert!(!is_deadlock_error(""));
1205    }
1206
1207    #[tokio::test]
1208    async fn test_m8_retry_on_deadlock_succeeds_first_attempt() -> Result<(), TxError> {
1209        use std::sync::atomic::{AtomicU32, Ordering};
1210
1211        let counter = Arc::new(AtomicU32::new(0));
1212        let counter_clone = counter.clone();
1213
1214        let result: Result<u32, TxError> =
1215            retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1216                let c = counter_clone.clone();
1217                async move {
1218                    c.fetch_add(1, Ordering::SeqCst);
1219                    Ok(42u32)
1220                }
1221            })
1222            .await;
1223
1224        assert_eq!(result?, 42);
1225        assert_eq!(counter.load(Ordering::SeqCst), 1);
1226        Ok(())
1227    }
1228
1229    #[tokio::test]
1230    async fn test_m8_retry_on_deadlock_retries_on_deadlock_error() -> Result<(), TxError> {
1231        use std::sync::atomic::{AtomicU32, Ordering};
1232
1233        let counter = Arc::new(AtomicU32::new(0));
1234        let counter_clone = counter.clone();
1235
1236        let result: Result<u32, TxError> =
1237            retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1238                let c = counter_clone.clone();
1239                async move {
1240                    let n = c.fetch_add(1, Ordering::SeqCst);
1241                    if n < 2 {
1242                        // 前两次返回死锁错误
1243                        Err(TxError::CommitFailed(
1244                            "Deadlock found when trying to get lock".to_string(),
1245                        ))
1246                    } else {
1247                        Ok(42u32)
1248                    }
1249                }
1250            })
1251            .await;
1252
1253        assert_eq!(result?, 42);
1254        assert_eq!(counter.load(Ordering::SeqCst), 3);
1255        Ok(())
1256    }
1257
1258    #[tokio::test]
1259    async fn test_m8_retry_on_deadlock_returns_error_after_max_attempts() {
1260        use std::sync::atomic::{AtomicU32, Ordering};
1261
1262        let counter = Arc::new(AtomicU32::new(0));
1263        let counter_clone = counter.clone();
1264
1265        let result: Result<u32, TxError> =
1266            retry_on_deadlock(2, Duration::from_millis(1), |_attempt| {
1267                let c = counter_clone.clone();
1268                async move {
1269                    c.fetch_add(1, Ordering::SeqCst);
1270                    Err(TxError::CommitFailed(
1271                        "Deadlock found when trying to get lock".to_string(),
1272                    ))
1273                }
1274            })
1275            .await;
1276
1277        // 应返回最后一次的死锁错误
1278        assert!(result.is_err());
1279        assert_eq!(counter.load(Ordering::SeqCst), 2);
1280    }
1281
1282    #[tokio::test]
1283    async fn test_m8_retry_on_deadlock_does_not_retry_non_deadlock_errors() {
1284        use std::sync::atomic::{AtomicU32, Ordering};
1285
1286        let counter = Arc::new(AtomicU32::new(0));
1287        let counter_clone = counter.clone();
1288
1289        let result: Result<u32, TxError> =
1290            retry_on_deadlock(3, Duration::from_millis(1), |_attempt| {
1291                let c = counter_clone.clone();
1292                async move {
1293                    c.fetch_add(1, Ordering::SeqCst);
1294                    Err(TxError::CommitFailed("syntax error".to_string()))
1295                }
1296            })
1297            .await;
1298
1299        // 非死锁错误应立即返回,不重试
1300        assert!(result.is_err());
1301        assert_eq!(counter.load(Ordering::SeqCst), 1);
1302    }
1303
1304    #[test]
1305    fn test_m8_deadlock_error_display() {
1306        let err = TxError::DeadlockDetected {
1307            attempt: 2,
1308            max_attempts: 3,
1309        };
1310        let msg = format!("{}", err);
1311        assert!(msg.contains("2"));
1312        assert!(msg.contains("3"));
1313        assert!(msg.contains("Deadlock"));
1314    }
1315}