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::materializer::MultiFormatWriter;
7use crate::testing::{assertions_from_model_tests, RecordBatchValidator};
8use anyhow::{Context, Result};
9use datafusion::datasource::MemTable;
10use datafusion::execution::context::SessionContext;
11use datafusion::physical_plan::SendableRecordBatchStream;
12use iceberg::Catalog;
13use iceberg_datafusion::IcebergCatalogProvider;
14use std::collections::HashSet;
15use std::path::{Path, PathBuf};
16use std::sync::Arc;
17
18pub use bronze::{
19    register_bronze_for_model, register_bronze_sources_for_dag, BronzeRegistrationMode,
20    BronzeSourceMeta, BronzeTableProvider,
21};
22
23/// Execution metric summary for a executed model DAG.
24#[derive(Debug, Clone)]
25pub struct DagExecutionSummary {
26    pub models_executed: usize,
27    pub total_rows_produced: usize,
28    pub bronze_sources_registered: usize,
29}
30
31/// Fluent Builder for configuring and launching `TransformationEngine` instances.
32#[derive(Default)]
33pub struct RbtEngineBuilder {
34    catalogs: Vec<(String, Arc<dyn Catalog>)>,
35}
36
37impl RbtEngineBuilder {
38    pub fn new() -> Self {
39        Self::default()
40    }
41
42    pub fn with_catalog(mut self, name: impl Into<String>, catalog: Arc<dyn Catalog>) -> Self {
43        self.catalogs.push((name.into(), catalog));
44        self
45    }
46
47    pub async fn build(self) -> Result<TransformationEngine> {
48        let engine = TransformationEngine::new();
49        for (name, cat) in self.catalogs {
50            engine.register_iceberg_catalog(&name, cat).await?;
51        }
52        Ok(engine)
53    }
54}
55
56pub struct TransformationEngine {
57    pub ctx: SessionContext,
58}
59
60impl Default for TransformationEngine {
61    fn default() -> Self {
62        Self::new()
63    }
64}
65
66impl TransformationEngine {
67    pub fn new() -> Self {
68        Self {
69            ctx: SessionContext::new(),
70        }
71    }
72
73    /// Registers an Apache Iceberg catalog directly into the DataFusion query context.
74    pub async fn register_iceberg_catalog(
75        &self,
76        catalog_name: &str,
77        catalog: Arc<dyn Catalog>,
78    ) -> Result<()> {
79        tracing::info!(
80            "Registering Iceberg catalog '{}' into DataFusion SessionContext",
81            catalog_name
82        );
83        let provider = IcebergCatalogProvider::try_new(catalog).await?;
84        self.ctx.register_catalog(catalog_name, Arc::new(provider));
85        Ok(())
86    }
87
88    /// Executes a SQL transform query against registered tables.
89    pub async fn execute_sql(&self, sql: &str) -> Result<SendableRecordBatchStream> {
90        tracing::info!(
91            "Executing SQL transform via Apache DataFusion engine: {}",
92            sql
93        );
94        let df = self.ctx.sql(sql).await?;
95        let stream = df.execute_stream().await?;
96        Ok(stream)
97    }
98
99    /// Executes a full pipeline DAG tier by tier.
100    ///
101    /// Before any model SQL runs, bronze sources declared in staging frontmatter are
102    /// registered via [`register_bronze_sources_for_dag`].
103    pub async fn execute_dag(
104        &self,
105        dag: &ModelDag,
106        project_dir: impl AsRef<Path>,
107        output_dir: impl AsRef<Path>,
108    ) -> Result<DagExecutionSummary> {
109        let project_dir = project_dir.as_ref();
110        let output_base = output_dir.as_ref();
111        tokio::fs::create_dir_all(output_base).await?;
112
113        let mut registered = HashSet::new();
114        let bronze_sources_registered =
115            register_bronze_sources_for_dag(&self.ctx, dag, project_dir, &mut registered)
116                .await
117                .context("frontmatter-driven bronze registration failed")?;
118
119        let tiers = dag.execution_tiers()?;
120        let mut models_executed = 0;
121        let mut total_rows_produced = 0;
122
123        for (tier_idx, tier) in tiers.iter().enumerate() {
124            tracing::info!(
125                "Executing DAG Tier {} with {} parallel models",
126                tier_idx,
127                tier.len()
128            );
129
130            for model in tier {
131                tracing::info!("Executing model '{}'...", model.name);
132
133                // Late-bind: if this model carries frontmatter not registered yet
134                register_bronze_for_model(&self.ctx, model, project_dir, &mut registered).await?;
135
136                let df = self.ctx.sql(&model.compiled_sql).await.with_context(|| {
137                    format!(
138                        "SQL execution failed for model '{}' (compiled: {})",
139                        model.name, model.compiled_sql
140                    )
141                })?;
142                let batches = df
143                    .collect()
144                    .await
145                    .with_context(|| format!("collect failed for model '{}'", model.name))?;
146                let row_count: usize = batches.iter().map(|b| b.num_rows()).sum();
147
148                let dest_path = model
149                    .output_path
150                    .as_ref()
151                    .map(PathBuf::from)
152                    .unwrap_or_else(|| match model.output_format {
153                        OutputFormat::Iceberg => output_base.join(&model.name),
154                        OutputFormat::Jsonl => output_base.join(format!("{}.jsonl", model.name)),
155                        OutputFormat::Csv => output_base.join(format!("{}.csv", model.name)),
156                        _ => output_base.join(format!("{}.parquet", model.name)),
157                    });
158
159                if let Some(parent) = dest_path.parent() {
160                    std::fs::create_dir_all(parent)?;
161                }
162
163                MultiFormatWriter::write_batches(&batches, &model.output_format, &dest_path)?;
164
165                // Frontmatter-declared tests (staging grain / not_null / unique_key).
166                if let Some(fm) = model.frontmatter.as_ref() {
167                    if let Some(tests) = fm.tests.as_ref() {
168                        if !tests.is_empty() {
169                            let unique = tests
170                                .unique
171                                .clone()
172                                .or_else(|| fm.unique_key.clone())
173                                .or_else(|| fm.grain.clone());
174                            let assertions = assertions_from_model_tests(
175                                tests.not_null.as_deref(),
176                                unique.as_deref(),
177                                tests.accepted_values.as_ref(),
178                            );
179                            if !assertions.is_empty() {
180                                let result =
181                                    RecordBatchValidator::validate_batches(&batches, &assertions);
182                                if result.failed_assertions > 0 {
183                                    let msg = format!(
184                                        "model '{}' failed {} test(s): {}",
185                                        model.name,
186                                        result.failed_assertions,
187                                        result.errors.join("; ")
188                                    );
189                                    if tests.should_fail_on_error() {
190                                        anyhow::bail!(msg);
191                                    }
192                                    tracing::warn!("{}", msg);
193                                } else {
194                                    tracing::info!(
195                                        "model '{}': {} assertion(s) passed ({} rows)",
196                                        model.name,
197                                        result.passed_assertions,
198                                        result.total_rows
199                                    );
200                                }
201                            }
202                        }
203                    } else if let Some(uk) = fm
204                        .unique_key
205                        .as_ref()
206                        .or(fm.grain.as_ref())
207                        .filter(|v| !v.is_empty())
208                    {
209                        // Implicit unique_key/grain check when no tests: block declared
210                        let assertions =
211                            assertions_from_model_tests(None, Some(uk.as_slice()), None);
212                        let result = RecordBatchValidator::validate_batches(&batches, &assertions);
213                        if result.failed_assertions > 0 {
214                            anyhow::bail!(
215                                "model '{}' grain/unique_key violated: {}",
216                                model.name,
217                                result.errors.join("; ")
218                            );
219                        }
220                    }
221                }
222
223                if !batches.is_empty() {
224                    let schema = batches[0].schema();
225                    let mem_table = MemTable::try_new(schema, vec![batches.clone()])?;
226                    // Allow re-runs in tests
227                    let _ = self.ctx.deregister_table(model.name.as_str());
228                    self.ctx
229                        .register_table(model.name.as_str(), Arc::new(mem_table))?;
230                }
231
232                models_executed += 1;
233                total_rows_produced += row_count;
234            }
235        }
236
237        Ok(DagExecutionSummary {
238            models_executed,
239            total_rows_produced,
240            bronze_sources_registered,
241        })
242    }
243}
244
245#[cfg(test)]
246mod tests {
247    use super::*;
248    use crate::core::dag::{Materialization, ModelDag, OutputFormat};
249
250    #[tokio::test]
251    async fn test_engine_initialization() -> Result<()> {
252        let engine = TransformationEngine::new();
253        let df = engine.ctx.sql("SELECT 1 AS col").await?;
254        let batches = df.collect().await?;
255        assert_eq!(batches.len(), 1);
256        assert_eq!(batches[0].num_rows(), 1);
257        Ok(())
258    }
259
260    #[tokio::test]
261    async fn test_dag_execution_multi_format() -> Result<()> {
262        let temp_dir = tempfile::tempdir()?;
263        let engine = TransformationEngine::new();
264
265        let mut dag = ModelDag::new();
266        dag.add_model_with_format(
267            "users",
268            "SELECT 1 AS id, 'Alice' AS name",
269            Materialization::Table,
270            OutputFormat::Jsonl,
271            None,
272            "",
273        )?;
274        dag.add_model_with_format(
275            "active_users",
276            "SELECT * FROM {{ ref('users') }} WHERE id = 1",
277            Materialization::Table,
278            OutputFormat::Parquet,
279            None,
280            "",
281        )?;
282        dag.build_graph()?;
283
284        let summary = engine
285            .execute_dag(&dag, temp_dir.path(), temp_dir.path())
286            .await?;
287        assert_eq!(summary.models_executed, 2);
288        assert_eq!(summary.total_rows_produced, 2);
289        assert!(temp_dir.path().join("users.jsonl").exists());
290        assert!(temp_dir.path().join("active_users.parquet").exists());
291        Ok(())
292    }
293
294    #[tokio::test]
295    async fn test_frontmatter_bronze_end_to_end() -> Result<()> {
296        let temp = tempfile::tempdir()?;
297        let bronze_dir = temp.path().join("lake/bronze");
298        std::fs::create_dir_all(&bronze_dir)?;
299        std::fs::write(
300            bronze_dir.join("raw_stock_trades.jsonl"),
301            r#"{"ticker":"NVDA","timestamp":"2026-07-24T09:30:01Z","price":125.5,"volume":100}
302{"ticker":"AAPL","timestamp":"2026-07-24T09:30:05Z","price":190.0,"volume":50}
303"#,
304        )?;
305
306        let sql = r#"---
307source_format: jsonl
308scan_path: "lake/bronze/raw_stock_trades.jsonl"
309---
310SELECT ticker, price, volume FROM {{ source('bronze', 'raw_stock_trades') }}
311"#;
312
313        let mut dag = ModelDag::new();
314        dag.add_model_with_format(
315            "stg_stock_trades",
316            sql,
317            Materialization::Table,
318            OutputFormat::Parquet,
319            Some(
320                temp.path()
321                    .join("lake/silver/stg_stock_trades.parquet")
322                    .to_string_lossy()
323                    .into(),
324            ),
325            "",
326        )?;
327        dag.build_graph()?;
328
329        let engine = TransformationEngine::new();
330        let summary = engine
331            .execute_dag(&dag, temp.path(), temp.path().join("out"))
332            .await?;
333        assert_eq!(summary.bronze_sources_registered, 1);
334        assert_eq!(summary.models_executed, 1);
335        assert_eq!(summary.total_rows_produced, 2);
336        assert!(temp
337            .path()
338            .join("lake/silver/stg_stock_trades.parquet")
339            .exists());
340        Ok(())
341    }
342}