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 super::dag::{Materialization, ModelDag, ModelLayer, OutputFormat};
9
10pub const DEFAULT_MEMTABLE_MAX_ROWS: usize = 50_000;
12
13#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
18#[serde(rename_all = "snake_case")]
19pub enum RefStrategy {
20 #[default]
22 #[serde(alias = "parquet_reread", alias = "lake", alias = "file")]
23 Parquet,
24 #[serde(alias = "mem_table", alias = "memory", alias = "arc")]
26 Memtable,
27}
28
29#[derive(Debug, Clone, Copy, PartialEq, Eq)]
31pub enum RefBackend {
32 MemTable,
34 LakeFile,
36}
37
38#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
42pub struct MaterializeConfig {
43 #[serde(default)]
45 pub ref_strategy: RefStrategy,
46 #[serde(default = "default_memtable_max_rows")]
48 pub memtable_max_rows: usize,
49}
50
51fn default_memtable_max_rows() -> usize {
52 DEFAULT_MEMTABLE_MAX_ROWS
53}
54
55impl Default for MaterializeConfig {
56 fn default() -> Self {
57 Self {
58 ref_strategy: RefStrategy::Parquet,
59 memtable_max_rows: DEFAULT_MEMTABLE_MAX_ROWS,
60 }
61 }
62}
63
64impl MaterializeConfig {
65 pub fn choose_ref_backend(&self, row_count: usize) -> RefBackend {
67 match self.ref_strategy {
68 RefStrategy::Parquet => RefBackend::LakeFile,
69 RefStrategy::Memtable if row_count < self.memtable_max_rows => RefBackend::MemTable,
70 RefStrategy::Memtable => RefBackend::LakeFile,
71 }
72 }
73}
74
75#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
77pub struct LayerConfig {
78 pub path: PathBuf,
79 pub target_path: PathBuf,
80 pub default_format: Option<String>,
81}
82
83#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
85pub struct RbtProjectConfig {
86 pub name: String,
87 pub version: String,
88 pub models_dir: PathBuf,
89 pub target_path: PathBuf,
90 #[serde(default)]
91 pub layers: HashMap<String, LayerConfig>,
92 #[serde(default)]
94 pub materialize: MaterializeConfig,
95}
96
97impl Default for RbtProjectConfig {
98 fn default() -> Self {
99 let mut layers = HashMap::new();
100 layers.insert(
101 "staging".to_string(),
102 LayerConfig {
103 path: PathBuf::from("models/staging"),
104 target_path: PathBuf::from("lake/silver"),
105 default_format: Some("parquet".to_string()),
106 },
107 );
108 layers.insert(
109 "transforms".to_string(),
110 LayerConfig {
111 path: PathBuf::from("models/transforms"),
112 target_path: PathBuf::from("lake/gold"),
113 default_format: Some("parquet".to_string()),
114 },
115 );
116 layers.insert(
117 "marts".to_string(),
118 LayerConfig {
119 path: PathBuf::from("models/marts"),
120 target_path: PathBuf::from("lake/gold"),
121 default_format: Some("parquet_and_iceberg".to_string()),
122 },
123 );
124
125 Self {
126 name: "rbt_project".to_string(),
127 version: "1.0.0".to_string(),
128 models_dir: PathBuf::from("models"),
129 target_path: PathBuf::from("lake/gold"),
130 layers,
131 materialize: MaterializeConfig::default(),
132 }
133 }
134}
135
136impl RbtProjectConfig {
137 pub fn load(project_dir: &Path) -> Result<Self> {
139 let project_file = project_dir.join("rbt_project.yml");
140 if project_file.exists() {
141 let content = fs::read_to_string(&project_file)?;
142 let mut config: RbtProjectConfig = serde_yaml::from_str(&content)
143 .with_context(|| format!("Failed to parse {}", project_file.display()))?;
144
145 let defaults = Self::default();
146 for (key, val) in defaults.layers {
147 config.layers.entry(key).or_insert(val);
148 }
149 Ok(config)
150 } else {
151 Ok(Self::default())
152 }
153 }
154
155 pub fn resolve_model_target_path(
157 &self,
158 project_dir: &Path,
159 model_name: &str,
160 layer: ModelLayer,
161 ext: &str,
162 ) -> PathBuf {
163 let layer_key = match layer {
164 ModelLayer::Staging => "staging",
165 ModelLayer::Transform => "transforms",
166 ModelLayer::Mart => "marts",
167 };
168
169 let target_dir = if let Some(layer_cfg) = self.layers.get(layer_key) {
170 project_dir.join(&layer_cfg.target_path)
171 } else {
172 project_dir.join(&self.target_path)
173 };
174
175 target_dir.join(format!("{}.{}", model_name, ext))
176 }
177
178 pub fn resolve_model_target_dir(
180 &self,
181 project_dir: &Path,
182 model_name: &str,
183 layer: ModelLayer,
184 ) -> PathBuf {
185 let layer_key = match layer {
186 ModelLayer::Staging => "staging",
187 ModelLayer::Transform => "transforms",
188 ModelLayer::Mart => "marts",
189 };
190 let target_dir = if let Some(layer_cfg) = self.layers.get(layer_key) {
191 project_dir.join(&layer_cfg.target_path)
192 } else {
193 project_dir.join(&self.target_path)
194 };
195 target_dir.join(model_name)
196 }
197
198 pub fn build_dag(
200 &self,
201 project_dir: &Path,
202 cli_format_override: Option<OutputFormat>,
203 ) -> Result<ModelDag> {
204 let models_dir = project_dir.join(&self.models_dir);
205 let mut dag = ModelDag::new();
206
207 if !models_dir.exists() {
208 let default_fmt = cli_format_override.unwrap_or(OutputFormat::Parquet);
209 dag.add_model_with_format(
210 "stg_users",
211 "SELECT 1 AS id, 'Alice' AS name, 'admin' AS role",
212 Materialization::Table,
213 default_fmt,
214 None,
215 "",
216 )?;
217 dag.build_graph()?;
218 return Ok(dag);
219 }
220
221 let mut model_count = 0;
222 for entry in WalkDir::new(&models_dir).into_iter().filter_map(|e| e.ok()) {
223 let path = entry.path();
224 if path.is_file() && path.extension().is_some_and(|ext| ext == "sql") {
225 let stem = path
226 .file_stem()
227 .and_then(|s| s.to_str())
228 .context("Invalid file stem")?;
229 let raw_sql = fs::read_to_string(path)?;
230
231 let layer = ModelLayer::from_name(stem);
232 let format = cli_format_override.clone().unwrap_or_else(|| {
233 let layer_key = match layer {
234 ModelLayer::Staging => "staging",
235 ModelLayer::Transform => "transforms",
236 ModelLayer::Mart => "marts",
237 };
238 if let Some(l_cfg) = self.layers.get(layer_key) {
239 match l_cfg.default_format.as_deref() {
240 Some("parquet") => OutputFormat::Parquet,
241 Some("jsonl") => OutputFormat::Jsonl,
242 Some("csv") => OutputFormat::Csv,
243 Some("iceberg") => OutputFormat::Iceberg,
244 Some("parquet_and_iceberg") => OutputFormat::ParquetAndIceberg,
245 _ => OutputFormat::Parquet,
246 }
247 } else {
248 OutputFormat::Parquet
249 }
250 });
251
252 let target_file_path = match format {
253 OutputFormat::Iceberg => {
254 self.resolve_model_target_dir(project_dir, stem, layer)
256 }
257 OutputFormat::Parquet
258 | OutputFormat::ParquetAndIceberg
259 | OutputFormat::ZeroCopyClone => {
260 self.resolve_model_target_path(project_dir, stem, layer, "parquet")
261 }
262 OutputFormat::Jsonl => {
263 self.resolve_model_target_path(project_dir, stem, layer, "jsonl")
264 }
265 OutputFormat::Csv => {
266 self.resolve_model_target_path(project_dir, stem, layer, "csv")
267 }
268 };
269
270 dag.add_model_with_format(
271 stem,
272 &raw_sql,
273 Materialization::Table,
274 format,
275 Some(target_file_path.to_string_lossy().to_string()),
276 "",
277 )?;
278 model_count += 1;
279 }
280 }
281
282 if model_count == 0 {
283 let default_fmt = cli_format_override.unwrap_or(OutputFormat::Parquet);
284 dag.add_model_with_format(
285 "stg_users",
286 "SELECT 1 AS id, 'Alice' AS name, 'admin' AS role",
287 Materialization::Table,
288 default_fmt,
289 None,
290 "",
291 )?;
292 }
293
294 dag.build_graph()?;
295 Ok(dag)
296 }
297}
298
299#[cfg(test)]
300mod tests {
301 use super::*;
302
303 #[test]
304 fn test_layer_target_path_resolution() -> Result<()> {
305 let config = RbtProjectConfig::default();
306 let project_dir = Path::new("/tmp/test_project");
307
308 let stg_path = config.resolve_model_target_path(
309 project_dir,
310 "stg_trades",
311 ModelLayer::Staging,
312 "parquet",
313 );
314 assert_eq!(stg_path, project_dir.join("lake/silver/stg_trades.parquet"));
315
316 let tf_path = config.resolve_model_target_path(
317 project_dir,
318 "tf_1m_bars",
319 ModelLayer::Transform,
320 "parquet",
321 );
322 assert_eq!(tf_path, project_dir.join("lake/gold/tf_1m_bars.parquet"));
323
324 let mart_path = config.resolve_model_target_path(
325 project_dir,
326 "fact_1d_bars",
327 ModelLayer::Mart,
328 "parquet",
329 );
330 assert_eq!(
331 mart_path,
332 project_dir.join("lake/gold/fact_1d_bars.parquet")
333 );
334
335 Ok(())
336 }
337
338 #[test]
339 fn materialize_defaults_to_parquet_reread() {
340 let cfg = MaterializeConfig::default();
341 assert_eq!(cfg.ref_strategy, RefStrategy::Parquet);
342 assert_eq!(cfg.memtable_max_rows, DEFAULT_MEMTABLE_MAX_ROWS);
343 assert_eq!(cfg.choose_ref_backend(0), RefBackend::LakeFile);
344 assert_eq!(cfg.choose_ref_backend(1_000_000), RefBackend::LakeFile);
345 }
346
347 #[test]
348 fn materialize_memtable_respects_cutoff() {
349 let cfg = MaterializeConfig {
350 ref_strategy: RefStrategy::Memtable,
351 memtable_max_rows: 50_000,
352 };
353 assert_eq!(cfg.choose_ref_backend(49_999), RefBackend::MemTable);
354 assert_eq!(cfg.choose_ref_backend(50_000), RefBackend::LakeFile);
355 assert_eq!(cfg.choose_ref_backend(50_001), RefBackend::LakeFile);
356 }
357
358 #[test]
359 fn parse_materialize_block_from_yaml() -> Result<()> {
360 let yml = r#"
361name: t
362version: "1"
363models_dir: models
364target_path: lake/gold
365materialize:
366 ref_strategy: memtable
367 memtable_max_rows: 10000
368"#;
369 let cfg: RbtProjectConfig = serde_yaml::from_str(yml)?;
370 assert_eq!(cfg.materialize.ref_strategy, RefStrategy::Memtable);
371 assert_eq!(cfg.materialize.memtable_max_rows, 10_000);
372 assert_eq!(
373 cfg.materialize.choose_ref_backend(9_999),
374 RefBackend::MemTable
375 );
376 assert_eq!(
377 cfg.materialize.choose_ref_backend(10_000),
378 RefBackend::LakeFile
379 );
380 Ok(())
381 }
382
383 #[test]
384 fn parse_project_without_materialize_uses_defaults() -> Result<()> {
385 let yml = r#"
386name: t
387version: "1"
388models_dir: models
389target_path: lake/gold
390"#;
391 let cfg: RbtProjectConfig = serde_yaml::from_str(yml)?;
392 assert_eq!(cfg.materialize, MaterializeConfig::default());
393 Ok(())
394 }
395
396 #[test]
397 fn parse_memtable_without_max_rows_defaults_cutoff() -> Result<()> {
398 let yml = r#"
399name: t
400version: "1"
401models_dir: models
402target_path: lake/gold
403materialize:
404 ref_strategy: memtable
405"#;
406 let cfg: RbtProjectConfig = serde_yaml::from_str(yml)?;
407 assert_eq!(cfg.materialize.ref_strategy, RefStrategy::Memtable);
408 assert_eq!(cfg.materialize.memtable_max_rows, DEFAULT_MEMTABLE_MAX_ROWS);
409 Ok(())
410 }
411}