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
12pub 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 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
40fn 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
106pub 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 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 let config = crate::config::load_config_or_default(Some(project_root.clone()))
124 .map_err(|e| format!("Failed to load config: {e}"))?;
125
126 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
139fn 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
206pub 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 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 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 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 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 #[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 let original = env::var("CARGO_MANIFEST_DIR").ok();
446
447 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 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 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 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 let result = load_models_from_dir(None);
658 let _ = result;
662 }
663
664 #[test]
665 #[serial]
666 fn test_load_models_at_compile_time() {
667 let result = load_models_at_compile_time();
671 let _ = result;
675 }
676}