1#![allow(missing_docs)]
2use crate::Value;
33use std::collections::HashMap;
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum EntityState {
38 Detached,
40 Unchanged,
42 Added,
44 Modified,
46 Deleted,
48}
49
50impl EntityState {
51 pub fn is_pending(&self) -> bool {
53 matches!(
54 self,
55 EntityState::Added | EntityState::Modified | EntityState::Deleted
56 )
57 }
58
59 pub fn as_sql_op(&self) -> &'static str {
61 match self {
62 EntityState::Added => "INSERT",
63 EntityState::Modified => "UPDATE",
64 EntityState::Deleted => "DELETE",
65 EntityState::Unchanged => "NOOP",
66 EntityState::Detached => "NOOP",
67 }
68 }
69}
70
71impl std::fmt::Display for EntityState {
72 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
73 write!(f, "{:?}", self)
74 }
75}
76
77#[derive(Debug, Clone, PartialEq, Eq, Hash)]
79pub struct EntityKey {
80 pub table: String,
81 pub id: String,
82}
83
84impl EntityKey {
85 pub fn new(table: &str, id: &str) -> Self {
86 Self {
87 table: table.to_string(),
88 id: id.to_string(),
89 }
90 }
91}
92
93#[derive(Debug, Clone)]
95pub struct EntityEntry {
96 pub key: EntityKey,
97 pub current: HashMap<String, Value>,
98 pub original: Option<HashMap<String, Value>>,
99 pub state: EntityState,
100}
101
102impl EntityEntry {
103 pub fn get_dirty_fields(&self) -> Vec<String> {
105 match &self.original {
106 Some(orig) => {
107 let mut dirty = Vec::new();
108 for (k, v) in &self.current {
109 match orig.get(k) {
110 Some(orig_v) if orig_v != v => dirty.push(k.clone()),
111 None => dirty.push(k.clone()),
112 _ => {}
113 }
114 }
115 for k in orig.keys() {
116 if !self.current.contains_key(k) {
117 dirty.push(k.clone());
118 }
119 }
120 dirty
121 }
122 None => self.current.keys().cloned().collect(),
123 }
124 }
125
126 pub fn is_dirty(&self) -> bool {
128 self.state.is_pending() && !self.get_dirty_fields().is_empty()
129 }
130}
131
132pub struct ChangeTracker {
136 entries: HashMap<EntityKey, EntityEntry>,
137}
138
139impl Default for ChangeTracker {
140 fn default() -> Self {
141 Self::new()
142 }
143}
144
145impl ChangeTracker {
146 pub fn new() -> Self {
147 Self {
148 entries: HashMap::new(),
149 }
150 }
151
152 pub fn track(
157 &mut self,
158 table: &str,
159 id: &str,
160 current: HashMap<String, Value>,
161 state: EntityState,
162 ) {
163 let key = EntityKey::new(table, id);
164 let original = match state {
165 EntityState::Added => None,
166 EntityState::Unchanged => Some(current.clone()),
167 EntityState::Modified => self
168 .entries
169 .get(&key)
170 .and_then(|e| e.original.clone())
171 .or_else(|| Some(current.clone())),
172 EntityState::Deleted => self
173 .entries
174 .get(&key)
175 .and_then(|e| e.original.clone())
176 .or_else(|| Some(current.clone())),
177 EntityState::Detached => None,
178 };
179
180 self.entries.insert(
181 key,
182 EntityEntry {
183 key: EntityKey::new(table, id),
184 current,
185 original,
186 state,
187 },
188 );
189 }
190
191 pub fn mark_added(&mut self, table: &str, id: &str, entity: HashMap<String, Value>) {
193 self.track(table, id, entity, EntityState::Added);
194 }
195
196 pub fn mark_unchanged(&mut self, table: &str, id: &str, entity: HashMap<String, Value>) {
198 self.track(table, id, entity, EntityState::Unchanged);
199 }
200
201 pub fn update(&mut self, table: &str, id: &str, entity: HashMap<String, Value>) {
203 let key = EntityKey::new(table, id);
204 if let Some(entry) = self.entries.get_mut(&key) {
205 entry.current = entity;
206 if entry.state == EntityState::Unchanged {
207 entry.state = EntityState::Modified;
208 }
209 } else {
210 self.track(table, id, entity, EntityState::Modified);
211 }
212 }
213
214 pub fn mark_deleted(&mut self, table: &str, id: &str) {
216 let key = EntityKey::new(table, id);
217 if let Some(entry) = self.entries.get_mut(&key) {
218 entry.state = EntityState::Deleted;
219 }
220 }
221
222 pub fn detach(&mut self, table: &str, id: &str) {
224 let key = EntityKey::new(table, id);
225 self.entries.remove(&key);
226 }
227
228 pub fn detect_changes(&mut self) {
232 for entry in self.entries.values_mut() {
233 if entry.state == EntityState::Unchanged && !entry.get_dirty_fields().is_empty() {
234 entry.state = EntityState::Modified;
235 }
236 }
237 }
238
239 pub fn get_pending_changes(&self) -> Vec<&EntityEntry> {
241 self.entries
242 .values()
243 .filter(|e| e.state.is_pending())
244 .collect()
245 }
246
247 pub fn get_pending_changes_by_table(&self) -> HashMap<String, Vec<&EntityEntry>> {
249 let mut result: HashMap<String, Vec<&EntityEntry>> = HashMap::new();
250 for entry in self.entries.values() {
251 if entry.state.is_pending() {
252 result
253 .entry(entry.key.table.clone())
254 .or_default()
255 .push(entry);
256 }
257 }
258 result
259 }
260
261 pub fn entry(&self, table: &str, id: &str) -> Option<&EntityEntry> {
263 self.entries.get(&EntityKey::new(table, id))
264 }
265
266 pub fn count(&self) -> usize {
268 self.entries.len()
269 }
270
271 pub fn pending_count(&self) -> usize {
273 self.entries
274 .values()
275 .filter(|e| e.state.is_pending())
276 .count()
277 }
278
279 pub fn accept_changes(&mut self) {
281 self.entries.retain(|_, entry| {
282 if entry.state == EntityState::Deleted {
283 false
284 } else {
285 entry.state = EntityState::Unchanged;
286 entry.original = Some(entry.current.clone());
287 true
288 }
289 });
290 }
291
292 pub fn build_sql_operations(&self) -> Vec<(String, Vec<Value>)> {
297 let mut ops = Vec::new();
298 for entry in self.entries.values() {
299 if !entry.state.is_pending() {
300 continue;
301 }
302 match entry.state {
303 EntityState::Added => {
304 let columns: Vec<&str> = entry.current.keys().map(|s| s.as_str()).collect();
305 let placeholders: Vec<&str> = columns.iter().map(|_| "?").collect();
306 let params: Vec<Value> = entry.current.values().cloned().collect();
307 let sql = format!(
308 "INSERT INTO {} ({}) VALUES ({})",
309 entry.key.table,
310 columns.join(", "),
311 placeholders.join(", ")
312 );
313 ops.push((sql, params));
314 }
315 EntityState::Modified => {
316 let dirty = entry.get_dirty_fields();
317 if dirty.is_empty() {
318 continue;
319 }
320 let set_clauses: Vec<String> =
321 dirty.iter().map(|c| format!("{} = ?", c)).collect();
322 let mut params: Vec<Value> = dirty
323 .iter()
324 .filter_map(|c| entry.current.get(c).cloned())
325 .collect();
326 params.push(Value::String(entry.key.id.clone()));
327 let sql = format!(
328 "UPDATE {} SET {} WHERE id = ?",
329 entry.key.table,
330 set_clauses.join(", ")
331 );
332 ops.push((sql, params));
333 }
334 EntityState::Deleted => {
335 let sql = format!("DELETE FROM {} WHERE id = ?", entry.key.table);
336 ops.push((sql, vec![Value::String(entry.key.id.clone())]));
337 }
338 _ => {}
339 }
340 }
341 ops
342 }
343}
344
345#[cfg(test)]
346mod tests {
347 use super::*;
348
349 fn make_entity(name: &str, age: i64) -> HashMap<String, Value> {
350 let mut m = HashMap::new();
351 m.insert("name".to_string(), Value::String(name.to_string()));
352 m.insert("age".to_string(), Value::I64(age));
353 m
354 }
355
356 #[test]
357 fn test_track_added() {
358 let mut tracker = ChangeTracker::new();
359 tracker.mark_added("users", "1", make_entity("alice", 25));
360 assert_eq!(tracker.pending_count(), 1);
361 let changes = tracker.get_pending_changes();
362 assert_eq!(changes[0].state, EntityState::Added);
363 }
364
365 #[test]
366 fn test_track_unchanged_then_detect_modified() {
367 let mut tracker = ChangeTracker::new();
368 tracker.mark_unchanged("users", "1", make_entity("alice", 25));
369 assert_eq!(tracker.pending_count(), 0);
370
371 tracker.update("users", "1", make_entity("alice", 26));
372 tracker.detect_changes();
373 assert_eq!(tracker.pending_count(), 1);
374 let entry = tracker.entry("users", "1").unwrap();
375 assert_eq!(entry.state, EntityState::Modified);
376 assert_eq!(entry.get_dirty_fields(), vec!["age"]);
377 }
378
379 #[test]
380 fn test_mark_deleted() {
381 let mut tracker = ChangeTracker::new();
382 tracker.mark_unchanged("users", "1", make_entity("alice", 25));
383 assert_eq!(tracker.pending_count(), 0);
384
385 tracker.mark_deleted("users", "1");
386 assert_eq!(tracker.pending_count(), 1);
387 let entry = tracker.entry("users", "1").unwrap();
388 assert_eq!(entry.state, EntityState::Deleted);
389 }
390
391 #[test]
392 fn test_accept_changes() {
393 let mut tracker = ChangeTracker::new();
394 tracker.mark_added("users", "1", make_entity("alice", 25));
395 tracker.mark_unchanged("users", "2", make_entity("bob", 30));
396 tracker.mark_unchanged("users", "3", make_entity("charlie", 35));
397 tracker.mark_deleted("users", "3");
398
399 assert_eq!(tracker.count(), 3);
400 tracker.accept_changes();
401
402 assert_eq!(tracker.count(), 2);
403 assert_eq!(tracker.pending_count(), 0);
404 assert!(tracker.entry("users", "3").is_none());
405 }
406
407 #[test]
408 fn test_pending_changes_by_table() {
409 let mut tracker = ChangeTracker::new();
410 tracker.mark_added("users", "1", make_entity("alice", 25));
411 tracker.mark_added("orders", "1", make_entity("order1", 100));
412
413 let by_table = tracker.get_pending_changes_by_table();
414 assert_eq!(by_table["users"].len(), 1);
415 assert_eq!(by_table["orders"].len(), 1);
416 }
417
418 #[test]
419 fn test_detach() {
420 let mut tracker = ChangeTracker::new();
421 tracker.mark_added("users", "1", make_entity("alice", 25));
422 assert_eq!(tracker.count(), 1);
423
424 tracker.detach("users", "1");
425 assert_eq!(tracker.count(), 0);
426 assert!(tracker.entry("users", "1").is_none());
427 }
428
429 #[test]
430 fn test_entity_state_is_pending() {
431 assert!(EntityState::Added.is_pending());
432 assert!(EntityState::Modified.is_pending());
433 assert!(EntityState::Deleted.is_pending());
434 assert!(!EntityState::Unchanged.is_pending());
435 assert!(!EntityState::Detached.is_pending());
436 }
437
438 #[test]
439 fn test_entity_state_as_sql_op() {
440 assert_eq!(EntityState::Added.as_sql_op(), "INSERT");
441 assert_eq!(EntityState::Modified.as_sql_op(), "UPDATE");
442 assert_eq!(EntityState::Deleted.as_sql_op(), "DELETE");
443 assert_eq!(EntityState::Unchanged.as_sql_op(), "NOOP");
444 }
445
446 #[test]
447 fn test_dirty_fields_detection() {
448 let mut tracker = ChangeTracker::new();
449 tracker.mark_unchanged("users", "1", make_entity("alice", 25));
450
451 let mut modified = make_entity("alice", 26);
452 modified.insert("email".to_string(), Value::String("new@email.com".into()));
453 tracker.update("users", "1", modified);
454
455 let entry = tracker.entry("users", "1").unwrap();
456 let dirty = entry.get_dirty_fields();
457 assert!(dirty.contains(&"age".to_string()));
458 assert!(dirty.contains(&"email".to_string()));
459 }
460
461 #[test]
462 fn test_multiple_entities_same_table() {
463 let mut tracker = ChangeTracker::new();
464 tracker.mark_added("users", "1", make_entity("alice", 25));
465 tracker.mark_added("users", "2", make_entity("bob", 30));
466 tracker.mark_added("users", "3", make_entity("charlie", 35));
467
468 assert_eq!(tracker.pending_count(), 3);
469 let by_table = tracker.get_pending_changes_by_table();
470 assert_eq!(by_table["users"].len(), 3);
471 }
472
473 #[test]
474 fn test_e2e_change_tracker_to_sql() {
475 let mut tracker = ChangeTracker::new();
476
477 tracker.mark_added("users", "1", make_entity("alice", 25));
478 tracker.mark_unchanged("users", "2", make_entity("bob", 30));
479 tracker.update("users", "2", make_entity("bob", 31));
480 tracker.mark_unchanged("users", "3", make_entity("charlie", 35));
481 tracker.mark_deleted("users", "3");
482
483 tracker.detect_changes();
484 let ops = tracker.build_sql_operations();
485 assert_eq!(ops.len(), 3);
486
487 let has_insert = ops
488 .iter()
489 .any(|(sql, _)| sql.starts_with("INSERT INTO users"));
490 let has_update = ops
491 .iter()
492 .any(|(sql, _)| sql.starts_with("UPDATE users SET"));
493 let has_delete = ops
494 .iter()
495 .any(|(sql, _)| sql.starts_with("DELETE FROM users"));
496 assert!(has_insert, "missing INSERT");
497 assert!(has_update, "missing UPDATE");
498 assert!(has_delete, "missing DELETE");
499
500 let update_op = ops
501 .iter()
502 .find(|(sql, _)| sql.starts_with("UPDATE"))
503 .unwrap();
504 assert!(update_op.0.contains("age = ?"));
505 assert_eq!(update_op.1.len(), 2);
506 }
507
508 #[test]
509 fn test_e2e_change_tracker_accept_then_clean() {
510 let mut tracker = ChangeTracker::new();
511 tracker.mark_added("users", "1", make_entity("alice", 25));
512 assert_eq!(tracker.pending_count(), 1);
513
514 let ops = tracker.build_sql_operations();
515 assert_eq!(ops.len(), 1);
516
517 tracker.accept_changes();
518 assert_eq!(tracker.pending_count(), 0);
519 let ops_after = tracker.build_sql_operations();
520 assert_eq!(ops_after.len(), 0);
521 }
522}