use crate::Dialect;
use crate::DbError;
use crate::Value;
use std::collections::HashMap;
pub trait OptimisticLock {
fn version_field() -> &'static str;
}
#[derive(Debug)]
pub enum LockError {
Conflict {
entity: String,
expected_version: i64,
},
MissingVersion {
field: &'static str,
},
InvalidVersion {
field: &'static str,
value: i64,
},
RetriesExhausted {
attempts: u32,
},
Other(DbError),
}
impl std::fmt::Display for LockError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
LockError::Conflict {
entity,
expected_version,
} => write!(
f,
"Optimistic lock conflict on {} (expected version {})",
entity, expected_version
),
LockError::MissingVersion { field } => {
write!(f, "Missing version value for field `{}`", field)
}
LockError::InvalidVersion { field, value } => {
write!(f, "Invalid version value for field `{}`: {}", field, value)
}
LockError::RetriesExhausted { attempts } => {
write!(f, "Retries exhausted after {} attempts", attempts)
}
LockError::Other(e) => write!(f, "Optimistic lock error: {}", e),
}
}
}
impl std::error::Error for LockError {}
impl From<DbError> for LockError {
fn from(e: DbError) -> Self {
LockError::Other(e)
}
}
pub type LockResult<T> = Result<T, LockError>;
pub fn build_update_with_lock(
dialect: &dyn Dialect,
table: &str,
pk_column: &str,
version_column: &str,
pk_value: &Value,
current_version: &Value,
data: &HashMap<String, Value>,
) -> String {
let quoted_table = dialect.quote(table);
let quoted_pk = dialect.quote(pk_column);
let quoted_version = dialect.quote(version_column);
let mut sets: Vec<String> = data
.iter()
.map(|(k, v)| {
format!(
"{} = {}",
dialect.quote(k),
v.to_param_with_dialect(dialect)
)
})
.collect();
sets.push(format!("{} = {} + 1", quoted_version, quoted_version));
let sets_sql = sets.join(", ");
format!(
"UPDATE {} SET {} WHERE {} = {} AND {} = {}",
quoted_table,
sets_sql,
quoted_pk,
pk_value.to_param_with_dialect(dialect),
quoted_version,
current_version.to_param_with_dialect(dialect),
)
}
pub fn build_delete_with_lock(
dialect: &dyn Dialect,
table: &str,
pk_column: &str,
version_column: &str,
pk_value: &Value,
current_version: &Value,
) -> String {
let quoted_table = dialect.quote(table);
let quoted_pk = dialect.quote(pk_column);
let quoted_version = dialect.quote(version_column);
format!(
"DELETE FROM {} WHERE {} = {} AND {} = {}",
quoted_table,
quoted_pk,
pk_value.to_param_with_dialect(dialect),
quoted_version,
current_version.to_param_with_dialect(dialect),
)
}
pub fn check_affected_rows(
affected: u64,
entity: impl Into<String>,
expected_version: i64,
) -> LockResult<()> {
if affected == 0 {
Err(LockError::Conflict {
entity: entity.into(),
expected_version,
})
} else {
Ok(())
}
}
pub fn extract_version(
row: &HashMap<String, Value>,
version_field: &'static str,
) -> LockResult<i64> {
match row.get(version_field) {
None => Err(LockError::MissingVersion {
field: version_field,
}),
Some(Value::I64(v)) => {
if *v < 0 {
Err(LockError::InvalidVersion {
field: version_field,
value: *v,
})
} else {
Ok(*v)
}
}
Some(Value::I32(v)) => {
if *v < 0 {
Err(LockError::InvalidVersion {
field: version_field,
value: *v as i64,
})
} else {
Ok(*v as i64)
}
}
Some(Value::U32(v)) => Ok(*v as i64),
Some(Value::U64(v)) => {
if *v > i64::MAX as u64 {
Err(LockError::InvalidVersion {
field: version_field,
value: *v as i64, })
} else {
Ok(*v as i64)
}
}
Some(other) => Err(LockError::InvalidVersion {
field: version_field,
value: other.as_i64().unwrap_or(-1),
}),
}
}
pub fn retry_on_conflict<F>(max_retries: u32, mut op: F) -> LockResult<()>
where
F: FnMut() -> LockResult<u64>,
{
let mut attempts = 0u32;
loop {
attempts += 1;
match op() {
Ok(affected) => {
if affected == 0 {
if attempts > max_retries {
return Err(LockError::RetriesExhausted { attempts });
}
continue;
}
return Ok(());
}
Err(LockError::Conflict { .. }) => {
if attempts > max_retries {
return Err(LockError::RetriesExhausted { attempts });
}
}
Err(e) => return Err(e),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::get_dialect;
use crate::DbType;
#[test]
fn test_build_update_with_lock_mysql() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut data = HashMap::new();
data.insert("name".to_string(), Value::String("alice".to_string()));
data.insert("age".to_string(), Value::I64(30));
let sql = build_update_with_lock(
&*dialect,
"users",
"id",
"version",
&Value::I64(1),
&Value::I64(5),
&data,
);
assert!(sql.starts_with("UPDATE `users` SET"));
assert!(sql.contains("`name` = 'alice'"));
assert!(sql.contains("`age` = 30"));
assert!(sql.contains("`version` = `version` + 1"));
assert!(sql.contains("WHERE `id` = 1 AND `version` = 5"));
}
#[test]
fn test_build_update_with_lock_postgres() {
let dialect = get_dialect(DbType::PostgreSQL).unwrap();
let mut data = HashMap::new();
data.insert("name".to_string(), Value::String("bob".to_string()));
let sql = build_update_with_lock(
&*dialect,
"products",
"id",
"version",
&Value::I64(42),
&Value::I64(3),
&data,
);
assert!(sql.contains("\"products\""));
assert!(sql.contains("\"name\" = 'bob'"));
assert!(sql.contains("\"version\" = \"version\" + 1"));
assert!(sql.contains("\"id\" = 42"));
assert!(sql.contains("\"version\" = 3"));
}
#[test]
fn test_build_update_with_lock_empty_data() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let data = HashMap::new();
let sql = build_update_with_lock(
&*dialect,
"users",
"id",
"version",
&Value::I64(1),
&Value::I64(0),
&data,
);
assert!(sql.contains("SET `version` = `version` + 1"));
assert!(sql.contains("`version` = 0"));
}
#[test]
fn test_build_update_with_lock_custom_version_field() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut data = HashMap::new();
data.insert("name".to_string(), Value::String("test".to_string()));
let sql = build_update_with_lock(
&*dialect,
"orders",
"order_id",
"lock_version",
&Value::I64(100),
&Value::I64(2),
&data,
);
assert!(sql.contains("`lock_version` = `lock_version` + 1"));
assert!(sql.contains("`order_id` = 100"));
assert!(sql.contains("`lock_version` = 2"));
}
#[test]
fn test_build_delete_with_lock_mysql() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let sql = build_delete_with_lock(
&*dialect,
"users",
"id",
"version",
&Value::I64(1),
&Value::I64(5),
);
assert_eq!(sql, "DELETE FROM `users` WHERE `id` = 1 AND `version` = 5");
}
#[test]
fn test_build_delete_with_lock_postgres() {
let dialect = get_dialect(DbType::PostgreSQL).unwrap();
let sql = build_delete_with_lock(
&*dialect,
"products",
"id",
"version",
&Value::I64(42),
&Value::I64(3),
);
assert_eq!(
sql,
"DELETE FROM \"products\" WHERE \"id\" = 42 AND \"version\" = 3"
);
}
#[test]
fn test_check_affected_rows_success() {
let result = check_affected_rows(1, "users#id=1", 5);
assert!(result.is_ok());
}
#[test]
fn test_check_affected_rows_conflict() {
let result = check_affected_rows(0, "users#id=1", 5);
assert!(matches!(
result,
Err(LockError::Conflict {
entity,
expected_version
}) if entity == "users#id=1" && expected_version == 5
));
}
#[test]
fn test_check_affected_rows_multi_rows_success() {
let result = check_affected_rows(5, "users#id=1", 5);
assert!(result.is_ok());
}
#[test]
fn test_extract_version_i64() {
let mut row = HashMap::new();
row.insert("version".to_string(), Value::I64(42));
let v = extract_version(&row, "version").unwrap();
assert_eq!(v, 42);
}
#[test]
fn test_extract_version_i32() {
let mut row = HashMap::new();
row.insert("version".to_string(), Value::I32(7));
let v = extract_version(&row, "version").unwrap();
assert_eq!(v, 7);
}
#[test]
fn test_extract_version_u32() {
let mut row = HashMap::new();
row.insert("version".to_string(), Value::U32(99));
let v = extract_version(&row, "version").unwrap();
assert_eq!(v, 99);
}
#[test]
fn test_extract_version_missing() {
let row = HashMap::new();
let result = extract_version(&row, "version");
assert!(matches!(result, Err(LockError::MissingVersion { field }) if field == "version"));
}
#[test]
fn test_extract_version_negative_invalid() {
let mut row = HashMap::new();
row.insert("version".to_string(), Value::I64(-1));
let result = extract_version(&row, "version");
assert!(matches!(
result,
Err(LockError::InvalidVersion { field, value }) if field == "version" && value == -1
));
}
#[test]
fn test_extract_version_wrong_type() {
let mut row = HashMap::new();
row.insert("version".to_string(), Value::String("abc".to_string()));
let result = extract_version(&row, "version");
assert!(matches!(result, Err(LockError::InvalidVersion { .. })));
}
#[test]
fn test_retry_on_conflict_immediate_success() {
let calls = std::cell::Cell::new(0u32);
let result: LockResult<()> = retry_on_conflict(3, || {
calls.set(calls.get() + 1);
Ok(1u64)
});
assert!(result.is_ok());
assert_eq!(calls.get(), 1);
}
#[test]
fn test_retry_on_conflict_after_one_failure() {
let calls = std::cell::Cell::new(0u32);
let result: LockResult<()> = retry_on_conflict(3, || {
calls.set(calls.get() + 1);
if calls.get() == 1 {
Err(LockError::Conflict {
entity: "x".to_string(),
expected_version: 1,
})
} else {
Ok(1u64)
}
});
assert!(result.is_ok());
assert_eq!(calls.get(), 2);
}
#[test]
fn test_retry_on_conflict_exhausted() {
let calls = std::cell::Cell::new(0u32);
let result: LockResult<()> = retry_on_conflict(2, || {
calls.set(calls.get() + 1);
Err(LockError::Conflict {
entity: "x".to_string(),
expected_version: 1,
})
});
assert!(matches!(result, Err(LockError::RetriesExhausted { .. })));
assert_eq!(calls.get(), 3);
}
#[test]
fn test_retry_on_conflict_zero_affected_treated_as_conflict() {
let calls = std::cell::Cell::new(0u32);
let result: LockResult<()> = retry_on_conflict(2, || {
calls.set(calls.get() + 1);
if calls.get() <= 1 {
Ok(0u64) } else {
Ok(1u64) }
});
assert!(result.is_ok());
assert_eq!(calls.get(), 2);
}
#[test]
fn test_retry_on_conflict_propagates_non_conflict_error() {
let calls = std::cell::Cell::new(0u32);
let result: LockResult<()> = retry_on_conflict(3, || {
calls.set(calls.get() + 1);
Err(LockError::MissingVersion { field: "version" })
});
assert!(matches!(result, Err(LockError::MissingVersion { .. })));
assert_eq!(calls.get(), 1); }
#[test]
fn test_lock_error_display_conflict() {
let e = LockError::Conflict {
entity: "users#id=1".to_string(),
expected_version: 5,
};
let s = format!("{}", e);
assert!(s.contains("Optimistic lock conflict"));
assert!(s.contains("users#id=1"));
assert!(s.contains("expected version 5"));
}
#[test]
fn test_lock_error_display_missing_version() {
let e = LockError::MissingVersion { field: "version" };
let s = format!("{}", e);
assert!(s.contains("Missing version value"));
assert!(s.contains("version"));
}
#[test]
fn test_lock_error_display_invalid_version() {
let e = LockError::InvalidVersion {
field: "version",
value: -1,
};
let s = format!("{}", e);
assert!(s.contains("Invalid version value"));
assert!(s.contains("-1"));
}
#[test]
fn test_lock_error_display_retries_exhausted() {
let e = LockError::RetriesExhausted { attempts: 5 };
let s = format!("{}", e);
assert!(s.contains("Retries exhausted"));
assert!(s.contains("5"));
}
struct Product {
_id: i64,
_version: i64,
}
impl OptimisticLock for Product {
fn version_field() -> &'static str {
"version"
}
}
#[test]
fn test_optimistic_lock_trait_implementable() {
assert_eq!(Product::version_field(), "version");
}
}