Skip to main content

sz_orm_core/
optimistic_lock.rs

1//! 乐观锁(Optimistic Locking)
2//!
3//! 对应文档 6.8 节改进项 20(乐观锁支持)。
4//!
5//! # 核心概念
6//!
7//! - **OptimisticLock**:Model trait,声明版本字段名
8//! - **build_update_with_lock**:生成 `UPDATE ... SET version = version + 1, ... WHERE pk = ? AND version = ?`
9//! - **LockError**:版本冲突错误
10//! - **retry_on_conflict**:冲突重试机制
11//!
12//! # 设计灵感
13//!
14//! - Hibernate `@Version`
15//! - Doctrine `@Version` / `LockMode::OPTIMISTIC`
16//! - Yii2 `OptimisticLockBehavior`
17//! - MyBatis-Plus `@Version`
18//!
19//! # 使用示例
20//!
21//! ```no_run
22//! use sz_orm_core::optimistic_lock::{OptimisticLock, build_update_with_lock};
23//! use sz_orm_core::{DbType, get_dialect, Value};
24//! use std::collections::HashMap;
25//!
26//! // 1. 定义带 version 字段的 Model
27//! struct Product {
28//!     id: i64,
29//!     name: String,
30//!     version: i64,
31//! }
32//!
33//! impl OptimisticLock for Product {
34//!     fn version_field() -> &'static str { "version" }
35//! }
36//!
37//! // 2. 生成乐观锁 UPDATE SQL
38//! let dialect = get_dialect(DbType::MySQL).unwrap();
39//! let mut data = HashMap::new();
40//! data.insert("name".to_string(), Value::String("new-name".to_string()));
41//! let sql = build_update_with_lock(&*dialect, "products", "id", "version", &Value::I64(1), &Value::I64(5), &data);
42//! // UPDATE `products` SET `name` = 'new-name', `version` = `version` + 1 WHERE `id` = 1 AND `version` = 5
43//! ```
44
45use crate::Dialect;
46use crate::DbError;
47use crate::Value;
48use std::collections::HashMap;
49
50// ============================================================================
51// OptimisticLock trait — Model 端声明
52// ============================================================================
53
54/// 乐观锁 trait — Model 实现此 trait 以声明版本字段名
55///
56/// 对应 Hibernate `@Version` / Doctrine `@Version` / MyBatis-Plus `@Version`。
57///
58/// # 示例
59///
60/// ```
61/// use sz_orm_core::optimistic_lock::OptimisticLock;
62///
63/// struct Product {
64///     id: i64,
65///     name: String,
66///     version: i64,
67/// }
68///
69/// impl OptimisticLock for Product {
70///     fn version_field() -> &'static str { "version" }
71/// }
72/// ```
73pub trait OptimisticLock {
74    /// 版本字段名(如 "version"、"lock_version"、"rev")
75    fn version_field() -> &'static str;
76}
77
78// ============================================================================
79// LockError — 乐观锁错误类型
80// ============================================================================
81
82/// 乐观锁错误类型
83#[derive(Debug)]
84pub enum LockError {
85    /// 版本冲突(受影响行数为 0)
86    ///
87    /// 携带 (entity, expected_version) 信息以便上层重试。
88    Conflict {
89        /// 实体描述(表名 + 主键,便于日志)
90        entity: String,
91        /// 期望的版本号
92        expected_version: i64,
93    },
94    /// 未指定版本号(必填)
95    MissingVersion {
96        /// 字段名
97        field: &'static str,
98    },
99    /// 版本号非法(负数或溢出)
100    InvalidVersion {
101        /// 字段名
102        field: &'static str,
103        /// 实际值
104        value: i64,
105    },
106    /// 重试次数耗尽
107    RetriesExhausted {
108        /// 已重试次数
109        attempts: u32,
110    },
111    /// 其他数据库错误
112    Other(DbError),
113}
114
115impl std::fmt::Display for LockError {
116    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
117        match self {
118            LockError::Conflict {
119                entity,
120                expected_version,
121            } => write!(
122                f,
123                "Optimistic lock conflict on {} (expected version {})",
124                entity, expected_version
125            ),
126            LockError::MissingVersion { field } => {
127                write!(f, "Missing version value for field `{}`", field)
128            }
129            LockError::InvalidVersion { field, value } => {
130                write!(f, "Invalid version value for field `{}`: {}", field, value)
131            }
132            LockError::RetriesExhausted { attempts } => {
133                write!(f, "Retries exhausted after {} attempts", attempts)
134            }
135            LockError::Other(e) => write!(f, "Optimistic lock error: {}", e),
136        }
137    }
138}
139
140impl std::error::Error for LockError {}
141
142impl From<DbError> for LockError {
143    fn from(e: DbError) -> Self {
144        LockError::Other(e)
145    }
146}
147
148/// 乐观锁结果
149pub type LockResult<T> = Result<T, LockError>;
150
151// ============================================================================
152// build_update_with_lock — 生成乐观锁 UPDATE SQL
153// ============================================================================
154
155/// 生成乐观锁 UPDATE SQL
156///
157/// 生成的 SQL 形如:
158/// ```sql
159/// UPDATE `table` SET col1 = ?, col2 = ?, `version` = `version` + 1
160/// WHERE `pk` = ? AND `version` = ?
161/// ```
162///
163/// # 参数
164/// - `dialect`:数据库方言
165/// - `table`:表名
166/// - `pk_column`:主键列名
167/// - `version_column`:版本列名
168/// - `pk_value`:主键值
169/// - `current_version`:当前版本号(从数据库读取)
170/// - `data`:要更新的字段(不含 version 字段,会自动追加)
171///
172/// # 示例
173///
174/// ```
175/// use sz_orm_core::optimistic_lock::build_update_with_lock;
176/// use sz_orm_core::{DbType, get_dialect, Value};
177/// use std::collections::HashMap;
178///
179/// let dialect = get_dialect(DbType::MySQL).unwrap();
180/// let mut data = HashMap::new();
181/// data.insert("name".to_string(), Value::String("alice".to_string()));
182/// let sql = build_update_with_lock(
183///     &*dialect, "users", "id", "version",
184///     &Value::I64(1), &Value::I64(5), &data,
185/// );
186/// assert!(sql.contains("UPDATE"));
187/// assert!(sql.contains("`users`"));
188/// assert!(sql.contains("`version` = `version` + 1"));
189/// assert!(sql.contains("`id` = 1"));
190/// assert!(sql.contains("`version` = 5"));
191/// ```
192pub fn build_update_with_lock(
193    dialect: &dyn Dialect,
194    table: &str,
195    pk_column: &str,
196    version_column: &str,
197    pk_value: &Value,
198    current_version: &Value,
199    data: &HashMap<String, Value>,
200) -> String {
201    let quoted_table = dialect.quote(table);
202    let quoted_pk = dialect.quote(pk_column);
203    let quoted_version = dialect.quote(version_column);
204
205    let mut sets: Vec<String> = data
206        .iter()
207        .map(|(k, v)| {
208            format!(
209                "{} = {}",
210                dialect.quote(k),
211                v.to_param_with_dialect(dialect)
212            )
213        })
214        .collect();
215    // version 字段自增(与数据中是否含 version 无关,强制使用 version = version + 1)
216    sets.push(format!("{} = {} + 1", quoted_version, quoted_version));
217
218    let sets_sql = sets.join(", ");
219
220    format!(
221        "UPDATE {} SET {} WHERE {} = {} AND {} = {}",
222        quoted_table,
223        sets_sql,
224        quoted_pk,
225        pk_value.to_param_with_dialect(dialect),
226        quoted_version,
227        current_version.to_param_with_dialect(dialect),
228    )
229}
230
231// ============================================================================
232// build_delete_with_lock — 生成乐观锁 DELETE SQL
233// ============================================================================
234
235/// 生成乐观锁 DELETE SQL(删除前校验版本号)
236///
237/// 生成的 SQL 形如:
238/// ```sql
239/// DELETE FROM `table` WHERE `pk` = ? AND `version` = ?
240/// ```
241pub fn build_delete_with_lock(
242    dialect: &dyn Dialect,
243    table: &str,
244    pk_column: &str,
245    version_column: &str,
246    pk_value: &Value,
247    current_version: &Value,
248) -> String {
249    let quoted_table = dialect.quote(table);
250    let quoted_pk = dialect.quote(pk_column);
251    let quoted_version = dialect.quote(version_column);
252
253    format!(
254        "DELETE FROM {} WHERE {} = {} AND {} = {}",
255        quoted_table,
256        quoted_pk,
257        pk_value.to_param_with_dialect(dialect),
258        quoted_version,
259        current_version.to_param_with_dialect(dialect),
260    )
261}
262
263// ============================================================================
264// check_affected_rows — 检查受影响行数判断是否冲突
265// ============================================================================
266
267/// 检查 UPDATE/DELETE 受影响行数,0 表示版本冲突
268///
269/// # 参数
270/// - `affected`:受影响行数
271/// - `entity`:实体描述(如 "products#id=1")
272/// - `expected_version`:期望的版本号
273pub fn check_affected_rows(
274    affected: u64,
275    entity: impl Into<String>,
276    expected_version: i64,
277) -> LockResult<()> {
278    if affected == 0 {
279        Err(LockError::Conflict {
280            entity: entity.into(),
281            expected_version,
282        })
283    } else {
284        Ok(())
285    }
286}
287
288// ============================================================================
289// extract_version — 从数据行中提取版本号
290// ============================================================================
291
292/// 从数据行(HashMap)中提取版本号
293///
294/// 若字段不存在或类型不匹配,返回 `LockError::MissingVersion` / `InvalidVersion`。
295pub fn extract_version(
296    row: &HashMap<String, Value>,
297    version_field: &'static str,
298) -> LockResult<i64> {
299    match row.get(version_field) {
300        None => Err(LockError::MissingVersion {
301            field: version_field,
302        }),
303        Some(Value::I64(v)) => {
304            if *v < 0 {
305                Err(LockError::InvalidVersion {
306                    field: version_field,
307                    value: *v,
308                })
309            } else {
310                Ok(*v)
311            }
312        }
313        Some(Value::I32(v)) => {
314            if *v < 0 {
315                Err(LockError::InvalidVersion {
316                    field: version_field,
317                    value: *v as i64,
318                })
319            } else {
320                Ok(*v as i64)
321            }
322        }
323        Some(Value::U32(v)) => Ok(*v as i64),
324        Some(Value::U64(v)) => {
325            if *v > i64::MAX as u64 {
326                Err(LockError::InvalidVersion {
327                    field: version_field,
328                    value: *v as i64, // 截断
329                })
330            } else {
331                Ok(*v as i64)
332            }
333        }
334        Some(other) => Err(LockError::InvalidVersion {
335            field: version_field,
336            value: other.as_i64().unwrap_or(-1),
337        }),
338    }
339}
340
341// ============================================================================
342// retry_on_conflict — 冲突重试机制
343// ============================================================================
344
345/// 在乐观锁冲突时自动重试
346///
347/// 重复调用 `op`,直到成功或重试次数耗尽。
348/// 每次 `op` 返回 `Err(LockError::Conflict)` 时,调用 `reload` 重新加载最新版本号后重试。
349///
350/// # 参数
351/// - `max_retries`:最大重试次数(不含首次调用)
352/// - `op`:执行更新操作,返回 `Result<u64, LockError>`(u64 = 受影响行数)
353///
354/// # 示例
355///
356/// ```no_run
357/// use sz_orm_core::optimistic_lock::{retry_on_conflict, LockResult, LockError};
358///
359/// let result: LockResult<()> = retry_on_conflict(3, || {
360///     // 模拟:第一次冲突,第二次成功
361///     static mut CALLS: u32 = 0;
362///     unsafe { CALLS += 1; }
363///     if unsafe { CALLS } == 1 {
364///         Err(LockError::Conflict { entity: "x".to_string(), expected_version: 1 })
365///     } else {
366///         Ok(1u64) // 1 行受影响
367///     }
368/// });
369/// ```
370pub fn retry_on_conflict<F>(max_retries: u32, mut op: F) -> LockResult<()>
371where
372    F: FnMut() -> LockResult<u64>,
373{
374    let mut attempts = 0u32;
375    loop {
376        attempts += 1;
377        match op() {
378            Ok(affected) => {
379                if affected == 0 {
380                    if attempts > max_retries {
381                        return Err(LockError::RetriesExhausted { attempts });
382                    }
383                    continue;
384                }
385                return Ok(());
386            }
387            Err(LockError::Conflict { .. }) => {
388                if attempts > max_retries {
389                    return Err(LockError::RetriesExhausted { attempts });
390                }
391                // 继续重试
392            }
393            Err(e) => return Err(e),
394        }
395    }
396}
397
398// ============================================================================
399// 单元测试
400// ============================================================================
401
402#[cfg(test)]
403mod tests {
404    use super::*;
405    use crate::get_dialect;
406    use crate::DbType;
407
408    // ===== build_update_with_lock 测试 =====
409
410    #[test]
411    fn test_build_update_with_lock_mysql() {
412        let dialect = get_dialect(DbType::MySQL).unwrap();
413        let mut data = HashMap::new();
414        data.insert("name".to_string(), Value::String("alice".to_string()));
415        data.insert("age".to_string(), Value::I64(30));
416
417        let sql = build_update_with_lock(
418            &*dialect,
419            "users",
420            "id",
421            "version",
422            &Value::I64(1),
423            &Value::I64(5),
424            &data,
425        );
426
427        // 应包含 UPDATE ... SET ... WHERE id = 1 AND version = 5
428        assert!(sql.starts_with("UPDATE `users` SET"));
429        assert!(sql.contains("`name` = 'alice'"));
430        assert!(sql.contains("`age` = 30"));
431        assert!(sql.contains("`version` = `version` + 1"));
432        assert!(sql.contains("WHERE `id` = 1 AND `version` = 5"));
433    }
434
435    #[test]
436    fn test_build_update_with_lock_postgres() {
437        let dialect = get_dialect(DbType::PostgreSQL).unwrap();
438        let mut data = HashMap::new();
439        data.insert("name".to_string(), Value::String("bob".to_string()));
440
441        let sql = build_update_with_lock(
442            &*dialect,
443            "products",
444            "id",
445            "version",
446            &Value::I64(42),
447            &Value::I64(3),
448            &data,
449        );
450
451        // PostgreSQL 使用双引号
452        assert!(sql.contains("\"products\""));
453        assert!(sql.contains("\"name\" = 'bob'"));
454        assert!(sql.contains("\"version\" = \"version\" + 1"));
455        assert!(sql.contains("\"id\" = 42"));
456        assert!(sql.contains("\"version\" = 3"));
457    }
458
459    #[test]
460    fn test_build_update_with_lock_empty_data() {
461        let dialect = get_dialect(DbType::MySQL).unwrap();
462        let data = HashMap::new();
463
464        let sql = build_update_with_lock(
465            &*dialect,
466            "users",
467            "id",
468            "version",
469            &Value::I64(1),
470            &Value::I64(0),
471            &data,
472        );
473
474        // 即使没有其他字段,也应包含 version 自增
475        assert!(sql.contains("SET `version` = `version` + 1"));
476        assert!(sql.contains("`version` = 0"));
477    }
478
479    #[test]
480    fn test_build_update_with_lock_custom_version_field() {
481        let dialect = get_dialect(DbType::MySQL).unwrap();
482        let mut data = HashMap::new();
483        data.insert("name".to_string(), Value::String("test".to_string()));
484
485        let sql = build_update_with_lock(
486            &*dialect,
487            "orders",
488            "order_id",
489            "lock_version",
490            &Value::I64(100),
491            &Value::I64(2),
492            &data,
493        );
494
495        assert!(sql.contains("`lock_version` = `lock_version` + 1"));
496        assert!(sql.contains("`order_id` = 100"));
497        assert!(sql.contains("`lock_version` = 2"));
498    }
499
500    // ===== build_delete_with_lock 测试 =====
501
502    #[test]
503    fn test_build_delete_with_lock_mysql() {
504        let dialect = get_dialect(DbType::MySQL).unwrap();
505        let sql = build_delete_with_lock(
506            &*dialect,
507            "users",
508            "id",
509            "version",
510            &Value::I64(1),
511            &Value::I64(5),
512        );
513
514        assert_eq!(sql, "DELETE FROM `users` WHERE `id` = 1 AND `version` = 5");
515    }
516
517    #[test]
518    fn test_build_delete_with_lock_postgres() {
519        let dialect = get_dialect(DbType::PostgreSQL).unwrap();
520        let sql = build_delete_with_lock(
521            &*dialect,
522            "products",
523            "id",
524            "version",
525            &Value::I64(42),
526            &Value::I64(3),
527        );
528
529        assert_eq!(
530            sql,
531            "DELETE FROM \"products\" WHERE \"id\" = 42 AND \"version\" = 3"
532        );
533    }
534
535    // ===== check_affected_rows 测试 =====
536
537    #[test]
538    fn test_check_affected_rows_success() {
539        let result = check_affected_rows(1, "users#id=1", 5);
540        assert!(result.is_ok());
541    }
542
543    #[test]
544    fn test_check_affected_rows_conflict() {
545        let result = check_affected_rows(0, "users#id=1", 5);
546        assert!(matches!(
547            result,
548            Err(LockError::Conflict {
549                entity,
550                expected_version
551            }) if entity == "users#id=1" && expected_version == 5
552        ));
553    }
554
555    #[test]
556    fn test_check_affected_rows_multi_rows_success() {
557        // 受影响多行也视为成功(虽然乐观锁通常一次只更新一行)
558        let result = check_affected_rows(5, "users#id=1", 5);
559        assert!(result.is_ok());
560    }
561
562    // ===== extract_version 测试 =====
563
564    #[test]
565    fn test_extract_version_i64() {
566        let mut row = HashMap::new();
567        row.insert("version".to_string(), Value::I64(42));
568        let v = extract_version(&row, "version").unwrap();
569        assert_eq!(v, 42);
570    }
571
572    #[test]
573    fn test_extract_version_i32() {
574        let mut row = HashMap::new();
575        row.insert("version".to_string(), Value::I32(7));
576        let v = extract_version(&row, "version").unwrap();
577        assert_eq!(v, 7);
578    }
579
580    #[test]
581    fn test_extract_version_u32() {
582        let mut row = HashMap::new();
583        row.insert("version".to_string(), Value::U32(99));
584        let v = extract_version(&row, "version").unwrap();
585        assert_eq!(v, 99);
586    }
587
588    #[test]
589    fn test_extract_version_missing() {
590        let row = HashMap::new();
591        let result = extract_version(&row, "version");
592        assert!(matches!(result, Err(LockError::MissingVersion { field }) if field == "version"));
593    }
594
595    #[test]
596    fn test_extract_version_negative_invalid() {
597        let mut row = HashMap::new();
598        row.insert("version".to_string(), Value::I64(-1));
599        let result = extract_version(&row, "version");
600        assert!(matches!(
601            result,
602            Err(LockError::InvalidVersion { field, value }) if field == "version" && value == -1
603        ));
604    }
605
606    #[test]
607    fn test_extract_version_wrong_type() {
608        let mut row = HashMap::new();
609        row.insert("version".to_string(), Value::String("abc".to_string()));
610        let result = extract_version(&row, "version");
611        assert!(matches!(result, Err(LockError::InvalidVersion { .. })));
612    }
613
614    // ===== retry_on_conflict 测试 =====
615
616    #[test]
617    fn test_retry_on_conflict_immediate_success() {
618        let calls = std::cell::Cell::new(0u32);
619        let result: LockResult<()> = retry_on_conflict(3, || {
620            calls.set(calls.get() + 1);
621            Ok(1u64)
622        });
623        assert!(result.is_ok());
624        assert_eq!(calls.get(), 1);
625    }
626
627    #[test]
628    fn test_retry_on_conflict_after_one_failure() {
629        let calls = std::cell::Cell::new(0u32);
630        let result: LockResult<()> = retry_on_conflict(3, || {
631            calls.set(calls.get() + 1);
632            if calls.get() == 1 {
633                Err(LockError::Conflict {
634                    entity: "x".to_string(),
635                    expected_version: 1,
636                })
637            } else {
638                Ok(1u64)
639            }
640        });
641        assert!(result.is_ok());
642        assert_eq!(calls.get(), 2);
643    }
644
645    #[test]
646    fn test_retry_on_conflict_exhausted() {
647        let calls = std::cell::Cell::new(0u32);
648        let result: LockResult<()> = retry_on_conflict(2, || {
649            calls.set(calls.get() + 1);
650            Err(LockError::Conflict {
651                entity: "x".to_string(),
652                expected_version: 1,
653            })
654        });
655        assert!(matches!(result, Err(LockError::RetriesExhausted { .. })));
656        // 1 initial + 2 retries = 3 calls
657        assert_eq!(calls.get(), 3);
658    }
659
660    #[test]
661    fn test_retry_on_conflict_zero_affected_treated_as_conflict() {
662        let calls = std::cell::Cell::new(0u32);
663        let result: LockResult<()> = retry_on_conflict(2, || {
664            calls.set(calls.get() + 1);
665            if calls.get() <= 1 {
666                Ok(0u64) // 0 行受影响 = 冲突
667            } else {
668                Ok(1u64) // 成功
669            }
670        });
671        assert!(result.is_ok());
672        assert_eq!(calls.get(), 2);
673    }
674
675    #[test]
676    fn test_retry_on_conflict_propagates_non_conflict_error() {
677        let calls = std::cell::Cell::new(0u32);
678        let result: LockResult<()> = retry_on_conflict(3, || {
679            calls.set(calls.get() + 1);
680            Err(LockError::MissingVersion { field: "version" })
681        });
682        assert!(matches!(result, Err(LockError::MissingVersion { .. })));
683        assert_eq!(calls.get(), 1); // 非 Conflict 错误立即返回,不重试
684    }
685
686    // ===== LockError Display 测试 =====
687
688    #[test]
689    fn test_lock_error_display_conflict() {
690        let e = LockError::Conflict {
691            entity: "users#id=1".to_string(),
692            expected_version: 5,
693        };
694        let s = format!("{}", e);
695        assert!(s.contains("Optimistic lock conflict"));
696        assert!(s.contains("users#id=1"));
697        assert!(s.contains("expected version 5"));
698    }
699
700    #[test]
701    fn test_lock_error_display_missing_version() {
702        let e = LockError::MissingVersion { field: "version" };
703        let s = format!("{}", e);
704        assert!(s.contains("Missing version value"));
705        assert!(s.contains("version"));
706    }
707
708    #[test]
709    fn test_lock_error_display_invalid_version() {
710        let e = LockError::InvalidVersion {
711            field: "version",
712            value: -1,
713        };
714        let s = format!("{}", e);
715        assert!(s.contains("Invalid version value"));
716        assert!(s.contains("-1"));
717    }
718
719    #[test]
720    fn test_lock_error_display_retries_exhausted() {
721        let e = LockError::RetriesExhausted { attempts: 5 };
722        let s = format!("{}", e);
723        assert!(s.contains("Retries exhausted"));
724        assert!(s.contains("5"));
725    }
726
727    // ===== OptimisticLock trait 测试 =====
728
729    struct Product {
730        _id: i64,
731        _version: i64,
732    }
733    impl OptimisticLock for Product {
734        fn version_field() -> &'static str {
735            "version"
736        }
737    }
738
739    #[test]
740    fn test_optimistic_lock_trait_implementable() {
741        // 验证 trait 可被实现
742        assert_eq!(Product::version_field(), "version");
743    }
744}