Skip to main content

rbt/engine/
mod.rs

1//! `rbt::engine`: Apache DataFusion query engine integration, bronze registration, and DAG execution.
2
3pub mod bronze;
4
5use crate::core::dag::{ModelDag, OutputFormat};
6use crate::core::project::{MaterializeConfig, RbtProjectConfig, RefBackend};
7use crate::materializer::{sibling_iceberg_dir, MultiFormatWriter};
8use crate::testing::{assertions_from_model_tests, RecordBatchValidator};
9use anyhow::{bail, Context, Result};
10use arrow::record_batch::RecordBatch;
11use datafusion::datasource::MemTable;
12use datafusion::execution::context::SessionContext;
13use datafusion::physical_plan::SendableRecordBatchStream;
14use datafusion::prelude::{CsvReadOptions, JsonReadOptions, ParquetReadOptions};
15use iceberg::Catalog;
16use iceberg_datafusion::IcebergCatalogProvider;
17use std::collections::HashSet;
18use std::path::{Path, PathBuf};
19use std::sync::Arc;
20
21pub use bronze::{
22    register_bronze_for_model, register_bronze_sources_for_dag, BronzeRegistrationMode,
23    BronzeSourceMeta, BronzeTableProvider,
24};
25
26/// Execution metric summary for a executed model DAG.
27#[derive(Debug, Clone)]
28pub struct DagExecutionSummary {
29    pub models_executed: usize,
30    pub total_rows_produced: usize,
31    pub bronze_sources_registered: usize,
32}
33
34/// Fluent Builder for configuring and launching `TransformationEngine` instances.
35#[derive(Default)]
36pub struct RbtEngineBuilder {
37    catalogs: Vec<(String, Arc<dyn Catalog>)>,
38}
39
40impl RbtEngineBuilder {
41    pub fn new() -> Self {
42        Self::default()
43    }
44
45    pub fn with_catalog(mut self, name: impl Into<String>, catalog: Arc<dyn Catalog>) -> Self {
46        self.catalogs.push((name.into(), catalog));
47        self
48    }
49
50    pub async fn build(self) -> Result<TransformationEngine> {
51        let engine = TransformationEngine::new();
52        for (name, cat) in self.catalogs {
53            engine.register_iceberg_catalog(&name, cat).await?;
54        }
55        Ok(engine)
56    }
57}
58
59pub struct TransformationEngine {
60    pub ctx: SessionContext,
61}
62
63impl Default for TransformationEngine {
64    fn default() -> Self {
65        Self::new()
66    }
67}
68
69impl TransformationEngine {
70    pub fn new() -> Self {
71        Self {
72            ctx: SessionContext::new(),
73        }
74    }
75
76    /// Registers an Apache Iceberg catalog directly into the DataFusion query context.
77    pub async fn register_iceberg_catalog(
78        &self,
79        catalog_name: &str,
80        catalog: Arc<dyn Catalog>,
81    ) -> Result<()> {
82        tracing::info!(
83            "Registering Iceberg catalog '{}' into DataFusion SessionContext",
84            catalog_name
85        );
86        let provider = IcebergCatalogProvider::try_new(catalog).await?;
87        self.ctx.register_catalog(catalog_name, Arc::new(provider));
88        Ok(())
89    }
90
91    /// Executes a SQL transform query against registered tables.
92    pub async fn execute_sql(&self, sql: &str) -> Result<SendableRecordBatchStream> {
93        tracing::info!(
94            "Executing SQL transform via Apache DataFusion engine: {}",
95            sql
96        );
97        let df = self.ctx.sql(sql).await?;
98        let stream = df.execute_stream().await?;
99        Ok(stream)
100    }
101
102    /// Executes a full pipeline DAG tier by tier.
103    ///
104    /// Loads `materialize:` policy from `rbt_project.yml` when present (defaults to
105    /// lake-as-truth Parquet re-read for `ref()`).
106    ///
107    /// Before any model SQL runs, bronze sources declared in staging frontmatter are
108    /// registered via [`register_bronze_sources_for_dag`].
109    pub async fn execute_dag(
110        &self,
111        dag: &ModelDag,
112        project_dir: impl AsRef<Path>,
113        output_dir: impl AsRef<Path>,
114    ) -> Result<DagExecutionSummary> {
115        let project_dir = project_dir.as_ref();
116        let materialize = RbtProjectConfig::load(project_dir)
117            .map(|c| c.materialize)
118            .unwrap_or_default();
119        self.execute_dag_with_materialize(dag, project_dir, output_dir, &materialize)
120            .await
121    }
122
123    /// Like [`execute_dag`] but with an explicit [`MaterializeConfig`] (tests / library).
124    pub async fn execute_dag_with_materialize(
125        &self,
126        dag: &ModelDag,
127        project_dir: impl AsRef<Path>,
128        output_dir: impl AsRef<Path>,
129        materialize: &MaterializeConfig,
130    ) -> Result<DagExecutionSummary> {
131        let project_dir = project_dir.as_ref();
132        let output_base = output_dir.as_ref();
133        tokio::fs::create_dir_all(output_base).await?;
134
135        let mut registered = HashSet::new();
136        let bronze_sources_registered =
137            register_bronze_sources_for_dag(&self.ctx, dag, project_dir, &mut registered)
138                .await
139                .context("frontmatter-driven bronze registration failed")?;
140
141        let tiers = dag.execution_tiers()?;
142        let mut models_executed = 0;
143        let mut total_rows_produced = 0;
144
145        for (tier_idx, tier) in tiers.iter().enumerate() {
146            tracing::info!(
147                "Executing DAG Tier {} with {} parallel models",
148                tier_idx,
149                tier.len()
150            );
151
152            for model in tier {
153                tracing::info!("Executing model '{}'...", model.name);
154
155                // Late-bind: if this model carries frontmatter not registered yet
156                register_bronze_for_model(&self.ctx, model, project_dir, &mut registered).await?;
157
158                let df = self.ctx.sql(&model.compiled_sql).await.with_context(|| {
159                    format!(
160                        "SQL execution failed for model '{}' (compiled: {})",
161                        model.name, model.compiled_sql
162                    )
163                })?;
164                let batches = df
165                    .collect()
166                    .await
167                    .with_context(|| format!("collect failed for model '{}'", model.name))?;
168                let row_count: usize = batches.iter().map(|b| b.num_rows()).sum();
169
170                let dest_path = model
171                    .output_path
172                    .as_ref()
173                    .map(PathBuf::from)
174                    .unwrap_or_else(|| match model.output_format {
175                        OutputFormat::Iceberg => output_base.join(&model.name),
176                        OutputFormat::Jsonl => output_base.join(format!("{}.jsonl", model.name)),
177                        OutputFormat::Csv => output_base.join(format!("{}.csv", model.name)),
178                        _ => output_base.join(format!("{}.parquet", model.name)),
179                    });
180
181                if let Some(parent) = dest_path.parent() {
182                    std::fs::create_dir_all(parent)?;
183                }
184
185                MultiFormatWriter::write_batches(&batches, &model.output_format, &dest_path)?;
186
187                // Frontmatter-declared tests (staging grain / not_null / unique_key).
188                if let Some(fm) = model.frontmatter.as_ref() {
189                    if let Some(tests) = fm.tests.as_ref() {
190                        if !tests.is_empty() {
191                            let unique = tests
192                                .unique
193                                .clone()
194                                .or_else(|| fm.unique_key.clone())
195                                .or_else(|| fm.grain.clone());
196                            let assertions = assertions_from_model_tests(
197                                tests.not_null.as_deref(),
198                                unique.as_deref(),
199                                tests.accepted_values.as_ref(),
200                            );
201                            if !assertions.is_empty() {
202                                let result =
203                                    RecordBatchValidator::validate_batches(&batches, &assertions);
204                                if result.failed_assertions > 0 {
205                                    let msg = format!(
206                                        "model '{}' failed {} test(s): {}",
207                                        model.name,
208                                        result.failed_assertions,
209                                        result.errors.join("; ")
210                                    );
211                                    if tests.should_fail_on_error() {
212                                        bail!(msg);
213                                    }
214                                    tracing::warn!("{}", msg);
215                                } else {
216                                    tracing::info!(
217                                        "model '{}': {} assertion(s) passed ({} rows)",
218                                        model.name,
219                                        result.passed_assertions,
220                                        result.total_rows
221                                    );
222                                }
223                            }
224                        }
225                    } else if let Some(uk) = fm
226                        .unique_key
227                        .as_ref()
228                        .or(fm.grain.as_ref())
229                        .filter(|v| !v.is_empty())
230                    {
231                        // Implicit unique_key/grain check when no tests: block declared
232                        let assertions =
233                            assertions_from_model_tests(None, Some(uk.as_slice()), None);
234                        let result = RecordBatchValidator::validate_batches(&batches, &assertions);
235                        if result.failed_assertions > 0 {
236                            bail!(
237                                "model '{}' grain/unique_key violated: {}",
238                                model.name,
239                                result.errors.join("; ")
240                            );
241                        }
242                    }
243                }
244
245                // Expose model for downstream {{ ref() }} per project materialize policy.
246                if !batches.is_empty() {
247                    let backend = materialize.choose_ref_backend(row_count);
248                    register_model_for_ref(
249                        &self.ctx,
250                        &model.name,
251                        &model.output_format,
252                        &dest_path,
253                        &batches,
254                        backend,
255                    )
256                    .await
257                    .with_context(|| {
258                        format!(
259                            "register model '{}' for ref() (backend={:?}, rows={})",
260                            model.name, backend, row_count
261                        )
262                    })?;
263                    tracing::debug!(
264                        model = %model.name,
265                        rows = row_count,
266                        ?backend,
267                        strategy = ?materialize.ref_strategy,
268                        "registered model for ref()"
269                    );
270                }
271
272                models_executed += 1;
273                total_rows_produced += row_count;
274            }
275        }
276
277        Ok(DagExecutionSummary {
278            models_executed,
279            total_rows_produced,
280            bronze_sources_registered,
281        })
282    }
283}
284
285/// Path used to re-read a model from the lake after materialize.
286fn lake_read_path(format: &OutputFormat, dest_path: &Path) -> PathBuf {
287    match format {
288        OutputFormat::Iceberg => dest_path.join("data/part-00000.parquet"),
289        OutputFormat::ParquetAndIceberg => {
290            // Flat parquet is the primary dual-write artifact for ref().
291            if dest_path.extension().and_then(|e| e.to_str()) == Some("parquet") {
292                dest_path.to_path_buf()
293            } else {
294                dest_path.with_extension("parquet")
295            }
296        }
297        _ => dest_path.to_path_buf(),
298    }
299}
300
301/// Register a completed model so later SQL `ref('name')` resolves.
302async fn register_model_for_ref(
303    ctx: &SessionContext,
304    name: &str,
305    format: &OutputFormat,
306    dest_path: &Path,
307    batches: &[RecordBatch],
308    backend: RefBackend,
309) -> Result<()> {
310    let _ = ctx.deregister_table(name);
311
312    match backend {
313        RefBackend::MemTable => {
314            let schema = batches[0].schema();
315            let mem_table = MemTable::try_new(schema, vec![batches.to_vec()])
316                .map_err(|e| anyhow::anyhow!("MemTable::try_new: {e}"))?;
317            ctx.register_table(name, Arc::new(mem_table))
318                .map_err(|e| anyhow::anyhow!("register_table MemTable: {e}"))?;
319        }
320        RefBackend::LakeFile => match format {
321            OutputFormat::Parquet
322            | OutputFormat::ZeroCopyClone
323            | OutputFormat::Iceberg
324            | OutputFormat::ParquetAndIceberg => {
325                let mut path = lake_read_path(format, dest_path);
326                if !path.exists() && matches!(format, OutputFormat::ParquetAndIceberg) {
327                    let alt = sibling_iceberg_dir(dest_path).join("data/part-00000.parquet");
328                    if alt.exists() {
329                        path = alt;
330                    }
331                }
332                if !path.exists() {
333                    bail!(
334                        "lake file missing for ref('{}'): expected {}",
335                        name,
336                        path.display()
337                    );
338                }
339                ctx.register_parquet(
340                    name,
341                    path.to_str().unwrap_or_default(),
342                    ParquetReadOptions::default(),
343                )
344                .await
345                .map_err(|e| anyhow::anyhow!("register_parquet {}: {e}", path.display()))?;
346            }
347            OutputFormat::Jsonl => {
348                let p = dest_path.to_str().unwrap_or_default();
349                let opts = JsonReadOptions::default()
350                    .file_extension(".jsonl")
351                    .newline_delimited(true);
352                if let Err(e) = ctx.register_json(name, p, opts).await {
353                    tracing::debug!("jsonl register failed ({e}); retry default");
354                    ctx.register_json(name, p, JsonReadOptions::default())
355                        .await
356                        .map_err(|e| anyhow::anyhow!("register_json: {e}"))?;
357                }
358            }
359            OutputFormat::Csv => {
360                ctx.register_csv(
361                    name,
362                    dest_path.to_str().unwrap_or_default(),
363                    CsvReadOptions::default(),
364                )
365                .await
366                .map_err(|e| anyhow::anyhow!("register_csv: {e}"))?;
367            }
368        },
369    }
370    Ok(())
371}
372
373#[cfg(test)]
374mod tests {
375    use super::*;
376    use crate::core::dag::{Materialization, ModelDag, OutputFormat};
377
378    #[tokio::test]
379    async fn test_engine_initialization() -> Result<()> {
380        let engine = TransformationEngine::new();
381        let df = engine.ctx.sql("SELECT 1 AS col").await?;
382        let batches = df.collect().await?;
383        assert_eq!(batches.len(), 1);
384        assert_eq!(batches[0].num_rows(), 1);
385        Ok(())
386    }
387
388    #[tokio::test]
389    async fn test_dag_execution_multi_format() -> Result<()> {
390        let temp_dir = tempfile::tempdir()?;
391        let engine = TransformationEngine::new();
392
393        let mut dag = ModelDag::new();
394        dag.add_model_with_format(
395            "users",
396            "SELECT 1 AS id, 'Alice' AS name",
397            Materialization::Table,
398            OutputFormat::Jsonl,
399            None,
400            "",
401        )?;
402        dag.add_model_with_format(
403            "active_users",
404            "SELECT * FROM {{ ref('users') }} WHERE id = 1",
405            Materialization::Table,
406            OutputFormat::Parquet,
407            None,
408            "",
409        )?;
410        dag.build_graph()?;
411
412        let summary = engine
413            .execute_dag(&dag, temp_dir.path(), temp_dir.path())
414            .await?;
415        assert_eq!(summary.models_executed, 2);
416        assert_eq!(summary.total_rows_produced, 2);
417        assert!(temp_dir.path().join("users.jsonl").exists());
418        assert!(temp_dir.path().join("active_users.parquet").exists());
419        Ok(())
420    }
421
422    #[tokio::test]
423    async fn test_frontmatter_bronze_end_to_end() -> Result<()> {
424        let temp = tempfile::tempdir()?;
425        let bronze_dir = temp.path().join("lake/bronze");
426        std::fs::create_dir_all(&bronze_dir)?;
427        std::fs::write(
428            bronze_dir.join("raw_stock_trades.jsonl"),
429            r#"{"ticker":"NVDA","timestamp":"2026-07-24T09:30:01Z","price":125.5,"volume":100}
430{"ticker":"AAPL","timestamp":"2026-07-24T09:30:05Z","price":190.0,"volume":50}
431"#,
432        )?;
433
434        let sql = r#"---
435source_format: jsonl
436scan_path: "lake/bronze/raw_stock_trades.jsonl"
437---
438SELECT ticker, price, volume FROM {{ source('bronze', 'raw_stock_trades') }}
439"#;
440
441        let mut dag = ModelDag::new();
442        dag.add_model_with_format(
443            "stg_stock_trades",
444            sql,
445            Materialization::Table,
446            OutputFormat::Parquet,
447            Some(
448                temp.path()
449                    .join("lake/silver/stg_stock_trades.parquet")
450                    .to_string_lossy()
451                    .into(),
452            ),
453            "",
454        )?;
455        dag.build_graph()?;
456
457        let engine = TransformationEngine::new();
458        let summary = engine
459            .execute_dag(&dag, temp.path(), temp.path().join("out"))
460            .await?;
461        assert_eq!(summary.bronze_sources_registered, 1);
462        assert_eq!(summary.models_executed, 1);
463        assert_eq!(summary.total_rows_produced, 2);
464        assert!(temp
465            .path()
466            .join("lake/silver/stg_stock_trades.parquet")
467            .exists());
468        Ok(())
469    }
470
471    #[tokio::test]
472    async fn test_ref_via_parquet_reread_default() -> Result<()> {
473        use crate::core::project::{MaterializeConfig, RefStrategy};
474
475        let temp = tempfile::tempdir()?;
476        let mut dag = ModelDag::new();
477        dag.add_model_with_format(
478            "stg_a",
479            "SELECT 1 AS id, 10 AS v UNION ALL SELECT 2, 20",
480            Materialization::Table,
481            OutputFormat::Parquet,
482            Some(temp.path().join("stg_a.parquet").to_string_lossy().into()),
483            "",
484        )?;
485        dag.add_model_with_format(
486            "tf_b",
487            "SELECT id, v * 2 AS v2 FROM {{ ref('stg_a') }}",
488            Materialization::Table,
489            OutputFormat::Parquet,
490            Some(temp.path().join("tf_b.parquet").to_string_lossy().into()),
491            "",
492        )?;
493        dag.build_graph()?;
494
495        let mat = MaterializeConfig {
496            ref_strategy: RefStrategy::Parquet,
497            memtable_max_rows: 50_000,
498        };
499        let engine = TransformationEngine::new();
500        let summary = engine
501            .execute_dag_with_materialize(&dag, temp.path(), temp.path(), &mat)
502            .await?;
503        assert_eq!(summary.models_executed, 2);
504        assert_eq!(summary.total_rows_produced, 4);
505        assert!(temp.path().join("tf_b.parquet").exists());
506        Ok(())
507    }
508
509    #[tokio::test]
510    async fn test_ref_via_memtable_when_configured() -> Result<()> {
511        use crate::core::project::{MaterializeConfig, RefStrategy};
512
513        let temp = tempfile::tempdir()?;
514        let mut dag = ModelDag::new();
515        dag.add_model_with_format(
516            "stg_a",
517            "SELECT 1 AS id UNION ALL SELECT 2",
518            Materialization::Table,
519            OutputFormat::Parquet,
520            Some(temp.path().join("stg_a.parquet").to_string_lossy().into()),
521            "",
522        )?;
523        dag.add_model_with_format(
524            "tf_b",
525            "SELECT count(*) AS c FROM {{ ref('stg_a') }}",
526            Materialization::Table,
527            OutputFormat::Parquet,
528            Some(temp.path().join("tf_b.parquet").to_string_lossy().into()),
529            "",
530        )?;
531        dag.build_graph()?;
532
533        let mat = MaterializeConfig {
534            ref_strategy: RefStrategy::Memtable,
535            memtable_max_rows: 50_000,
536        };
537        let engine = TransformationEngine::new();
538        let summary = engine
539            .execute_dag_with_materialize(&dag, temp.path(), temp.path(), &mat)
540            .await?;
541        assert_eq!(summary.models_executed, 2);
542        assert!(temp.path().join("tf_b.parquet").exists());
543        Ok(())
544    }
545
546    #[tokio::test]
547    async fn test_memtable_falls_back_to_lake_above_cutoff() -> Result<()> {
548        use crate::core::project::{MaterializeConfig, RefStrategy};
549
550        // Cutoff 1 → 2-row model must use lake re-read.
551        let temp = tempfile::tempdir()?;
552        let mut dag = ModelDag::new();
553        dag.add_model_with_format(
554            "stg_a",
555            "SELECT 1 AS id UNION ALL SELECT 2",
556            Materialization::Table,
557            OutputFormat::Parquet,
558            Some(temp.path().join("stg_a.parquet").to_string_lossy().into()),
559            "",
560        )?;
561        dag.add_model_with_format(
562            "tf_b",
563            "SELECT * FROM {{ ref('stg_a') }}",
564            Materialization::Table,
565            OutputFormat::Parquet,
566            Some(temp.path().join("tf_b.parquet").to_string_lossy().into()),
567            "",
568        )?;
569        dag.build_graph()?;
570
571        let mat = MaterializeConfig {
572            ref_strategy: RefStrategy::Memtable,
573            memtable_max_rows: 1,
574        };
575        let engine = TransformationEngine::new();
576        let summary = engine
577            .execute_dag_with_materialize(&dag, temp.path(), temp.path(), &mat)
578            .await?;
579        assert_eq!(summary.models_executed, 2);
580        assert_eq!(summary.total_rows_produced, 4);
581        Ok(())
582    }
583}