Skip to content
Merged
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
98 changes: 98 additions & 0 deletions datafusion/core/tests/physical_optimizer/enforce_distribution.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1030,6 +1030,104 @@ fn range_right_mark_hash_join_reuses_range_partitioning() -> Result<()> {
Ok(())
}

#[test]
fn range_hash_join_repartitions_unsatisfied_side_to_match_range() -> Result<()> {
let left = parquet_exec_with_output_partitioning(range_partitioning(
"b",
[100, 200, 300],
SortOptions::default(),
)?);
let right = parquet_exec_with_output_partitioning(range_partitioning(
"a",
[10, 20, 30],
SortOptions::default(),
)?);
let join_on = vec![(
Arc::new(Column::new_with_schema("a", &left.schema())?) as _,
Arc::new(Column::new_with_schema("a", &right.schema())?) as _,
)];
let join = hash_join_exec(left, right, &join_on, &JoinType::Inner);

let plan = TestConfig::default()
.with_query_execution_partitions(4)
.to_plan(join, &DISTRIB_DISTRIB_SORT);

assert_plan!(
plan,
@r"
HashJoinExec: mode=Partitioned, join_type=Inner, on=[(a@0, a@0)]
RepartitionExec: partitioning=Range([a@0 ASC], [(10), (20), (30)], 4), input_partitions=4
DataSourceExec: file_groups={4 groups: [[p0], [p1], [p2], [p3]]}, projection=[a, b, c, d, e], output_partitioning=Range([b@1 ASC], [(100), (200), (300)], 4), file_type=parquet
DataSourceExec: file_groups={4 groups: [[p0], [p1], [p2], [p3]]}, projection=[a, b, c, d, e], output_partitioning=Range([a@0 ASC], [(10), (20), (30)], 4), file_type=parquet
"
);

Ok(())
}

#[test]
fn range_hash_join_repartitions_unpartitioned_side_to_match_range() -> Result<()> {
let left = parquet_exec();
let right = parquet_exec_with_output_partitioning(range_partitioning(
"a",
[10, 20, 30],
SortOptions::default(),
)?);
let join_on = vec![(
Arc::new(Column::new_with_schema("a", &left.schema())?) as _,
Arc::new(Column::new_with_schema("a", &right.schema())?) as _,
)];
let join = hash_join_exec(left, right, &join_on, &JoinType::Inner);

let plan = TestConfig::default()
.with_query_execution_partitions(4)
.to_plan(join, &DISTRIB_DISTRIB_SORT);

assert_plan!(
plan,
@r"
HashJoinExec: mode=Partitioned, join_type=Inner, on=[(a@0, a@0)]
RepartitionExec: partitioning=Range([a@0 ASC], [(10), (20), (30)], 4), input_partitions=1
DataSourceExec: file_groups={1 group: [[x]]}, projection=[a, b, c, d, e], file_type=parquet
DataSourceExec: file_groups={4 groups: [[p0], [p1], [p2], [p3]]}, projection=[a, b, c, d, e], output_partitioning=Range([a@0 ASC], [(10), (20), (30)], 4), file_type=parquet
"
);

Ok(())
}

#[test]
fn range_hash_join_rehashes_incompatible_data_type() -> Result<()> {
let left = parquet_exec();
let right = parquet_exec_with_output_partitioning(range_partitioning(
"a",
[10, 20, 30],
SortOptions::default(),
)?);
let join_on = vec![(
Arc::new(Column::new_with_schema("d", &left.schema())?) as _,
Arc::new(Column::new_with_schema("a", &right.schema())?) as _,
)];
let join = hash_join_exec(left, right, &join_on, &JoinType::Inner);

let plan = TestConfig::default()
.with_query_execution_partitions(4)
.to_plan(join, &DISTRIB_DISTRIB_SORT);

assert_plan!(
plan,
@r"
HashJoinExec: mode=Partitioned, join_type=Inner, on=[(d@3, a@0)]
RepartitionExec: partitioning=Hash([d@3], 4), input_partitions=1
DataSourceExec: file_groups={1 group: [[x]]}, projection=[a, b, c, d, e], file_type=parquet
RepartitionExec: partitioning=Hash([a@0], 4), input_partitions=4
DataSourceExec: file_groups={4 groups: [[p0], [p1], [p2], [p3]]}, projection=[a, b, c, d, e], output_partitioning=Range([a@0 ASC], [(10), (20), (30)], 4), file_type=parquet
"
);

Ok(())
}

#[test]
fn range_right_semi_hash_join_rehashes_incompatible_sort_options() -> Result<()> {
let left = parquet_exec_with_output_partitioning(range_partitioning(
Expand Down
186 changes: 186 additions & 0 deletions datafusion/physical-expr/src/partitioning.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ use crate::{
EquivalenceProperties, PhysicalExpr, equivalence::ProjectionMapping,
expressions::UnKnownColumn, physical_exprs_contains, physical_exprs_equal,
};
use arrow::datatypes::Schema;
pub use datafusion_common::SplitPoint;
use datafusion_common::{Result, validate_range_split_points};
use datafusion_physical_expr_common::physical_expr::format_physical_expr_list;
Expand Down Expand Up @@ -281,6 +282,49 @@ impl RangePartitioning {
split_points: self.split_points.clone(),
})
}

/// Checks whether the types of the given expressions match the data types of the split points in this range partitioning.
pub fn is_compatible_with_expressions(
&self,
exprs: &[Arc<dyn PhysicalExpr>],
schema: &Schema,
) -> bool {
if self.ordering.len() != exprs.len() {
return false;
}
if let Some(first_split) = self.split_points.first() {
exprs.iter().zip(first_split.values()).all(|(expr, val)| {
expr.data_type(schema)
.map(|dt| dt == val.data_type())
.unwrap_or(false)
})
} else {
true
}
}

/// Adapts this range partitioning to the given expressions, preserving split points and sort options.
/// Returns `None` if `exprs` count doesn't match ordering length or expression types don't match split points.
pub fn adapt(
&self,
exprs: &[Arc<dyn PhysicalExpr>],
schema: &Schema,
) -> Option<Self> {
if !self.is_compatible_with_expressions(exprs, schema) {
return None;
}
let new_ordering = LexOrdering::new(
exprs
.iter()
.zip(&self.ordering)
.map(|(expr, sort_expr)| PhysicalSortExpr {
expr: Arc::clone(expr),
options: sort_expr.options,
})
.collect::<Vec<_>>(),
)?;
Self::try_new(new_ordering, self.split_points.clone()).ok()
}
}

impl Display for RangePartitioning {
Expand Down Expand Up @@ -517,6 +561,37 @@ impl Partitioning {
}
}
}

/// Adapts this partitioning scheme to satisfy a required [`Distribution`] on the given schema.
///
/// - For `Partitioning::Hash`: creates `Partitioning::Hash(exprs, partition_count)`.
/// - For `Partitioning::Range`: adapts the range partitioning to the requirement's expressions using [`RangePartitioning::adapt`].
/// - For other partitioning schemes: returns `None`.
#[expect(
deprecated,
reason = "HashPartitioned is accepted during the KeyPartitioned migration"
)]
pub fn adapt(
&self,
child_requirement: &Distribution,
child_schema: &Schema,
) -> Option<Self> {
let (Distribution::HashPartitioned(exprs) | Distribution::KeyPartitioned(exprs)) =
child_requirement
else {
return None;
};

match self {
Partitioning::Range(ref_range) => ref_range
.adapt(exprs, child_schema)
.map(Partitioning::Range),
Partitioning::Hash(_, ref_count) => {
Some(Partitioning::Hash(exprs.to_vec(), *ref_count))
}
_ => None,
}
}
}

/// Protobuf conversions for [`Partitioning`].
Expand Down Expand Up @@ -1291,6 +1366,117 @@ mod tests {

Ok(())
}

#[test]
fn test_range_partitioning_adapt() -> Result<()> {
let fixture = PartitioningTestFixture::new(vec![
("a", DataType::Int32),
("b", DataType::Int64),
("c", DataType::Int32),
])?;

let range = fixture.range(
[0],
vec![
SplitPoint::new(vec![ScalarValue::Int32(Some(10))]),
SplitPoint::new(vec![ScalarValue::Int32(Some(20))]),
],
);

// Adapting to col_c (same type Int32) succeeds
let adapted = range.adapt(&[fixture.col(2)], &fixture.schema).unwrap();
assert_eq!(adapted.ordering().len(), 1);
assert!(adapted.ordering()[0].expr.eq(&fixture.col(2)));
assert_eq!(adapted.partition_count(), 3);

// Adapting to col_b (different type Int64) fails
assert!(range.adapt(&[fixture.col(1)], &fixture.schema).is_none());

// Adapting to empty or mismatch count fails
assert!(range.adapt(&[], &fixture.schema).is_none());
assert!(
range
.adapt(&fixture.cols([0, 2]), &fixture.schema)
.is_none()
);

// Partitioning::adapt works with Distribution::KeyPartitioned
let part = Partitioning::Range(range);
assert!(
part.adapt(&fixture.key_distribution([1]), &fixture.schema)
.is_none()
);

let adapted_part = part
.adapt(&fixture.key_distribution([2]), &fixture.schema)
.unwrap();
match adapted_part {
Partitioning::Range(r) => assert!(r.ordering()[0].expr.eq(&fixture.col(2))),
_ => panic!("expected Range partitioning"),
}

// Partitioning::Hash adaptation
let hash_part = fixture.hash_partitioning([1], 4);
let adapted_hash = hash_part
.adapt(&fixture.key_distribution([2]), &fixture.schema)
.unwrap();
match adapted_hash {
Partitioning::Hash(exprs, count) => {
assert_eq!(count, 4);
assert_eq!(exprs.len(), 1);
assert!(exprs[0].eq(&fixture.col(2)));
}
_ => panic!("expected Hash partitioning"),
}

Ok(())
}

#[test]
fn test_range_partitioning_adapt_multi_key() -> Result<()> {
let fixture = PartitioningTestFixture::new(vec![
("k1", DataType::Int32),
("k2", DataType::Utf8),
("t1", DataType::Int32),
("t2", DataType::Utf8),
])?;

let opt_k1 = SortOptions {
descending: true,
nulls_first: false,
};
let opt_k2 = SortOptions {
descending: false,
nulls_first: true,
};

let ordering = LexOrdering::new(vec![
fixture.range_sort_expr(0, opt_k1),
fixture.range_sort_expr(1, opt_k2),
])
.unwrap();

let split_points = vec![
SplitPoint::new(vec![ScalarValue::Int32(Some(20)), ScalarValue::Utf8(None)]),
SplitPoint::new(vec![
ScalarValue::Int32(Some(10)),
ScalarValue::Utf8(Some("foo".to_string())),
]),
];

let range = RangePartitioning::try_new(ordering, split_points.clone())?;
let adapted = range.adapt(&fixture.cols([2, 3]), &fixture.schema).unwrap();

assert_eq!(adapted.ordering().len(), 2);
assert!(adapted.ordering()[0].expr.eq(&fixture.col(2)));
assert_eq!(adapted.ordering()[0].options, opt_k1);
assert!(adapted.ordering()[1].expr.eq(&fixture.col(3)));
assert_eq!(adapted.ordering()[1].options, opt_k2);
assert_eq!(adapted.split_points(), &split_points);
assert_eq!(adapted.partition_count(), 3);

Ok(())
}
}

#[cfg(all(test, feature = "proto"))]
Expand Down
Loading