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
325 changes: 320 additions & 5 deletions datafusion/core/tests/physical_optimizer/enforce_distribution.rs
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@ use crate::physical_optimizer::test_utils::{
sort_merge_join_exec, sort_preserving_merge_exec, union_exec,
};

use arrow::array::{RecordBatch, UInt8Array, UInt64Array};
use arrow::array::{Int64Array, RecordBatch, UInt8Array, UInt64Array};
use arrow::compute::SortOptions;
use arrow_schema::{DataType, Field, Schema, SchemaRef};
use datafusion::config::ConfigOptions;
Expand All @@ -47,6 +47,7 @@ use datafusion_common::tree_node::{
};
use datafusion_datasource::file_groups::FileGroup;
use datafusion_datasource::file_scan_config::FileScanConfigBuilder;
use datafusion_datasource::memory::MemorySourceConfig;
use datafusion_expr::{JoinType, Operator};
use datafusion_functions_aggregate::count::count_udaf;
use datafusion_physical_expr::aggregate::AggregateExprBuilder;
Expand Down Expand Up @@ -76,11 +77,12 @@ use datafusion_physical_plan::joins::utils::JoinOn;
use datafusion_physical_plan::limit::{GlobalLimitExec, LocalLimitExec};
use datafusion_physical_plan::projection::{ProjectionExec, ProjectionExpr};
use datafusion_physical_plan::repartition::RepartitionExec;
use datafusion_physical_plan::sorts::sort::SortExec;
use datafusion_physical_plan::sorts::sort_preserving_merge::SortPreservingMergeExec;
use datafusion_physical_plan::union::{InterleaveExec, UnionExec};
use datafusion_physical_plan::{
ChildrenPropertiesMode, DisplayAs, DisplayFormatType, ExecutionPlanProperties,
PlanProperties, ReplaceChildrenOptions, displayable,
PlanProperties, ReplaceChildrenOptions, collect, displayable,
};
use insta::Settings;

Expand Down Expand Up @@ -4844,11 +4846,12 @@ fn test_replace_order_preserving_variants_with_fetch() -> Result<()> {
// Apply the function
let result = replace_order_preserving_variants(dist_context)?;

// Verify the plan was transformed to CoalescePartitionsExec
// A fetched ordered merge must still select the TopK rows.
let result = check_integrity(result)?;
result
.plan
.downcast_ref::<CoalescePartitionsExec>()
.expect("Expected CoalescePartitionsExec");
.downcast_ref::<SortExec>()
.expect("Expected a TopK SortExec");

// Verify fetch was preserved
assert_eq!(
Expand All @@ -4860,6 +4863,318 @@ fn test_replace_order_preserving_variants_with_fetch() -> Result<()> {
Ok(())
}

#[test]
fn preserve_fetch_when_reoptimizing_ordered_merge() -> Result<()> {
let schema = schema();
let sort_key: LexOrdering =
[PhysicalSortExpr::new_default(col("c", &schema)?)].into();
let input = parquet_exec_multiple_sorted(vec![sort_key.clone()]);
let plan: Arc<dyn ExecutionPlan> =
Arc::new(SortPreservingMergeExec::new(sort_key, input).with_fetch(Some(5)));

let optimized =
EnsureRequirements::new().optimize(plan, &test_suite_default_config_options())?;
let plan = displayable(optimized.as_ref()).indent(true).to_string();

assert!(
plan.contains("SortPreservingMergeExec: [c@2 ASC], fetch=5"),
"expected the optimizer to preserve fetch:\n{plan}"
);

Ok(())
}

#[test]
fn preserve_fetch_when_reoptimizing_coalesce_partitions() -> Result<()> {
let input = parquet_exec_multiple();
let plan: Arc<dyn ExecutionPlan> =
Arc::new(CoalescePartitionsExec::new(input).with_fetch(Some(5)));

let optimized =
EnsureRequirements::new().optimize(plan, &test_suite_default_config_options())?;

assert_eq!(optimized.fetch(), Some(5));
optimized
.downcast_ref::<CoalescePartitionsExec>()
.expect("expected CoalescePartitionsExec");

Ok(())
}

#[tokio::test]
async fn move_fetch_to_replacement_sort() -> Result<()> {
for (options, partitions, expected) in [
(
SortOptions::default(),
[
vec![None, Some(1), Some(1), Some(6)],
vec![None, Some(1), Some(2), Some(7)],
],
vec![None, None, Some(1), Some(1), Some(1)],
),
(
SortOptions {
descending: true,
nulls_first: false,
},
[vec![Some(7), Some(1), None], vec![Some(6), Some(1), None]],
vec![Some(7), Some(6), Some(1), Some(1), None],
),
] {
let (input, sort_key) = sorted_memory_input(partitions, options)?;
let merge: Arc<dyn ExecutionPlan> = Arc::new(
SortPreservingMergeExec::new(sort_key.clone(), input).with_fetch(Some(5)),
);
assert_eq!(fetch_test_values(Arc::clone(&merge)).await?, expected);
let plan = sort_required_exec_with_req(merge, sort_key);
let optimized = ensure_distribution_helper(plan, 10, false)?;
let replacement = Arc::clone(optimized.children()[0]);
let sort = replacement
.downcast_ref::<SortExec>()
.expect("expected a replacement sort");
assert_eq!(sort.fetch(), Some(5));
assert_eq!(fetch_test_values(replacement).await?, expected);
}
Ok(())
}

#[tokio::test]
async fn preserve_fetch_in_nested_distribution_operators() -> Result<()> {
for outer_fetch in [0, 3, 10] {
let (input, sort_key) = sorted_memory_input(
[0, 1].map(|start| (start..10).step_by(2).map(Some).collect()),
SortOptions::default(),
)?;
let merge: Arc<dyn ExecutionPlan> =
Arc::new(SortPreservingMergeExec::new(sort_key, input).with_fetch(Some(5)));
let plan: Arc<dyn ExecutionPlan> =
Arc::new(CoalescePartitionsExec::new(merge).with_fetch(Some(outer_fetch)));
let expected = (0..outer_fetch.min(5))
.map(|value| Some(value as i64))
.collect::<Vec<_>>();
assert_reoptimized_fetch_values(plan, &expected).await?;
}
Ok(())
}

#[tokio::test]
async fn preserve_topk_when_parent_changes_ordering() -> Result<()> {
let (input, sort_key) = sorted_memory_input(
[0, 1].map(|start| (start..10).step_by(2).map(Some).collect()),
SortOptions::default(),
)?;
let descending = [PhysicalSortExpr::new(
col("c", &input.schema())?,
SortOptions {
descending: true,
nulls_first: false,
},
)]
.into();
let merge: Arc<dyn ExecutionPlan> =
Arc::new(SortPreservingMergeExec::new(sort_key, input).with_fetch(Some(5)));
let plan: Arc<dyn ExecutionPlan> = Arc::new(SortExec::new(descending, merge));
assert_reoptimized_fetch_values(plan, &[Some(4), Some(3), Some(2), Some(1), Some(0)])
.await
}

#[tokio::test]
async fn preserve_fetch_when_parallelizing_sort_above_filter() -> Result<()> {
let (input, sort_key) = sorted_memory_input(
[
vec![Some(-4), Some(-2), Some(2), Some(4), Some(6)],
vec![Some(-3), Some(-1), Some(3), Some(5), Some(7)],
],
SortOptions::default(),
)?;
let predicate = Arc::new(BinaryExpr::new(
col("c", &input.schema())?,
Operator::Gt,
lit(0_i64),
));
let coalesce: Arc<dyn ExecutionPlan> =
Arc::new(CoalescePartitionsExec::new(input).with_fetch(Some(5)));
let filter: Arc<dyn ExecutionPlan> =
Arc::new(FilterExec::try_new(predicate, coalesce)?);
let mut plan: Arc<dyn ExecutionPlan> = Arc::new(SortExec::new(sort_key, filter));
let mut config = test_suite_default_config_options();
config.optimizer.enable_round_robin_repartition = false;
config.optimizer.repartition_sorts = true;
for iteration in 0..3 {
if iteration > 0 {
plan = EnsureRequirements::new().optimize(plan, &config)?;
}
// Either input batch can arrive first. Both contain three positive
// rows, so keeping the limit below the filter always returns three.
assert_eq!(
fetch_test_values(Arc::clone(&plan)).await?.len(),
3,
"iteration {iteration}:\n{}",
displayable(plan.as_ref()).indent(true)
);
}
Ok(())
}

fn sorted_memory_input(
partitions: [Vec<Option<i64>>; 2],
options: SortOptions,
) -> Result<(Arc<dyn ExecutionPlan>, LexOrdering)> {
let schema = Arc::new(Schema::new(vec![Field::new("c", DataType::Int64, true)]));
let order: LexOrdering = [PhysicalSortExpr::new(col("c", &schema)?, options)].into();
let partitions = partitions
.into_iter()
.map(|values| {
RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(values))],
)
.map(|batch| vec![batch])
})
.collect::<std::result::Result<Vec<_>, _>>()?;
let source = MemorySourceConfig::try_new(&partitions, schema, None)?
.try_with_sort_information(vec![order.clone()])?;
Ok((DataSourceExec::from_data_source(source), order))
}

async fn fetch_test_values(plan: Arc<dyn ExecutionPlan>) -> Result<Vec<Option<i64>>> {
let batches = collect(plan, SessionContext::new().task_ctx()).await?;
Ok(batches
.iter()
.flat_map(|batch| {
batch
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap()
.iter()
})
.collect())
}

async fn assert_reoptimized_fetch_values(
plan: Arc<dyn ExecutionPlan>,
expected: &[Option<i64>],
) -> Result<()> {
for repartition_sorts in [false, true] {
let mut optimized = Arc::clone(&plan);
let mut config = test_suite_default_config_options();
config.optimizer.enable_round_robin_repartition = false;
config.optimizer.repartition_sorts = repartition_sorts;
for iteration in 0..3 {
if iteration > 0 {
let distribution =
DistributionContext::new_default(Arc::clone(&optimized))
.transform_up(|context| ensure_distribution(context, &config))?
.data;
check_integrity(distribution)?;
optimized = EnsureRequirements::new().optimize(optimized, &config)?;
}
assert_eq!(
fetch_test_values(Arc::clone(&optimized)).await?,
expected,
"iteration {iteration}, repartition_sorts={repartition_sorts}:\n{}",
displayable(optimized.as_ref()).indent(true)
);
}
}
Ok(())
}

#[tokio::test]
async fn preserve_fetch_below_filter_when_reoptimizing() -> Result<()> {
check_fetch_below_filter(
Operator::Gt,
[vec![-2, 0, 2, 4], vec![-1, 1, 3, 5]],
&[1, 2],
)
.await
}

#[tokio::test]
async fn preserve_fetch_below_filter_with_constant_ordering() -> Result<()> {
check_fetch_below_filter(
Operator::Eq,
[vec![-2, 0, 0, 0], vec![-1, 0, 0, 0]],
&[0, 0, 0],
)
.await
}

async fn check_fetch_below_filter(
op: Operator,
partitions: [Vec<i64>; 2],
expected: &[i64],
) -> Result<()> {
let schema = Arc::new(Schema::new(vec![Field::new("c", DataType::Int64, false)]));
let sort_key: LexOrdering =
[PhysicalSortExpr::new_default(col("c", &schema)?)].into();
let partitions = partitions
.into_iter()
.map(|values| {
RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(Int64Array::from(values))],
)
.map(|batch| vec![batch])
})
.collect::<std::result::Result<Vec<_>, _>>()?;
let source = MemorySourceConfig::try_new(&partitions, Arc::clone(&schema), None)?
.try_with_sort_information(vec![sort_key.clone()])?;
let merge: Arc<dyn ExecutionPlan> = Arc::new(
SortPreservingMergeExec::new(
sort_key.clone(),
DataSourceExec::from_data_source(source),
)
.with_fetch(Some(5)),
);
let predicate = Arc::new(BinaryExpr::new(col("c", &schema)?, op, lit(0_i64)));
let filter: Arc<dyn ExecutionPlan> = Arc::new(FilterExec::try_new(predicate, merge)?);
let mut plan = sort_required_exec_with_req(filter, sort_key);
let mut config = test_suite_default_config_options();
config.optimizer.enable_round_robin_repartition = false;
let task_context = SessionContext::new().task_ctx();

// The test operator only declares ordering requirements. Execute its child
// to compare query results before optimization and after repeated passes.
for iteration in 0..3 {
if iteration > 0 {
let distribution = DistributionContext::new_default(Arc::clone(&plan))
.transform_up(|context| ensure_distribution(context, &config))?
.data;
check_integrity(distribution)?;
plan = EnsureRequirements::new().optimize(plan, &config)?;
}
let input = Arc::clone(plan.children()[0]);
let batches = collect(input, Arc::clone(&task_context)).await?;
let values = batches
.iter()
.flat_map(|batch| {
batch
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap()
.values()
.iter()
.copied()
})
.collect::<Vec<_>>();
assert_eq!(
values,
expected,
"iteration {iteration}:\n{}",
displayable(plan.as_ref()).indent(true)
);
assert!(
plan.children()[0].is::<FilterExec>(),
"fetch must stay below the filter:\n{}",
displayable(plan.as_ref()).indent(true)
);
}
Ok(())
}

/// When a parent requires SinglePartition and maintains input order, order-preserving
/// variants (e.g. SortPreservingMergeExec) should be kept so that ordering can
/// propagate to ancestors. Replacing them with CoalescePartitionsExec would destroy
Expand Down
4 changes: 3 additions & 1 deletion datafusion/core/tests/physical_optimizer/enforce_sorting.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2327,7 +2327,9 @@ async fn test_remove_unnecessary_spm2() -> Result<()> {
DataSourceExec: partitions=1, partition_sizes=[0]

Optimized Plan:
DataSourceExec: partitions=1, partition_sizes=[0]
LocalLimitExec: fetch=100
SortExec: expr=[non_nullable_col@1 ASC], preserve_partitioning=[false]
DataSourceExec: partitions=1, partition_sizes=[0]
");

Ok(())
Expand Down
Loading