Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 58 additions & 16 deletions datafusion/expr/src/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -616,28 +616,45 @@ pub fn compare_sort_expr(
Ordering::Equal
}

/// Group a slice of window expression expr by their order by expressions
/// Group window expressions by their sort keys, preserving any outer alias.
pub fn group_window_expr_by_sort_keys(
window_expr: impl IntoIterator<Item = Expr>,
) -> Result<Vec<(WindowSortKey, Vec<Expr>)>> {
let mut result = vec![];
window_expr.into_iter().try_for_each(|expr| match &expr {
Expr::WindowFunction(window_fun) => {
let WindowFunctionParams{ partition_by, order_by, ..} = &window_fun.as_ref().params;
let sort_key = generate_sort_key(partition_by, order_by)?;
if let Some((_, values)) = result.iter_mut().find(
|group: &&mut (WindowSortKey, Vec<Expr>)| matches!(group, (key, _) if *key == sort_key),
) {
values.push(expr);
} else {
result.push((sort_key, vec![expr]))
}
Ok(())

window_expr.into_iter().try_for_each(|expr| {
// Read the window's settings through one alias.
// Keep `expr` intact so its output name is preserved.
let inner = match &expr {
Expr::Alias(alias) => alias.expr.as_ref(),
_ => &expr,
};

let Expr::WindowFunction(window_fun) = inner else {
return internal_err!("Impossibly got non-window expr {expr:?}");
};

let WindowFunctionParams {
partition_by,
order_by,
..
} = &window_fun.as_ref().params;

let sort_key = generate_sort_key(partition_by, order_by)?;

if let Some((_, values)) = result.iter_mut().find(
|group: &&mut (WindowSortKey, Vec<Expr>)| {
matches!(group, (key, _) if *key == sort_key)
},
) {
values.push(expr);
} else {
result.push((sort_key, vec![expr]));
}
other => internal_err!(
"Impossibly got non-window expr {other:?}"
),

Ok(())
})?;

Ok(result)
}

Expand Down Expand Up @@ -1550,6 +1567,31 @@ mod tests {
Ok(())
}

#[test]
fn test_group_window_expr_by_sort_keys_aliased_window_expr() -> Result<()> {
let age_asc = Sort::new(col("age"), true, true);
let max1 = Expr::from(WindowFunction::new(
WindowFunctionDefinition::AggregateUDF(max_udaf()),
vec![col("name")],
))
.order_by(vec![age_asc.clone()])
.build()
.unwrap();
// The same window function under an alias, as the Substrait consumer
// produces when a window column's default name collides with an
// input column. It must be grouped by the inner function's sort key
// and kept aliased.
let max1_aliased = max1.clone().alias("max_name");

let result =
group_window_expr_by_sort_keys(vec![max1.clone(), max1_aliased.clone()])?;

let key = vec![(age_asc, false)];
let expected: Vec<(WindowSortKey, Vec<Expr>)> =
vec![(key, vec![max1, max1_aliased])];
assert_eq!(expected, result);
Ok(())
}
#[test]
fn test_group_window_expr_by_sort_keys() -> Result<()> {
let age_asc = Sort::new(col("age"), true, true);
Expand Down
60 changes: 60 additions & 0 deletions datafusion/substrait/src/logical_plan/consumer/rel/project_rel.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,8 +18,10 @@
use crate::logical_plan::consumer::SubstraitConsumer;
use crate::logical_plan::consumer::utils::NameTracker;
use async_recursion::async_recursion;
use datafusion::common::tree_node::{Transformed, TreeNode, TreeNodeRecursion};
use datafusion::common::{Column, not_impl_err};
use datafusion::logical_expr::builder::project;
use datafusion::logical_expr::expr_rewriter::NamePreserver;
use datafusion::logical_expr::utils::find_window_exprs;
use datafusion::logical_expr::{Expr, LogicalPlan, LogicalPlanBuilder};
use std::collections::HashSet;
Expand Down Expand Up @@ -65,6 +67,37 @@ pub async fn from_project_rel(
}

let input = if !window_exprs.is_empty() {
// Window outputs must have unique names across the input schema
// and the new window expressions.
let mut window_names = NameTracker::new();
window_names.reserve_schema(&original_schema);

let mut aliased_columns: Vec<(Expr, Expr)> = vec![];
let window_exprs = window_exprs
.into_iter()
.map(|window_expr| {
let named =
window_names.get_uniquely_named_expr(window_expr.clone())?;

if let Expr::Alias(alias) = &named {
aliased_columns.push((
window_expr,
Expr::Column(Column::from_name(&alias.name)),
));
}

Ok(named)
})
.collect::<datafusion::common::Result<Vec<_>>>()?;

// References to renamed windows must point to their new output columns.
if !aliased_columns.is_empty() {
explicit_exprs = explicit_exprs
.into_iter()
.map(|expr| reference_aliased_windows(expr, &aliased_columns))
.collect::<datafusion::common::Result<Vec<_>>>()?;
}

LogicalPlanBuilder::window_plan(input, window_exprs)?
} else {
input
Expand All @@ -81,3 +114,30 @@ pub async fn from_project_rel(
not_impl_err!("Projection without an input is not supported")
}
}

/// Reference renamed window outputs while preserving the projection's
/// original output name.
fn reference_aliased_windows(
expr: Expr,
aliased_columns: &[(Expr, Expr)],
) -> datafusion::common::Result<Expr> {
let saved_name = NamePreserver::new_for_projection().save(&expr);

let rewritten = expr
.transform_down(|node| {
match aliased_columns
.iter()
.find(|(window_expr, _)| *window_expr == node)
{
Some((_, column)) => Ok(Transformed::new(
column.clone(),
true,
TreeNodeRecursion::Jump,
)),
None => Ok(Transformed::no(node)),
}
})?
.data;

Ok(saved_name.restore(rewritten))
}
8 changes: 7 additions & 1 deletion datafusion/substrait/src/logical_plan/consumer/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
use crate::logical_plan::consumer::SubstraitConsumer;
use datafusion::arrow::datatypes::{DataType, Field, Schema, TimeUnit, UnionFields};
use datafusion::common::{
DFSchema, DFSchemaRef, TableReference, exec_err, not_impl_err,
Column, DFSchema, DFSchemaRef, TableReference, exec_err, not_impl_err,
substrait_datafusion_err, substrait_err,
};
use datafusion::logical_expr::expr::Sort;
Expand Down Expand Up @@ -424,6 +424,12 @@ impl NameTracker {
}
}

pub(super) fn reserve_schema(&mut self, schema: &DFSchema) {
for (qualifier, field) in schema.iter() {
self.insert(&Expr::Column(Column::from((qualifier, field))));
}
}

/// Check if the expression would cause a conflict either in:
/// 1. validate_unique_names (duplicate schema_name)
/// 2. DFSchema::check_names (ambiguous reference)
Expand Down
57 changes: 57 additions & 0 deletions datafusion/substrait/tests/cases/roundtrip_logical_plan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1684,6 +1684,63 @@ async fn simple_window_function() -> Result<()> {
roundtrip("SELECT RANK() OVER (PARTITION BY a ORDER BY b), d, sum(b) OVER (PARTITION BY a) FROM data;").await
}

#[tokio::test]
async fn stacked_windows_with_same_default_name_via_builder() -> Result<()> {
// Substrait projections are positional and drop DataFusion's aliases, so the
// inherited `rn1` column and the new window both come back under the default
// name `row_number() ROWS BETWEEN ...`. The consumer must keep them apart.
// See https://github.com/apache/datafusion/issues/23007
use datafusion::functions_window::expr_fn::row_number;

let ctx = create_context().await?;
let scan = ctx.table("data").await?.into_optimized_plan()?;
let plan = LogicalPlanBuilder::from(scan)
.window(vec![row_number().alias("rn1")])?
.window(vec![row_number().alias("rn2")])?
.build()?;

let plan2 = substrait_roundtrip(&plan, &ctx).await?;
// Compare output fields; qualifiers and functional dependencies
// can differ after a Substrait round trip.
assert_eq!(plan.schema().as_arrow(), plan2.schema().as_arrow());
DataFrame::new(ctx.state(), plan2).show().await?;
Ok(())
}

#[tokio::test]
async fn chained_windows_with_same_default_name() -> Result<()> {
// SQL form from https://github.com/apache/datafusion/issues/23007. The two
// `avg` windows differ only by aliases that Substrait does not carry. As in
// roundtrip_self_join, the consumer must synthesize an alias, so the plan
// text differs; verify schema and executability instead.
let ctx = create_context().await?;
let plan = ctx
.sql(
"SELECT a, b, avg1, avg2 FROM (
SELECT a, b, avg1,
row_number() OVER () AS seq2,
avg(b) OVER (PARTITION BY a ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW) AS avg2
FROM (
SELECT a, b, avg1 FROM (
SELECT a, b,
row_number() OVER () AS seq1,
avg(b) OVER (PARTITION BY a ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW) AS avg1
FROM data
) t1 ORDER BY seq1
) t2
) t3 ORDER BY seq2",
)
.await?
.into_optimized_plan()?;

let plan2 = substrait_roundtrip(&plan, &ctx).await?;
// Compare output fields; qualifiers and functional dependencies
// can differ after a Substrait round trip.
assert_eq!(plan.schema().as_arrow(), plan2.schema().as_arrow());
DataFrame::new(ctx.state(), plan2).show().await?;
Ok(())
}

#[tokio::test]
async fn window_with_rows() -> Result<()> {
roundtrip("SELECT sum(b) OVER (PARTITION BY a ROWS BETWEEN 2 PRECEDING AND CURRENT ROW) FROM data;").await?;
Expand Down