Skip to content
119 changes: 68 additions & 51 deletions datafusion/optimizer/src/push_down_filter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1135,8 +1135,8 @@ impl OptimizerRule for PushDownFilter {
// multiple window functions, each with potentially different partition keys.
// Therefore, we need to ensure that any potential partition key returned is used in
// ALL window functions. Otherwise, filters cannot be pushed by through that column.
fn extract_partition_keys(func: &WindowFunction) -> HashSet<Column> {
expr_columns(&func.params.partition_by)
fn extract_partition_keys(func: &WindowFunction) -> HashSet<&Expr> {
func.params.partition_by.iter().collect()
}

let potential_partition_keys = window
Expand Down Expand Up @@ -1172,21 +1172,23 @@ impl OptimizerRule for PushDownFilter {
let mut keep_predicates = vec![];
let mut push_predicates = vec![];
for expr in predicates {
let cols = expr.column_refs();
if cols.iter().all(|c| potential_partition_keys.contains(c)) {
// A volatile predicate has to stay above the window: pushing it
// changes which rows the window function sees.
if !expr.is_volatile()
&& reads_only_partition_keys(&expr, &potential_partition_keys)?
{
push_predicates.push(expr);
} else {
keep_predicates.push(expr);
}
}

// Unlike with aggregations, there are no cases where we have to replace, e.g.,
// `a+b` with Column(a)+Column(b). This is because partition expressions are not
// available as standalone columns to the user. For example, while an aggregation on
// `a+b` becomes Column(a + b), in a window partition it becomes
// `func() PARTITION BY [a + b] ...`. Thus, filters on expressions always remain in
// place, so we can use `push_predicates` directly. This is consistent with other
// optimizers, such as the one used by Postgres.
// `a+b` with Column(a+b). This is because partition expressions are not available
// as standalone columns to the user: while an aggregation on `a+b` becomes
// Column(a + b), in a window partition it stays `func() PARTITION BY [a + b] ...`.
// That is why the predicate is matched against the key expressions themselves and
// can be pushed unchanged.

// If we have a filter to push, we push it down to the input of the aggregate
let result = if let Some(predicate) = conjunction(push_predicates) {
Expand Down Expand Up @@ -1585,6 +1587,54 @@ fn with_filters(predicates: Vec<Expr>, plan: LogicalPlan) -> LogicalPlan {
}
}

/// Can `expr` be evaluated below a window with these `PARTITION BY` keys?
///
/// A predicate that reads only the partition keys is constant within each
/// partition, so filtering before the window drops whole partitions and leaves
/// the surviving rows' window values unchanged. "Reads only the keys" is checked
/// structurally: every column reference must sit inside a subtree that is equal
/// to one of the keys. Given `PARTITION BY a, b + c`:
///
/// * `a < 5`, `b + c = 4` and `(b + c) + 1 > 10` can be pushed down

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

👍

/// * `d < 5` and `b < 5` cannot (`b` on its own is not a key), and neither can
/// `c + b = 4` (the match is structural, `c + b` is not `b + c`)
///
/// A predicate containing a subquery is never pushed: what the subquery reads is
/// not visible from the expression tree.
fn reads_only_partition_keys(
expr: &Expr,
partition_keys: &HashSet<&Expr>,
) -> Result<bool> {
let mut reads_something_else = false;
expr.apply(|node| {
Ok(if partition_keys.contains(&node) {
// the whole key was matched, so whatever it reads is accounted for
TreeNodeRecursion::Jump
} else if reads_beyond_this_node(node) {
reads_something_else = true;
TreeNodeRecursion::Stop
} else {
TreeNodeRecursion::Continue
})
})?;
Ok(!reads_something_else)
}

/// Does this node read data that walking its children cannot account for? A
/// column reads itself; a subquery reads whatever its plan reads, which
/// `Expr::apply` does not visit.
fn reads_beyond_this_node(node: &Expr) -> bool {
matches!(
node,
Expr::Column(_)
| Expr::OuterReferenceColumn(..)
| Expr::ScalarSubquery(_)
| Expr::Exists(_)
| Expr::InSubquery(_)
| Expr::SetComparison(_)
)
}

fn expr_columns(exprs: &[Expr]) -> HashSet<Column> {
exprs
.iter()
Expand Down Expand Up @@ -1956,10 +2006,12 @@ mod tests {
)
}

/// verifies that filters on partition expressions are not pushed, as the single expression
/// column is not available to the user, unlike with aggregations
/// verifies that a filter on an expression partition key is pushed; the
/// remaining shapes (mixed keys, operand order, subqueries, volatile
/// predicates, several windows) are covered in
/// `push_down_filter_regression.slt`
#[test]
fn filter_expression_keep_window() -> Result<()> {
fn filter_expression_move_window() -> Result<()> {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

the filter is moved!

let table_scan = test_table_scan()?;

let window = Expr::from(WindowFunction::new(
Expand All @@ -1975,49 +2027,14 @@ mod tests {

let plan = LogicalPlanBuilder::from(table_scan)
.window(vec![window])?
// unlike with aggregations, single partition column "test.a + test.b" is not available
// to the plan, so we use multiple columns when filtering
.filter(add(col("a"), col("b")).gt(lit(10i64)))?
.build()?;

assert_optimized_plan_equal!(
plan,
@r"
Filter: test.a + test.b > Int64(10)
WindowAggr: windowExpr=[[rank() PARTITION BY [test.a + test.b] ORDER BY [test.c ASC NULLS FIRST] ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW]]
TableScan: test
"
)
}

/// verifies that filters are not pushed on order by columns (that are not used in partitioning)
#[test]
fn filter_order_keep_window() -> Result<()> {
let table_scan = test_table_scan()?;

let window = Expr::from(WindowFunction::new(
WindowFunctionDefinition::WindowUDF(
datafusion_functions_window::rank::rank_udwf(),
),
vec![],
))
.partition_by(vec![col("a")])
.order_by(vec![col("c").sort(true, true)])
.build()
.unwrap();

let plan = LogicalPlanBuilder::from(table_scan)
.window(vec![window])?
.filter(col("c").gt(lit(10i64)))?
.filter(add(col("a"), col("b")).gt(lit(10i64)))? // a + b > 10
.build()?;
assert_plan_not_transformed!(plan.clone());

assert_optimized_plan_equal!(
plan,
@r"
Filter: test.c > Int64(10)
WindowAggr: windowExpr=[[rank() PARTITION BY [test.a] ORDER BY [test.c ASC NULLS FIRST] ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW]]
TableScan: test
WindowAggr: windowExpr=[[rank() PARTITION BY [test.a + test.b] ORDER BY [test.c ASC NULLS FIRST] ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW]]
TableScan: test, full_filters=[test.a + test.b > Int64(10)]
"
)
}
Expand Down
Loading
Loading