1use super::{FixtureTemplate, SeedError};
4use std::collections::{HashMap, HashSet};
5use std::time::Duration;
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum SeedMode {
10 Upsert,
12 TruncateInsert,
14}
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum SeedEnv {
19 Dev,
21 Test,
23 Staging,
25 Production,
27}
28
29#[derive(Debug, Clone)]
31pub struct SeedFile {
32 pub version: String,
34 pub description: String,
36 pub dependencies: Vec<String>,
38 pub template: FixtureTemplate,
40}
41
42#[derive(Debug, Clone)]
44pub struct SeedReport {
45 pub executed_seeds: Vec<String>,
47 pub total_rows: u64,
49 pub total_duration: Duration,
51 pub idempotent: bool,
53 pub env: SeedEnv,
55}
56
57pub 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 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 pub fn allow_production(mut self) -> Self {
80 self.allow_production = true;
81 self
82 }
83
84 pub fn add_seed(&mut self, seed: SeedFile) {
86 self.seeds.push(seed);
87 }
88
89 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 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 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 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 pub fn executed_versions(&self) -> &HashSet<String> {
195 &self.executed_versions
196 }
197
198 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}