Skip to main content

rbt/core/
project.rs

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
10/// Default MemTable row cutoff when `ref_strategy: memtable` and max rows omitted.
11pub const DEFAULT_MEMTABLE_MAX_ROWS: usize = 50_000;
12
13/// How completed models are exposed to downstream `{{ ref() }}` in the same run.
14///
15/// Default is lake-as-truth **parquet / file re-read** (no long-lived MemTable).
16/// Opt into MemTable via `materialize.ref_strategy: memtable` in `rbt_project.yml`.
17#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
18#[serde(rename_all = "snake_case")]
19pub enum RefStrategy {
20    /// Always re-read the written lake file for `ref()` (default).
21    #[default]
22    #[serde(alias = "parquet_reread", alias = "lake", alias = "file")]
23    Parquet,
24    /// Keep an in-memory `MemTable` when `row_count < memtable_max_rows`; else re-read file.
25    #[serde(alias = "mem_table", alias = "memory", alias = "arc")]
26    Memtable,
27}
28
29/// Chosen backend after applying strategy + row cutoff.
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
31pub enum RefBackend {
32    /// DataFusion `MemTable` holding Arrow batches (Arc-shared).
33    MemTable,
34    /// Re-read model output from the lake path (`register_parquet` / json / csv).
35    LakeFile,
36}
37
38/// Optional materialization / `ref()` registration policy (`materialize:` in yml).
39///
40/// All fields are optional; omitting the whole block keeps lake-as-truth Parquet re-read.
41#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
42pub struct MaterializeConfig {
43    /// `parquet` (default) | `memtable`
44    #[serde(default)]
45    pub ref_strategy: RefStrategy,
46    /// Used only when `ref_strategy: memtable`. Defaults to [`DEFAULT_MEMTABLE_MAX_ROWS`].
47    #[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    /// Decide MemTable vs lake file for a model that produced `row_count` rows.
66    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/// Layer-specific target storage & path configuration.
76#[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/// Project-wide `rbt_project.yml` configuration schema.
84#[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    /// Optional; defaults to lake-as-truth Parquet re-read for `ref()`.
93    #[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    /// Loads `rbt_project.yml` from project directory or returns default configuration.
138    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    /// Resolves destination output directory path for a model based on its layer configuration.
156    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    /// Directory target for Iceberg-style table roots (no file extension).
179    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    /// Discovers all `.sql` models under `models/` directory, resolves layer target paths, and constructs `ModelDag`.
199    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                        // Directory table root: lake/.../model_name/
255                        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}