1use std::collections::HashMap;
10use std::path::{Path, PathBuf};
11use std::process::Command;
12
13use prax_query::dialect::{Mysql, SqlDialect};
14use prax_query::sql::escape_identifier;
15use serde::{Deserialize, Serialize};
16
17use crate::config::Config;
18use crate::error::{CliError, CliResult};
19use crate::output;
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum SeedFileType {
24 Rust,
26 Sql,
28 Json,
30 Toml,
32}
33
34impl SeedFileType {
35 pub fn from_path(path: &Path) -> Option<Self> {
37 match path.extension()?.to_str()? {
38 "rs" => Some(Self::Rust),
39 "sql" => Some(Self::Sql),
40 "json" => Some(Self::Json),
41 "toml" => Some(Self::Toml),
42 _ => None,
43 }
44 }
45}
46
47#[derive(Debug, Clone)]
49pub struct SeedRunner {
50 pub seed_path: PathBuf,
52 pub file_type: SeedFileType,
54 pub database_url: String,
56 pub provider: String,
58 pub cwd: PathBuf,
60 pub environment: String,
62 pub reset_before_seed: bool,
64}
65
66impl SeedRunner {
67 pub fn new(
69 seed_path: PathBuf,
70 database_url: String,
71 provider: String,
72 cwd: PathBuf,
73 ) -> CliResult<Self> {
74 let file_type = SeedFileType::from_path(&seed_path).ok_or_else(|| {
75 CliError::Config(format!(
76 "Unsupported seed file type: {}. Supported: .rs, .sql, .json, .toml",
77 seed_path.display()
78 ))
79 })?;
80
81 Ok(Self {
82 seed_path,
83 file_type,
84 database_url,
85 provider,
86 cwd,
87 environment: std::env::var("PRAX_ENV").unwrap_or_else(|_| "development".to_string()),
88 reset_before_seed: false,
89 })
90 }
91
92 pub fn with_environment(mut self, env: impl Into<String>) -> Self {
94 self.environment = env.into();
95 self
96 }
97
98 pub fn with_reset(mut self, reset: bool) -> Self {
100 self.reset_before_seed = reset;
101 self
102 }
103
104 pub async fn run(&self) -> CliResult<SeedResult> {
106 if self.reset_before_seed {
107 self.reset_tables().await?;
108 }
109
110 match self.file_type {
111 SeedFileType::Rust => self.run_rust_seed().await,
112 SeedFileType::Sql => self.run_sql_seed().await,
113 SeedFileType::Json => self.run_json_seed().await,
114 SeedFileType::Toml => self.run_toml_seed().await,
115 }
116 }
117
118 async fn reset_tables(&self) -> CliResult<()> {
120 let tables = self.seed_table_names()?;
121 if tables.is_empty() {
122 return Ok(());
123 }
124
125 output::list_item(&format!(
126 "Resetting {} table(s) before seeding...",
127 tables.len()
128 ));
129
130 let sql = self.generate_reset_sql(&tables);
131 self.execute_sql(&sql).await?;
132 Ok(())
133 }
134
135 fn seed_table_names(&self) -> CliResult<Vec<String>> {
140 match self.file_type {
141 SeedFileType::Json => {
142 let content = std::fs::read_to_string(&self.seed_path)?;
143 let seed_data: SeedData =
144 serde_json::from_str(&content).map_err(|e| CliError::Config(e.to_string()))?;
145 Ok(seed_data.tables.keys().cloned().collect())
146 }
147 SeedFileType::Toml => {
148 let content = std::fs::read_to_string(&self.seed_path)?;
149 let seed_data: SeedData =
150 toml::from_str(&content).map_err(|e| CliError::Config(e.to_string()))?;
151 Ok(seed_data.tables.keys().cloned().collect())
152 }
153 SeedFileType::Rust | SeedFileType::Sql => Err(CliError::Config(
154 "--reset requires table metadata; not supported yet".to_string(),
155 )),
156 }
157 }
158
159 async fn run_rust_seed(&self) -> CliResult<SeedResult> {
161 output::step(1, 4, "Compiling seed script...");
162
163 let cargo_toml = self.cwd.join("Cargo.toml");
165 if !cargo_toml.exists() {
166 return Err(CliError::Config(
167 "No Cargo.toml found. Rust seed scripts require a Rust project.".to_string(),
168 ));
169 }
170
171 let seed_name = self
173 .seed_path
174 .file_stem()
175 .and_then(|s| s.to_str())
176 .unwrap_or("seed");
177
178 let has_bin_target = self.check_bin_target(seed_name)?;
180
181 let mut records_affected = 0u64;
182
183 if has_bin_target {
184 output::step(2, 4, &format!("Building seed binary '{}'...", seed_name));
186
187 let build_status = Command::new("cargo")
188 .args(["build", "--bin", seed_name, "--release"])
189 .current_dir(&self.cwd)
190 .env("DATABASE_URL", &self.database_url)
191 .env("PRAX_ENV", &self.environment)
192 .status()?;
193
194 if !build_status.success() {
195 return Err(CliError::Command("Failed to build seed binary".to_string()));
196 }
197
198 output::step(3, 4, "Running seed...");
199
200 let run_output = Command::new("cargo")
201 .args(["run", "--bin", seed_name, "--release"])
202 .current_dir(&self.cwd)
203 .env("DATABASE_URL", &self.database_url)
204 .env("PRAX_ENV", &self.environment)
205 .output()?;
206
207 if !run_output.status.success() {
208 let stderr = String::from_utf8_lossy(&run_output.stderr);
209 return Err(CliError::Command(format!("Seed failed: {}", stderr)));
210 }
211
212 let stdout = String::from_utf8_lossy(&run_output.stdout);
214 for line in stdout.lines() {
215 output::list_item(line);
216 if let Some(count) = parse_seed_output(line) {
218 records_affected += count;
219 }
220 }
221
222 output::step(4, 4, "Verifying seed data...");
223 } else {
224 output::step(2, 4, "Compiling standalone seed script...");
226
227 let temp_dir = std::env::temp_dir().join("prax_seed");
229 std::fs::create_dir_all(&temp_dir)?;
230
231 let output_binary = temp_dir.join(seed_name);
232
233 let seed_content = std::fs::read_to_string(&self.seed_path)?;
235
236 if seed_content.contains("use prax") || seed_content.contains("#[tokio::main]") {
237 output::list_item("Creating temporary build environment...");
239
240 let temp_project = temp_dir.join("seed_project");
241 std::fs::create_dir_all(temp_project.join("src"))?;
242
243 std::fs::copy(&self.seed_path, temp_project.join("src/main.rs"))?;
245
246 let seed_cargo = create_seed_cargo_toml(&self.cwd)?;
248 std::fs::write(temp_project.join("Cargo.toml"), seed_cargo)?;
249
250 let build_status = Command::new("cargo")
252 .args(["build", "--release"])
253 .current_dir(&temp_project)
254 .env("DATABASE_URL", &self.database_url)
255 .env("PRAX_ENV", &self.environment)
256 .status()?;
257
258 if !build_status.success() {
259 return Err(CliError::Command(
260 "Failed to compile seed script".to_string(),
261 ));
262 }
263
264 let built_binary = temp_project.join("target/release/seed");
266 if built_binary.exists() {
267 std::fs::copy(&built_binary, &output_binary)?;
268 }
269 } else {
270 return Err(CliError::Config(
271 "Seed script must be a valid Rust file with a main function".to_string(),
272 ));
273 }
274
275 output::step(3, 4, "Running seed...");
276
277 let run_output = Command::new(&output_binary)
278 .current_dir(&self.cwd)
279 .env("DATABASE_URL", &self.database_url)
280 .env("PRAX_ENV", &self.environment)
281 .output()?;
282
283 if !run_output.status.success() {
284 let stderr = String::from_utf8_lossy(&run_output.stderr);
285 return Err(CliError::Command(format!("Seed failed: {}", stderr)));
286 }
287
288 let stdout = String::from_utf8_lossy(&run_output.stdout);
289 for line in stdout.lines() {
290 output::list_item(line);
291 if let Some(count) = parse_seed_output(line) {
292 records_affected += count;
293 }
294 }
295
296 output::step(4, 4, "Verifying seed data...");
297 }
298
299 Ok(SeedResult {
300 file_type: self.file_type,
301 records_affected,
302 tables_seeded: Vec::new(),
303 duration: std::time::Duration::from_secs(0),
304 })
305 }
306
307 async fn run_sql_seed(&self) -> CliResult<SeedResult> {
309 output::step(1, 3, "Reading SQL seed file...");
310
311 let sql_content = std::fs::read_to_string(&self.seed_path)?;
312
313 let statements: Vec<&str> = sql_content
315 .split(';')
316 .map(|s| s.trim())
317 .filter(|s| !s.is_empty() && !s.starts_with("--"))
318 .collect();
319
320 output::list_item(&format!("Found {} SQL statements", statements.len()));
321
322 output::step(2, 3, "Executing SQL...");
323
324 let records = self.execute_sql(&sql_content).await?;
326
327 output::step(3, 3, "Verifying seed data...");
328
329 Ok(SeedResult {
330 file_type: self.file_type,
331 records_affected: records,
332 tables_seeded: Vec::new(),
333 duration: std::time::Duration::from_secs(0),
334 })
335 }
336
337 async fn run_json_seed(&self) -> CliResult<SeedResult> {
339 output::step(1, 4, "Reading JSON seed file...");
340
341 let json_content = std::fs::read_to_string(&self.seed_path)?;
342 let seed_data: SeedData =
343 serde_json::from_str(&json_content).map_err(|e| CliError::Config(e.to_string()))?;
344
345 output::step(2, 4, "Validating seed data...");
346 output::list_item(&format!("Found {} tables to seed", seed_data.tables.len()));
347
348 output::step(3, 4, "Inserting seed data...");
349
350 let mut total_records = 0u64;
351 let mut tables_seeded = Vec::new();
352
353 for (table_name, records) in &seed_data.tables {
354 let sql = self.generate_insert_sql(table_name, records)?;
355 let count = self.execute_sql(&sql).await?;
356 output::list_item(&format!(" {} - {} records", table_name, records.len()));
357 total_records += count;
358 tables_seeded.push(table_name.clone());
359 }
360
361 output::step(4, 4, "Verifying seed data...");
362
363 Ok(SeedResult {
364 file_type: self.file_type,
365 records_affected: total_records,
366 tables_seeded,
367 duration: std::time::Duration::from_secs(0),
368 })
369 }
370
371 async fn run_toml_seed(&self) -> CliResult<SeedResult> {
373 output::step(1, 4, "Reading TOML seed file...");
374
375 let toml_content = std::fs::read_to_string(&self.seed_path)?;
376 let seed_data: SeedData =
377 toml::from_str(&toml_content).map_err(|e| CliError::Config(e.to_string()))?;
378
379 output::step(2, 4, "Validating seed data...");
380 output::list_item(&format!("Found {} tables to seed", seed_data.tables.len()));
381
382 output::step(3, 4, "Inserting seed data...");
383
384 let mut total_records = 0u64;
385 let mut tables_seeded = Vec::new();
386
387 for (table_name, records) in &seed_data.tables {
388 let sql = self.generate_insert_sql(table_name, records)?;
389 let count = self.execute_sql(&sql).await?;
390 output::list_item(&format!(" {} - {} records", table_name, records.len()));
391 total_records += count;
392 tables_seeded.push(table_name.clone());
393 }
394
395 output::step(4, 4, "Verifying seed data...");
396
397 Ok(SeedResult {
398 file_type: self.file_type,
399 records_affected: total_records,
400 tables_seeded,
401 duration: std::time::Duration::from_secs(0),
402 })
403 }
404
405 fn check_bin_target(&self, name: &str) -> CliResult<bool> {
407 let cargo_toml = self.cwd.join("Cargo.toml");
408 let content = std::fs::read_to_string(&cargo_toml)?;
409
410 Ok(content.contains(&format!("name = \"{}\"", name))
412 || content.contains(&format!("name = '{}'", name)))
413 }
414
415 fn generate_insert_sql(
417 &self,
418 table: &str,
419 records: &[HashMap<String, serde_json::Value>],
420 ) -> CliResult<String> {
421 if records.is_empty() {
422 return Ok(String::new());
423 }
424
425 let mut sql = String::new();
426
427 let columns: Vec<&String> = records[0].keys().collect();
429 let column_names = columns
430 .iter()
431 .map(|c| escape_identifier(c))
432 .collect::<Vec<_>>()
433 .join(", ");
434
435 for record in records {
436 let values = columns
437 .iter()
438 .map(|col| {
439 record
440 .get(*col)
441 .map(|v| self.value_to_sql(v))
442 .unwrap_or_else(|| "NULL".to_string())
443 })
444 .collect::<Vec<_>>()
445 .join(", ");
446
447 sql.push_str(&format!(
448 "INSERT INTO {} ({}) VALUES ({});\n",
449 escape_identifier(table),
450 column_names,
451 values
452 ));
453 }
454
455 Ok(sql)
456 }
457
458 fn generate_reset_sql(&self, tables: &[String]) -> String {
460 let quote: fn(&str) -> String = match self.provider.as_str() {
464 "mysql" => |t| Mysql.quote_ident(t),
465 _ => escape_identifier,
466 };
467 match self.provider.as_str() {
468 "postgresql" | "postgres" => {
469 let names = tables
470 .iter()
471 .map(|t| quote(t))
472 .collect::<Vec<_>>()
473 .join(", ");
474 format!("TRUNCATE TABLE {} CASCADE;", names)
475 }
476 "mysql" => tables
477 .iter()
478 .map(|t| format!("TRUNCATE TABLE {};", quote(t)))
479 .collect::<Vec<_>>()
480 .join("\n"),
481 "sqlite" => tables
482 .iter()
483 .map(|t| format!("DELETE FROM {};", quote(t)))
484 .collect::<Vec<_>>()
485 .join("\n"),
486 _ => String::new(),
487 }
488 }
489
490 fn value_to_sql(&self, value: &serde_json::Value) -> String {
492 match value {
493 serde_json::Value::Null => "NULL".to_string(),
494 serde_json::Value::Bool(b) => {
495 if *b {
496 "TRUE".to_string()
497 } else {
498 "FALSE".to_string()
499 }
500 }
501 serde_json::Value::Number(n) => n.to_string(),
502 serde_json::Value::String(s) => {
503 match s.as_str() {
505 "now()" | "NOW()" => match self.provider.as_str() {
506 "postgresql" => "CURRENT_TIMESTAMP".to_string(),
507 "mysql" => "NOW()".to_string(),
508 "sqlite" => "datetime('now')".to_string(),
509 _ => "CURRENT_TIMESTAMP".to_string(),
510 },
511 "uuid()" | "UUID()" => match self.provider.as_str() {
512 "postgresql" => "gen_random_uuid()".to_string(),
513 "mysql" => "UUID()".to_string(),
514 "sqlite" => format!("'{}'", uuid::Uuid::new_v4()),
515 _ => "gen_random_uuid()".to_string(),
516 },
517 _ => format!("'{}'", s.replace('\'', "''")),
518 }
519 }
520 serde_json::Value::Array(arr) => {
521 let items = arr
523 .iter()
524 .map(|v| self.value_to_sql(v))
525 .collect::<Vec<_>>()
526 .join(", ");
527 format!("ARRAY[{}]", items)
528 }
529 serde_json::Value::Object(_) => {
530 format!("'{}'", value)
532 }
533 }
534 }
535
536 async fn execute_sql(&self, sql: &str) -> CliResult<u64> {
538 match self.provider.as_str() {
540 "postgresql" | "postgres" => self.execute_postgres_sql(sql).await,
541 "mysql" => self.execute_mysql_sql(sql).await,
542 "sqlite" => self.execute_sqlite_sql(sql).await,
543 _ => Err(CliError::Database(format!(
544 "Unsupported database provider: {}",
545 self.provider
546 ))),
547 }
548 }
549
550 async fn execute_postgres_sql(&self, sql: &str) -> CliResult<u64> {
552 let psql_result = Command::new("psql")
554 .args(["-d", &self.database_url, "-c", sql])
555 .output();
556
557 match psql_result {
558 Ok(output) if output.status.success() => {
559 let stdout = String::from_utf8_lossy(&output.stdout);
561 Ok(parse_affected_rows(&stdout).unwrap_or(0))
562 }
563 Ok(output) => {
564 let stderr = String::from_utf8_lossy(&output.stderr);
565 if stderr.contains("not found") || stderr.contains("No such file") {
567 Err(CliError::Command(
568 "psql not found. Install PostgreSQL client tools or use a Rust seed script.".to_string()
569 ))
570 } else {
571 Err(CliError::Database(format!(
572 "SQL execution failed: {}",
573 stderr
574 )))
575 }
576 }
577 Err(e) => Err(CliError::Command(format!(
578 "Failed to execute SQL. Install psql (PostgreSQL client tools) or use a Rust seed script: {}",
579 e
580 ))),
581 }
582 }
583
584 async fn execute_mysql_sql(&self, sql: &str) -> CliResult<u64> {
586 let url = url::Url::parse(&self.database_url)
588 .map_err(|e| CliError::Config(format!("Invalid MySQL URL: {}", e)))?;
589
590 let host = url.host_str().unwrap_or("localhost");
591 let port = url.port().unwrap_or(3306);
592 let user = url.username();
593 let password = url.password().unwrap_or("");
594 let database = url.path().trim_start_matches('/');
595
596 let mut cmd = Command::new("mysql");
597 cmd.args(["-h", host, "-P", &port.to_string(), "-u", user]);
598
599 if !password.is_empty() {
600 cmd.arg(format!("-p{}", password));
601 }
602
603 cmd.args(["-D", database, "-e", sql]);
604
605 let output = cmd.output()?;
606
607 if output.status.success() {
608 let stdout = String::from_utf8_lossy(&output.stdout);
609 Ok(parse_affected_rows(&stdout).unwrap_or(0))
610 } else {
611 let stderr = String::from_utf8_lossy(&output.stderr);
612 if stderr.contains("not found") || stderr.contains("No such file") {
613 Err(CliError::Command(
614 "mysql client not found. Install MySQL client tools or use a Rust seed script."
615 .to_string(),
616 ))
617 } else {
618 Err(CliError::Database(format!(
619 "SQL execution failed: {}",
620 stderr
621 )))
622 }
623 }
624 }
625
626 async fn execute_sqlite_sql(&self, sql: &str) -> CliResult<u64> {
628 let db_path = self
630 .database_url
631 .strip_prefix("sqlite://")
632 .or_else(|| self.database_url.strip_prefix("sqlite:"))
633 .unwrap_or(&self.database_url);
634
635 let output = Command::new("sqlite3").args([db_path, sql]).output()?;
636
637 if output.status.success() {
638 let stdout = String::from_utf8_lossy(&output.stdout);
639 Ok(parse_affected_rows(&stdout).unwrap_or(0))
640 } else {
641 let stderr = String::from_utf8_lossy(&output.stderr);
642 if stderr.contains("not found") || stderr.contains("No such file") {
643 Err(CliError::Command(
644 "sqlite3 not found. Install SQLite tools or use a Rust seed script."
645 .to_string(),
646 ))
647 } else {
648 Err(CliError::Database(format!(
649 "SQL execution failed: {}",
650 stderr
651 )))
652 }
653 }
654 }
655}
656
657#[derive(Debug)]
659pub struct SeedResult {
660 pub file_type: SeedFileType,
662 pub records_affected: u64,
664 pub tables_seeded: Vec<String>,
666 pub duration: std::time::Duration,
668}
669
670#[derive(Debug, Clone, Deserialize, Serialize)]
672pub struct SeedData {
673 #[serde(default)]
675 pub tables: HashMap<String, Vec<HashMap<String, serde_json::Value>>>,
676
677 #[serde(default)]
679 pub order: Vec<String>,
680
681 #[serde(default)]
683 pub truncate: bool,
684
685 #[serde(default)]
687 pub disable_fk_checks: bool,
688}
689
690pub fn find_seed_file(cwd: &Path, config: &Config) -> Option<PathBuf> {
696 if let Some(ref seed_path) = config.database.seed_path
698 && seed_path.exists()
699 {
700 return Some(seed_path.clone());
701 }
702
703 let candidates = [
705 cwd.join("seed.rs"),
706 cwd.join("seed.sql"),
707 cwd.join("seed.json"),
708 cwd.join("seed.toml"),
709 cwd.join("prax/seed.rs"),
710 cwd.join("prax/seed.sql"),
711 cwd.join("prax/seed.json"),
712 cwd.join("prax/seed.toml"),
713 cwd.join("prisma/seed.rs"),
714 cwd.join("src/seed.rs"),
715 cwd.join("seeds/seed.rs"),
716 cwd.join("seeds/seed.sql"),
717 ];
718
719 candidates.into_iter().find(|p| p.exists())
720}
721
722pub fn get_database_url(config: &Config) -> CliResult<String> {
724 if let Some(ref url) = config.database.url {
726 let expanded = expand_env_var(url);
728 if !expanded.is_empty() && !expanded.contains("${") {
729 return Ok(expanded);
730 }
731 }
732
733 std::env::var("DATABASE_URL").map_err(|_| {
735 CliError::Config(
736 "Database URL not found. Set DATABASE_URL environment variable or configure in prax.toml"
737 .to_string(),
738 )
739 })
740}
741
742fn expand_env_var(s: &str) -> String {
744 let mut result = s.to_string();
745
746 let re = regex_lite::Regex::new(r"\$\{([^}]+)\}").unwrap();
748 for cap in re.captures_iter(s) {
749 let var_name = &cap[1];
750 if let Ok(value) = std::env::var(var_name) {
751 result = result.replace(&cap[0], &value);
752 }
753 }
754
755 let re2 = regex_lite::Regex::new(r"\$([A-Z_][A-Z0-9_]*)").unwrap();
757 for cap in re2.captures_iter(&result.clone()) {
758 let var_name = &cap[1];
759 if let Ok(value) = std::env::var(var_name) {
760 result = result.replace(&cap[0], &value);
761 }
762 }
763
764 result
765}
766
767fn parse_seed_output(line: &str) -> Option<u64> {
769 let patterns = [
774 r"(?i)created\s+(\d+)",
775 r"(?i)seeded\s+(\d+)",
776 r"(?i)inserted[:\s]+(\d+)",
777 r"(?i)(\d+)\s+records?",
778 r"(?i)(\d+)\s+rows?",
779 ];
780
781 for pattern in patterns {
782 if let Ok(re) = regex_lite::Regex::new(pattern)
783 && let Some(caps) = re.captures(line)
784 && let Some(m) = caps.get(1)
785 && let Ok(n) = m.as_str().parse()
786 {
787 return Some(n);
788 }
789 }
790
791 None
792}
793
794fn parse_affected_rows(output: &str) -> Option<u64> {
796 let patterns = [
801 r"INSERT\s+\d+\s+(\d+)",
802 r"UPDATE\s+(\d+)",
803 r"DELETE\s+(\d+)",
804 r"(\d+)\s+rows?\s+affected",
805 ];
806
807 let mut total = 0u64;
808
809 for pattern in patterns {
810 if let Ok(re) = regex_lite::Regex::new(pattern) {
811 for caps in re.captures_iter(output) {
812 if let Some(m) = caps.get(1)
813 && let Ok(n) = m.as_str().parse::<u64>()
814 {
815 total += n;
816 }
817 }
818 }
819 }
820
821 if total > 0 { Some(total) } else { None }
822}
823
824fn create_seed_cargo_toml(project_root: &Path) -> CliResult<String> {
826 let workspace_cargo = project_root.join("Cargo.toml");
828 let prax_version = if workspace_cargo.exists() {
829 let content = std::fs::read_to_string(&workspace_cargo)?;
830 extract_prax_version(&content).unwrap_or_else(|| env!("CARGO_PKG_VERSION").to_string())
832 } else {
833 env!("CARGO_PKG_VERSION").to_string()
834 };
835
836 Ok(format!(
837 r#"[package]
838name = "seed"
839version = "0.1.0"
840edition = "2024"
841
842[dependencies]
843prax-orm = "{}"
844tokio = {{ version = "1", features = ["full"] }}
845"#,
846 prax_version
847 ))
848}
849
850fn extract_prax_version(content: &str) -> Option<String> {
852 let simple_re = regex_lite::Regex::new(r#"prax-orm\s*=\s*"([^"]+)""#).ok()?;
854 if let Some(caps) = simple_re.captures(content) {
855 return Some(caps.get(1)?.as_str().to_string());
856 }
857
858 let complex_re =
859 regex_lite::Regex::new(r#"prax-orm\s*=\s*\{[^}]*version\s*=\s*"([^"]+)""#).ok()?;
860 if let Some(caps) = complex_re.captures(content) {
861 return Some(caps.get(1)?.as_str().to_string());
862 }
863
864 None
865}
866
867#[cfg(test)]
872mod tests {
873 use super::*;
874
875 #[test]
876 fn test_seed_file_type_detection() {
877 assert_eq!(
878 SeedFileType::from_path(Path::new("seed.rs")),
879 Some(SeedFileType::Rust)
880 );
881 assert_eq!(
882 SeedFileType::from_path(Path::new("seed.sql")),
883 Some(SeedFileType::Sql)
884 );
885 assert_eq!(
886 SeedFileType::from_path(Path::new("data.json")),
887 Some(SeedFileType::Json)
888 );
889 assert_eq!(
890 SeedFileType::from_path(Path::new("data.toml")),
891 Some(SeedFileType::Toml)
892 );
893 assert_eq!(SeedFileType::from_path(Path::new("seed.txt")), None);
894 }
895
896 #[test]
897 fn test_parse_seed_output() {
898 assert_eq!(parse_seed_output("Created 10 users"), Some(10));
899 assert_eq!(parse_seed_output("Seeded 100 records"), Some(100));
900 assert_eq!(parse_seed_output("Inserted: 50"), Some(50));
901 assert_eq!(parse_seed_output("5 rows affected"), Some(5));
902 assert_eq!(parse_seed_output("no numbers here"), None);
903 }
904
905 #[test]
906 fn test_parse_affected_rows() {
907 assert_eq!(parse_affected_rows("INSERT 0 5"), Some(5));
908 assert_eq!(parse_affected_rows("UPDATE 3"), Some(3));
909 assert_eq!(parse_affected_rows("Query OK, 10 rows affected"), Some(10));
910 }
911
912 #[test]
913 fn test_expand_env_var() {
914 unsafe {
916 std::env::set_var("TEST_VAR", "test_value");
917 }
918 assert_eq!(expand_env_var("${TEST_VAR}"), "test_value");
919 assert_eq!(expand_env_var("$TEST_VAR"), "test_value");
920 assert_eq!(
921 expand_env_var("postgres://${TEST_VAR}@localhost"),
922 "postgres://test_value@localhost"
923 );
924 unsafe {
926 std::env::remove_var("TEST_VAR");
927 }
928 }
929
930 fn make_runner(seed_path: &str, provider: &str) -> SeedRunner {
931 SeedRunner::new(
932 PathBuf::from(seed_path),
933 "sqlite://test.db".to_string(),
934 provider.to_string(),
935 PathBuf::from("."),
936 )
937 .expect("supported seed file type")
938 }
939
940 #[test]
941 fn test_find_seed_file_skips_unsupported_ts() {
942 let dir = tempfile::tempdir().unwrap();
943 let config = Config::default();
944
945 let prisma_dir = dir.path().join("prisma");
947 std::fs::create_dir_all(&prisma_dir).unwrap();
948 std::fs::write(prisma_dir.join("seed.ts"), "export default {}").unwrap();
949 assert!(find_seed_file(dir.path(), &config).is_none());
950
951 let src_dir = dir.path().join("src");
953 std::fs::create_dir_all(&src_dir).unwrap();
954 std::fs::write(src_dir.join("seed.rs"), "fn main() {}").unwrap();
955 assert_eq!(
956 find_seed_file(dir.path(), &config),
957 Some(src_dir.join("seed.rs"))
958 );
959 }
960
961 #[test]
962 fn test_create_seed_cargo_toml_version_fallback() {
963 let dir = tempfile::tempdir().unwrap();
965 let cargo_toml = create_seed_cargo_toml(dir.path()).unwrap();
966 assert!(
967 cargo_toml.contains(&format!("prax-orm = \"{}\"", env!("CARGO_PKG_VERSION"))),
968 "unexpected seed Cargo.toml:\n{cargo_toml}"
969 );
970
971 let dir = tempfile::tempdir().unwrap();
973 std::fs::write(dir.path().join("Cargo.toml"), "[package]\nname = \"app\"\n").unwrap();
974 let cargo_toml = create_seed_cargo_toml(dir.path()).unwrap();
975 assert!(
976 cargo_toml.contains(&format!("prax-orm = \"{}\"", env!("CARGO_PKG_VERSION"))),
977 "unexpected seed Cargo.toml:\n{cargo_toml}"
978 );
979 }
980
981 #[test]
982 fn test_create_seed_cargo_toml_uses_project_version() {
983 let dir = tempfile::tempdir().unwrap();
984 std::fs::write(
985 dir.path().join("Cargo.toml"),
986 "[dependencies]\nprax-orm = \"0.7.3\"\n",
987 )
988 .unwrap();
989 let cargo_toml = create_seed_cargo_toml(dir.path()).unwrap();
990 assert!(cargo_toml.contains("prax-orm = \"0.7.3\""));
991 }
992
993 #[test]
994 fn test_generate_insert_sql_escapes_identifiers() {
995 let runner = make_runner("seed.json", "sqlite");
996
997 let mut record = HashMap::new();
998 record.insert(
999 "na\"me".to_string(),
1000 serde_json::Value::String("va\"l".to_string()),
1001 );
1002
1003 let sql = runner.generate_insert_sql("we\"ird", &[record]).unwrap();
1004
1005 assert!(sql.contains("INSERT INTO \"we\"\"ird\""));
1007 assert!(sql.contains("(\"na\"\"me\")"));
1008 assert!(sql.contains("'va\"l'"));
1010 }
1011
1012 #[test]
1013 fn test_generate_reset_sql_per_provider() {
1014 let tables = vec!["users".to_string(), "posts".to_string()];
1015
1016 let pg = make_runner("seed.json", "postgresql");
1017 assert_eq!(
1018 pg.generate_reset_sql(&tables),
1019 "TRUNCATE TABLE \"users\", \"posts\" CASCADE;"
1020 );
1021
1022 let mysql = make_runner("seed.json", "mysql");
1023 assert_eq!(
1024 mysql.generate_reset_sql(&tables),
1025 "TRUNCATE TABLE `users`;\nTRUNCATE TABLE `posts`;"
1026 );
1027
1028 let sqlite = make_runner("seed.json", "sqlite");
1029 assert_eq!(
1030 sqlite.generate_reset_sql(&tables),
1031 "DELETE FROM \"users\";\nDELETE FROM \"posts\";"
1032 );
1033 }
1034
1035 #[test]
1036 fn test_generate_reset_sql_escapes_identifiers() {
1037 let sqlite = make_runner("seed.json", "sqlite");
1038 let tables = vec!["we\"ird".to_string()];
1039 assert_eq!(
1040 sqlite.generate_reset_sql(&tables),
1041 "DELETE FROM \"we\"\"ird\";"
1042 );
1043 }
1044
1045 #[test]
1046 fn test_seed_table_names_from_json() {
1047 let dir = tempfile::tempdir().unwrap();
1048 let seed_path = dir.path().join("seed.json");
1049 std::fs::write(
1050 &seed_path,
1051 r#"{"tables": {"users": [{"id": 1}], "posts": [{"id": 2}]}}"#,
1052 )
1053 .unwrap();
1054
1055 let runner = SeedRunner::new(
1056 seed_path,
1057 "sqlite://test.db".to_string(),
1058 "sqlite".to_string(),
1059 dir.path().to_path_buf(),
1060 )
1061 .unwrap();
1062
1063 let mut names = runner.seed_table_names().unwrap();
1064 names.sort();
1065 assert_eq!(names, vec!["posts".to_string(), "users".to_string()]);
1066 }
1067
1068 #[test]
1069 fn test_seed_table_names_from_toml() {
1070 let dir = tempfile::tempdir().unwrap();
1071 let seed_path = dir.path().join("seed.toml");
1072 std::fs::write(
1073 &seed_path,
1074 "[[tables.users]]\nid = 1\n\n[[tables.posts]]\nid = 2\n",
1075 )
1076 .unwrap();
1077
1078 let runner = SeedRunner::new(
1079 seed_path,
1080 "sqlite://test.db".to_string(),
1081 "sqlite".to_string(),
1082 dir.path().to_path_buf(),
1083 )
1084 .unwrap();
1085
1086 let mut names = runner.seed_table_names().unwrap();
1087 names.sort();
1088 assert_eq!(names, vec!["posts".to_string(), "users".to_string()]);
1089 }
1090
1091 #[tokio::test]
1092 async fn test_reset_without_table_metadata_errors() {
1093 for seed in ["seed.rs", "seed.sql"] {
1095 let runner = make_runner(seed, "sqlite").with_reset(true);
1096 let err = runner.run().await.unwrap_err();
1097 assert!(
1098 err.to_string().contains("--reset requires table metadata"),
1099 "expected metadata error for {seed}, got: {err}"
1100 );
1101 }
1102 }
1103}