Skip to main content

sz_orm_core/
change_tracker.rs

1#![allow(missing_docs)]
2//! 实体变更跟踪器(Change Tracker)
3//!
4//! 对标 EF Core `ChangeTracker` / Hibernate `PersistenceContext`。
5//!
6//! 跟踪实体的生命周期状态(Added / Modified / Deleted / Unchanged / Detached),
7//! 在 `SaveChanges` 时自动生成相应的 INSERT / UPDATE / DELETE 语句。
8//!
9//! # 使用示例
10//!
11//! ```
12//! use sz_orm_core::change_tracker::{ChangeTracker, EntityEntry, EntityState};
13//! use sz_orm_core::Value;
14//! use std::collections::HashMap;
15//!
16//! let mut tracker = ChangeTracker::new();
17//!
18//! // 添加新实体
19//! let mut user: HashMap<String, Value> = HashMap::new();
20//! user.insert("name".to_string(), Value::String("alice".into()));
21//! tracker.track("users", "1", user.clone(), EntityState::Added);
22//!
23//! // 修改实体
24//! user.insert("name".to_string(), Value::String("bob".into()));
25//! tracker.track("users", "1", user, EntityState::Modified);
26//!
27//! let changes = tracker.get_pending_changes();
28//! assert_eq!(changes.len(), 1);
29//! assert_eq!(changes[0].state, EntityState::Modified);
30//! ```
31
32use crate::Value;
33use std::collections::HashMap;
34
35/// 实体状态
36#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum EntityState {
38    /// 未跟踪
39    Detached,
40    /// 未变更
41    Unchanged,
42    /// 新增
43    Added,
44    /// 已修改
45    Modified,
46    /// 已删除
47    Deleted,
48}
49
50impl EntityState {
51    /// 是否需要生成 SQL
52    pub fn is_pending(&self) -> bool {
53        matches!(
54            self,
55            EntityState::Added | EntityState::Modified | EntityState::Deleted
56        )
57    }
58
59    /// 对应的 SQL 操作名
60    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/// 实体键(表名 + 主键值)
78#[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/// 实体跟踪条目
94#[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    /// 获取脏字段(current 与 original 的差异)
104    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    /// 是否有变更
127    pub fn is_dirty(&self) -> bool {
128        self.state.is_pending() && !self.get_dirty_fields().is_empty()
129    }
130}
131
132/// 变更跟踪器
133///
134/// 管理多个实体的状态,提供批量变更检测和变更集生成。
135pub 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    /// 跟踪实体
153    ///
154    /// 如果状态为 `Unchanged`,会保存当前值作为原始值快照。
155    /// 如果状态为 `Modified`,会保留之前的原始值快照。
156    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    /// 标记为新增
192    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    /// 标记为未变更(从数据库加载后调用)
197    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    /// 标记为已修改
202    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    /// 标记为已删除
215    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    /// 分离实体(停止跟踪)
223    pub fn detach(&mut self, table: &str, id: &str) {
224        let key = EntityKey::new(table, id);
225        self.entries.remove(&key);
226    }
227
228    /// 自动检测变更
229    ///
230    /// 遍历所有 `Unchanged` 实体,如果 current 与 original 不同,自动转为 `Modified`。
231    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    /// 获取所有待提交的变更
240    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    /// 按表分组获取待提交的变更
248    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    /// 获取条目
262    pub fn entry(&self, table: &str, id: &str) -> Option<&EntityEntry> {
263        self.entries.get(&EntityKey::new(table, id))
264    }
265
266    /// 跟踪的实体数量
267    pub fn count(&self) -> usize {
268        self.entries.len()
269    }
270
271    /// 待提交的变更数量
272    pub fn pending_count(&self) -> usize {
273        self.entries
274            .values()
275            .filter(|e| e.state.is_pending())
276            .count()
277    }
278
279    /// 清除所有跟踪状态(SaveChanges 成功后调用)
280    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    /// 将待提交的变更转换为 SQL 语句列表(生产接线点)
293    ///
294    /// 这是 ChangeTracker 与 SQL 生成的集成点。
295    /// 生成的 SQL 使用参数化占位符 `?`,参数通过返回值一并给出。
296    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}