Skip to main content

vespertide_loader/
models.rs

1use std::fs;
2use std::path::{Path, PathBuf};
3
4use anyhow::{Context, Result};
5use rayon::prelude::*;
6use vespertide_config::VespertideConfig;
7use vespertide_core::TableDef;
8use vespertide_planner::validate_schema;
9
10use crate::parallel_config::{LOAD_FILES_PAR_MIN_LEN, LOAD_FILES_PAR_THRESHOLD};
11
12/// Load all model definitions from the models directory (recursively).
13pub fn load_models(config: &VespertideConfig) -> Result<Vec<TableDef>> {
14    let models_dir = config.models_dir();
15    if !models_dir.exists() {
16        return Ok(Vec::new());
17    }
18
19    let mut tables = Vec::new();
20    load_models_recursive(models_dir, &mut tables)?;
21
22    // Validate schema integrity using normalized version
23    // But return the original tables to preserve inline constraints
24    if !tables.is_empty() {
25        let normalized_tables: Vec<TableDef> = tables
26            .iter()
27            .map(|t| {
28                t.normalize()
29                    .map_err(|e| anyhow::anyhow!("Failed to normalize table '{}': {}", t.name, e))
30            })
31            .collect::<Result<Vec<_>, _>>()?;
32
33        validate_schema(&normalized_tables)
34            .map_err(|e| anyhow::anyhow!("schema validation failed: {e}"))?;
35    }
36
37    Ok(tables)
38}
39
40/// Recursively walk directory and load model files.
41fn load_models_recursive(dir: &Path, tables: &mut Vec<TableDef>) -> Result<()> {
42    let paths = collect_model_paths(dir)?;
43    let results: Vec<Result<TableDef>> = if paths.len() < LOAD_FILES_PAR_THRESHOLD {
44        paths.iter().map(|path| load_model_file(path)).collect()
45    } else {
46        paths
47            .par_iter()
48            .with_min_len(LOAD_FILES_PAR_MIN_LEN)
49            .map(|path| load_model_file(path))
50            .collect()
51    };
52
53    for result in results {
54        tables.push(result?);
55    }
56
57    Ok(())
58}
59
60fn collect_model_paths(dir: &Path) -> Result<Vec<PathBuf>> {
61    let entries =
62        fs::read_dir(dir).with_context(|| format!("read models directory: {}", dir.display()))?;
63    let mut paths = Vec::new();
64
65    for entry in entries {
66        let entry = entry.context("read directory entry")?;
67        let path = entry.path();
68
69        if path.is_dir() {
70            paths.extend(collect_model_paths(&path)?);
71        } else if path.is_file() && has_model_extension(&path) {
72            paths.push(path);
73        }
74    }
75
76    Ok(paths)
77}
78
79fn has_model_extension(path: &Path) -> bool {
80    matches!(
81        path.extension().and_then(|s| s.to_str()),
82        Some("json" | "yaml" | "yml")
83    )
84}
85
86fn load_model_file(path: &Path) -> Result<TableDef> {
87    let ext = path.extension().and_then(|s| s.to_str());
88    let content =
89        fs::read_to_string(path).with_context(|| format!("read model file: {}", path.display()))?;
90
91    let table: TableDef = if ext == Some("json") {
92        serde_json::from_str(&content)
93            .with_context(|| format!("parse JSON model: {}", path.display()))?
94    } else {
95        serde_yaml::from_str(&content)
96            .with_context(|| format!("parse YAML model: {}", path.display()))?
97    };
98
99    table
100        .validate_unique_column_names()
101        .with_context(|| format!("validate model: {}", path.display()))?;
102
103    Ok(table)
104}
105
106/// Load models from a specific directory (for compile-time use in macros).
107pub fn load_models_from_dir(
108    project_root: Option<std::path::PathBuf>,
109) -> Result<Vec<TableDef>, Box<dyn std::error::Error>> {
110    use std::env;
111
112    // Locate project root from CARGO_MANIFEST_DIR or use provided path
113    let project_root = if let Some(root) = project_root {
114        root
115    } else {
116        std::path::PathBuf::from(
117            env::var("CARGO_MANIFEST_DIR")
118                .context("CARGO_MANIFEST_DIR environment variable not set")?,
119        )
120    };
121
122    // Read vespertide.json or use defaults
123    let config = crate::config::load_config_or_default(Some(project_root.clone()))
124        .map_err(|e| format!("Failed to load config: {e}"))?;
125
126    // Read models directory
127    let models_dir = project_root.join(config.models_dir());
128    if !models_dir.exists() {
129        return Ok(Vec::new());
130    }
131
132    let mut tables = Vec::new();
133    load_models_recursive_internal(&models_dir, &mut tables)
134        .map_err(|e| format!("Failed to load models: {e}"))?;
135
136    Ok(tables)
137}
138
139/// Internal recursive function for loading models (used by both runtime and compile-time).
140fn load_models_recursive_internal(
141    dir: &Path,
142    tables: &mut Vec<TableDef>,
143) -> Result<(), Box<dyn std::error::Error>> {
144    let paths = collect_model_paths_internal(dir)?;
145    let results: Vec<Result<TableDef, String>> = if paths.len() < LOAD_FILES_PAR_THRESHOLD {
146        paths
147            .iter()
148            .map(|path| load_normalized_model_file_internal(path))
149            .collect()
150    } else {
151        paths
152            .par_iter()
153            .with_min_len(LOAD_FILES_PAR_MIN_LEN)
154            .map(|path| load_normalized_model_file_internal(path))
155            .collect()
156    };
157
158    for result in results {
159        tables.push(result.map_err(|e| -> Box<dyn std::error::Error> { e.into() })?);
160    }
161
162    Ok(())
163}
164
165fn collect_model_paths_internal(dir: &Path) -> Result<Vec<PathBuf>, String> {
166    let entries = fs::read_dir(dir)
167        .map_err(|e| format!("Failed to read models directory {}: {}", dir.display(), e))?;
168    let mut paths = Vec::new();
169
170    for entry in entries {
171        let entry = entry.map_err(|e| format!("Failed to read directory entry: {e}"))?;
172        let path = entry.path();
173
174        if path.is_dir() {
175            paths.extend(collect_model_paths_internal(&path)?);
176        } else if path.is_file() && has_model_extension(&path) {
177            paths.push(path);
178        }
179    }
180
181    Ok(paths)
182}
183
184fn load_normalized_model_file_internal(path: &Path) -> Result<TableDef, String> {
185    let ext = path.extension().and_then(|s| s.to_str());
186    let content = fs::read_to_string(path)
187        .map_err(|e| format!("Failed to read model file {}: {}", path.display(), e))?;
188
189    let table: TableDef = if ext == Some("json") {
190        serde_json::from_str(&content)
191            .map_err(|e| format!("Failed to parse JSON model {}: {}", path.display(), e))?
192    } else {
193        serde_yaml::from_str(&content)
194            .map_err(|e| format!("Failed to parse YAML model {}: {}", path.display(), e))?
195    };
196
197    table
198        .validate_unique_column_names()
199        .map_err(|e| format!("Failed to validate model {}: {}", path.display(), e))?;
200
201    table
202        .normalize()
203        .map_err(|e| format!("Failed to normalize table '{}': {}", table.name, e))
204}
205
206/// Load models at compile time (for macro use).
207pub fn load_models_at_compile_time() -> Result<Vec<TableDef>, Box<dyn std::error::Error>> {
208    load_models_from_dir(None)
209}
210
211#[cfg(test)]
212mod tests {
213    use super::*;
214    use serial_test::serial;
215    use std::fs;
216    use tempfile::tempdir;
217    use vespertide_core::{
218        ColumnDef, ColumnType, SimpleColumnType, TableConstraint,
219        schema::foreign_key::ForeignKeySyntax,
220    };
221
222    struct CwdGuard {
223        original: std::path::PathBuf,
224    }
225
226    impl CwdGuard {
227        fn new(dir: &std::path::PathBuf) -> Self {
228            let original = std::env::current_dir().unwrap();
229            std::env::set_current_dir(dir).unwrap();
230            Self { original }
231        }
232    }
233
234    impl Drop for CwdGuard {
235        fn drop(&mut self) {
236            let _ = std::env::set_current_dir(&self.original);
237        }
238    }
239
240    fn write_config() {
241        let cfg = VespertideConfig::default();
242        let text = serde_json::to_string_pretty(&cfg).unwrap();
243        fs::write("vespertide.json", text).unwrap();
244    }
245
246    #[test]
247    #[serial]
248    fn load_models_returns_empty_when_no_models_dir() {
249        let tmp = tempdir().unwrap();
250        let _guard = CwdGuard::new(&tmp.path().to_path_buf());
251        write_config();
252
253        // Don't create models directory
254        let models = load_models(&VespertideConfig::default()).unwrap();
255        assert_eq!(models.len(), 0);
256    }
257
258    #[test]
259    #[serial]
260    fn load_models_reads_yaml_and_validates() {
261        let tmp = tempdir().unwrap();
262        let _guard = CwdGuard::new(&tmp.path().to_path_buf());
263        write_config();
264
265        fs::create_dir_all("models").unwrap();
266        let table = TableDef {
267            name: "users".into(),
268            description: None,
269            columns: vec![ColumnDef {
270                name: "id".into(),
271                r#type: ColumnType::Simple(SimpleColumnType::Integer),
272                nullable: false,
273                default: None,
274                comment: None,
275                primary_key: None,
276                unique: None,
277                index: None,
278                foreign_key: None,
279            }],
280            constraints: vec![TableConstraint::PrimaryKey {
281                auto_increment: false,
282                columns: vec!["id".into()],
283                strategy: vespertide_core::PrimaryKeyAdditionStrategy::default(),
284            }],
285        };
286        fs::write("models/users.yaml", serde_yaml::to_string(&table).unwrap()).unwrap();
287
288        let models = load_models(&VespertideConfig::default()).unwrap();
289        assert_eq!(models.len(), 1);
290        assert_eq!(models[0].name, "users");
291    }
292
293    #[test]
294    #[serial]
295    fn load_models_recursive_processes_subdirectories() {
296        let tmp = tempdir().unwrap();
297        let _guard = CwdGuard::new(&tmp.path().to_path_buf());
298        write_config();
299
300        fs::create_dir_all("models/subdir").unwrap();
301
302        // Create model in subdirectory
303        let table = TableDef {
304            name: "subtable".into(),
305            description: None,
306            columns: vec![ColumnDef {
307                name: "id".into(),
308                r#type: ColumnType::Simple(SimpleColumnType::Integer),
309                nullable: false,
310                default: None,
311                comment: None,
312                primary_key: None,
313                unique: None,
314                index: None,
315                foreign_key: None,
316            }],
317            constraints: vec![TableConstraint::PrimaryKey {
318                auto_increment: false,
319                columns: vec!["id".into()],
320                strategy: vespertide_core::PrimaryKeyAdditionStrategy::default(),
321            }],
322        };
323        let content = serde_json::to_string_pretty(&table).unwrap();
324        fs::write("models/subdir/subtable.json", content).unwrap();
325
326        let models = load_models(&VespertideConfig::default()).unwrap();
327        assert_eq!(models.len(), 1);
328        assert_eq!(models[0].name, "subtable");
329    }
330
331    #[test]
332    #[serial]
333    fn load_models_fails_on_invalid_fk_format() {
334        let tmp = tempdir().unwrap();
335        let _guard = CwdGuard::new(&tmp.path().to_path_buf());
336        write_config();
337
338        fs::create_dir_all("models").unwrap();
339
340        // Create a model with invalid FK string format (missing dot separator)
341        let table = TableDef {
342            name: "orders".into(),
343            description: None,
344            columns: vec![ColumnDef {
345                name: "user_id".into(),
346                r#type: ColumnType::Simple(SimpleColumnType::Integer),
347                nullable: false,
348                default: None,
349                comment: None,
350                primary_key: None,
351                unique: None,
352                index: None,
353                // Invalid FK format: should be "table.column" but missing the dot
354                foreign_key: Some(ForeignKeySyntax::String("invalid_format".into())),
355            }],
356            constraints: vec![],
357        };
358        fs::write(
359            "models/orders.json",
360            serde_json::to_string_pretty(&table).unwrap(),
361        )
362        .unwrap();
363
364        let result = load_models(&VespertideConfig::default());
365        assert!(result.is_err());
366        let err_msg = result.unwrap_err().to_string();
367        assert!(err_msg.contains("Failed to normalize table 'orders'"));
368    }
369
370    // A non-model-extension file (e.g. `.txt`) must be ignored by the
371    // collector. Pins `path.is_file() && has_model_extension(&path)`: a
372    // `&& -> ||` mutant would pick up the `.txt`, then fail to parse it as a
373    // model and surface an error instead of an empty load.
374    #[test]
375    #[serial]
376    fn load_models_ignores_non_model_extension_files() {
377        let tmp = tempdir().unwrap();
378        let _guard = CwdGuard::new(&tmp.path().to_path_buf());
379        write_config();
380        fs::create_dir_all("models").unwrap();
381        fs::write("models/README.txt", "not a model: {{{ invalid").unwrap();
382
383        let models = load_models(&VespertideConfig::default()).unwrap();
384        assert_eq!(models.len(), 0, "the .txt file must be skipped");
385    }
386
387    #[test]
388    #[serial]
389    fn load_models_from_dir_ignores_non_model_extension_files() {
390        let temp_dir = tempdir().unwrap();
391        let models_dir = temp_dir.path().join("models");
392        fs::create_dir_all(&models_dir).unwrap();
393        fs::write(models_dir.join("notes.txt"), "not a model: {{{ invalid").unwrap();
394
395        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
396        assert!(
397            result.is_ok(),
398            "the .txt file must be skipped, not parsed: {result:?}"
399        );
400        assert_eq!(result.unwrap().len(), 0);
401    }
402
403    #[test]
404    #[serial]
405    fn test_load_models_from_dir_with_root() {
406        let temp_dir = tempdir().unwrap();
407        let models_dir = temp_dir.path().join("models");
408        fs::create_dir_all(&models_dir).unwrap();
409
410        let table = TableDef {
411            name: "users".into(),
412            description: None,
413            columns: vec![ColumnDef {
414                name: "id".into(),
415                r#type: ColumnType::Simple(SimpleColumnType::Integer),
416                nullable: false,
417                default: None,
418                comment: None,
419                primary_key: None,
420                unique: None,
421                index: None,
422                foreign_key: None,
423            }],
424            constraints: vec![],
425        };
426        fs::write(
427            models_dir.join("users.json"),
428            serde_json::to_string_pretty(&table).unwrap(),
429        )
430        .unwrap();
431
432        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
433        assert!(result.is_ok());
434        let models = result.unwrap();
435        assert_eq!(models.len(), 1);
436        assert_eq!(models[0].name, "users");
437    }
438
439    #[test]
440    #[serial]
441    fn test_load_models_from_dir_without_root() {
442        use std::env;
443
444        // Save the original value
445        let original = env::var("CARGO_MANIFEST_DIR").ok();
446
447        // Remove CARGO_MANIFEST_DIR to test the error path
448        unsafe {
449            env::remove_var("CARGO_MANIFEST_DIR");
450        }
451
452        let result = load_models_from_dir(None);
453        assert!(result.is_err());
454        let err_msg = result.unwrap_err().to_string();
455        assert!(err_msg.contains("CARGO_MANIFEST_DIR environment variable not set"));
456
457        // Restore the original value if it existed
458        if let Some(val) = original {
459            unsafe {
460                env::set_var("CARGO_MANIFEST_DIR", val);
461            }
462        }
463    }
464
465    #[test]
466    #[serial]
467    fn test_load_models_from_dir_no_models_dir() {
468        let temp_dir = tempdir().unwrap();
469        // Don't create models directory
470
471        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
472        assert!(result.is_ok());
473        let models = result.unwrap();
474        assert_eq!(models.len(), 0);
475    }
476
477    #[test]
478    #[serial]
479    fn test_load_models_from_dir_with_yaml() {
480        let temp_dir = tempdir().unwrap();
481        let models_dir = temp_dir.path().join("models");
482        fs::create_dir_all(&models_dir).unwrap();
483
484        let table = TableDef {
485            name: "users".into(),
486            description: None,
487            columns: vec![ColumnDef {
488                name: "id".into(),
489                r#type: ColumnType::Simple(SimpleColumnType::Integer),
490                nullable: false,
491                default: None,
492                comment: None,
493                primary_key: None,
494                unique: None,
495                index: None,
496                foreign_key: None,
497            }],
498            constraints: vec![],
499        };
500        fs::write(
501            models_dir.join("users.yaml"),
502            serde_yaml::to_string(&table).unwrap(),
503        )
504        .unwrap();
505
506        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
507        assert!(result.is_ok());
508        let models = result.unwrap();
509        assert_eq!(models.len(), 1);
510        assert_eq!(models[0].name, "users");
511    }
512
513    #[test]
514    #[serial]
515    fn test_load_models_from_dir_with_yml() {
516        let temp_dir = tempdir().unwrap();
517        let models_dir = temp_dir.path().join("models");
518        fs::create_dir_all(&models_dir).unwrap();
519
520        let table = TableDef {
521            name: "users".into(),
522            description: None,
523            columns: vec![ColumnDef {
524                name: "id".into(),
525                r#type: ColumnType::Simple(SimpleColumnType::Integer),
526                nullable: false,
527                default: None,
528                comment: None,
529                primary_key: None,
530                unique: None,
531                index: None,
532                foreign_key: None,
533            }],
534            constraints: vec![],
535        };
536        fs::write(
537            models_dir.join("users.yml"),
538            serde_yaml::to_string(&table).unwrap(),
539        )
540        .unwrap();
541
542        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
543        assert!(result.is_ok());
544        let models = result.unwrap();
545        assert_eq!(models.len(), 1);
546        assert_eq!(models[0].name, "users");
547    }
548
549    #[test]
550    #[serial]
551    fn test_load_models_from_dir_recursive() {
552        let temp_dir = tempdir().unwrap();
553        let models_dir = temp_dir.path().join("models");
554        let subdir = models_dir.join("subdir");
555        fs::create_dir_all(&subdir).unwrap();
556
557        let table = TableDef {
558            name: "subtable".into(),
559            description: None,
560            columns: vec![ColumnDef {
561                name: "id".into(),
562                r#type: ColumnType::Simple(SimpleColumnType::Integer),
563                nullable: false,
564                default: None,
565                comment: None,
566                primary_key: None,
567                unique: None,
568                index: None,
569                foreign_key: None,
570            }],
571            constraints: vec![],
572        };
573        fs::write(
574            subdir.join("subtable.json"),
575            serde_json::to_string_pretty(&table).unwrap(),
576        )
577        .unwrap();
578
579        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
580        assert!(result.is_ok());
581        let models = result.unwrap();
582        assert_eq!(models.len(), 1);
583        assert_eq!(models[0].name, "subtable");
584    }
585
586    #[test]
587    #[serial]
588    fn test_load_models_from_dir_with_invalid_json() {
589        let temp_dir = tempdir().unwrap();
590        let models_dir = temp_dir.path().join("models");
591        fs::create_dir_all(&models_dir).unwrap();
592
593        fs::write(models_dir.join("invalid.json"), r#"{"invalid": json}"#).unwrap();
594
595        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
596        assert!(result.is_err());
597        let err_msg = result.unwrap_err().to_string();
598        assert!(err_msg.contains("Failed to parse JSON model"));
599    }
600
601    #[test]
602    #[serial]
603    fn test_load_models_from_dir_with_invalid_yaml() {
604        let temp_dir = tempdir().unwrap();
605        let models_dir = temp_dir.path().join("models");
606        fs::create_dir_all(&models_dir).unwrap();
607
608        fs::write(models_dir.join("invalid.yaml"), r"invalid: [yaml").unwrap();
609
610        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
611        assert!(result.is_err());
612        let err_msg = result.unwrap_err().to_string();
613        assert!(err_msg.contains("Failed to parse YAML model"));
614    }
615
616    #[test]
617    #[serial]
618    fn test_load_models_from_dir_normalization_error() {
619        let temp_dir = tempdir().unwrap();
620        let models_dir = temp_dir.path().join("models");
621        fs::create_dir_all(&models_dir).unwrap();
622
623        // Create a model with invalid FK format
624        let table = TableDef {
625            name: "orders".into(),
626            description: None,
627            columns: vec![ColumnDef {
628                name: "user_id".into(),
629                r#type: ColumnType::Simple(SimpleColumnType::Integer),
630                nullable: false,
631                default: None,
632                comment: None,
633                primary_key: None,
634                unique: None,
635                index: None,
636                foreign_key: Some(ForeignKeySyntax::String("invalid_format".into())),
637            }],
638            constraints: vec![],
639        };
640        fs::write(
641            models_dir.join("orders.json"),
642            serde_json::to_string_pretty(&table).unwrap(),
643        )
644        .unwrap();
645
646        let result = load_models_from_dir(Some(temp_dir.path().to_path_buf()));
647        assert!(result.is_err());
648        let err_msg = result.unwrap_err().to_string();
649        assert!(err_msg.contains("Failed to normalize table 'orders'"));
650    }
651
652    #[test]
653    #[serial]
654    fn test_load_models_from_dir_with_cargo_manifest_dir() {
655        // Test the path where CARGO_MANIFEST_DIR is set (line 87)
656        // In cargo test environment, CARGO_MANIFEST_DIR is usually set
657        let result = load_models_from_dir(None);
658        // This might succeed if CARGO_MANIFEST_DIR is set (like in cargo test)
659        // or fail if it's not set
660        // Either way, we're testing the code path including line 87
661        let _ = result;
662    }
663
664    #[test]
665    #[serial]
666    fn test_load_models_at_compile_time() {
667        // This function just calls load_models_from_dir(None)
668        // We can't easily test it without CARGO_MANIFEST_DIR, but we can verify
669        // it doesn't panic
670        let result = load_models_at_compile_time();
671        // This might succeed if CARGO_MANIFEST_DIR is set (like in cargo test)
672        // or fail if it's not set
673        // Either way, we're testing the code path
674        let _ = result;
675    }
676}