Skip to main content

prax_cli/commands/
seed.rs

1//! Database seeding implementation.
2//!
3//! Supports multiple seed file types:
4//! - `.rs` - Rust seed scripts (compiled and executed)
5//! - `.sql` - Raw SQL files (executed directly)
6//! - `.json` - JSON data files (declarative seeding)
7//! - `.toml` - TOML data files (declarative seeding)
8
9use 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/// Seed file types
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum SeedFileType {
24    /// Rust seed script (.rs)
25    Rust,
26    /// SQL seed file (.sql)
27    Sql,
28    /// JSON seed data (.json)
29    Json,
30    /// TOML seed data (.toml)
31    Toml,
32}
33
34impl SeedFileType {
35    /// Detect seed file type from path extension
36    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/// Seed runner configuration
48#[derive(Debug, Clone)]
49pub struct SeedRunner {
50    /// Path to the seed file
51    pub seed_path: PathBuf,
52    /// Seed file type
53    pub file_type: SeedFileType,
54    /// Database URL for execution
55    pub database_url: String,
56    /// Database provider (postgresql, mysql, sqlite)
57    pub provider: String,
58    /// Current working directory
59    pub cwd: PathBuf,
60    /// Environment name (development, staging, production)
61    pub environment: String,
62    /// Whether to reset database before seeding
63    pub reset_before_seed: bool,
64}
65
66impl SeedRunner {
67    /// Create a new seed runner
68    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    /// Set environment
93    pub fn with_environment(mut self, env: impl Into<String>) -> Self {
94        self.environment = env.into();
95        self
96    }
97
98    /// Set reset before seed
99    pub fn with_reset(mut self, reset: bool) -> Self {
100        self.reset_before_seed = reset;
101        self
102    }
103
104    /// Run the seed
105    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    /// Reset the tables this runner knows about before seeding.
119    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    /// Table names this runner can reset, from declarative seed metadata.
136    ///
137    /// `.rs` and `.sql` seeds carry no table metadata, so `--reset` is
138    /// rejected for them rather than silently skipped.
139    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    /// Run a Rust seed script
160    async fn run_rust_seed(&self) -> CliResult<SeedResult> {
161        output::step(1, 4, "Compiling seed script...");
162
163        // Check if we're in a Cargo workspace
164        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        // Create a temporary bin target or use cargo run
172        let seed_name = self
173            .seed_path
174            .file_stem()
175            .and_then(|s| s.to_str())
176            .unwrap_or("seed");
177
178        // Check if there's a [[bin]] entry for the seed, or we need to compile manually
179        let has_bin_target = self.check_bin_target(seed_name)?;
180
181        let mut records_affected = 0u64;
182
183        if has_bin_target {
184            // Use cargo run directly
185            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            // Parse output for record count if available
213            let stdout = String::from_utf8_lossy(&run_output.stdout);
214            for line in stdout.lines() {
215                output::list_item(line);
216                // Try to parse seed output for counts
217                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            // Compile and run as a standalone script using rustc
225            output::step(2, 4, "Compiling standalone seed script...");
226
227            // Create temp directory for compiled seed
228            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            // Try to compile with cargo if it looks like a full Rust file
234            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                // This is a standalone Rust file - we'll create a temporary Cargo project
238                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                // Copy seed file
244                std::fs::copy(&self.seed_path, temp_project.join("src/main.rs"))?;
245
246                // Create Cargo.toml for the seed
247                let seed_cargo = create_seed_cargo_toml(&self.cwd)?;
248                std::fs::write(temp_project.join("Cargo.toml"), seed_cargo)?;
249
250                // Build
251                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                // Copy binary
265                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    /// Run a SQL seed file
308    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        // Count statements for progress
314        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        // Execute SQL based on provider
325        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    /// Run a JSON seed file (declarative)
338    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    /// Run a TOML seed file (declarative)
372    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    /// Check if there's a bin target in Cargo.toml
406    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        // Simple check - look for [[bin]] with our name
411        Ok(content.contains(&format!("name = \"{}\"", name))
412            || content.contains(&format!("name = '{}'", name)))
413    }
414
415    /// Generate INSERT SQL from seed records
416    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        // Get columns from first record
428        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    /// Generate provider-specific reset SQL for the given tables
459    fn generate_reset_sql(&self, tables: &[String]) -> String {
460        // MySQL without ANSI_QUOTES mode parses `"users"` as a string
461        // literal, so identifiers need backtick quoting there; other
462        // providers take standard double-quoted identifiers.
463        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    /// Convert JSON value to SQL literal
491    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                // Check for special functions
504                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                // PostgreSQL array literal
522                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                // JSON/JSONB
531                format!("'{}'", value)
532            }
533        }
534    }
535
536    /// Execute SQL against the database
537    async fn execute_sql(&self, sql: &str) -> CliResult<u64> {
538        // Use command-line tools based on provider
539        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    /// Execute SQL using psql
551    async fn execute_postgres_sql(&self, sql: &str) -> CliResult<u64> {
552        // First try using psql
553        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                // Try to parse affected rows from output
560                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 psql not found, suggest alternative
566                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    /// Execute SQL using mysql client
585    async fn execute_mysql_sql(&self, sql: &str) -> CliResult<u64> {
586        // Parse MySQL URL to extract components
587        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    /// Execute SQL using sqlite3
627    async fn execute_sqlite_sql(&self, sql: &str) -> CliResult<u64> {
628        // Extract database path from URL
629        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/// Seed execution result
658#[derive(Debug)]
659pub struct SeedResult {
660    /// Type of seed file that was executed
661    pub file_type: SeedFileType,
662    /// Number of records affected
663    pub records_affected: u64,
664    /// Tables that were seeded
665    pub tables_seeded: Vec<String>,
666    /// Execution duration
667    pub duration: std::time::Duration,
668}
669
670/// Declarative seed data structure
671#[derive(Debug, Clone, Deserialize, Serialize)]
672pub struct SeedData {
673    /// Tables to seed, keyed by table name
674    #[serde(default)]
675    pub tables: HashMap<String, Vec<HashMap<String, serde_json::Value>>>,
676
677    /// Seed order (optional - tables will be seeded in this order)
678    #[serde(default)]
679    pub order: Vec<String>,
680
681    /// Truncate tables before seeding
682    #[serde(default)]
683    pub truncate: bool,
684
685    /// Disable foreign key checks during seeding
686    #[serde(default)]
687    pub disable_fk_checks: bool,
688}
689
690// =============================================================================
691// Helper Functions
692// =============================================================================
693
694/// Find seed file in common locations
695pub fn find_seed_file(cwd: &Path, config: &Config) -> Option<PathBuf> {
696    // Check config first
697    if let Some(ref seed_path) = config.database.seed_path
698        && seed_path.exists()
699    {
700        return Some(seed_path.clone());
701    }
702
703    // Common locations
704    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
722/// Get database URL from config or environment
723pub fn get_database_url(config: &Config) -> CliResult<String> {
724    // Try config first
725    if let Some(ref url) = config.database.url {
726        // Expand environment variables
727        let expanded = expand_env_var(url);
728        if !expanded.is_empty() && !expanded.contains("${") {
729            return Ok(expanded);
730        }
731    }
732
733    // Try environment variable
734    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
742/// Expand environment variables in a string
743fn expand_env_var(s: &str) -> String {
744    let mut result = s.to_string();
745
746    // Match ${VAR} pattern
747    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    // Also match $VAR pattern (no braces)
756    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
767/// Parse seed output for record counts
768fn parse_seed_output(line: &str) -> Option<u64> {
769    // Common patterns:
770    // "Created 10 users"
771    // "Seeded 100 records"
772    // "Inserted: 50"
773    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
794/// Parse affected rows from database output
795fn parse_affected_rows(output: &str) -> Option<u64> {
796    // PostgreSQL: "INSERT 0 5" or "UPDATE 3"
797    // MySQL: "Query OK, 5 rows affected"
798    // SQLite: no standard format
799
800    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
824/// Create a Cargo.toml for standalone seed script
825fn create_seed_cargo_toml(project_root: &Path) -> CliResult<String> {
826    // Try to read the workspace Cargo.toml to get prax version
827    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        // Try to extract prax version from dependencies
831        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
850/// Extract prax-orm version from Cargo.toml
851fn extract_prax_version(content: &str) -> Option<String> {
852    // Look for prax-orm = "x.y.z" or prax-orm = { version = "x.y.z" }
853    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// =============================================================================
868// Tests
869// =============================================================================
870
871#[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        // SAFETY: Single-threaded test environment
915        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        // SAFETY: Single-threaded test environment
925        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        // Only an unsupported .ts seed exists -> nothing found
946        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        // A supported file alongside the .ts one is used instead of erroring
952        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        // No Cargo.toml at all -> fall back to the CLI's own (workspace) version
964        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        // Cargo.toml without a prax-orm dependency -> same fallback
972        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        // Embedded double quotes in identifiers are doubled
1006        assert!(sql.contains("INSERT INTO \"we\"\"ird\""));
1007        assert!(sql.contains("(\"na\"\"me\")"));
1008        // Values are string literals - double quotes need no escaping there
1009        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        // .rs and .sql seeds carry no table metadata: --reset must fail loudly
1094        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}