use crate::dialect::Dialect;
use crate::Value;
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct DirtyTracker {
original: HashMap<String, Value>,
current: HashMap<String, Value>,
}
impl DirtyTracker {
pub fn new(initial: HashMap<String, Value>) -> Self {
let original = initial.clone();
Self {
original,
current: initial,
}
}
pub fn empty() -> Self {
Self {
original: HashMap::new(),
current: HashMap::new(),
}
}
pub fn set(&mut self, field: impl Into<String>, value: Value) {
self.current.insert(field.into(), value);
}
pub fn set_many(&mut self, fields: HashMap<String, Value>) {
for (k, v) in fields {
self.current.insert(k, v);
}
}
pub fn get(&self, field: &str) -> Option<&Value> {
self.current.get(field)
}
pub fn get_original(&self, field: &str) -> Option<&Value> {
self.original.get(field)
}
pub fn current(&self) -> &HashMap<String, Value> {
&self.current
}
pub fn original(&self) -> &HashMap<String, Value> {
&self.original
}
pub fn is_dirty(&self) -> bool {
self.original.len() != self.current.len() || self.dirty_fields_iter().next().is_some()
}
pub fn is_field_dirty(&self, field: &str) -> bool {
match (self.original.get(field), self.current.get(field)) {
(None, None) => false,
(None, Some(_)) => true, (Some(_), None) => true, (Some(o), Some(c)) => o != c,
}
}
pub fn get_dirty_fields(&self) -> Vec<String> {
let mut dirty: Vec<String> = self.dirty_fields_iter().cloned().collect();
dirty.sort();
dirty
}
pub fn get_dirty_attributes(&self) -> HashMap<String, Value> {
let mut result = HashMap::new();
for field in self.dirty_fields_iter() {
if let Some(v) = self.current.get(field) {
result.insert(field.clone(), v.clone());
}
}
result
}
pub fn mark_clean(&mut self) {
self.original = self.current.clone();
}
pub fn rollback(&mut self) {
self.current = self.original.clone();
}
pub fn clear(&mut self) {
self.original.clear();
self.current.clear();
}
fn dirty_fields_iter(&self) -> impl Iterator<Item = &String> {
let keys: Vec<&String> = self.current.keys().collect();
keys.into_iter()
.filter(move |k| match self.original.get(*k) {
None => true, Some(o) => self.current.get(*k).map(|c| c != o).unwrap_or(true),
})
.chain(
self.original
.keys()
.filter(move |k| !self.current.contains_key(*k)),
)
}
}
pub fn build_dynamic_update(
dialect: &dyn Dialect,
table: &str,
pk_column: &str,
pk_value: &Value,
tracker: &DirtyTracker,
) -> Option<String> {
let dirty = tracker.get_dirty_attributes();
if dirty.is_empty() {
return None;
}
let quoted_table = dialect.quote(table);
let quoted_pk = dialect.quote(pk_column);
let mut fields: Vec<&String> = dirty.keys().collect();
fields.sort();
let sets: Vec<String> = fields
.iter()
.map(|k| {
format!(
"{} = {}",
dialect.quote(k),
dirty[*k].to_param_with_dialect(dialect)
)
})
.collect();
let sets_sql = sets.join(", ");
Some(format!(
"UPDATE {} SET {} WHERE {} = {}",
quoted_table,
sets_sql,
quoted_pk,
pk_value.to_param_with_dialect(dialect),
))
}
pub fn build_dynamic_insert(
dialect: &dyn Dialect,
table: &str,
data: &HashMap<String, Value>,
) -> Option<String> {
let non_null: Vec<(&String, &Value)> = data
.iter()
.filter(|(_, v)| !matches!(v, Value::Null))
.collect();
if non_null.is_empty() {
return None;
}
let mut sorted = non_null.clone();
sorted.sort_by(|a, b| a.0.cmp(b.0));
let quoted_table = dialect.quote(table);
let columns: Vec<String> = sorted.iter().map(|(k, _)| dialect.quote(k)).collect();
let values: Vec<String> = sorted
.iter()
.map(|(_, v)| v.to_param_with_dialect(dialect).to_string())
.collect();
Some(format!(
"INSERT INTO {} ({}) VALUES ({})",
quoted_table,
columns.join(", "),
values.join(", "),
))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::get_dialect;
use crate::DbType;
#[test]
fn test_new_tracker_no_dirty() {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
let tracker = DirtyTracker::new(row);
assert!(!tracker.is_dirty());
assert!(tracker.get_dirty_fields().is_empty());
}
#[test]
fn test_empty_tracker() {
let tracker = DirtyTracker::empty();
assert!(!tracker.is_dirty());
assert!(tracker.get_dirty_fields().is_empty());
}
#[test]
fn test_set_existing_field_makes_dirty() {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("bob".to_string()));
assert!(tracker.is_dirty());
assert!(tracker.is_field_dirty("name"));
assert!(!tracker.is_field_dirty("id"));
assert_eq!(tracker.get_dirty_fields(), vec!["name"]);
}
#[test]
fn test_set_new_field_makes_dirty() {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("alice".to_string()));
assert!(tracker.is_dirty());
assert!(tracker.is_field_dirty("name"));
assert_eq!(tracker.get_dirty_fields(), vec!["name"]);
}
#[test]
fn test_set_same_value_not_dirty() {
let mut row = HashMap::new();
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("alice".to_string()));
assert!(!tracker.is_dirty());
}
#[test]
fn test_set_int_value_not_dirty_when_same() {
let mut row = HashMap::new();
row.insert("age".to_string(), Value::I64(25));
let mut tracker = DirtyTracker::new(row);
tracker.set("age", Value::I64(25));
assert!(!tracker.is_dirty());
tracker.set("age", Value::I64(26));
assert!(tracker.is_dirty());
}
#[test]
fn test_set_null_makes_dirty_when_was_value() {
let mut row = HashMap::new();
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::Null);
assert!(tracker.is_dirty());
assert!(tracker.is_field_dirty("name"));
}
#[test]
fn test_get_original() {
let mut row = HashMap::new();
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("bob".to_string()));
assert_eq!(
tracker.get_original("name"),
Some(&Value::String("alice".to_string()))
);
assert_eq!(tracker.get("name"), Some(&Value::String("bob".to_string())));
}
#[test]
fn test_get_original_nonexistent() {
let tracker = DirtyTracker::empty();
assert_eq!(tracker.get_original("foo"), None);
}
#[test]
fn test_get_dirty_attributes() {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
row.insert("age".to_string(), Value::I64(25));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("bob".to_string()));
tracker.set("age", Value::I64(26));
let dirty = tracker.get_dirty_attributes();
assert_eq!(dirty.len(), 2);
assert_eq!(dirty.get("name"), Some(&Value::String("bob".to_string())));
assert_eq!(dirty.get("age"), Some(&Value::I64(26)));
}
#[test]
fn test_mark_clean() {
let mut row = HashMap::new();
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("bob".to_string()));
assert!(tracker.is_dirty());
tracker.mark_clean();
assert!(!tracker.is_dirty());
assert_eq!(
tracker.get_original("name"),
Some(&Value::String("bob".to_string()))
);
}
#[test]
fn test_rollback() {
let mut row = HashMap::new();
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("bob".to_string()));
assert!(tracker.is_dirty());
tracker.rollback();
assert!(!tracker.is_dirty());
assert_eq!(
tracker.get("name"),
Some(&Value::String("alice".to_string()))
);
}
#[test]
fn test_clear() {
let mut row = HashMap::new();
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.clear();
assert!(!tracker.is_dirty());
assert!(tracker.current().is_empty());
assert!(tracker.original().is_empty());
}
#[test]
fn test_set_many() {
let mut row = HashMap::new();
row.insert("a".to_string(), Value::I64(1));
let mut tracker = DirtyTracker::new(row);
let mut updates = HashMap::new();
updates.insert("a".to_string(), Value::I64(2));
updates.insert("b".to_string(), Value::I64(3));
tracker.set_many(updates);
assert!(tracker.is_dirty());
let dirty = tracker.get_dirty_fields();
assert!(dirty.contains(&"a".to_string()));
assert!(dirty.contains(&"b".to_string()));
}
#[test]
fn test_multiple_dirty_fields_sorted() {
let mut row = HashMap::new();
row.insert("z".to_string(), Value::I64(1));
row.insert("a".to_string(), Value::I64(1));
row.insert("m".to_string(), Value::I64(1));
let mut tracker = DirtyTracker::new(row);
tracker.set("z", Value::I64(2));
tracker.set("a", Value::I64(2));
tracker.set("m", Value::I64(2));
assert_eq!(tracker.get_dirty_fields(), vec!["a", "m", "z"]);
}
#[test]
fn test_remove_field_makes_dirty() {
let row = HashMap::new();
let tracker = DirtyTracker::new(row);
assert!(!tracker.is_field_dirty("name"));
}
#[test]
fn test_build_dynamic_update_with_dirty_fields() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
row.insert("age".to_string(), Value::I64(25));
let mut tracker = DirtyTracker::new(row);
tracker.set("age", Value::I64(26));
tracker.set("name", Value::String("bob".to_string()));
let sql = build_dynamic_update(&*dialect, "users", "id", &Value::I64(1), &tracker).unwrap();
assert!(sql.starts_with("UPDATE `users` SET"));
assert!(sql.contains("`age` = 26"));
assert!(sql.contains("`name` = 'bob'"));
assert!(sql.contains("WHERE `id` = 1"));
let set_clause = sql.split("WHERE").next().unwrap();
assert!(!set_clause.contains("`id` ="));
}
#[test]
fn test_build_dynamic_update_no_dirty_returns_none() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
let tracker = DirtyTracker::new(row);
let result = build_dynamic_update(&*dialect, "users", "id", &Value::I64(1), &tracker);
assert!(result.is_none());
}
#[test]
fn test_build_dynamic_update_postgres() {
let dialect = get_dialect(DbType::PostgreSQL).unwrap();
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("bob".to_string()));
let sql = build_dynamic_update(&*dialect, "users", "id", &Value::I64(1), &tracker).unwrap();
assert!(sql.contains("\"users\""));
assert!(sql.contains("\"name\" = 'bob'"));
assert!(sql.contains("\"id\" = 1"));
}
#[test]
fn test_build_dynamic_update_single_dirty_field() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
row.insert("age".to_string(), Value::I64(25));
let mut tracker = DirtyTracker::new(row);
tracker.set("age", Value::I64(26));
let sql = build_dynamic_update(&*dialect, "users", "id", &Value::I64(1), &tracker).unwrap();
assert!(sql.contains("`age` = 26"));
assert!(!sql.contains("`name`"));
}
#[test]
fn test_build_dynamic_update_after_mark_clean() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("bob".to_string()));
tracker.mark_clean();
let result = build_dynamic_update(&*dialect, "users", "id", &Value::I64(1), &tracker);
assert!(result.is_none());
}
#[test]
fn test_build_dynamic_insert_filters_null() {
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(25));
data.insert("bio".to_string(), Value::Null);
let sql = build_dynamic_insert(&*dialect, "users", &data).unwrap();
assert!(sql.starts_with("INSERT INTO `users`"));
assert!(sql.contains("`name`"));
assert!(sql.contains("`age`"));
assert!(!sql.contains("`bio`"));
assert!(sql.contains("'alice'"));
assert!(sql.contains("25"));
}
#[test]
fn test_build_dynamic_insert_all_null_returns_none() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut data = HashMap::new();
data.insert("a".to_string(), Value::Null);
data.insert("b".to_string(), Value::Null);
let result = build_dynamic_insert(&*dialect, "users", &data);
assert!(result.is_none());
}
#[test]
fn test_build_dynamic_insert_empty_data_returns_none() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let data = HashMap::new();
let result = build_dynamic_insert(&*dialect, "users", &data);
assert!(result.is_none());
}
#[test]
fn test_build_dynamic_insert_postgres() {
let dialect = get_dialect(DbType::PostgreSQL).unwrap();
let mut data = HashMap::new();
data.insert("name".to_string(), Value::String("alice".to_string()));
data.insert("age".to_string(), Value::I64(25));
let sql = build_dynamic_insert(&*dialect, "users", &data).unwrap();
assert!(sql.contains("INSERT INTO \"users\""));
assert!(sql.contains("\"name\""));
assert!(sql.contains("\"age\""));
assert!(sql.contains("'alice'"));
assert!(sql.contains("25"));
}
#[test]
fn test_build_dynamic_insert_columns_and_values_aligned() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut data = HashMap::new();
data.insert("a".to_string(), Value::I64(1));
data.insert("b".to_string(), Value::I64(2));
data.insert("c".to_string(), Value::I64(3));
let sql = build_dynamic_insert(&*dialect, "test", &data).unwrap();
let cols_start = sql.find('(').unwrap();
let cols_end = sql.find(") VALUES").unwrap();
let cols = &sql[cols_start + 1..cols_end];
let vals_start = sql.rfind('(').unwrap();
let vals_end = sql.rfind(')').unwrap();
let vals = &sql[vals_start + 1..vals_end];
let col_count = cols.split(',').count();
let val_count = vals.split(',').count();
assert_eq!(col_count, val_count);
assert_eq!(col_count, 3);
}
#[test]
fn test_build_dynamic_insert_with_bool() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut data = HashMap::new();
data.insert("active".to_string(), Value::Bool(true));
data.insert("name".to_string(), Value::String("alice".to_string()));
let sql = build_dynamic_insert(&*dialect, "users", &data).unwrap();
assert!(sql.contains("TRUE"));
assert!(sql.contains("'alice'"));
}
#[test]
fn test_workflow_load_modify_save() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
row.insert("name".to_string(), Value::String("alice".to_string()));
row.insert("age".to_string(), Value::I64(25));
row.insert(
"updated_at".to_string(),
Value::String("2026-01-01".to_string()),
);
let mut tracker = DirtyTracker::new(row);
tracker.set("age", Value::I64(26));
tracker.set("updated_at", Value::String("2026-07-19".to_string()));
let sql = build_dynamic_update(&*dialect, "users", "id", &Value::I64(1), &tracker).unwrap();
assert!(sql.contains("`age` = 26"));
assert!(sql.contains("`updated_at` = '2026-07-19'"));
let set_clause = sql.split("WHERE").next().unwrap();
assert!(!set_clause.contains("`name`"));
assert!(!set_clause.contains("`id` ="));
tracker.mark_clean();
assert!(!tracker.is_dirty());
tracker.set("name", Value::String("bob".to_string()));
assert!(tracker.is_dirty());
assert_eq!(tracker.get_dirty_fields(), vec!["name"]);
}
#[test]
fn test_workflow_insert_with_optional_fields() {
let dialect = get_dialect(DbType::MySQL).unwrap();
let mut data = HashMap::new();
data.insert("name".to_string(), Value::String("alice".to_string()));
data.insert(
"email".to_string(),
Value::String("alice@example.com".to_string()),
);
data.insert("bio".to_string(), Value::Null);
data.insert("age".to_string(), Value::I64(25));
let sql = build_dynamic_insert(&*dialect, "users", &data).unwrap();
assert!(!sql.contains("`bio`"));
assert!(sql.contains("`name`"));
assert!(sql.contains("`email`"));
assert!(sql.contains("`age`"));
}
#[test]
fn test_workflow_rollback_on_failure() {
let mut row = HashMap::new();
row.insert("name".to_string(), Value::String("alice".to_string()));
row.insert("age".to_string(), Value::I64(25));
let mut tracker = DirtyTracker::new(row);
tracker.set("name", Value::String("bob".to_string()));
tracker.set("age", Value::I64(99));
tracker.rollback();
assert!(!tracker.is_dirty());
assert_eq!(
tracker.get("name"),
Some(&Value::String("alice".to_string()))
);
assert_eq!(tracker.get("age"), Some(&Value::I64(25)));
}
}