Skip to content
310 changes: 309 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,300 @@ 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.
///
/// The build side reports larger statistics than the probe side so that
/// `should_swap_join_order` returns true: without the dynamic-filter guard,
/// `JoinSelection` would attempt `swap_inputs` and fail.
#[tokio::test]
async fn test_join_selection_skips_real_dynamic_filter_pushdown() -> Result<()> {
let build_schema = Arc::new(Schema::new(vec![
Field::new("build_a", DataType::Utf8, false),
Field::new("build_b", DataType::Utf8, false),
]));
let build_scan = TestScanBuilder::new(Arc::clone(&build_schema))
.with_support(true)
.with_batches(vec![
record_batch!(
("build_a", Utf8, ["aa", "ab"]),
("build_b", Utf8, ["ba", "bb"])
)
.unwrap(),
])
.build();

let probe_schema = Arc::new(Schema::new(vec![
Field::new("probe_a", DataType::Utf8, false),
Field::new("probe_b", DataType::Utf8, false),
]));
let probe_scan = TestScanBuilder::new(Arc::clone(&probe_schema))
.with_support(true)
.with_batches(vec![
record_batch!(
("probe_a", Utf8, ["aa", "ab", "ac"]),
("probe_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("build_a", &build_schema)?,
col("build_b", &build_schema)?,
],
partition_count,
),
)
.unwrap(),
);
let probe_repartition = Arc::new(
RepartitionExec::try_new(
probe_scan,
Partitioning::Hash(
vec![
col("probe_a", &probe_schema)?,
col("probe_b", &probe_schema)?,
],
partition_count,
),
)
.unwrap(),
);

let on = vec![
(
col("build_a", &build_schema)?,
col("probe_a", &probe_schema)?,
),
(
col("build_b", &build_schema)?,
col("probe_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"
);
let orig_left = Arc::clone(join_with_filter.left());
let orig_right = Arc::clone(join_with_filter.right());

// Both `TestScanBuilder` inputs report the same 123-byte size, so force the
// build side to look larger. Without the dynamic-filter guard this makes
// `JoinSelection` attempt the swap and fail in `swap_inputs`.
let mut registry = StatisticsRegistry::new();
registry.register(Arc::new(ClosureStatisticsProvider::with_matches(
|plan| {
plan.schema()
.fields()
.iter()
.any(|f| f.name().starts_with("build_"))
},
|_plan, _child_stats| Ok(StatisticsResult::Computed(big_statistics().into())),
)));

struct ContextWithRegistry {
config: ConfigOptions,
registry: StatisticsRegistry,
}

impl PhysicalOptimizerContext for ContextWithRegistry {
fn config_options(&self) -> &ConfigOptions {
&self.config
}

fn statistics_registry(&self) -> Option<&StatisticsRegistry> {
Some(&self.registry)
}
}

let context = ContextWithRegistry {
config: config.clone(),
registry,
};

// JoinSelection must not fail in `swap_inputs` and must leave the join as-is
let after_selection =
JoinSelection::new().optimize_with_context(with_filter, &context)?;
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);
assert!(
Arc::ptr_eq(final_join.left(), &orig_left),
"build side must not be swapped"
);
assert!(
Arc::ptr_eq(final_join.right(), &orig_right),
"probe side must not be swapped"
);

// 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(())
}
Loading