diff --git a/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs b/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs index b48905500d54..3009acfcc847 100644 --- a/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs +++ b/datafusion/physical-plan/src/joins/sort_merge_join/exec.rs @@ -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}; @@ -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, @@ -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, diff --git a/datafusion/physical-plan/src/joins/sort_merge_join/tests.rs b/datafusion/physical-plan/src/joins/sort_merge_join/tests.rs index 91d1b893f1b2..d811b0e489d3 100644 --- a/datafusion/physical-plan/src/joins/sort_merge_join/tests.rs +++ b/datafusion/physical-plan/src/joins/sort_merge_join/tests.rs @@ -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; @@ -361,6 +362,118 @@ async fn join_collect_batch_size_equals_two( Ok((columns, batches)) } +fn join_and_projection_for_pushdown( + filter: Option, +) -> Result<(Arc, 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 = 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::() + .expect("swapped plan should be a SortMergeJoinExec"); + + let (left_on, right_on) = &swapped.on()[0]; + assert_eq!(left_on.downcast_ref::().unwrap().index(), 1); + assert_eq!(right_on.downcast_ref::().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::() + .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( diff --git a/datafusion/sqllogictest/test_files/joins.slt b/datafusion/sqllogictest/test_files/joins.slt index 7a706836f44d..496255b4c957 100644 --- a/datafusion/sqllogictest/test_files/joins.slt +++ b/datafusion/sqllogictest/test_files/joins.slt @@ -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 ####