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
28 changes: 23 additions & 5 deletions datafusion/physical-plan/src/joins/sort_merge_join/exec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ use crate::joins::utils::{
use crate::metrics::{ExecutionPlanMetricsSet, MetricsSet, SpillMetrics};
use crate::projection::{
ProjectionExec, join_allows_pushdown, join_table_borders, new_join_children,
physical_to_column_exprs, update_join_on,
physical_to_column_exprs, update_join_filter, update_join_on,
};
use crate::spill::spill_manager::SpillManager;
use crate::statistics::{ChildStats, StatisticsArgs};
Expand Down Expand Up @@ -651,15 +651,33 @@ impl ExecutionPlan for SortMergeJoinExec {
return Ok(None);
}

let left_field_size = self.left().schema().fields().len();
let left_projection = &projection_as_columns[0..=far_right_left_col_ind as usize];
let right_projection = &projection_as_columns[far_left_right_col_ind as usize..];

let Some(new_on) = update_join_on(
&projection_as_columns[0..=far_right_left_col_ind as _],
&projection_as_columns[far_left_right_col_ind as _..],
left_projection,
right_projection,
self.on(),
self.left().schema().fields().len(),
left_field_size,
) else {
return Ok(None);
};

let new_filter = if let Some(filter) = self.filter() {
let Some(filter) = update_join_filter(
left_projection,
right_projection,
filter,
left_field_size,
) else {
return Ok(None);
};
Some(filter)
} else {
None
};

let (new_left, new_right) = new_join_children(
&projection_as_columns,
far_right_left_col_ind,
Expand All @@ -672,7 +690,7 @@ impl ExecutionPlan for SortMergeJoinExec {
Arc::new(new_left),
Arc::new(new_right),
new_on,
self.filter.clone(),
new_filter,
self.join_type,
self.sort_options.clone(),
self.null_equality,
Expand Down
113 changes: 113 additions & 0 deletions datafusion/physical-plan/src/joins/sort_merge_join/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ use super::bitwise_stream::BitwiseSortMergeJoinStream;
use crate::joins::utils::{ColumnIndex, JoinFilter, JoinOn};
use crate::joins::{HashJoinExec, PartitionMode, SortMergeJoinExec};
use crate::metrics::{ExecutionPlanMetricsSet, SpillMetrics};
use crate::projection::{ProjectionExec, ProjectionExpr};
use crate::spill::spill_manager::SpillManager;
use crate::test::TestMemoryExec;
use crate::test::exec::BarrierExec;
Expand Down Expand Up @@ -361,6 +362,118 @@ async fn join_collect_batch_size_equals_two(
Ok((columns, batches))
}

fn join_and_projection_for_pushdown(
filter: Option<JoinFilter>,
) -> Result<(Arc<SortMergeJoinExec>, ProjectionExec)> {
let left = build_table(("a1", &vec![1]), ("b1", &vec![2]), ("c1", &vec![3]));
let right = build_table(("a2", &vec![4]), ("b2", &vec![2]), ("c2", &vec![5]));
let on = vec![(
Arc::new(Column::new("b1", 1)) as _,
Arc::new(Column::new("b2", 1)) as _,
)];
let join = Arc::new(SortMergeJoinExec::try_new(
left,
right,
on,
filter,
Inner,
vec![SortOptions::default()],
NullEquality::NullEqualsNothing,
)?);
let input: Arc<dyn ExecutionPlan> = Arc::clone(&join) as _;
let projection = ProjectionExec::try_new(
[
ProjectionExpr {
expr: Arc::new(Column::new("c1", 2)),
alias: "c1".to_string(),
},
ProjectionExpr {
expr: Arc::new(Column::new("b1", 1)),
alias: "b1".to_string(),
},
ProjectionExpr {
expr: Arc::new(Column::new("c2", 5)),
alias: "c2".to_string(),
},
ProjectionExpr {
expr: Arc::new(Column::new("b2", 4)),
alias: "b2".to_string(),
},
],
input,
)?;

Ok((join, projection))
}

#[test]
fn projection_pushdown_remaps_filter() -> Result<()> {
let filter = JoinFilter::new(
Arc::new(BinaryExpr::new(
Arc::new(Column::new("c1", 0)),
Operator::Lt,
Arc::new(Column::new("c2", 1)),
)),
vec![
ColumnIndex {
index: 2,
side: JoinSide::Left,
},
ColumnIndex {
index: 2,
side: JoinSide::Right,
},
],
Arc::new(Schema::new(vec![
Field::new("c1", DataType::Int32, false),
Field::new("c2", DataType::Int32, false),
])),
);
let (join, projection) = join_and_projection_for_pushdown(Some(filter))?;

let swapped = join
.try_swapping_with_projection(&projection)?
.expect("projection should be pushed below the join");
let swapped = swapped
.downcast_ref::<SortMergeJoinExec>()
.expect("swapped plan should be a SortMergeJoinExec");

let (left_on, right_on) = &swapped.on()[0];
assert_eq!(left_on.downcast_ref::<Column>().unwrap().index(), 1);
assert_eq!(right_on.downcast_ref::<Column>().unwrap().index(), 1);
assert_eq!(
swapped.filter().as_ref().unwrap().column_indices(),
&[
ColumnIndex {
index: 0,
side: JoinSide::Left,
},
ColumnIndex {
index: 0,
side: JoinSide::Right,
},
]
);

Ok(())
}

#[test]
fn projection_pushdown_without_filter() -> Result<()> {
let (join, projection) = join_and_projection_for_pushdown(None)?;

let swapped = join
.try_swapping_with_projection(&projection)?
.expect("projection should be pushed below the join");
let swapped = swapped
.downcast_ref::<SortMergeJoinExec>()
.expect("swapped plan should be a SortMergeJoinExec");

assert!(swapped.filter().is_none());

Ok(())
}

#[tokio::test]
async fn join_inner_one() -> Result<()> {
let left = build_table(
Expand Down
38 changes: 38 additions & 0 deletions datafusion/sqllogictest/test_files/joins.slt
Original file line number Diff line number Diff line change
Expand Up @@ -2858,6 +2858,44 @@ NULL 1970-01-04T00:00:00 789 ghi 1970-01-04 NULL 789 qwe
NULL NULL NULL NULL NULL 1970-01-04T00:00:00 0 qwerty
NULL NULL NULL NULL NULL NULL 100000 abcdefg

# Regression test: projection pushdown through SortMergeJoinExec must not reuse a
# JoinFilter after its input columns have been projected away.
statement ok
set datafusion.optimizer.repartition_joins = true;

query TT
explain
select t1.column2 as left_b1, t2.column2 as right_b1
from (values (100, 1, 0)) t1
join (values (10, 1)) t2
on t1.column2 = t2.column2
and t1.column1 > t2.column1;
----
logical_plan
01)Projection: t1.column2 AS left_b1, t2.column2 AS right_b1
02)--Inner Join: t1.column2 = t2.column2 Filter: t1.column1 > t2.column1
03)----SubqueryAlias: t1
04)------Projection: column1, column2
05)--------Values: (Int64(100), Int64(1), Int64(0))
06)----SubqueryAlias: t2
07)------Values: (Int64(10), Int64(1))
physical_plan
01)ProjectionExec: expr=[column2@1 as left_b1, column2@3 as right_b1]
02)--SortMergeJoinExec: join_type=Inner, on=[(column2@1, column2@1)], filter=column1@0 > column1@1
03)----SortExec: expr=[column2@1 ASC], preserve_partitioning=[false]
04)------DataSourceExec: partitions=1, partition_sizes=[1]
05)----SortExec: expr=[column2@1 ASC], preserve_partitioning=[false]
06)------DataSourceExec: partitions=1, partition_sizes=[1]

query II
select t1.column2 as left_b1, t2.column2 as right_b1
from (values (100, 1, 0)) t1
join (values (10, 1)) t2
on t1.column2 = t2.column2
and t1.column1 > t2.column1;
----
1 1

####
# Config teardown
####
Expand Down