1use crate::Dialect;
46use crate::DbError;
47use crate::Value;
48use std::collections::HashMap;
49
50pub trait OptimisticLock {
74 fn version_field() -> &'static str;
76}
77
78#[derive(Debug)]
84pub enum LockError {
85 Conflict {
89 entity: String,
91 expected_version: i64,
93 },
94 MissingVersion {
96 field: &'static str,
98 },
99 InvalidVersion {
101 field: &'static str,
103 value: i64,
105 },
106 RetriesExhausted {
108 attempts: u32,
110 },
111 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
148pub type LockResult<T> = Result<T, LockError>;
150
151pub 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 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
231pub 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
263pub 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
288pub 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, })
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
341pub 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 }
393 Err(e) => return Err(e),
394 }
395 }
396}
397
398#[cfg(test)]
403mod tests {
404 use super::*;
405 use crate::get_dialect;
406 use crate::DbType;
407
408 #[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 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 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 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 #[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 #[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 let result = check_affected_rows(5, "users#id=1", 5);
559 assert!(result.is_ok());
560 }
561
562 #[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 #[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 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) } else {
668 Ok(1u64) }
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); }
685
686 #[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 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 assert_eq!(Product::version_field(), "version");
743 }
744}