#![cfg(feature = "native")]
use knut_thund::backend::{ExecBackend, native::NativeBackend};
use knut_thund::ir::{Dataset, Flow, OutputType, Pipeline};
use datafusion::arrow::array::{Array, Int64Array, RecordBatch, StringArray, StringViewArray};
use std::io::Write;
fn str_col(b: &RecordBatch, name: &str) -> Vec<String> {
let col = b
.column_by_name(name)
.unwrap_or_else(|| panic!("column `{name}` present"))
.as_any();
if let Some(a) = col.downcast_ref::<StringArray>() {
(0..a.len()).map(|i| a.value(i).to_string()).collect()
} else if let Some(a) = col.downcast_ref::<StringViewArray>() {
(0..a.len()).map(|i| a.value(i).to_string()).collect()
} else {
panic!("column `{name}` is neither Utf8 nor Utf8View");
}
}
fn rollup_pipeline() -> Pipeline {
Pipeline::new("orders_rollup")
.with_dataset(Dataset::new("by_customer", OutputType::MaterializedView))
.with_flow(
Flow::batch("agg", "by_customer", ["orders"])
.with_query("SELECT customer, SUM(amount) AS total FROM orders GROUP BY customer"),
)
}
fn sorted_totals(batches: &[RecordBatch]) -> Vec<(String, i64)> {
let mut out: Vec<(String, i64)> = Vec::new();
for b in batches {
let col = b.column(0).as_any();
let cust: Vec<String> = if let Some(a) = col.downcast_ref::<StringArray>() {
(0..a.len()).map(|i| a.value(i).to_string()).collect()
} else if let Some(a) = col.downcast_ref::<StringViewArray>() {
(0..a.len()).map(|i| a.value(i).to_string()).collect()
} else {
panic!("customer column is neither Utf8 nor Utf8View");
};
let total = b
.column(1)
.as_any()
.downcast_ref::<Int64Array>()
.expect("total is Int64");
for (i, c) in cust.into_iter().enumerate() {
out.push((c, total.value(i)));
}
}
out.sort();
out
}
#[test]
fn csv_file_source_batch_rollup() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("orders.csv");
{
let mut f = std::fs::File::create(&path).expect("create csv");
writeln!(f, "customer,amount").unwrap();
writeln!(f, "alice,10").unwrap();
writeln!(f, "bob,5").unwrap();
writeln!(f, "alice,20").unwrap();
writeln!(f, "bob,7").unwrap();
writeln!(f, "alice,3").unwrap();
}
let run = NativeBackend::new()
.with_file_input("orders", path.to_str().unwrap(), "csv")
.run(&rollup_pipeline())
.expect("native run over the on-disk CSV completes");
let totals = sorted_totals(run.output("by_customer").expect("materialised output"));
assert_eq!(
totals,
vec![("alice".to_string(), 33), ("bob".to_string(), 12)],
"the batch flow scanned the local CSV and rolled it up per customer"
);
knut_thund::functional_status(
"knut-thund/native_file_source",
"csv_batch",
totals == vec![("alice".to_string(), 33), ("bob".to_string(), 12)],
"orders.csv -> registered `orders` -> batch rollup",
);
}
#[test]
fn parquet_file_source_batch_rollup() {
use datafusion::arrow::array::{Int64Array as I64, StringArray as Utf8};
use datafusion::arrow::datatypes::{DataType, Field, Schema};
use datafusion::parquet::arrow::ArrowWriter;
use std::sync::Arc;
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("orders.parquet");
let schema = Arc::new(Schema::new(vec![
Field::new("customer", DataType::Utf8, false),
Field::new("amount", DataType::Int64, false),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Utf8::from(vec!["alice", "bob", "alice", "bob", "alice"])),
Arc::new(I64::from(vec![10, 5, 20, 7, 3])),
],
)
.expect("build orders batch");
{
let file = std::fs::File::create(&path).expect("create parquet");
let mut w = ArrowWriter::try_new(file, schema, None).expect("arrow writer");
w.write(&batch).expect("write parquet batch");
w.close().expect("close parquet");
}
let run = NativeBackend::new()
.with_file_input("orders", path.to_str().unwrap(), "parquet")
.run(&rollup_pipeline())
.expect("native run over the on-disk parquet completes");
let totals = sorted_totals(run.output("by_customer").expect("materialised output"));
assert_eq!(
totals,
vec![("alice".to_string(), 33), ("bob".to_string(), 12)],
"the batch flow scanned the local parquet and rolled it up per customer"
);
knut_thund::functional_status(
"knut-thund/native_file_source",
"parquet_batch",
totals == vec![("alice".to_string(), 33), ("bob".to_string(), 12)],
"orders.parquet -> registered `orders` -> batch rollup",
);
}
fn sorted_regional(batches: &[RecordBatch]) -> Vec<(String, String, i64)> {
let mut out: Vec<(String, String, i64)> = Vec::new();
for b in batches {
let region = str_col(b, "region");
let customer = str_col(b, "customer");
let total = b
.column_by_name("total")
.expect("total column present")
.as_any()
.downcast_ref::<Int64Array>()
.expect("total is Int64");
for i in 0..b.num_rows() {
out.push((region[i].clone(), customer[i].clone(), total.value(i)));
}
}
out.sort();
out
}
fn write_regional_orders_csv(path: &std::path::Path) {
let mut f = std::fs::File::create(path).expect("create orders.csv");
writeln!(f, "region,customer,amount").unwrap();
writeln!(f, "US,alice,10").unwrap();
writeln!(f, "EU,bob,5").unwrap();
writeln!(f, "US,alice,20").unwrap();
writeln!(f, "EU,bob,7").unwrap();
writeln!(f, "US,carol,3").unwrap();
}
#[test]
fn partitioned_parquet_source_round_trip() {
let dir = tempfile::tempdir().expect("tempdir");
let src = dir.path().join("orders.csv");
write_regional_orders_csv(&src);
let root = dir.path().join("warehouse").join("by_region");
let write = NativeBackend::new()
.with_file_input("orders", src.to_str().unwrap(), "csv")
.with_file_output("by_customer", root.to_str().unwrap(), "parquet")
.run(
&Pipeline::new("orders_by_region")
.with_dataset(
Dataset::new("by_customer", OutputType::Sink).with_partition_cols(["region"]),
)
.with_flow(Flow::batch("agg", "by_customer", ["orders"]).with_query(
"SELECT region, customer, SUM(amount) AS total \
FROM orders GROUP BY region, customer",
)),
)
.expect("partitioned sink write completes");
assert_eq!(
write
.sink("by_customer")
.expect("sink report")
.partition_cols,
vec!["region".to_string()],
"the sink partitioned by region"
);
assert!(root.is_dir(), "the sink produced a Hive directory root");
let back = NativeBackend::new()
.with_partitioned_file_input("t", root.to_str().unwrap(), "parquet", ["region"])
.run(
&Pipeline::new("read_partitioned")
.with_dataset(Dataset::new("t_out", OutputType::MaterializedView))
.with_flow(
Flow::batch("scan", "t_out", ["t"])
.with_query("SELECT region, customer, total FROM t"),
),
)
.expect("partitioned read-back completes (region recovered from the path)");
let got = sorted_regional(back.output("t_out").expect("read-back output"));
let expected = vec![
("EU".to_string(), "bob".to_string(), 12),
("US".to_string(), "alice".to_string(), 30),
("US".to_string(), "carol".to_string(), 3),
];
assert_eq!(
got, expected,
"the partition column `region` is recovered from the path and the rows survive the round trip"
);
knut_thund::functional_status(
"knut-thund/native_file_source",
"partitioned_parquet",
got == expected,
"region=US/ , region=EU/ (Hive parquet dir) -> partitioned source -> region recovered from path",
);
}
#[test]
fn partitioned_csv_source_round_trip() {
let dir = tempfile::tempdir().expect("tempdir");
let src = dir.path().join("orders.csv");
write_regional_orders_csv(&src);
let root = dir.path().join("csv_warehouse");
NativeBackend::new()
.with_file_input("orders", src.to_str().unwrap(), "csv")
.with_file_output("by_customer", root.to_str().unwrap(), "csv")
.run(
&Pipeline::new("orders_by_region_csv")
.with_dataset(
Dataset::new("by_customer", OutputType::Sink).with_partition_cols(["region"]),
)
.with_flow(Flow::batch("agg", "by_customer", ["orders"]).with_query(
"SELECT region, customer, SUM(amount) AS total \
FROM orders GROUP BY region, customer",
)),
)
.expect("partitioned csv sink write completes");
let back = NativeBackend::new()
.with_partitioned_file_input("t", root.to_str().unwrap(), "csv", ["region"])
.run(
&Pipeline::new("read_partitioned_csv")
.with_dataset(Dataset::new("t_out", OutputType::MaterializedView))
.with_flow(
Flow::batch("scan", "t_out", ["t"])
.with_query("SELECT region, customer, total FROM t"),
),
)
.expect("partitioned csv read-back completes");
let got = sorted_regional(back.output("t_out").expect("read-back output"));
let expected = vec![
("EU".to_string(), "bob".to_string(), 12),
("US".to_string(), "alice".to_string(), 30),
("US".to_string(), "carol".to_string(), 3),
];
assert_eq!(
got, expected,
"partitioned CSV round-trips with region from the path"
);
knut_thund::functional_status(
"knut-thund/native_file_source",
"partitioned_csv",
got == expected,
"region=US/ , region=EU/ (Hive csv dir) -> partitioned source -> region recovered from path",
);
}