1pub 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#[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#[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 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 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 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 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 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 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 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 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
285fn 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 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
301async 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 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}