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
166 changes: 156 additions & 10 deletions datafusion/physical-expr/src/equivalence/class.rs

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 public EquivalenceGroup::join (no schema) keeps the old unsound outer-join behavior with no doc warning or deprecation, so external callers still get wrong equivalences.

Original file line number Diff line number Diff line change
Expand Up @@ -22,14 +22,15 @@ use std::sync::Arc;
use std::vec::IntoIter;

use super::ProjectionMapping;
use crate::expressions::Literal;
use crate::expressions::{Column, Literal};
use crate::physical_expr::add_offset_to_expr;
use crate::projection::ProjectionTargets;
use crate::{PhysicalExpr, PhysicalExprRef, PhysicalSortExpr, PhysicalSortRequirement};

use datafusion_common::tree_node::{Transformed, TransformedResult, TreeNode};
use datafusion_common::{JoinType, Result, ScalarValue};
use datafusion_physical_expr_common::physical_expr::format_physical_expr_list;
use datafusion_physical_expr_common::physical_expr::is_volatile;

use indexmap::{IndexMap, IndexSet};

Expand Down Expand Up @@ -809,23 +810,54 @@ impl EquivalenceGroup {
}

/// Combine equivalence groups of the given join children.
///
/// For outer joins, equivalences are conservatively reduced to direct
/// columns on the null-extended side and expressions independent of that
/// side. This avoids evaluating arbitrary expressions during planning.
pub fn join(
&self,
right_equivalences: &Self,
join_type: &JoinType,
left_size: usize,
on: &[(PhysicalExprRef, PhysicalExprRef)],
) -> Result<Self> {
self.join_equivalences(right_equivalences, join_type, left_size, on)
}

fn join_equivalences(
&self,
right_equivalences: &Self,
join_type: &JoinType,
left_size: usize,
on: &[(PhysicalExprRef, PhysicalExprRef)],
) -> Result<Self> {
let group = match join_type {
JoinType::Inner | JoinType::Left | JoinType::Full | JoinType::Right => {
let mut result = Self::new(
self.iter().cloned().chain(
right_equivalences
.iter()
.map(|cls| cls.try_with_offset(left_size as _))
.collect::<Result<Vec<_>>>()?,
),
let mut left = self.clone();
let mut right = Self::new(
right_equivalences
.iter()
.map(|cls| cls.try_with_offset(left_size as _))
.collect::<Result<Vec<_>>>()?,
);

// Outer joins add rows where the nullable side is NULL. Keep
// direct columns (which all become NULL together) and
// expressions independent of that side. Drop other expressions
// conservatively; evaluating arbitrary ScalarUDFs during
// planning can panic or have side effects.
match join_type {
JoinType::Left => right = right.with_null_extended_side(),
JoinType::Right => left = left.with_null_extended_side(),
JoinType::Full => {
left = left.with_null_extended_side();
right = right.with_null_extended_side();
}
_ => {}
}

let mut result =
Self::new(left.iter().cloned().chain(right.iter().cloned()));
// In we have an inner join, expressions in the "on" condition
// are equal in the resulting table.
if join_type == &JoinType::Inner {
Expand All @@ -847,6 +879,42 @@ impl EquivalenceGroup {
Ok(group)
}

/// Conservatively retain equivalences valid when this side is null-extended.
fn with_null_extended_side(self) -> Self {
let classes = self.classes.into_iter().flat_map(|class| {
let mut null_equivalences: Vec<(PhysicalExprRef, Vec<PhysicalExprRef>)> =
vec![];
let mut unchanged = Vec::new();
for expr in class {
if is_volatile(&expr) {
continue;
}
if contains_column_reference(&expr) {
if let Some(null_expr) = null_extended_expr(&expr) {
if let Some((_, equivalent_exprs)) =
null_equivalences.iter_mut().find(|(existing, _)| {
existing.as_ref().eq(null_expr.as_ref())
})
{
equivalent_exprs.push(expr);
} else {
null_equivalences.push((null_expr, vec![expr]));
}
}
} else {
unchanged.push(expr);
}
}
null_equivalences
.into_iter()
.map(|(_, exprs)| exprs)
.chain(std::iter::once(unchanged))
.map(EquivalenceClass::new)
.filter(|class| !class.is_trivial())
});
Self::new(classes)
}

/// Checks if two expressions are equal directly or through equivalence
/// classes. For complex expressions (e.g. `a + b`), checks that the
/// expression trees are structurally identical and their leaf nodes are
Expand Down Expand Up @@ -902,6 +970,34 @@ impl EquivalenceGroup {
}
}

/// Whether an expression tree references any input column.
fn contains_column_reference(expr: &PhysicalExprRef) -> bool {
if expr.downcast_ref::<Column>().is_some() {
return true;
}
let children = expr.children();
if children.is_empty() {
return expr.downcast_ref::<Literal>().is_none();
}
children.into_iter().any(contains_column_reference)
}

/// Replace input columns with NULL without evaluating the expression.
fn null_extended_expr(expr: &PhysicalExprRef) -> Option<PhysicalExprRef> {
if expr.downcast_ref::<Column>().is_some() {
return Some(Arc::new(Literal::new(ScalarValue::Null)));
}
let children = expr.children();
if children.is_empty() {
return expr.downcast_ref::<Literal>().map(|_| Arc::clone(expr));
}
let children = children
.into_iter()
.map(null_extended_expr)
.collect::<Option<Vec<_>>>()?;
Arc::clone(expr).with_new_children(children).ok()
}

impl Deref for EquivalenceGroup {
type Target = [EquivalenceClass];

Expand Down Expand Up @@ -954,11 +1050,61 @@ impl From<Vec<EquivalenceClass>> for EquivalenceGroup {
mod tests {
use super::*;
use crate::equivalence::tests::create_test_params;
use crate::expressions::{BinaryExpr, Column, binary, col, lit};
use crate::expressions::{BinaryExpr, Column, binary, col, is_null, lit};
use arrow::datatypes::{DataType, Field, Schema, SchemaRef};

use datafusion_expr::Operator;

#[test]
fn outer_joins_preserve_column_equivalences_on_both_sides() -> Result<()> {
let schema = abc_schema();
let left = group_of(&schema, &[("a", "b")])?;
let right = group_of(&schema, &[("a", "b")])?;

let joined = left.join(&right, &JoinType::Right, 4, &[])?;

assert_eq!(joined.len(), 2);
assert!(joined.exprs_equal(&col("a", &schema)?, &col("b", &schema)?));
assert!(joined.exprs_equal(
&(Arc::new(Column::new("a", 4)) as _),
&(Arc::new(Column::new("b", 5)) as _),
));
Ok(())
}

#[test]
fn null_extension_keeps_expressions_with_matching_non_null_results() -> Result<()> {
let schema = abc_schema();
let expr_a = is_null(col("a", &schema)?)?;
let expr_b = is_null(col("b", &schema)?)?;
let group = EquivalenceGroup::new(vec![EquivalenceClass::new([
Arc::clone(&expr_a),
Arc::clone(&expr_b),
])]);

let joined = group.with_null_extended_side();

assert!(joined.exprs_equal(&expr_a, &expr_b));
Ok(())
}

#[test]
fn cross_join_path_preserves_both_inputs() -> Result<()> {
let schema = abc_schema();
let left = group_of(&schema, &[("a", "b")])?;
let right = group_of(&schema, &[("a", "b")])?;

// A cross join uses the inner-join equivalence path with no join keys.
let joined = left.join(&right, &JoinType::Inner, 4, &[])?;

assert_eq!(joined.len(), 2);
assert!(joined.exprs_equal(&col("a", &schema)?, &col("b", &schema)?));
assert!(joined.exprs_equal(
&(Arc::new(Column::new("a", 4)) as _),
&(Arc::new(Column::new("b", 5)) as _),
));
Ok(())
}

#[test]
fn test_bridge_groups() -> Result<()> {
// First entry in the tuple is argument, second entry is the bridged result
Expand Down
10 changes: 4 additions & 6 deletions datafusion/physical-expr/src/equivalence/properties/joins.rs
Original file line number Diff line number Diff line change
Expand Up @@ -32,13 +32,11 @@ pub fn join_equivalence_properties(
on: &[(PhysicalExprRef, PhysicalExprRef)],
) -> Result<EquivalenceProperties> {
let left_size = left.schema.fields.len();
let equivalences =
left.eq_group()
.join(right.eq_group(), join_type, left_size, on)?;
let mut result = EquivalenceProperties::new(join_schema);
result.add_equivalence_group(left.eq_group().join(
right.eq_group(),
join_type,
left_size,
on,
)?)?;
result.add_equivalence_group(equivalences)?;

let EquivalenceProperties {
oeq_class: left_oeq_class,
Expand Down
49 changes: 49 additions & 0 deletions datafusion/sqllogictest/test_files/joins.slt
Original file line number Diff line number Diff line change
Expand Up @@ -2074,6 +2074,55 @@ set datafusion.explain.logical_plan_only = false;
statement ok
set datafusion.execution.target_partitions = 4;

statement ok
set datafusion.execution.target_partitions = 1;

statement ok
CREATE TABLE outer_join_sort_left (id BIGINT) AS VALUES (1), (2);

statement ok
CREATE TABLE outer_join_sort_right (a BIGINT, b BIGINT) AS VALUES (2, 2), (3, 3);

query III

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 regression tests cover only LEFT and FULL OUTER JOIN; the new JoinType::Right branch (class.rs:874) and the cross-join path are never exercised.

SELECT * FROM outer_join_sort_left
LEFT JOIN (
SELECT a, b FROM outer_join_sort_right WHERE b = coalesce(a, 0)
) ON id = b
ORDER BY coalesce(a, 0), id;
----
1 NULL NULL
2 2 2

query III
SELECT * FROM outer_join_sort_left
LEFT JOIN (
SELECT a, b FROM outer_join_sort_right WHERE b = coalesce(a, 0)
) ON id = b
ORDER BY coalesce(a, 0), id
LIMIT 1;
----
1 NULL NULL

query III
SELECT * FROM outer_join_sort_left
FULL OUTER JOIN (
SELECT a, b FROM outer_join_sort_right WHERE b = coalesce(a, 0)
) ON id = b
ORDER BY coalesce(a, 0), id;
----
1 NULL NULL
2 2 2
NULL 3 3

statement ok
DROP TABLE outer_join_sort_left;

statement ok
DROP TABLE outer_join_sort_right;

statement ok
set datafusion.execution.target_partitions = 4;

# Planning inner nested loop join
# inputs are swapped due to inexact statistics + join reordering caused additional projection

Expand Down