Skip to main content

sz_orm_core/seeding/
manager.rs

1//! SeedManager — 种子版本管理 + 依赖排序 + 幂等执行 + 环境隔离
2
3use super::{FixtureTemplate, SeedError};
4use std::collections::{HashMap, HashSet};
5use std::time::Duration;
6
7/// 幂等执行模式
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum SeedMode {
10    /// INSERT ... ON CONFLICT UPDATE
11    Upsert,
12    /// TRUNCATE + INSERT
13    TruncateInsert,
14}
15
16/// 执行环境
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum SeedEnv {
19    /// 开发环境
20    Dev,
21    /// 测试环境
22    Test,
23    /// 预发布环境
24    Staging,
25    /// 生产环境
26    Production,
27}
28
29/// 种子文件
30#[derive(Debug, Clone)]
31pub struct SeedFile {
32    /// 种子版本号
33    pub version: String,
34    /// 种子描述
35    pub description: String,
36    /// 依赖的种子版本列表
37    pub dependencies: Vec<String>,
38    /// fixture 模板
39    pub template: FixtureTemplate,
40}
41
42/// 执行报告
43#[derive(Debug, Clone)]
44pub struct SeedReport {
45    /// 已执行的种子版本列表
46    pub executed_seeds: Vec<String>,
47    /// 总插入行数
48    pub total_rows: u64,
49    /// 总耗时
50    pub total_duration: Duration,
51    /// 是否幂等执行
52    pub idempotent: bool,
53    /// 执行环境
54    pub env: SeedEnv,
55}
56
57/// 种子管理器
58pub struct SeedManager {
59    seeds: Vec<SeedFile>,
60    mode: SeedMode,
61    env: SeedEnv,
62    allow_production: bool,
63    executed_versions: HashSet<String>,
64}
65
66impl SeedManager {
67    /// 创建新的 SeedManager
68    pub fn new(mode: SeedMode, env: SeedEnv) -> Self {
69        Self {
70            seeds: Vec::new(),
71            mode,
72            env,
73            allow_production: false,
74            executed_versions: HashSet::new(),
75        }
76    }
77
78    /// 允许生产环境执行
79    pub fn allow_production(mut self) -> Self {
80        self.allow_production = true;
81        self
82    }
83
84    /// 添加种子文件
85    pub fn add_seed(&mut self, seed: SeedFile) {
86        self.seeds.push(seed);
87    }
88
89    /// 从目录加载种子文件
90    pub fn load_seeds(&mut self, templates: Vec<FixtureTemplate>) -> Result<(), SeedError> {
91        for (i, template) in templates.into_iter().enumerate() {
92            self.seeds.push(SeedFile {
93                version: format!("v{}", i + 1),
94                description: format!("seed for {}", template.table),
95                dependencies: Vec::new(),
96                template,
97            });
98        }
99        Ok(())
100    }
101
102    /// 拓扑排序(按依赖关系)
103    pub fn topological_sort(&self) -> Result<Vec<&SeedFile>, SeedError> {
104        let mut sorted = Vec::new();
105        let mut visited = HashSet::new();
106        let mut visiting = HashSet::new();
107        let seed_map: HashMap<&str, &SeedFile> =
108            self.seeds.iter().map(|s| (s.version.as_str(), s)).collect();
109        for seed in &self.seeds {
110            self.visit(
111                seed,
112                &seed_map,
113                &mut sorted,
114                &mut visited,
115                &mut visiting,
116                Vec::new(),
117            )?;
118        }
119        Ok(sorted)
120    }
121
122    fn visit<'a>(
123        &'a self,
124        seed: &'a SeedFile,
125        seed_map: &HashMap<&str, &'a SeedFile>,
126        sorted: &mut Vec<&'a SeedFile>,
127        visited: &mut HashSet<String>,
128        visiting: &mut HashSet<String>,
129        path: Vec<String>,
130    ) -> Result<(), SeedError> {
131        if visited.contains(&seed.version) {
132            return Ok(());
133        }
134        if visiting.contains(&seed.version) {
135            let chain = path.join(" <- ");
136            return Err(SeedError::DependencyCycle { chain });
137        }
138        visiting.insert(seed.version.clone());
139        let mut new_path = path.clone();
140        new_path.push(seed.version.clone());
141        for dep in &seed.dependencies {
142            if let Some(dep_seed) = seed_map.get(dep.as_str()) {
143                self.visit(
144                    dep_seed,
145                    seed_map,
146                    sorted,
147                    visited,
148                    visiting,
149                    new_path.clone(),
150                )?;
151            }
152        }
153        visiting.remove(&seed.version);
154        visited.insert(seed.version.clone());
155        sorted.push(seed);
156        Ok(())
157    }
158
159    /// 环境隔离检查
160    pub fn check_env(&self) -> Result<(), SeedError> {
161        if self.env == SeedEnv::Production && !self.allow_production {
162            return Err(SeedError::EnvForbidden);
163        }
164        Ok(())
165    }
166
167    /// 执行 seeding(编排:环境检查 → 拓扑排序 → 执行 → 记录)
168    pub fn seed(&mut self) -> Result<SeedReport, SeedError> {
169        self.check_env()?;
170        let sorted = self.topological_sort()?;
171        let to_execute: Vec<(String, usize)> = sorted
172            .iter()
173            .filter(|s| !self.executed_versions.contains(&s.version))
174            .map(|s| (s.version.clone(), s.template.records.len()))
175            .collect();
176        let start = std::time::Instant::now();
177        let mut executed = Vec::new();
178        let mut total_rows = 0u64;
179        for (version, row_count) in to_execute {
180            total_rows += row_count as u64;
181            executed.push(version.clone());
182            self.executed_versions.insert(version);
183        }
184        Ok(SeedReport {
185            executed_seeds: executed,
186            total_rows,
187            total_duration: start.elapsed(),
188            idempotent: matches!(self.mode, SeedMode::Upsert),
189            env: self.env,
190        })
191    }
192
193    /// 获取已执行版本
194    pub fn executed_versions(&self) -> &HashSet<String> {
195        &self.executed_versions
196    }
197
198    /// 获取种子数量
199    pub fn seed_count(&self) -> usize {
200        self.seeds.len()
201    }
202}
203
204#[cfg(test)]
205mod tests {
206    use super::*;
207    use crate::seeding::FixtureTemplate;
208
209    fn make_template(table: &str) -> FixtureTemplate {
210        FixtureTemplate {
211            table: table.to_string(),
212            records: vec![serde_json::Map::new()],
213            count: 1,
214            references: Vec::new(),
215            extends: None,
216        }
217    }
218
219    fn make_seed(version: &str, deps: Vec<&str>) -> SeedFile {
220        SeedFile {
221            version: version.to_string(),
222            description: format!("seed {}", version),
223            dependencies: deps.iter().map(|s| s.to_string()).collect(),
224            template: make_template(&format!("table_{}", version)),
225        }
226    }
227
228    #[test]
229    fn test_topological_sort() {
230        let mut manager = SeedManager::new(SeedMode::Upsert, SeedEnv::Test);
231        manager.add_seed(make_seed("A", vec![]));
232        manager.add_seed(make_seed("B", vec!["A"]));
233        manager.add_seed(make_seed("C", vec!["B"]));
234        let sorted = manager.topological_sort().unwrap();
235        assert_eq!(sorted.len(), 3);
236        assert_eq!(sorted[0].version, "A");
237        assert_eq!(sorted[1].version, "B");
238        assert_eq!(sorted[2].version, "C");
239    }
240
241    #[test]
242    fn test_dependency_cycle_detection() {
243        let mut manager = SeedManager::new(SeedMode::Upsert, SeedEnv::Test);
244        manager.add_seed(make_seed("A", vec!["B"]));
245        manager.add_seed(make_seed("B", vec!["A"]));
246        let result = manager.topological_sort();
247        assert!(result.is_err());
248        assert!(matches!(
249            result.unwrap_err(),
250            SeedError::DependencyCycle { .. }
251        ));
252    }
253
254    #[test]
255    fn test_env_forbidden() {
256        let mut manager = SeedManager::new(SeedMode::Upsert, SeedEnv::Production);
257        manager.add_seed(make_seed("A", vec![]));
258        let result = manager.seed();
259        assert!(result.is_err());
260        assert!(matches!(result.unwrap_err(), SeedError::EnvForbidden));
261    }
262
263    #[test]
264    fn test_env_allowed_with_flag() {
265        let mut manager =
266            SeedManager::new(SeedMode::Upsert, SeedEnv::Production).allow_production();
267        manager.add_seed(make_seed("A", vec![]));
268        let result = manager.seed();
269        assert!(result.is_ok());
270    }
271
272    #[test]
273    fn test_idempotent_execution() {
274        let mut manager = SeedManager::new(SeedMode::Upsert, SeedEnv::Test);
275        manager.add_seed(make_seed("A", vec![]));
276        let report1 = manager.seed().unwrap();
277        assert_eq!(report1.executed_seeds.len(), 1);
278        let report2 = manager.seed().unwrap();
279        assert_eq!(report2.executed_seeds.len(), 0);
280        assert!(report2.total_rows == 0);
281    }
282
283    #[test]
284    fn test_truncate_insert_mode() {
285        let mut manager = SeedManager::new(SeedMode::TruncateInsert, SeedEnv::Test);
286        manager.add_seed(make_seed("A", vec![]));
287        manager.add_seed(make_seed("B", vec![]));
288        let report = manager.seed().unwrap();
289        assert_eq!(report.executed_seeds.len(), 2);
290        assert!(!report.idempotent);
291    }
292
293    #[test]
294    fn test_seed_report_env() {
295        let mut manager = SeedManager::new(SeedMode::Upsert, SeedEnv::Staging);
296        manager.add_seed(make_seed("A", vec![]));
297        let report = manager.seed().unwrap();
298        assert_eq!(report.env, SeedEnv::Staging);
299    }
300}