1use anyhow::{Context, Result};
2use serde::{Deserialize, Serialize};
3use std::collections::HashMap;
4use std::fs;
5use std::path::{Path, PathBuf};
6use walkdir::WalkDir;
7
8use crate::dag::{Materialization, ModelDag, ModelLayer, OutputFormat};
9
10#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
12pub struct LayerConfig {
13 pub path: PathBuf,
14 pub target_path: PathBuf,
15 pub default_format: Option<String>,
16}
17
18#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
20pub struct RbtProjectConfig {
21 pub name: String,
22 pub version: String,
23 pub models_dir: PathBuf,
24 pub target_path: PathBuf,
25 #[serde(default)]
26 pub layers: HashMap<String, LayerConfig>,
27}
28
29impl Default for RbtProjectConfig {
30 fn default() -> Self {
31 let mut layers = HashMap::new();
32 layers.insert(
33 "staging".to_string(),
34 LayerConfig {
35 path: PathBuf::from("models/staging"),
36 target_path: PathBuf::from("lake/silver"),
37 default_format: Some("parquet".to_string()),
38 },
39 );
40 layers.insert(
41 "transforms".to_string(),
42 LayerConfig {
43 path: PathBuf::from("models/transforms"),
44 target_path: PathBuf::from("lake/gold"),
45 default_format: Some("parquet".to_string()),
46 },
47 );
48 layers.insert(
49 "marts".to_string(),
50 LayerConfig {
51 path: PathBuf::from("models/marts"),
52 target_path: PathBuf::from("lake/gold"),
53 default_format: Some("parquet_and_iceberg".to_string()),
54 },
55 );
56
57 Self {
58 name: "rbt_project".to_string(),
59 version: "1.0.0".to_string(),
60 models_dir: PathBuf::from("models"),
61 target_path: PathBuf::from("lake/gold"),
62 layers,
63 }
64 }
65}
66
67impl RbtProjectConfig {
68 pub fn load(project_dir: &Path) -> Result<Self> {
70 let project_file = project_dir.join("rbt_project.yml");
71 if project_file.exists() {
72 let content = fs::read_to_string(&project_file)?;
73 let mut config: RbtProjectConfig = serde_yaml::from_str(&content)
74 .with_context(|| format!("Failed to parse {}", project_file.display()))?;
75
76 let defaults = Self::default();
77 for (key, val) in defaults.layers {
78 config.layers.entry(key).or_insert(val);
79 }
80 Ok(config)
81 } else {
82 Ok(Self::default())
83 }
84 }
85
86 pub fn resolve_model_target_path(
88 &self,
89 project_dir: &Path,
90 model_name: &str,
91 layer: ModelLayer,
92 ext: &str,
93 ) -> PathBuf {
94 let layer_key = match layer {
95 ModelLayer::Staging => "staging",
96 ModelLayer::Transform => "transforms",
97 ModelLayer::Mart => "marts",
98 };
99
100 let target_dir = if let Some(layer_cfg) = self.layers.get(layer_key) {
101 project_dir.join(&layer_cfg.target_path)
102 } else {
103 project_dir.join(&self.target_path)
104 };
105
106 target_dir.join(format!("{}.{}", model_name, ext))
107 }
108
109 pub fn resolve_model_target_dir(
111 &self,
112 project_dir: &Path,
113 model_name: &str,
114 layer: ModelLayer,
115 ) -> PathBuf {
116 let layer_key = match layer {
117 ModelLayer::Staging => "staging",
118 ModelLayer::Transform => "transforms",
119 ModelLayer::Mart => "marts",
120 };
121 let target_dir = if let Some(layer_cfg) = self.layers.get(layer_key) {
122 project_dir.join(&layer_cfg.target_path)
123 } else {
124 project_dir.join(&self.target_path)
125 };
126 target_dir.join(model_name)
127 }
128
129 pub fn build_dag(
131 &self,
132 project_dir: &Path,
133 cli_format_override: Option<OutputFormat>,
134 ) -> Result<ModelDag> {
135 let models_dir = project_dir.join(&self.models_dir);
136 let mut dag = ModelDag::new();
137
138 if !models_dir.exists() {
139 let default_fmt = cli_format_override.unwrap_or(OutputFormat::Parquet);
140 dag.add_model_with_format(
141 "stg_users",
142 "SELECT 1 AS id, 'Alice' AS name, 'admin' AS role",
143 Materialization::Table,
144 default_fmt,
145 None,
146 "",
147 )?;
148 dag.build_graph()?;
149 return Ok(dag);
150 }
151
152 let mut model_count = 0;
153 for entry in WalkDir::new(&models_dir).into_iter().filter_map(|e| e.ok()) {
154 let path = entry.path();
155 if path.is_file() && path.extension().map_or(false, |ext| ext == "sql") {
156 let stem = path
157 .file_stem()
158 .and_then(|s| s.to_str())
159 .context("Invalid file stem")?;
160 let raw_sql = fs::read_to_string(path)?;
161
162 let layer = ModelLayer::from_name(stem);
163 let format = cli_format_override.clone().unwrap_or_else(|| {
164 let layer_key = match layer {
165 ModelLayer::Staging => "staging",
166 ModelLayer::Transform => "transforms",
167 ModelLayer::Mart => "marts",
168 };
169 if let Some(l_cfg) = self.layers.get(layer_key) {
170 match l_cfg.default_format.as_deref() {
171 Some("parquet") => OutputFormat::Parquet,
172 Some("jsonl") => OutputFormat::Jsonl,
173 Some("csv") => OutputFormat::Csv,
174 Some("iceberg") => OutputFormat::Iceberg,
175 Some("parquet_and_iceberg") => OutputFormat::ParquetAndIceberg,
176 _ => OutputFormat::Parquet,
177 }
178 } else {
179 OutputFormat::Parquet
180 }
181 });
182
183 let target_file_path = match format {
184 OutputFormat::Iceberg => {
185 self.resolve_model_target_dir(project_dir, stem, layer)
187 }
188 OutputFormat::Parquet
189 | OutputFormat::ParquetAndIceberg
190 | OutputFormat::ZeroCopyClone => {
191 self.resolve_model_target_path(project_dir, stem, layer, "parquet")
192 }
193 OutputFormat::Jsonl => {
194 self.resolve_model_target_path(project_dir, stem, layer, "jsonl")
195 }
196 OutputFormat::Csv => {
197 self.resolve_model_target_path(project_dir, stem, layer, "csv")
198 }
199 };
200
201 dag.add_model_with_format(
202 stem,
203 &raw_sql,
204 Materialization::Table,
205 format,
206 Some(target_file_path.to_string_lossy().to_string()),
207 "",
208 )?;
209 model_count += 1;
210 }
211 }
212
213 if model_count == 0 {
214 let default_fmt = cli_format_override.unwrap_or(OutputFormat::Parquet);
215 dag.add_model_with_format(
216 "stg_users",
217 "SELECT 1 AS id, 'Alice' AS name, 'admin' AS role",
218 Materialization::Table,
219 default_fmt,
220 None,
221 "",
222 )?;
223 }
224
225 dag.build_graph()?;
226 Ok(dag)
227 }
228}
229
230#[cfg(test)]
231mod tests {
232 use super::*;
233
234 #[test]
235 fn test_layer_target_path_resolution() -> Result<()> {
236 let config = RbtProjectConfig::default();
237 let project_dir = Path::new("/tmp/test_project");
238
239 let stg_path = config.resolve_model_target_path(project_dir, "stg_trades", ModelLayer::Staging, "parquet");
240 assert_eq!(stg_path, project_dir.join("lake/silver/stg_trades.parquet"));
241
242 let tf_path = config.resolve_model_target_path(project_dir, "tf_1m_bars", ModelLayer::Transform, "parquet");
243 assert_eq!(tf_path, project_dir.join("lake/gold/tf_1m_bars.parquet"));
244
245 let mart_path = config.resolve_model_target_path(project_dir, "fact_1d_bars", ModelLayer::Mart, "parquet");
246 assert_eq!(mart_path, project_dir.join("lake/gold/fact_1d_bars.parquet"));
247
248 Ok(())
249 }
250}