use datafusion::arrow::datatypes::{DataType, Field, Schema};
use datafusion::physical_expr::expressions::Column;
use datafusion::physical_plan::ExecutionPlan;
use datafusion::physical_plan::joins::{HashJoinExec, PartitionMode};
use datafusion::physical_plan::repartition::RepartitionExec;
use datafusion::physical_plan::Partitioning;
use datafusion::physical_optimizer::PhysicalOptimizerRule;
use krishiv_sql::distributed_plan::{ShuffleReadExec, planning_session_context_with_options};
use std::sync::Arc;
fn wide_schema() -> Arc<Schema> {
Arc::new(Schema::new(vec![
Field::new("c_custkey", DataType::Int64, false),
Field::new("c_name", DataType::Utf8, false),
Field::new("c_address", DataType::Utf8, false),
Field::new("c_phone", DataType::Utf8, false),
Field::new("c_comment", DataType::Utf8, false),
]))
}
fn narrow_schema() -> Arc<Schema> {
Arc::new(Schema::new(vec![
Field::new("o_custkey", DataType::Int64, false),
Field::new("o_orderkey", DataType::Int64, false),
]))
}
fn build_side(rows: usize) -> Arc<dyn ExecutionPlan> {
Arc::new(
ShuffleReadExec::new(0, 4, 1, wide_schema(), None).with_upstream_estimate(Some(rows), None),
)
}
fn probe_side() -> datafusion::error::Result<Arc<dyn ExecutionPlan>> {
let read: Arc<dyn ExecutionPlan> = Arc::new(
ShuffleReadExec::new(1, 4, 4, narrow_schema(), None)
.with_upstream_estimate(Some(5_000_000), Some(80_000_000)),
);
Ok(Arc::new(RepartitionExec::try_new(
read,
Partitioning::RoundRobinBatch(4),
)?))
}
fn chosen_mode(build_rows: usize) -> datafusion::error::Result<PartitionMode> {
let ctx = planning_session_context_with_options(4, None, None);
let state = ctx.state();
let opts = state.config().options();
let join = HashJoinExec::try_new(
build_side(build_rows),
probe_side()?,
vec![(
Arc::new(Column::new("c_custkey", 0)),
Arc::new(Column::new("o_custkey", 0)),
)],
None,
&datafusion::common::JoinType::Inner,
None,
PartitionMode::Auto,
datafusion::common::NullEquality::NullEqualsNothing,
false,
)?;
let optimized = datafusion::physical_optimizer::join_selection::JoinSelection::new()
.optimize(Arc::new(join), opts)?;
let hj = optimized.downcast_ref::<HashJoinExec>().ok_or_else(|| {
datafusion::error::DataFusionError::Internal(String::from(
"JoinSelection replaced the hash join with another operator",
))
})?;
Ok(*hj.partition_mode())
}
#[test]
fn a_wide_sub_ceiling_build_side_with_no_byte_estimate_is_broadcast() {
let mode = chosen_mode(900_000).expect("join selection");
assert_eq!(
mode,
PartitionMode::CollectLeft,
"MECHANISM CHECK: if this is no longer CollectLeft, the width-blind row \
ceiling is not how q10 broadcasts a wide build side, and the fix must \
be re-derived before touching join_estimates"
);
}
#[test]
fn the_wide_broadcast_is_converted_to_a_partitioned_join() {
let join = wide_join(900_000).expect("wide join");
let converted = krishiv_sql::distributed_plan::redistribute_unsplittable_broadcast_joins(
Arc::clone(&join),
)
.expect("conversion");
let mode = converted
.downcast_ref::<HashJoinExec>()
.map(|hj| *hj.partition_mode());
assert_eq!(
mode,
Some(PartitionMode::Partitioned),
"a build side the ROW ceiling admitted but the BYTE ceiling would not \
must be hash-partitioned, not copied to every task"
);
}
#[test]
fn a_narrow_sub_ceiling_build_side_stays_broadcast() {
let build: Arc<dyn ExecutionPlan> = Arc::new(
ShuffleReadExec::new(0, 4, 1, narrow_schema(), None)
.with_upstream_estimate(Some(900_000), None),
);
let join: Arc<dyn ExecutionPlan> = Arc::new(
HashJoinExec::try_new(
build,
probe_side().expect("probe side"),
vec![(
Arc::new(Column::new("o_custkey", 0)),
Arc::new(Column::new("o_custkey", 0)),
)],
None,
&datafusion::common::JoinType::Inner,
None,
PartitionMode::CollectLeft,
datafusion::common::NullEquality::NullEqualsNothing,
false,
)
.expect("hash join"),
);
let converted =
krishiv_sql::distributed_plan::redistribute_unsplittable_broadcast_joins(Arc::clone(&join))
.expect("conversion");
assert_eq!(
converted
.downcast_ref::<HashJoinExec>()
.map(|hj| *hj.partition_mode()),
Some(PartitionMode::CollectLeft),
"a genuinely small build side must still broadcast; converting these is \
the regression join_estimates was written about"
);
}
fn wide_join(rows: usize) -> datafusion::error::Result<Arc<dyn ExecutionPlan>> {
Ok(Arc::new(
HashJoinExec::try_new(
build_side(rows),
probe_side()?,
vec![(
Arc::new(Column::new("c_custkey", 0)),
Arc::new(Column::new("o_custkey", 0)),
)],
None,
&datafusion::common::JoinType::Inner,
None,
PartitionMode::CollectLeft,
datafusion::common::NullEquality::NullEqualsNothing,
false,
)?,
))
}
#[test]
fn the_q8_shape_stays_partitioned_and_must_keep_doing_so() {
let mode = chosen_mode(4_000_000).expect("join selection");
assert_ne!(
mode,
PartitionMode::CollectLeft,
"a 4M-row build side must not broadcast; this is the shape whose \
conversion regressed four queries"
);
}