1pub 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#[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#[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 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 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 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 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 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 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 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}