Skip to content
245 changes: 244 additions & 1 deletion datafusion/core/tests/physical_optimizer/join_selection.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,8 +22,11 @@ use std::{
task::{Context, Poll},
};

use arrow::array::record_batch;
use arrow::datatypes::{DataType, Field, Schema, SchemaRef};
use arrow::record_batch::RecordBatch;
use datafusion::datasource::object_store::ObjectStoreUrl;
use datafusion::prelude::{SessionConfig, SessionContext};
use datafusion_common::config::ConfigOptions;
use datafusion_common::tree_node::TreeNodeRecursion;
use datafusion_common::{ColumnStatistics, JoinType, ScalarValue, stats::Precision};
Expand All @@ -33,13 +36,17 @@ use datafusion_execution::{RecordBatchStream, SendableRecordBatchStream, TaskCon
use datafusion_expr::Operator;
use datafusion_physical_expr::PhysicalExprRef;
use datafusion_physical_expr::expressions::col;
use datafusion_physical_expr::expressions::{BinaryExpr, Column, NegativeExpr};
use datafusion_physical_expr::expressions::{
BinaryExpr, Column, DynamicFilterPhysicalExpr, NegativeExpr, lit,
};
use datafusion_physical_expr::intervals::utils::check_support;
use datafusion_physical_expr::{EquivalenceProperties, Partitioning, PhysicalExpr};
use datafusion_physical_expr_common::sort_expr::PhysicalSortExpr;
use datafusion_physical_optimizer::PhysicalOptimizerContext;
use datafusion_physical_optimizer::PhysicalOptimizerRule;
use datafusion_physical_optimizer::filter_pushdown::FilterPushdown;
use datafusion_physical_optimizer::join_selection::JoinSelection;
use datafusion_physical_plan::collect;
use datafusion_physical_plan::displayable;
use datafusion_physical_plan::joins::utils::ColumnIndex;
use datafusion_physical_plan::joins::utils::JoinFilter;
Expand All @@ -48,6 +55,7 @@ use datafusion_physical_plan::operator_statistics::{
ClosureStatisticsProvider, StatisticsRegistry, StatisticsResult,
};
use datafusion_physical_plan::projection::ProjectionExec;
use datafusion_physical_plan::repartition::RepartitionExec;
use datafusion_physical_plan::sorts::sort_preserving_merge::SortPreservingMergeExec;
use datafusion_physical_plan::{
ChildrenPropertiesMode, ExecutionPlanProperties, ReplaceChildrenOptions,
Expand All @@ -59,8 +67,11 @@ use datafusion_physical_plan::{
};

use futures::Stream;
use object_store::memory::InMemory;
use rstest::rstest;

use super::pushdown_utils::TestScanBuilder;

/// Return statistics for empty table
fn empty_statistics() -> Statistics {
Statistics {
Expand Down Expand Up @@ -1969,3 +1980,235 @@ fn test_join_with_maybe_swap_unbounded_case(t: TestCase) -> Result<()> {
}
Ok(())
}

#[rstest]

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.

Could we add a regression that lets FilterPushdown::new_post_optimization() create the dynamic filter and probe-side consumer, then reruns JoinSelection and executes the plan? The current tests correctly exercise both guards, but manually calling with_dynamic_filter_expr does not verify the producer/consumer connection or query results.

#[case(PartitionMode::CollectLeft)]
#[case(PartitionMode::Partitioned)]
#[tokio::test]
async fn test_join_selection_skips_hash_join_with_dynamic_filter(
#[case] partition_mode: PartitionMode,
) -> Result<()> {
// Left has larger statistics than right, which would normally trigger swap
let (big, small) = create_big_and_small();
let on = vec![(
Arc::new(Column::new_with_schema("big_col", &big.schema())?) as PhysicalExprRef,
Arc::new(Column::new_with_schema("small_col", &small.schema())?)
as PhysicalExprRef,
)];

let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new(
vec![Arc::clone(&on[0].1)],
lit(true),
));

#[expect(deprecated)]
let join = Arc::new(
HashJoinExec::try_new(
Arc::clone(&big),
Arc::clone(&small),
on,
None,
&JoinType::Inner,
None,
partition_mode,
NullEquality::NullEqualsNothing,
false,
)?
.with_dynamic_filter_expr(dynamic_filter)?,
);

let original_schema = join.schema();

// JoinSelection must not fail and must leave the join unchanged
let optimized = JoinSelection::new().optimize(join, &ConfigOptions::new())?;
let optimized_join = optimized
.downcast_ref::<HashJoinExec>()
.expect("join should remain HashJoinExec without wrapping projection");

assert_eq!(*optimized_join.partition_mode(), partition_mode);
assert_eq!(*optimized_join.join_type(), JoinType::Inner);
assert!(Arc::ptr_eq(optimized_join.left(), &big));
assert!(Arc::ptr_eq(optimized_join.right(), &small));
assert_eq!(optimized_join.schema(), original_schema);
assert_eq!(optimized_join.dynamic_expressions_produced().len(), 1);

Ok(())
}

#[tokio::test]
async fn test_join_selection_skips_unbounded_hash_join_with_dynamic_filter() -> Result<()>
{
let left_exec: Arc<dyn ExecutionPlan> = Arc::new(UnboundedExec::new(
None,
RecordBatch::new_empty(Arc::new(Schema::new(vec![Field::new(
"a",
DataType::Int32,
false,
)]))),
2,
));
let right_exec: Arc<dyn ExecutionPlan> = Arc::new(UnboundedExec::new(
Some(1),
RecordBatch::new_empty(Arc::new(Schema::new(vec![Field::new(
"b",
DataType::Int32,
false,
)]))),
2,
));

let on = vec![(
col("a", &left_exec.schema())?,
col("b", &right_exec.schema())?,
)];

let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new(
vec![Arc::clone(&on[0].1)],
lit(true),
));

#[expect(deprecated)]
let join = Arc::new(
HashJoinExec::try_new(
Arc::clone(&left_exec),
Arc::clone(&right_exec),
on,
None,
&JoinType::Inner,
None,
PartitionMode::Partitioned,
NullEquality::NullEqualsNothing,
false,
)?
.with_dynamic_filter_expr(dynamic_filter)?,
);

let original_schema = join.schema();

// hash_join_swap_subrule would normally swap unbounded left with bounded right,
// but must skip this join because it has a dynamic filter.
let optimized = JoinSelection::new().optimize(join, &ConfigOptions::new())?;
let optimized_join = optimized
.downcast_ref::<HashJoinExec>()
.expect("join should remain HashJoinExec without wrapping projection");

assert_eq!(*optimized_join.partition_mode(), PartitionMode::Partitioned);
assert_eq!(*optimized_join.join_type(), JoinType::Inner);
assert!(Arc::ptr_eq(optimized_join.left(), &left_exec));
assert!(Arc::ptr_eq(optimized_join.right(), &right_exec));
assert_eq!(optimized_join.schema(), original_schema);
assert_eq!(optimized_join.dynamic_expressions_produced().len(), 1);

Ok(())
}

/// End-to-end regression for #26106: let `FilterPushdown::new_post_optimization()`
/// wire up a real dynamic filter, then ensure `JoinSelection` leaves the join
/// untouched and the plan still executes with correct results.
#[tokio::test]
async fn test_join_selection_skips_real_dynamic_filter_pushdown() -> Result<()> {
let build_schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::Utf8, false),
Field::new("b", DataType::Utf8, false),
]));
let build_scan = TestScanBuilder::new(Arc::clone(&build_schema))
.with_support(true)
.with_batches(vec![
record_batch!(("a", Utf8, ["aa", "ab"]), ("b", Utf8, ["ba", "bb"])).unwrap(),
])
.build();

let probe_schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::Utf8, false),
Field::new("b", DataType::Utf8, false),
]));
let probe_scan = TestScanBuilder::new(Arc::clone(&probe_schema))
.with_support(true)
.with_batches(vec![
record_batch!(
("a", Utf8, ["aa", "ab", "ac"]),
("b", Utf8, ["ba", "bb", "bc"])
)
.unwrap(),
])
.build();

let partition_count = 4;
let build_repartition = Arc::new(
RepartitionExec::try_new(
build_scan,
Partitioning::Hash(
vec![col("a", &build_schema)?, col("b", &build_schema)?],
partition_count,
),
)
.unwrap(),
);
let probe_repartition = Arc::new(
RepartitionExec::try_new(
probe_scan,
Partitioning::Hash(
vec![col("a", &probe_schema)?, col("b", &probe_schema)?],
partition_count,
),
)
.unwrap(),
);

let on = vec![
(col("a", &build_schema)?, col("a", &probe_schema)?),
(col("b", &build_schema)?, col("b", &probe_schema)?),
];
let join = Arc::new(
HashJoinExec::try_new(
build_repartition,
probe_repartition,
on,
None,
&JoinType::Inner,
None,
PartitionMode::Partitioned,
NullEquality::NullEqualsNothing,
false,
)
.unwrap(),
) as Arc<dyn ExecutionPlan>;

let mut config = ConfigOptions::new();
config.execution.parquet.pushdown_filters = true;

// Real optimizer wiring (not manual `with_dynamic_filter_expr`)
let with_filter = FilterPushdown::new_post_optimization().optimize(join, &config)?;
let join_with_filter = with_filter
.downcast_ref::<HashJoinExec>()
.expect("plan should still be a HashJoinExec after pushdown");
assert!(
!join_with_filter.dynamic_expressions_produced().is_empty(),
"FilterPushdown should have created a dynamic filter"
);

// JoinSelection must not panic in `swap_inputs` and must leave the join as-is
let after_selection = JoinSelection::new().optimize(with_filter, &config)?;
let final_join = after_selection
.downcast_ref::<HashJoinExec>()
.expect("JoinSelection must leave dynamic-filter join as HashJoinExec");
assert!(!final_join.dynamic_expressions_produced().is_empty());
assert_eq!(*final_join.partition_mode(), PartitionMode::Partitioned);

// The wired plan must still execute and return the 2 matching rows
let session_config = SessionConfig::new();
let session_ctx = SessionContext::new_with_config(session_config);
session_ctx.register_object_store(
ObjectStoreUrl::parse("test://").unwrap().as_ref(),
Arc::new(InMemory::new()),
);
let task_ctx = session_ctx.task_ctx();
let batches = collect(after_selection, task_ctx).await?;
let total_rows: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(
total_rows, 2,
"expected 2 inner-join matches, got {batches:?}"
);

Ok(())
}
59 changes: 34 additions & 25 deletions datafusion/physical-optimizer/src/join_selection.rs
Original file line number Diff line number Diff line change
Expand Up @@ -301,32 +301,40 @@ fn statistical_join_selection_subrule(
context: &dyn PhysicalOptimizerContext,
) -> Result<Transformed<Arc<dyn ExecutionPlan>>> {
let transformed = if let Some(hash_join) = plan.downcast_ref::<HashJoinExec>() {
match hash_join.partition_mode() {
PartitionMode::Auto => try_collect_left(hash_join, false, context)?
.map_or_else(
|| partitioned_hash_join(hash_join, context).map(Some),
|v| Ok(Some(v)),
)?,
PartitionMode::CollectLeft => try_collect_left(hash_join, true, context)?
.map_or_else(
|| partitioned_hash_join(hash_join, context).map(Some),
|v| Ok(Some(v)),
)?,
PartitionMode::Partitioned => {
let left = hash_join.left();
let right = hash_join.right();
if can_swap_hash_join(hash_join)
&& should_swap_join_order(&**left, &**right, context)?
{
// Null-aware RightAnti only supports CollectLeft
let partition_mode = if hash_join.null_aware {
PartitionMode::CollectLeft
if !hash_join.dynamic_expressions_produced().is_empty() {
// Once a HashJoinExec carries a dynamic filter, its build side
// has already been determined and the dynamic filter has been wired
// to the probe side. Reordering inputs would invalidate the dynamic
// filter, so skip this join and leave it unchanged.
None
} else {
match hash_join.partition_mode() {
PartitionMode::Auto => try_collect_left(hash_join, false, context)?
.map_or_else(
|| partitioned_hash_join(hash_join, context).map(Some),
|v| Ok(Some(v)),
)?,
PartitionMode::CollectLeft => try_collect_left(hash_join, true, context)?
.map_or_else(
|| partitioned_hash_join(hash_join, context).map(Some),
|v| Ok(Some(v)),
)?,
PartitionMode::Partitioned => {
let left = hash_join.left();
let right = hash_join.right();
if can_swap_hash_join(hash_join)
&& should_swap_join_order(&**left, &**right, context)?
{
// Null-aware RightAnti only supports CollectLeft
let partition_mode = if hash_join.null_aware {
PartitionMode::CollectLeft
} else {
PartitionMode::Partitioned
};
hash_join.swap_inputs(partition_mode).map(Some)?
} else {
PartitionMode::Partitioned
};
hash_join.swap_inputs(partition_mode).map(Some)?
} else {
None
None
}
}
}
}
Expand Down Expand Up @@ -524,6 +532,7 @@ pub fn hash_join_swap_subrule(
_config_options: &ConfigOptions,
) -> Result<Arc<dyn ExecutionPlan>> {
if let Some(hash_join) = input.downcast_ref::<HashJoinExec>()
&& hash_join.dynamic_expressions_produced().is_empty()
&& hash_join.left.boundedness().is_unbounded()
&& !hash_join.right.boundedness().is_unbounded()
&& !hash_join.null_aware // Don't swap null-aware anti joins
Expand Down