use datafusion::common::tree_node::Transformed;
use datafusion::common::{Column, DFSchema, Dependency, NullEquality, Result, TableReference};
use datafusion::logical_expr::{
Aggregate, Expr, LogicalPlan, LogicalPlanBuilder, Projection, SubqueryAlias, TableScan,
};
use datafusion::optimizer::{ApplyOrder, OptimizerConfig, OptimizerRule};
use std::collections::HashSet;
use std::sync::Arc;
pub const LATE_MATERIALIZATION_ENV: &str = "KRISHIV_LATE_MATERIALIZATION";
const TOPN_ALIAS: &str = "__krishiv_lm";
const MAX_LATE_MATERIALIZE_FETCH: usize = 10_000;
const MAX_DIM_CHAIN: usize = 4;
pub fn late_materialization_enabled() -> bool {
enabled_from(&std::env::var(LATE_MATERIALIZATION_ENV).unwrap_or_default())
}
fn enabled_from(value: &str) -> bool {
!matches!(
value.trim().to_ascii_lowercase().as_str(),
"0" | "off" | "false" | "no"
)
}
#[derive(Debug, Default)]
pub struct LateMaterializeTopKAggregate {
forced: bool,
}
impl LateMaterializeTopKAggregate {
#[must_use]
pub fn forced() -> Self {
Self { forced: true }
}
}
impl OptimizerRule for LateMaterializeTopKAggregate {
fn name(&self) -> &str {
"late_materialize_topk_aggregate"
}
fn apply_order(&self) -> Option<ApplyOrder> {
Some(ApplyOrder::TopDown)
}
fn rewrite(
&self,
plan: LogicalPlan,
_config: &dyn OptimizerConfig,
) -> Result<Transformed<LogicalPlan>> {
if !self.forced && !late_materialization_enabled() {
return Ok(Transformed::no(plan));
}
let LogicalPlan::Sort(sort) = &plan else {
return Ok(Transformed::no(plan));
};
let Some(fetch) = sort.fetch.filter(|n| *n <= MAX_LATE_MATERIALIZE_FETCH) else {
return Ok(Transformed::no(plan));
};
let mut sort_exprs: Vec<Expr> = sort.expr.iter().map(|s| s.expr.clone()).collect();
let mut node: &LogicalPlan = &sort.input;
let mut projections: Vec<&Projection> = Vec::new();
loop {
match node {
LogicalPlan::Projection(proj) => {
let Some(lowered) = lower_exprs_through(proj, &sort_exprs) else {
return Ok(Transformed::no(plan));
};
sort_exprs = lowered;
projections.push(proj);
node = &proj.input;
}
LogicalPlan::Aggregate(_) => break,
_ => return Ok(Transformed::no(plan)),
}
}
let LogicalPlan::Aggregate(agg) = node else {
return Ok(Transformed::no(plan));
};
let Some(rewritten) = rewrite_aggregate(agg, &sort_exprs, sort, fetch)? else {
return Ok(Transformed::no(plan));
};
let mut rebuilt = rewritten;
for proj in projections.into_iter().rev() {
rebuilt = LogicalPlan::Projection(Projection::try_new(
proj.expr.clone(),
Arc::new(rebuilt),
)?);
}
Ok(Transformed::yes(LogicalPlan::Sort(
datafusion::logical_expr::Sort {
expr: sort.expr.clone(),
input: Arc::new(rebuilt),
fetch: sort.fetch,
},
)))
}
}
fn rewrite_aggregate(
agg: &Aggregate,
sort_exprs: &[Expr],
sort: &datafusion::logical_expr::Sort,
fetch: usize,
) -> Result<Option<LogicalPlan>> {
let mut group_cols = Vec::with_capacity(agg.group_expr.len());
for expr in &agg.group_expr {
let Expr::Column(col) = expr else {
return Ok(None);
};
group_cols.push(col.clone());
}
if group_cols.len() < 2 {
return Ok(None);
}
let Some(facts) = InputFacts::collect(&agg.input) else {
return Ok(None);
};
if facts.tables.is_empty() {
return Ok(None);
}
let mut key = facts.minimal_key(&group_cols);
for expr in sort_exprs {
for col in expr.column_refs() {
if group_cols.contains(col) && !key.contains(col) {
key.push(col.clone());
}
}
}
key.sort_by_key(|col| group_cols.iter().position(|g| g == col).unwrap_or(usize::MAX));
let deferred: Vec<Column> = group_cols
.iter()
.filter(|col| !key.contains(col))
.cloned()
.collect();
if deferred.is_empty() || key.is_empty() {
return Ok(None);
}
let agg_schema = agg.schema.as_ref();
for expr in sort_exprs {
for col in expr.column_refs() {
let is_aggregate_output = agg_schema
.index_of_column(col)
.is_ok_and(|idx| idx >= group_cols.len());
if !key.contains(col) && !is_aggregate_output {
return Ok(None);
}
}
}
let Some(chain) = facts.dim_chain(&key, &deferred) else {
return Ok(None);
};
let key_exprs: Vec<Expr> = key.iter().cloned().map(Expr::Column).collect();
let narrow = LogicalPlan::Aggregate(Aggregate::try_new(
Arc::clone(&agg.input),
key_exprs,
agg.aggr_expr.clone(),
)?);
{
let mut names = HashSet::new();
if !narrow
.schema()
.fields()
.iter()
.all(|field| names.insert(field.name().clone()))
{
return Ok(None);
}
}
let narrow_sort = LogicalPlan::Sort(datafusion::logical_expr::Sort {
expr: sort
.expr
.iter()
.zip(sort_exprs)
.map(|(original, lowered)| datafusion::logical_expr::SortExpr {
expr: lowered.clone(),
asc: original.asc,
nulls_first: original.nulls_first,
})
.collect(),
input: Arc::new(narrow),
fetch: Some(fetch),
});
let topn = LogicalPlan::SubqueryAlias(SubqueryAlias::try_new(
Arc::new(narrow_sort),
TableReference::bare(TOPN_ALIAS),
)?);
let mut joined = topn;
for link in &chain {
let resolves = link
.probe_keys
.iter()
.all(|col| index_of(joined.schema(), col).is_some())
&& link
.key_columns
.iter()
.all(|col| index_of(link.scan.schema(), col).is_some());
if !resolves {
return Ok(None);
}
joined = LogicalPlanBuilder::from(joined)
.join_detailed(
link.scan.clone(),
datafusion::common::JoinType::Inner,
(link.probe_keys.clone(), link.key_columns.clone()),
None,
NullEquality::NullEqualsNothing,
)?
.build()?;
}
let mut exprs = Vec::with_capacity(agg_schema.fields().len());
for (idx, (qualifier, field)) in agg_schema.iter().enumerate() {
let target = Column::new(qualifier.cloned(), field.name());
let source = match group_cols.get(idx).filter(|col| deferred.contains(col)) {
Some(col) => col.clone(),
None => Column::new(Some(TableReference::bare(TOPN_ALIAS)), field.name()),
};
exprs.push(if source == target {
Expr::Column(source)
} else {
Expr::Column(source).alias_qualified(qualifier.cloned(), field.name())
});
}
Ok(Some(LogicalPlan::Projection(Projection::try_new(
exprs,
Arc::new(joined),
)?)))
}
#[derive(Debug)]
struct DimLink {
scan: LogicalPlan,
probe_keys: Vec<Column>,
key_columns: Vec<Column>,
}
#[derive(Debug, Clone)]
struct DimTable {
scan: LogicalPlan,
columns: Vec<Column>,
primary_key: Vec<Column>,
}
#[derive(Debug, Default)]
struct InputFacts {
tables: Vec<DimTable>,
equalities: Vec<(Column, Column)>,
}
impl InputFacts {
fn collect(plan: &LogicalPlan) -> Option<Self> {
let mut facts = Self::default();
facts.walk(plan)?;
let mut seen = HashSet::new();
for table in &facts.tables {
for col in &table.columns {
if !seen.insert(col.flat_name()) {
return None;
}
}
}
Some(facts)
}
fn walk(&mut self, plan: &LogicalPlan) -> Option<()> {
match plan {
LogicalPlan::TableScan(scan) => {
self.tables.push(dim_table(scan));
Some(())
}
LogicalPlan::Filter(filter) => self.walk(&filter.input),
LogicalPlan::Projection(proj) => {
projection_preserves_names(proj).then_some(())?;
self.walk(&proj.input)
}
LogicalPlan::Join(join) => {
(join.join_type == datafusion::common::JoinType::Inner).then_some(())?;
for (left, right) in &join.on {
if let (Expr::Column(l), Expr::Column(r)) = (left, right) {
self.equalities.push((l.clone(), r.clone()));
}
}
self.walk(&join.left)?;
self.walk(&join.right)
}
_ => None,
}
}
fn reach_one(&self, available: &HashSet<String>, skip: &[usize]) -> Option<(usize, DimLink)> {
for (index, table) in self.tables.iter().enumerate() {
if skip.contains(&index) || table.primary_key.is_empty() {
continue;
}
let mut probe_keys = Vec::with_capacity(table.primary_key.len());
let resolved = table.primary_key.iter().all(|key_col| {
if available.contains(&key_col.flat_name()) {
probe_keys.push(key_col.clone());
return true;
}
for (left, right) in &self.equalities {
for (near, far) in [(left, right), (right, left)] {
if near == key_col && available.contains(&far.flat_name()) {
probe_keys.push(far.clone());
return true;
}
}
}
false
});
if !resolved {
continue;
}
if table
.columns
.iter()
.all(|col| available.contains(&col.flat_name()))
{
continue;
}
return Some((
index,
DimLink {
scan: table.scan.clone(),
probe_keys,
key_columns: table.primary_key.clone(),
},
));
}
None
}
fn columns_of(&self, index: usize) -> &[Column] {
self.tables.get(index).map_or(&[], |t| t.columns.as_slice())
}
fn closure(&self, seed: &[Column]) -> HashSet<String> {
let mut available: HashSet<String> = seed.iter().map(Column::flat_name).collect();
let mut used: Vec<usize> = Vec::new();
while let Some((index, _)) = self.reach_one(&available, &used) {
for col in self.columns_of(index) {
available.insert(col.flat_name());
}
used.push(index);
}
available
}
fn minimal_key(&self, group_cols: &[Column]) -> Vec<Column> {
let mut retained: Vec<Column> = group_cols.to_vec();
let mut index = 0;
while let Some(candidate) = retained.get(index).cloned() {
let mut without = retained.clone();
without.remove(index);
if self.closure(&without).contains(&candidate.flat_name()) {
retained = without;
} else {
index += 1;
}
}
retained
}
fn dim_chain(&self, key: &[Column], deferred: &[Column]) -> Option<Vec<DimLink>> {
let mut available: HashSet<String> = key.iter().map(Column::flat_name).collect();
let mut materialised: HashSet<String> = HashSet::new();
let mut links: Vec<(usize, DimLink)> = Vec::new();
let mut used: Vec<usize> = Vec::new();
while !deferred
.iter()
.all(|col| available.contains(&col.flat_name()))
{
if links.len() >= MAX_DIM_CHAIN {
return None;
}
let (index, mut link) = self.reach_one(&available, &used)?;
for probe in &mut link.probe_keys {
if !materialised.contains(&probe.flat_name()) {
*probe = Column::new(
Some(TableReference::bare(TOPN_ALIAS)),
probe.name.clone(),
);
}
}
for col in self.columns_of(index) {
available.insert(col.flat_name());
materialised.insert(col.flat_name());
}
used.push(index);
links.push((index, link));
}
let mut wanted: HashSet<String> = deferred.iter().map(Column::flat_name).collect();
let mut keep: Vec<bool> = links
.iter()
.enumerate()
.rev()
.map(|(_, (index, link))| {
let supplies = self
.columns_of(*index)
.iter()
.any(|col| wanted.contains(&col.flat_name()));
if supplies {
for probe in &link.probe_keys {
wanted.insert(probe.flat_name());
}
}
supplies
})
.collect();
keep.reverse();
let pruned: Vec<DimLink> = links
.into_iter()
.zip(keep)
.filter_map(|((_, link), keep)| keep.then_some(link))
.collect();
(!pruned.is_empty()).then_some(pruned)
}
}
fn dim_table(scan: &TableScan) -> DimTable {
let schema = scan.projected_schema.as_ref();
let columns: Vec<Column> = schema
.iter()
.map(|(qualifier, field)| Column::new(qualifier.cloned(), field.name()))
.collect();
let primary_key = schema
.functional_dependencies()
.iter()
.find(|dep| dep.mode == Dependency::Single && !dep.nullable)
.map(|dep| {
dep.source_indices
.iter()
.filter_map(|idx| columns.get(*idx).cloned())
.collect::<Vec<Column>>()
})
.filter(|key| !key.is_empty())
.unwrap_or_default();
DimTable {
scan: LogicalPlan::TableScan(scan.clone()),
columns,
primary_key,
}
}
fn projection_preserves_names(proj: &Projection) -> bool {
let below: HashSet<String> = proj.input.schema().field_names().into_iter().collect();
proj.schema
.iter()
.enumerate()
.all(|(idx, (qualifier, field))| {
let out = Column::new(qualifier.cloned(), field.name());
if !below.contains(&out.flat_name()) {
return true;
}
matches!(proj.expr.get(idx), Some(Expr::Column(col)) if *col == out)
})
}
fn lower_exprs_through(proj: &Projection, exprs: &[Expr]) -> Option<Vec<Expr>> {
use datafusion::common::tree_node::TreeNode;
let mut lowered = Vec::with_capacity(exprs.len());
for expr in exprs {
let mut unfollowable = false;
let rewritten = expr
.clone()
.transform(|node| {
if let Expr::Column(col) = &node {
return match index_of(&proj.schema, col).and_then(|idx| proj.expr.get(idx)) {
Some(inner) => match unalias(inner) {
Some(inner) => Ok(Transformed::yes(inner)),
None => {
unfollowable = true;
Ok(Transformed::no(node))
}
},
None => Ok(Transformed::no(node)),
};
}
Ok(Transformed::no(node))
})
.ok()?;
if unfollowable {
return None;
}
lowered.push(rewritten.data);
}
Some(lowered)
}
fn unalias(expr: &Expr) -> Option<Expr> {
match expr {
Expr::Column(_) => Some(expr.clone()),
Expr::Alias(alias) => unalias(&alias.expr),
_ => None,
}
}
fn index_of(schema: &DFSchema, col: &Column) -> Option<usize> {
schema.index_of_column(col).ok()
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use datafusion::arrow::array::{Array as _, Decimal128Array, Int64Array, StringArray};
use datafusion::common::get_required_group_by_exprs_indices;
use datafusion::arrow::datatypes::{DataType, Field, Schema};
use datafusion::arrow::record_batch::RecordBatch;
use datafusion::common::{Constraint, Constraints};
use datafusion::datasource::MemTable;
use datafusion::execution::session_state::SessionStateBuilder;
use datafusion::prelude::SessionContext;
fn customer_table(with_key: bool) -> Arc<MemTable> {
let schema = Arc::new(Schema::new(vec![
Field::new("c_custkey", DataType::Int64, false),
Field::new("c_name", DataType::Utf8, false),
Field::new("c_address", DataType::Utf8, false),
Field::new("c_nationkey", DataType::Int64, false),
Field::new("c_phone", DataType::Utf8, false),
Field::new("c_acctbal", DataType::Decimal128(15, 2), false),
Field::new("c_comment", DataType::Utf8, false),
]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int64Array::from(vec![1i64, 2, 3, 4])),
Arc::new(StringArray::from(vec!["alice", "bob", "cara", "dan"])),
Arc::new(StringArray::from(vec!["a1", "a2", "a3", "a4"])),
Arc::new(Int64Array::from(vec![7i64, 7, 8, 8])),
Arc::new(StringArray::from(vec!["p1", "p2", "p3", "p4"])),
Arc::new(
Decimal128Array::from(vec![100i128, 200, 300, 400])
.with_precision_and_scale(15, 2)
.unwrap(),
),
Arc::new(StringArray::from(vec!["k1", "k2", "k3", "k4"])),
],
)
.unwrap();
let table = MemTable::try_new(schema, vec![vec![batch]]).unwrap();
Arc::new(if with_key {
table.with_constraints(Constraints::new_unverified(vec![Constraint::PrimaryKey(
vec![0],
)]))
} else {
table
})
}
fn nation_table(with_key: bool) -> Arc<MemTable> {
let schema = Arc::new(Schema::new(vec![
Field::new("n_nationkey", DataType::Int64, false),
Field::new("n_name", DataType::Utf8, false),
]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int64Array::from(vec![7i64, 8])),
Arc::new(StringArray::from(vec!["GERMANY", "FRANCE"])),
],
)
.unwrap();
let table = MemTable::try_new(schema, vec![vec![batch]]).unwrap();
Arc::new(if with_key {
table.with_constraints(Constraints::new_unverified(vec![Constraint::PrimaryKey(
vec![0],
)]))
} else {
table
})
}
fn orders_table() -> Arc<MemTable> {
let schema = Arc::new(Schema::new(vec![
Field::new("o_orderkey", DataType::Int64, false),
Field::new("o_custkey", DataType::Int64, false),
Field::new("o_totalprice", DataType::Decimal128(15, 2), false),
]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int64Array::from(vec![10i64, 11, 12, 13, 14, 15])),
Arc::new(Int64Array::from(vec![1i64, 1, 2, 3, 3, 3])),
Arc::new(
Decimal128Array::from(vec![500i128, 700, 900, 100, 200, 300])
.with_precision_and_scale(15, 2)
.unwrap(),
),
],
)
.unwrap();
let table = MemTable::try_new(schema, vec![vec![batch]]).unwrap();
Arc::new(table.with_constraints(Constraints::new_unverified(vec![
Constraint::PrimaryKey(vec![0]),
])))
}
fn context(with_rule: bool, with_keys: bool) -> SessionContext {
let mut builder = SessionStateBuilder::new().with_default_features();
if with_rule {
builder =
builder.with_optimizer_rule(Arc::new(LateMaterializeTopKAggregate::forced()));
}
let ctx = SessionContext::new_with_state(builder.build());
ctx.register_table("customer", customer_table(with_keys))
.unwrap();
ctx.register_table("nation", nation_table(with_keys))
.unwrap();
ctx.register_table("orders", orders_table()).unwrap();
ctx
}
async fn rows(ctx: &SessionContext, sql: &str) -> Vec<String> {
let batches = ctx.sql(sql).await.unwrap().collect().await.unwrap();
let mut out = Vec::new();
for batch in &batches {
for row in 0..batch.num_rows() {
let mut cells = Vec::new();
for col in 0..batch.num_columns() {
let casted =
datafusion::arrow::compute::cast(batch.column(col), &DataType::Utf8)
.unwrap();
let array = datafusion::common::cast::as_string_array(&casted).unwrap();
cells.push(if array.is_null(row) {
String::from("NULL")
} else {
array.value(row).to_string()
});
}
out.push(cells.join("|"));
}
}
out
}
async fn plan_of(ctx: &SessionContext, sql: &str) -> String {
format!(
"{}",
ctx.sql(sql)
.await
.unwrap()
.into_optimized_plan()
.unwrap()
.display_indent()
)
}
const Q10_SHAPE: &str = "SELECT c_custkey, c_name, sum(o_totalprice) AS revenue, \
c_acctbal, n_name, c_address, c_phone, c_comment \
FROM customer, orders, nation \
WHERE c_custkey = o_custkey AND c_nationkey = n_nationkey \
GROUP BY c_custkey, c_name, c_acctbal, c_phone, n_name, c_address, c_comment \
ORDER BY revenue DESC LIMIT 20";
#[tokio::test]
async fn the_group_by_narrows_to_the_key_alone() {
let plan = plan_of(&context(true, true), Q10_SHAPE).await;
assert!(
plan.contains("groupBy=[[customer.c_custkey]]"),
"expected a single-column group by:\n{plan}"
);
assert!(
plan.contains(TOPN_ALIAS),
"expected the bounded top-N branch:\n{plan}"
);
}
#[tokio::test]
async fn the_wide_columns_leave_the_aggregate_branch() {
let plan = plan_of(&context(true, true), Q10_SHAPE).await;
let narrow_branch = subtree_under(&plan, &format!("SubqueryAlias: {TOPN_ALIAS}"));
for wide in ["c_comment", "c_address", "c_phone", "c_name"] {
assert!(
!narrow_branch.contains(wide),
"{wide} still crosses the joins under the aggregate:\n\
--- narrowed branch ---\n{narrow_branch}\n--- whole plan ---\n{plan}"
);
}
assert!(
narrow_branch.contains("TableScan: customer projection=[c_custkey, c_nationkey]"),
"the narrowed branch should scan only the keys:\n{narrow_branch}"
);
assert!(
plan.contains("TableScan: customer projection=[c_custkey, c_name, c_address"),
"the join-back must still fetch the deferred columns:\n{plan}"
);
}
fn subtree_under(plan: &str, header: &str) -> String {
let indent = |line: &str| line.len() - line.trim_start().len();
let mut lines = plan.lines().skip_while(|line| !line.contains(header));
let root = lines.next().expect("subtree root not found");
std::iter::once(root)
.chain(lines.take_while(|line| indent(line) > indent(root)))
.collect::<Vec<_>>()
.join("\n")
}
#[tokio::test]
async fn results_are_identical_with_and_without_the_rule() {
for sql in [
Q10_SHAPE,
"SELECT c_custkey, c_name, count(*) AS n, c_comment FROM customer, orders \
WHERE c_custkey = o_custkey GROUP BY c_custkey, c_name, c_comment \
ORDER BY n DESC, c_custkey LIMIT 3",
"SELECT c_custkey, c_name, max(c_comment) AS m, sum(o_totalprice) AS t \
FROM customer, orders WHERE c_custkey = o_custkey \
GROUP BY c_custkey, c_name ORDER BY t DESC LIMIT 2",
"SELECT c_custkey, n_name, sum(o_totalprice) AS t FROM customer, orders, nation \
WHERE c_custkey = o_custkey AND c_nationkey = n_nationkey \
GROUP BY c_custkey, n_name ORDER BY t DESC, c_custkey LIMIT 4",
"SELECT c_custkey, c_name, c_address, sum(o_totalprice) AS t \
FROM customer, orders WHERE c_custkey = o_custkey \
GROUP BY c_custkey, c_name, c_address ORDER BY t DESC LIMIT 1",
"SELECT c_custkey, c_name, c_address, sum(o_totalprice) AS t \
FROM customer, orders WHERE c_custkey = o_custkey \
GROUP BY c_custkey, c_name, c_address ORDER BY t DESC LIMIT 100",
] {
let with = rows(&context(true, true), sql).await;
let without = rows(&context(false, true), sql).await;
assert_eq!(with, without, "results diverged for:\n{sql}");
assert!(!with.is_empty(), "test query returned nothing: {sql}");
}
}
#[tokio::test]
async fn the_ordering_survives_the_join_back() {
let out = rows(&context(true, true), Q10_SHAPE).await;
let revenues: Vec<f64> = out
.iter()
.map(|row| row.split('|').nth(2).unwrap().parse().unwrap())
.collect();
assert!(
revenues.windows(2).all(|w| w[0] >= w[1]),
"rows came back out of order: {revenues:?}"
);
assert_eq!(
out,
rows(&context(false, true), Q10_SHAPE).await,
"the ordered result must match the unrewritten plan row for row"
);
}
#[tokio::test]
async fn an_undeclared_key_is_not_assumed() {
let plan = plan_of(&context(true, false), Q10_SHAPE).await;
assert!(
!plan.contains(TOPN_ALIAS),
"no declared key means no rewrite:\n{plan}"
);
assert_eq!(
rows(&context(true, false), Q10_SHAPE).await,
rows(&context(false, false), Q10_SHAPE).await
);
}
#[tokio::test]
async fn an_unbounded_aggregate_is_left_alone() {
let sql = "SELECT c_custkey, c_name, sum(o_totalprice) AS t FROM customer, orders \
WHERE c_custkey = o_custkey GROUP BY c_custkey, c_name ORDER BY t DESC";
let plan = plan_of(&context(true, true), sql).await;
assert!(
!plan.contains(TOPN_ALIAS),
"no fetch means no bound to exploit:\n{plan}"
);
}
#[tokio::test]
async fn an_outer_join_below_refuses_the_rewrite() {
let sql = "SELECT c_custkey, c_name, c_comment, sum(o_totalprice) AS t \
FROM customer LEFT JOIN orders ON c_custkey = o_custkey \
GROUP BY c_custkey, c_name, c_comment ORDER BY t DESC LIMIT 5";
let plan = plan_of(&context(true, true), sql).await;
assert!(
!plan.contains(TOPN_ALIAS),
"must not rewrite under an outer join:\n{plan}"
);
assert_eq!(
rows(&context(true, true), sql).await,
rows(&context(false, true), sql).await
);
}
#[tokio::test]
async fn a_sort_on_a_determined_column_keeps_it_in_the_key() {
let sql = "SELECT c_custkey, c_name, c_comment, sum(o_totalprice) AS t \
FROM customer, orders WHERE c_custkey = o_custkey \
GROUP BY c_custkey, c_name, c_comment ORDER BY c_name DESC LIMIT 3";
let plan = plan_of(&context(true, true), sql).await;
assert!(
plan.contains("groupBy=[[customer.c_custkey, customer.c_name]]"),
"the sorted column must stay in the key:\n{plan}"
);
assert_eq!(
rows(&context(true, true), sql).await,
rows(&context(false, true), sql).await
);
}
#[tokio::test]
async fn the_rewrite_is_applied_exactly_once() {
let plan = plan_of(&context(true, true), Q10_SHAPE).await;
assert_eq!(
plan.matches(TOPN_ALIAS).count(),
plan.matches(&format!("{TOPN_ALIAS}.")).count() + 1,
"expected a single aliased branch:\n{plan}"
);
assert_eq!(
plan.matches("SubqueryAlias").count(),
1,
"expected exactly one rewrite:\n{plan}"
);
}
#[tokio::test]
async fn the_join_back_does_not_duplicate_rows() {
let sql = "SELECT c_custkey, c_name, c_comment, count(*) AS n \
FROM customer, orders WHERE c_custkey = o_custkey \
GROUP BY c_custkey, c_name, c_comment ORDER BY n DESC, c_custkey LIMIT 10";
let with = rows(&context(true, true), sql).await;
assert_eq!(with, rows(&context(false, true), sql).await);
assert_eq!(with.len(), 3, "expected one row per customer: {with:?}");
}
#[tokio::test]
async fn a_self_join_is_refused() {
let sql = "SELECT a.c_custkey, a.c_name, a.c_comment, count(*) AS n \
FROM customer a, customer b \
WHERE a.c_nationkey = b.c_nationkey \
GROUP BY a.c_custkey, a.c_name, a.c_comment ORDER BY n DESC LIMIT 5";
assert_eq!(
rows(&context(true, true), sql).await,
rows(&context(false, true), sql).await,
"a self-join must not change the answer"
);
}
async fn physical_plan_of(ctx: &SessionContext, sql: &str) -> String {
let logical = ctx.sql(sql).await.unwrap().into_optimized_plan().unwrap();
let physical = ctx.state().create_physical_plan(&logical).await.unwrap();
format!(
"{}",
datafusion::physical_plan::displayable(physical.as_ref()).indent(false)
)
}
#[tokio::test]
async fn the_join_back_is_a_real_equi_join() {
for sql in [
Q10_SHAPE,
"SELECT c_custkey, c_name, c_comment, sum(o_totalprice) AS t \
FROM customer, orders WHERE c_custkey = o_custkey \
GROUP BY c_custkey, c_name, c_comment ORDER BY t DESC LIMIT 5",
] {
let plan = physical_plan_of(&context(true, true), sql).await;
assert!(
!plan.contains("NestedLoopJoin") && !plan.contains("CrossJoin"),
"the join-back lost its keys for:\n{sql}\n\n{plan}"
);
}
}
#[tokio::test]
async fn the_inner_bound_survives_physical_planning() {
let plan = physical_plan_of(&context(true, true), Q10_SHAPE).await;
let bounded = plan
.lines()
.filter(|line| line.contains("Sort") && line.contains("fetch=20"))
.count();
assert!(
bounded >= 2,
"expected a bounded sort on the narrowed branch as well as on top, \
found {bounded}:\n{plan}"
);
}
#[test]
fn the_env_switch_is_honoured() {
for off in ["off", "OFF", "0", "false", "no", " off "] {
assert!(!enabled_from(off), "{off:?} should disable the rule");
}
for on in ["", "on", "1", "true", "anything-else"] {
assert!(enabled_from(on), "{on:?} should leave the rule enabled");
}
}
#[tokio::test]
async fn datafusion_alone_does_not_reach_through_the_second_table() {
let ctx = context(false, true);
let plan = ctx.sql(Q10_SHAPE).await.unwrap().into_optimized_plan().unwrap();
let agg = find_aggregate(&plan)
.unwrap_or_else(|| panic!("no aggregate in:\n{}", plan.display_indent()));
let names: Vec<String> = agg
.group_expr
.iter()
.map(|e| e.schema_name().to_string())
.collect();
let minimal = get_required_group_by_exprs_indices(agg.input.schema(), &names)
.expect("customer's declared key must be visible");
let kept: Vec<&String> = minimal.iter().map(|i| &names[*i]).collect();
assert!(
kept.iter().any(|n| n.ends_with("n_name")),
"DataFusion is expected to keep n_name; it kept {kept:?}"
);
}
fn find_aggregate(plan: &LogicalPlan) -> Option<&Aggregate> {
if let LogicalPlan::Aggregate(agg) = plan {
return Some(agg);
}
plan.inputs().into_iter().find_map(find_aggregate)
}
}