diff --git a/datafusion/core/tests/physical_optimizer/enforce_distribution.rs b/datafusion/core/tests/physical_optimizer/enforce_distribution.rs index 4aeee422394ca..9df5aa1f7e22f 100644 --- a/datafusion/core/tests/physical_optimizer/enforce_distribution.rs +++ b/datafusion/core/tests/physical_optimizer/enforce_distribution.rs @@ -26,16 +26,18 @@ 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::{ArrayRef, Int32Array, RecordBatch, UInt8Array, UInt64Array}; use arrow::compute::SortOptions; use arrow_schema::{DataType, Field, Schema, SchemaRef}; use datafusion::config::ConfigOptions; use datafusion::datasource::MemTable; use datafusion::datasource::file_format::file_compression_type::FileCompressionType; use datafusion::datasource::listing::PartitionedFile; +use datafusion::datasource::memory::MemorySourceConfig; use datafusion::datasource::object_store::ObjectStoreUrl; use datafusion::datasource::physical_plan::{CsvSource, ParquetSource}; -use datafusion::datasource::source::DataSourceExec; +use datafusion::datasource::source::{DataSource, DataSourceExec}; +use datafusion::physical_planner::DefaultPhysicalPlanner; use datafusion::prelude::{SessionConfig, SessionContext}; use datafusion_common::ScalarValue; use datafusion_common::Statistics; @@ -80,7 +82,7 @@ use datafusion_physical_plan::sorts::sort_preserving_merge::SortPreservingMergeE use datafusion_physical_plan::union::{InterleaveExec, UnionExec}; use datafusion_physical_plan::{ ChildrenPropertiesMode, DisplayAs, DisplayFormatType, ExecutionPlanProperties, - PlanProperties, ReplaceChildrenOptions, displayable, + PlanProperties, ReplaceChildrenOptions, collect_partitioned, displayable, }; use insta::Settings; @@ -4794,6 +4796,98 @@ async fn test_distribute_sort_memtable() -> Result<()> { Ok(()) } +/// The full physical optimizer must not lose a source fetch while increasing source parallelism. +#[tokio::test] +async fn memory_source_repartition_preserves_fetch_through_optimizer() -> Result<()> { + let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])); + let batches = vec![vec![ + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![0, 1])) as ArrayRef, + )])?, + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![2, 3])) as ArrayRef, + )])?, + ]]; + let source = MemorySourceConfig::try_new(&batches, Arc::clone(&schema), None)? + .with_limit(Some(1)); + assert_eq!(source.fetch(), Some(1)); + let source = DataSourceExec::from_data_source(source); + let predicate = Arc::new(BinaryExpr::new( + col("a", &schema)?, + Operator::GtEq, + Arc::new(Literal::new(ScalarValue::Int32(Some(0)))), + )); + let original: Arc = + Arc::new(FilterExec::try_new(predicate, source)?); + + let session_config = SessionConfig::new() + .with_target_partitions(2) + .with_batch_size(1) + .with_repartition_file_scans(true); + let ctx = SessionContext::new_with_config(session_config); + let state = ctx.state(); + let config = state.config_options(); + + let values = |partitions: Vec>| { + partitions + .iter() + .flatten() + .flat_map(|batch| { + batch + .column(0) + .as_any() + .downcast_ref::() + .unwrap() + .values() + .iter() + .copied() + }) + .collect::>() + }; + let original_values = + values(collect_partitioned(Arc::clone(&original), ctx.task_ctx()).await?); + + // Exercise EnsureRequirements independently from the default full pipeline. + let ensure = EnsureRequirements::new().optimize(Arc::clone(&original), config)?; + let ensure_values = + values(collect_partitioned(Arc::clone(&ensure), ctx.task_ctx()).await?); + let ensure_plan = displayable(ensure.as_ref()).indent(true).to_string(); + + // Run the actual default physical-optimizer pipeline from the original plan, + // then run it once more to cover a subsequent physical-optimization pass. + let planner = DefaultPhysicalPlanner::default(); + let full = + planner.optimize_physical_plan(Arc::clone(&original), &state, |_, _| {})?; + let full_values = + values(collect_partitioned(Arc::clone(&full), ctx.task_ctx()).await?); + let full_plan = displayable(full.as_ref()).indent(true).to_string(); + + let full_second = planner.optimize_physical_plan(full, &state, |_, _| {})?; + let full_second_values = + values(collect_partitioned(Arc::clone(&full_second), ctx.task_ctx()).await?); + let full_second_plan = displayable(full_second.as_ref()).indent(true).to_string(); + + assert_eq!(original_values, [0]); + for plan in [&ensure_plan, &full_plan, &full_second_plan] { + assert!( + plan.contains( + "RepartitionExec: partitioning=RoundRobinBatch(2), input_partitions=1" + ), + "expected ordinary repartition fallback:\n{plan}" + ); + assert!( + plan.contains("DataSourceExec: partitions=1") && plan.contains("fetch=1"), + "source partition count and fetch must be retained:\n{plan}" + ); + } + assert_eq!(ensure_values, [0]); + assert_eq!(full_values, [0]); + assert_eq!(full_second_values, [0]); + Ok(()) +} + /// Create a [`MemTable`] with 100 batches of 8192 rows each, in 1 partition fn create_memtable() -> Result { let mut batches = Vec::with_capacity(100); diff --git a/datafusion/datasource/src/memory.rs b/datafusion/datasource/src/memory.rs index 4c79cf4a9851d..682e107398ac1 100644 --- a/datafusion/datasource/src/memory.rs +++ b/datafusion/datasource/src/memory.rs @@ -159,6 +159,11 @@ impl DataSource for MemorySourceConfig { _repartition_file_min_size: usize, output_ordering: Option, ) -> Result>> { + // Fetch is an original per-partition cap, which repartitioning would change. + if self.fetch.is_some() { + return Ok(None); + } + if self.partitions.is_empty() || self.partitions.len() >= target_partitions // if have no partitions, or already have more partitions than desired, do not repartition { @@ -1055,6 +1060,7 @@ mod tests { use crate::tests::{aggr_test_schema, make_partition}; use arrow::array::{ArrayRef, Int32Array, Int64Array, StringArray}; + use arrow::compute::SortOptions; use arrow::datatypes::{DataType, Field}; use datafusion_common::assert_batches_eq; use datafusion_common::stats::{ColumnStatistics, Precision}; @@ -1062,7 +1068,9 @@ mod tests { use datafusion_physical_plan::expressions::lit; use datafusion_physical_plan::statistics::{StatisticsArgs, StatisticsContext}; - use datafusion_physical_plan::ExecutionPlan; + use datafusion_physical_plan::{ + ExecutionPlan, ExecutionPlanProperties, collect, collect_partitioned, + }; #[tokio::test] async fn exec_with_limit() -> Result<()> { @@ -1090,6 +1098,182 @@ mod tests { Ok(()) } + #[tokio::test] + async fn repartition_refuses_existing_fetch() -> Result<()> { + let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])); + let batches = vec![ + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![0, 1])) as ArrayRef, + )])?, + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![2, 3])) as ArrayRef, + )])?, + ]; + let task_ctx = Arc::new(TaskContext::default()); + + let uncapped = + MemorySourceConfig::try_new_from_batches(Arc::clone(&schema), batches)?; + let uncapped_repartitioned = uncapped + .repartitioned(2, &datafusion_common::config::ConfigOptions::default())? + .expect("uncapped source can be repartitioned"); + assert_eq!( + uncapped_repartitioned + .output_partitioning() + .partition_count(), + 2 + ); + let uncapped_rows: usize = + collect_partitioned(uncapped_repartitioned, Arc::clone(&task_ctx)) + .await? + .iter() + .flatten() + .map(RecordBatch::num_rows) + .sum(); + assert_eq!(uncapped_rows, 4); + + let capped = MemorySourceConfig::try_new_from_batches( + Arc::clone(&schema), + vec![ + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![0, 1])) as ArrayRef, + )])?, + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![2, 3])) as ArrayRef, + )])?, + ], + )? + .with_fetch(Some(0)) + .unwrap(); + assert!( + capped + .repartitioned(2, &datafusion_common::config::ConfigOptions::default())? + .is_none() + ); + assert!(collect(capped, Arc::clone(&task_ctx)).await?.is_empty()); + + let partitions = vec![ + vec![ + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![0, 1])) as ArrayRef, + )])?, + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![2, 3])) as ArrayRef, + )])?, + ], + vec![ + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![10, 11])) as ArrayRef, + )])?, + RecordBatch::try_from_iter(vec![( + "a", + Arc::new(Int32Array::from(vec![12, 13])) as ArrayRef, + )])?, + ], + ]; + let capped = DataSourceExec::from_data_source( + MemorySourceConfig::try_new(&partitions, schema, None)?.with_limit(Some(1)), + ); + assert!( + capped + .repartitioned(3, &datafusion_common::config::ConfigOptions::default())? + .is_none() + ); + let values: Vec = collect_partitioned(capped, task_ctx) + .await? + .iter() + .flatten() + .flat_map(|batch| { + batch + .column(0) + .as_any() + .downcast_ref::() + .unwrap() + .values() + .iter() + .copied() + }) + .collect(); + assert_eq!(values, [0, 10]); + Ok(()) + } + + #[tokio::test] + async fn repartition_refuses_fetch_with_projection_and_sort_metadata() -> Result<()> { + let schema = Arc::new(Schema::new(vec![ + Field::new("a", DataType::Int32, false), + Field::new("b", DataType::Int32, false), + ])); + let partitions = vec![vec![ + RecordBatch::try_from_iter(vec![ + ("a", Arc::new(Int32Array::from(vec![1, 2])) as ArrayRef), + ("b", Arc::new(Int32Array::from(vec![10, 20])) as ArrayRef), + ])?, + RecordBatch::try_from_iter(vec![ + ("a", Arc::new(Int32Array::from(vec![3, 4])) as ArrayRef), + ("b", Arc::new(Int32Array::from(vec![30, 40])) as ArrayRef), + ])?, + ]]; + let ordering: LexOrdering = [PhysicalSortExpr { + expr: col("b", &schema)?, + options: SortOptions::default(), + }] + .into(); + let projected_schema = + Arc::new(Schema::new(vec![Field::new("b", DataType::Int32, false)])); + let projected_ordering: LexOrdering = [PhysicalSortExpr { + expr: col("b", &projected_schema)?, + options: SortOptions::default(), + }] + .into(); + let source = + MemorySourceConfig::try_new(&partitions, Arc::clone(&schema), Some(vec![1]))? + .try_with_sort_information(vec![ordering])? + .with_limit(Some(1)); + + assert_eq!(source.sort_information(), from_ref(&projected_ordering)); + assert!(source.repartitioned(2, usize::MAX, None)?.is_none()); + assert!( + source + .repartitioned(2, usize::MAX, source.sort_information().first().cloned())? + .is_none() + ); + + let projected = DataSourceExec::from_data_source(source); + assert_eq!( + projected.schema(), + Arc::new(Schema::new(vec![Field::new("b", DataType::Int32, false)])) + ); + assert_eq!( + projected.properties().output_ordering(), + Some(&projected_ordering) + ); + let batches = + collect_partitioned(projected, Arc::new(TaskContext::default())).await?; + let values: Vec = batches + .iter() + .flatten() + .flat_map(|batch| { + batch + .column(0) + .as_any() + .downcast_ref::() + .unwrap() + .values() + .iter() + .copied() + }) + .collect(); + assert_eq!(values, [10]); + Ok(()) + } + /// Test that `try_swapping_with_projection` preserves the `fetch` limit. /// Regression test for #[test]