diff --git a/datafusion/core/tests/physical_optimizer/filter_pushdown.rs b/datafusion/core/tests/physical_optimizer/filter_pushdown.rs index a98c1b7bcf98b..340e78d161f45 100644 --- a/datafusion/core/tests/physical_optimizer/filter_pushdown.rs +++ b/datafusion/core/tests/physical_optimizer/filter_pushdown.rs @@ -18,11 +18,13 @@ use std::sync::{Arc, LazyLock}; use arrow::{ - array::{RecordBatch, record_batch}, - datatypes::{DataType, Field, Schema, SchemaRef}, + array::{RecordBatch, TimestampNanosecondArray, record_batch}, + datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit}, util::pretty::pretty_format_batches, }; use arrow_schema::SortOptions; +use bytes::Bytes; +use datafusion::datasource::physical_plan::ParquetSource; use datafusion::{ assert_batches_eq, logical_expr::Operator, @@ -36,6 +38,7 @@ use datafusion::{ use datafusion_catalog::memory::DataSourceExec; use datafusion_common::{ JoinType, + arrow::datatypes::IntervalMonthDayNano, config::ConfigOptions, tree_node::{TreeNode, TreeNodeRecursion}, }; @@ -44,6 +47,7 @@ use datafusion_datasource::{ }; use datafusion_execution::object_store::ObjectStoreUrl; use datafusion_expr::ScalarUDF; +use datafusion_expr::physical_planning_context::{ScalarSubqueryResults, SubqueryIndex}; use datafusion_functions::math::random::RandomFunc; use datafusion_functions_aggregate::{ count::count_udaf, @@ -52,6 +56,7 @@ use datafusion_functions_aggregate::{ use datafusion_physical_expr::{ LexOrdering, PhysicalSortExpr, expressions::{DynamicFilterPhysicalExpr, col}, + scalar_subquery::ScalarSubqueryExpr, utils::conjunction, }; use datafusion_physical_expr::{ @@ -61,6 +66,7 @@ use datafusion_physical_expr::{ use datafusion_physical_optimizer::{ PhysicalOptimizerRule, filter_pushdown::FilterPushdown, }; +use datafusion_physical_plan::scalar_subquery::{ScalarSubqueryExec, ScalarSubqueryLink}; use datafusion_physical_plan::{ ExecutionPlan, aggregates::{AggregateExec, AggregateMode, PhysicalGroupBy}, @@ -72,6 +78,8 @@ use datafusion_physical_plan::{ repartition::RepartitionExec, sorts::sort::SortExec, }; +use object_store::{ObjectStoreExt, path::Path}; +use parquet::{arrow::ArrowWriter, file::properties::WriterProperties}; use super::pushdown_utils::{ OptimizationTest, TestNode, TestScanBuilder, TestSource, format_plan_for_test, @@ -3064,6 +3072,174 @@ fn test_hashjoin_dynamic_filter_pushdown_is_used() { } } +/// A completed scalar-subquery filter is retained for Parquet statistics pruning even +/// when Parquet row-filter pushdown is disabled. +#[tokio::test] +async fn scalar_subquery_dynamic_filter_parquet_pruning_with_pushdown_disabled() { + let timestamp = DataType::Timestamp(TimeUnit::Nanosecond, None); + let schema = Arc::new(Schema::new(vec![Field::new( + "ts", + timestamp.clone(), + false, + )])); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(TimestampNanosecondArray::from([0, 2_000_000_000]))], + ) + .unwrap(); + + let mut parquet_bytes = Vec::new(); + let properties = WriterProperties::builder() + .set_max_row_group_row_count(Some(1)) + .build(); + let mut writer = + ArrowWriter::try_new(&mut parquet_bytes, Arc::clone(&schema), Some(properties)) + .unwrap(); + writer.write(&batch).unwrap(); + writer.close().unwrap(); + + let object_store = Arc::new(InMemory::new()); + let path = Path::from("scalar-subquery.parquet"); + let file_size = parquet_bytes.len() as u64; + object_store + .put(&path, Bytes::from(parquet_bytes).into()) + .await + .unwrap(); + let scan = DataSourceExec::from_data_source( + FileScanConfigBuilder::new( + ObjectStoreUrl::parse("test://").unwrap(), + Arc::new(ParquetSource::new(Arc::clone(&schema))), + ) + .with_file(PartitionedFile::new("scalar-subquery.parquet", file_size)) + .build(), + ) as Arc; + + let results = ScalarSubqueryResults::new(1); + let producer_batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(TimestampNanosecondArray::from([1_500_000_000]))], + ) + .unwrap(); + let producer = datafusion::datasource::memory::MemorySourceConfig::try_new_exec( + &[vec![producer_batch]], + Arc::clone(&schema), + None, + ) + .unwrap(); + let scalar = Arc::new(ScalarSubqueryExpr::new( + timestamp, + false, + SubqueryIndex::new(0), + results.clone(), + )); + let predicate = Arc::new(BinaryExpr::new( + col("ts", &schema).unwrap(), + Operator::GtEq, + Arc::new(BinaryExpr::new( + scalar, + Operator::Minus, + Arc::new(Literal::new(ScalarValue::IntervalMonthDayNano(Some( + IntervalMonthDayNano { + months: 0, + days: 0, + nanoseconds: 500_000_000, + }, + )))), + )), + )) as Arc; + let plan = Arc::new(ScalarSubqueryExec::new( + Arc::new(FilterExec::try_new(predicate, scan).unwrap()), + vec![ScalarSubqueryLink { + plan: producer, + index: SubqueryIndex::new(0), + }], + results, + )) as Arc; + + let mut config = ConfigOptions::default(); + config.execution.parquet.pushdown_filters = false; + config.execution.parquet.pruning = true; + config.optimizer.enable_dynamic_filter_pushdown = true; + let optimized = FilterPushdown::new_post_optimization() + .optimize(plan, &config) + .unwrap(); + let scalar_exec = optimized + .downcast_ref::() + .expect("optimized root should retain ScalarSubqueryExec"); + assert_eq!(scalar_exec.subqueries().len(), 1); + let expression_id = scalar_exec + .dynamic_expressions_produced() + .first() + .expect("scalar subquery producer should retain its dynamic filter") + .expression_id() + .expect("dynamic filter should have an expression ID"); + assert!(scalar_exec.input().downcast_ref::().is_some()); + + let mut scan_consumer = None; + scalar_exec + .input() + .apply(|node| { + if let Some(scan) = node.downcast_ref::() + && let Some((_, parquet)) = + scan.downcast_to_file_source::() + && let Some(predicate) = parquet.filter() + { + predicate.apply(|expr| { + if expr.expression_id() == Some(expression_id) { + scan_consumer = Some(Arc::clone(expr)); + Ok(TreeNodeRecursion::Stop) + } else { + Ok(TreeNodeRecursion::Continue) + } + })?; + } + Ok(TreeNodeRecursion::Continue) + }) + .unwrap(); + let scan_consumer = + scan_consumer.expect("Parquet predicate should retain the dynamic filter ID"); + + let context = SessionContext::new_with_config(SessionConfig::from(config)); + context.register_object_store( + ObjectStoreUrl::parse("test://").unwrap().as_ref(), + object_store, + ); + let batches = collect(optimized.clone(), context.task_ctx()) + .await + .unwrap(); + assert_batches_eq!( + &[ + "+---------------------+", + "| ts |", + "+---------------------+", + "| 1970-01-01T00:00:02 |", + "+---------------------+", + ], + &batches + ); + + let consumer = scan_consumer + .downcast_ref::() + .expect("scan consumer should be a dynamic filter"); + assert_ne!(consumer.current().unwrap().to_string(), "true"); + consumer.wait_complete().await; + + let metrics = + datafusion::test_util::parquet::TestParquetFile::parquet_metrics(&optimized) + .expect("Parquet scan should have metrics"); + let pruned = match metrics.sum_by_name("row_groups_pruned_statistics").unwrap() { + datafusion::physical_plan::metrics::MetricValue::PruningMetrics { + pruning_metrics, + .. + } => pruning_metrics.pruned(), + metric => panic!("unexpected statistics pruning metric: {metric:?}"), + }; + assert!( + pruned >= 1, + "expected statistics pruning, observed {pruned}" + ); +} + /// Regression test for https://github.com/apache/datafusion/issues/20109. /// /// Not portable to sqllogictest: the regression specifically targets the diff --git a/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs b/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs index eb3d457de82ad..a712677c4c0b9 100644 --- a/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs +++ b/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs @@ -183,6 +183,16 @@ impl Display for DynamicFilterPhysicalExpr { } impl DynamicFilterPhysicalExpr { + /// Returns `true` if `other` shares this filter's runtime state. + /// + /// Derived filters can have distinct outer objects while sharing the same + /// state. Conversely, filters reconstructed independently (for example by + /// decoding) do not share state even when their expression IDs match. + #[doc(hidden)] + pub fn shares_runtime_state(&self, other: &Self) -> bool { + Arc::ptr_eq(&self.inner, &other.inner) + } + /// Create a new [`DynamicFilterPhysicalExpr`] /// from an initial expression and a list of children. /// The list of children is provided separately because @@ -314,15 +324,7 @@ impl DynamicFilterPhysicalExpr { /// - When we've computed the probe side's hash table in a HashJoinExec /// - After every batch is processed if we update the TopK heap in a SortExec using a TopK approach. pub fn update(&self, new_expr: Arc) -> Result<()> { - // Remap the children of the new expression to match the original children - // We still do this again in `current()` but doing it preventively here - // reduces the work needed in some cases if `current()` is called multiple times - // and the same externally facing `PhysicalExpr` is used for both `with_new_children` and `update()`.` - let new_expr = Self::remap_children( - &self.children, - self.remapped_children.as_ref(), - new_expr, - )?; + // Store the canonical expression; `current()` remaps it for each derived consumer. // Load the current inner, increment generation, and store the new one let mut current = self.inner.write(); diff --git a/datafusion/physical-expr/src/scalar_subquery.rs b/datafusion/physical-expr/src/scalar_subquery.rs index 473b52a5cb45c..4f1a837300cce 100644 --- a/datafusion/physical-expr/src/scalar_subquery.rs +++ b/datafusion/physical-expr/src/scalar_subquery.rs @@ -29,6 +29,8 @@ use datafusion_expr_common::columnar_value::ColumnarValue; use datafusion_expr_common::sort_properties::{ExprProperties, SortProperties}; use datafusion_physical_expr_common::physical_expr::PhysicalExpr; +use crate::expressions::Literal; + /// A physical expression whose value is provided by a scalar subquery. /// /// Subquery execution is handled by `ScalarSubqueryExec`, which stores the @@ -133,6 +135,15 @@ impl PhysicalExpr for ScalarSubqueryExpr { Ok(ColumnarValue::Scalar(value)) } + fn snapshot(&self) -> Result>> { + let value = self.results.get(self.index).ok_or_else(|| { + internal_datafusion_err!( + "ScalarSubqueryExpr snapshotted before the subquery was executed" + ) + })?; + Ok(Some(Arc::new(Literal::new(value)))) + } + fn children(&self) -> Vec<&Arc> { vec![] } @@ -274,6 +285,56 @@ mod tests { assert!(result.is_err()); } + #[test] + fn test_snapshot_pending_and_populated_values() -> Result<()> { + let pending_results = ScalarSubqueryResults::new(1); + let pending = ScalarSubqueryExpr::new( + DataType::Int32, + true, + SubqueryIndex::new(0), + pending_results, + ); + assert!(pending.snapshot().is_err()); + + let results = ScalarSubqueryResults::new(2); + let non_null = ScalarSubqueryExpr::new( + DataType::Int32, + false, + SubqueryIndex::new(0), + results.clone(), + ); + let null = ScalarSubqueryExpr::new( + DataType::Utf8, + true, + SubqueryIndex::new(1), + results.clone(), + ); + results.set(SubqueryIndex::new(0), ScalarValue::Int32(Some(42)))?; + results.set(SubqueryIndex::new(1), ScalarValue::Utf8(None))?; + + let non_null_snapshot = non_null + .snapshot()? + .expect("populated scalar subquery should snapshot"); + assert_eq!( + non_null_snapshot + .downcast_ref::() + .expect("snapshot should be a literal") + .value(), + &ScalarValue::Int32(Some(42)) + ); + let null_snapshot = null + .snapshot()? + .expect("populated scalar subquery should snapshot"); + assert_eq!( + null_snapshot + .downcast_ref::() + .expect("snapshot should be a literal") + .value(), + &ScalarValue::Utf8(None) + ); + Ok(()) + } + #[test] fn test_identity_equality() { let results = make_results(vec![None, None]); diff --git a/datafusion/physical-plan/src/proto.rs b/datafusion/physical-plan/src/proto.rs index 7640d76c3e010..57fd35fdb83b0 100644 --- a/datafusion/physical-plan/src/proto.rs +++ b/datafusion/physical-plan/src/proto.rs @@ -135,6 +135,15 @@ pub trait ExecutionPlanDecode { input_schema: &Schema, ) -> Result>; + /// Deserialize a physical expression with `results` active for scalar + /// subquery expressions in that expression's subtree. + fn decode_expr_with_scalar_subquery_results( + &self, + node: &PhysicalExprNode, + input_schema: &Schema, + results: ScalarSubqueryResults, + ) -> Result>; + /// The session task context, used by plans that need the function registry /// or session configuration. Never exposes the proto extension codec. fn task_ctx(&self) -> &TaskContext; @@ -294,6 +303,18 @@ impl<'a> ExecutionPlanDecodeCtx<'a> { self.decoder.decode_expr(node, input_schema) } + /// Deserialize a physical expression with scalar subquery results scoped + /// to this ScalarSubqueryExec. + pub fn decode_expr_with_scalar_subquery_results( + &self, + node: &PhysicalExprNode, + input_schema: &Schema, + results: ScalarSubqueryResults, + ) -> Result> { + self.decoder + .decode_expr_with_scalar_subquery_results(node, input_schema, results) + } + /// Deserialize a required physical expression against `input_schema`. pub fn decode_required_expr( &self, diff --git a/datafusion/physical-plan/src/scalar_subquery.rs b/datafusion/physical-plan/src/scalar_subquery.rs index f2b7c5e0b53e9..5ed0eb6608982 100644 --- a/datafusion/physical-plan/src/scalar_subquery.rs +++ b/datafusion/physical-plan/src/scalar_subquery.rs @@ -25,15 +25,30 @@ //! [`ScalarSubqueryExpr`]: datafusion_physical_expr::scalar_subquery::ScalarSubqueryExpr use std::fmt; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; -use datafusion_common::tree_node::TreeNodeRecursion; -use datafusion_common::{Result, ScalarValue, Statistics, exec_err, internal_err}; +use datafusion_common::config::ConfigOptions; +use datafusion_common::tree_node::{TreeNode, TreeNodeRecursion}; +use datafusion_common::{ + DataFusionError, Result, ScalarValue, Statistics, exec_err, internal_err, +}; use datafusion_execution::TaskContext; use datafusion_expr::physical_planning_context::{ScalarSubqueryResults, SubqueryIndex}; -use datafusion_physical_expr::PhysicalExpr; +use datafusion_physical_expr::expressions::{ + BinaryExpr, Column, DynamicFilterPhysicalExpr, Literal, lit, +}; +use datafusion_physical_expr::scalar_subquery::ScalarSubqueryExpr; +use datafusion_physical_expr_common::physical_expr::{ + PhysicalExpr, snapshot_physical_expr, +}; -use crate::execution_plan::{CardinalityEffect, ExecutionPlan, PlanProperties}; +use crate::execution_plan::{ + CardinalityEffect, ExecutionPlan, PlanProperties, plan_contains_expression_id, +}; +use crate::filter_pushdown::{ + ChildFilterDescription, ChildPushdownResult, FilterDescription, FilterPushdownPhase, + FilterPushdownPropagation, PushedDown, +}; use crate::joins::utils::{OnceAsync, OnceFut}; use crate::statistics::{ChildStats, StatisticsArgs}; use crate::stream::RecordBatchStreamAdapter; @@ -92,10 +107,25 @@ pub struct ScalarSubqueryExec { /// Shared results container; the corresponding `ScalarSubqueryExpr` /// nodes in the input plan hold the same underlying container. results: ScalarSubqueryResults, + /// Dynamic filters associated with predicates in the main input. + /// This state is separate from the input so it survives child replacement. + bindings: Arc>>, /// Cached plan properties (copied from input). cache: Arc, } +#[derive(Debug, Clone)] +struct ScalarSubqueryBinding { + /// The predicate produced by this scalar subquery. + predicate: Arc, + /// Every concrete dynamic-filter consumer for this binding's ID. + /// + /// Optimized plans commonly share one filter instance between the + /// producer and its consumers. Proto decoding, however, can produce + /// multiple independent filter instances with the same expression ID. + consumers: Vec>, +} + impl ScalarSubqueryExec { pub fn new( input: Arc, @@ -108,6 +138,7 @@ impl ScalarSubqueryExec { subqueries, subquery_future: Arc::default(), results, + bindings: Arc::default(), cache, } } @@ -131,6 +162,61 @@ impl ScalarSubqueryExec { .chain(std::iter::repeat_n(false, self.subqueries.len())) .collect() } + + fn discover_binding( + &self, + predicate: &Arc, + ) -> Option { + let binary = predicate.downcast_ref::()?; + if *binary.op() != datafusion_expr::Operator::GtEq { + return None; + } + let left = binary.left().downcast_ref::()?; + let subtraction = binary.right().downcast_ref::()?; + if *subtraction.op() != datafusion_expr::Operator::Minus + || subtraction.right().downcast_ref::().is_none() + { + return None; + } + let scalar = subtraction.left().downcast_ref::()?; + if !ScalarSubqueryResults::ptr_eq(scalar.results(), &self.results) { + return None; + } + let filter = Arc::new(DynamicFilterPhysicalExpr::new( + vec![Arc::new(left.clone()) as Arc], + lit(true), + )); + Some(ScalarSubqueryBinding { + predicate: Arc::clone(predicate), + consumers: vec![filter as Arc], + }) + } + + fn discover_bindings(&self) -> Result<()> { + let mut discovered = Vec::new(); + self.input.apply(|plan| { + plan.apply_expressions(&mut |root| { + root.apply(&mut |expr: &Arc| { + if let Some(binding) = self.discover_binding(expr) { + discovered.push(binding); + } + Ok(TreeNodeRecursion::Continue) + })?; + Ok(TreeNodeRecursion::Continue) + })?; + Ok(TreeNodeRecursion::Continue) + })?; + let mut bindings = self.bindings.lock().unwrap(); + for binding in discovered { + if !bindings + .iter() + .any(|existing| Arc::ptr_eq(&existing.predicate, &binding.predicate)) + { + bindings.push(binding); + } + } + Ok(()) + } } impl DisplayAs for ScalarSubqueryExec { @@ -183,11 +269,10 @@ impl ExecutionPlan for ScalarSubqueryExec { index: sq.index, }) .collect(); - Ok(Arc::new(ScalarSubqueryExec::new( - input, - subqueries, - self.results.clone(), - ))) + let mut new_node = + ScalarSubqueryExec::new(input, subqueries, self.results.clone()); + new_node.bindings = Arc::clone(&self.bindings); + Ok(Arc::new(new_node)) } fn with_new_children( @@ -201,16 +286,83 @@ impl ExecutionPlan for ScalarSubqueryExec { } fn reset_state(self: Arc) -> Result> { + // DynamicFilterPhysicalExpr currently has no supported reset operation. + // Do not silently reuse a completed filter on a later execution. + if !self.bindings.lock().unwrap().is_empty() { + return internal_err!( + "cannot reset ScalarSubqueryExec with dynamic-filter bindings: DynamicFilterPhysicalExpr has no reset API" + ); + } self.results.clear(); Ok(Arc::new(ScalarSubqueryExec { input: Arc::clone(&self.input), subqueries: self.subqueries.clone(), subquery_future: Arc::default(), results: self.results.clone(), + bindings: Arc::clone(&self.bindings), cache: Arc::clone(&self.cache), })) } + fn gather_filters_for_pushdown( + &self, + phase: FilterPushdownPhase, + parent_filters: Vec>, + _config: &ConfigOptions, + ) -> Result { + if phase == FilterPushdownPhase::Post { + self.discover_bindings()?; + } + let mut main = ChildFilterDescription::from_child(&parent_filters, self.input())?; + if phase == FilterPushdownPhase::Post { + for binding in self.bindings.lock().unwrap().iter() { + main = main.with_self_filter(Arc::clone(&binding.consumers[0])); + } + } + let mut description = FilterDescription::new().with_child(main); + for subquery in &self.subqueries { + description = description + .with_child(ChildFilterDescription::all_unsupported(&parent_filters)); + let _ = subquery; + } + Ok(description) + } + + fn handle_child_pushdown_result( + &self, + phase: FilterPushdownPhase, + child_pushdown_result: ChildPushdownResult, + _config: &ConfigOptions, + ) -> Result>> { + let result = FilterPushdownPropagation::if_any(child_pushdown_result.clone()); + if phase == FilterPushdownPhase::Post { + let accepted = child_pushdown_result + .self_filters + .first() + .into_iter() + .flatten() + .filter(|predicate| matches!(predicate.discriminant, PushedDown::Yes)) + .filter_map(|predicate| predicate.predicate.expression_id()) + .collect::>(); + let mut bindings = self.bindings.lock().unwrap(); + let mut retained = Vec::with_capacity(bindings.len()); + for binding in bindings.drain(..) { + let keep = match binding.consumers[0].expression_id() { + Some(id) => { + accepted.contains(&id) + || plan_contains_expression_id(&self.input, id)? + } + None => false, + }; + if keep { + retained.push(binding); + } + } + *bindings = retained; + } + Ok(result) + } + fn execute( &self, partition: usize, @@ -218,9 +370,12 @@ impl ExecutionPlan for ScalarSubqueryExec { ) -> Result { let subqueries = self.subqueries.clone(); let results = self.results.clone(); + let bindings = Arc::clone(&self.bindings); let planning_ctx = Arc::clone(&context); let mut subquery_future = self.subquery_future.try_once(move || { - Ok(async move { execute_subqueries(subqueries, results, planning_ctx).await }) + Ok(async move { + execute_subqueries(subqueries, results, bindings, planning_ctx).await + }) })?; let input = Arc::clone(&self.input); let schema = self.schema(); @@ -242,9 +397,27 @@ impl ExecutionPlan for ScalarSubqueryExec { fn apply_expressions( &self, - _f: &mut dyn FnMut(&Arc) -> Result, + f: &mut dyn FnMut(&Arc) -> Result, ) -> Result { - Ok(TreeNodeRecursion::Continue) + let bindings = self.bindings.lock().unwrap(); + crate::apply_expression_roots( + bindings.iter().flat_map(|binding| { + [ + Arc::clone(&binding.predicate), + Arc::clone(&binding.consumers[0]), + ] + }), + f, + ) + } + + fn dynamic_expressions_produced(&self) -> Vec> { + self.bindings + .lock() + .unwrap() + .iter() + .map(|binding| Arc::clone(&binding.consumers[0])) + .collect() } fn maintains_input_order(&self) -> Vec { @@ -285,16 +458,43 @@ impl ExecutionPlan for ScalarSubqueryExec { ) -> Result> { use datafusion_proto_models::protobuf; - let input = ctx.encode_child(self.input())?; + let ScalarSubqueryExec { + input: input_plan, + subqueries, + subquery_future: _, + results: _, + bindings, + cache: _, + } = self; + let input = ctx.encode_child(input_plan)?; // Subquery indices are positional and recovered during decoding. let subqueries = - ctx.encode_children(self.subqueries().iter().map(|subquery| &subquery.plan))?; + ctx.encode_children(subqueries.iter().map(|subquery| &subquery.plan))?; + let dynamic_filter_bindings = bindings + .lock() + .unwrap() + .iter() + .map(|binding| { + Ok(protobuf::ScalarSubqueryDynamicFilterBindingNode { + predicate: Some(ctx.encode_expr(&binding.predicate)?), + dynamic_filter_id: Some( + binding.consumers[0].expression_id().ok_or_else(|| { + DataFusionError::Internal( + "ScalarSubquery dynamic filter is missing expression_id" + .to_string(), + ) + })?, + ), + }) + }) + .collect::>>()?; Ok(Some(protobuf::PhysicalPlanNode { physical_plan_type: Some( protobuf::physical_plan_node::PhysicalPlanType::ScalarSubquery(Box::new( protobuf::ScalarSubqueryExecNode { input: Some(Box::new(input)), subqueries, + dynamic_filter_bindings, }, )), ), @@ -318,13 +518,61 @@ impl ScalarSubqueryExec { ); let results = ScalarSubqueryResults::new(scalar_subquery.subqueries.len()); let input_node = scalar_subquery.input.as_deref().ok_or_else(|| { - datafusion_common::internal_datafusion_err!( - "ScalarSubqueryExec is missing required field 'input'" + DataFusionError::Internal( + "ScalarSubqueryExec is missing required field 'input'".to_string(), ) })?; - // The input's ScalarSubqueryExpr nodes must share this results container. let input = ctx.decode_child_with_scalar_subquery_results(input_node, results.clone())?; + + let mut dynamic_filters = + std::collections::HashMap::>>::new(); + input.apply(|plan| { + plan.apply_expressions(&mut |root| { + root.apply(&mut |expr: &Arc| { + if let Some(filter) = expr.downcast_ref::() + { + if let Some(id) = filter.expression_id() { + dynamic_filters + .entry(id) + .or_default() + .push(Arc::clone(expr)); + } + } + Ok(TreeNodeRecursion::Continue) + }) + })?; + Ok(TreeNodeRecursion::Continue) + })?; + + let mut bindings = + Vec::with_capacity(scalar_subquery.dynamic_filter_bindings.len()); + for binding in &scalar_subquery.dynamic_filter_bindings { + let predicate = ctx.decode_expr_with_scalar_subquery_results( + binding.predicate.as_ref().ok_or_else(|| { + DataFusionError::Internal( + "ScalarSubqueryDynamicFilterBindingNode is missing required field 'predicate'".to_string(), + ) + })?, + input.schema().as_ref(), + results.clone(), + )?; + let id = binding.dynamic_filter_id.ok_or_else(|| { + DataFusionError::Internal( + "ScalarSubqueryDynamicFilterBindingNode is missing required field 'dynamic_filter_id'".to_string(), + ) + })?; + let filters = dynamic_filters.get(&id).ok_or_else(|| { + DataFusionError::Internal(format!( + "ScalarSubquery dynamic filter binding references missing expression_id {id}" + )) + })?; + bindings.push(ScalarSubqueryBinding { + predicate, + consumers: filters.clone(), + }); + } + let subqueries = scalar_subquery .subqueries .iter() @@ -336,8 +584,9 @@ impl ScalarSubqueryExec { }) }) .collect::>>()?; - - Ok(Arc::new(Self::new(input, subqueries, results))) + let exec = Self::new(input, subqueries, results); + exec.bindings.lock().unwrap().extend(bindings); + Ok(Arc::new(exec)) } } @@ -350,6 +599,7 @@ async fn wait_for_subqueries(fut: &mut OnceFut<()>) -> Result<()> { async fn execute_subqueries( subqueries: Vec, results: ScalarSubqueryResults, + bindings: Arc>>, context: Arc, ) -> Result<()> { // Evaluate subqueries in parallel; wait for them all to finish evaluation @@ -366,6 +616,32 @@ async fn execute_subqueries( } }); futures::future::try_join_all(futures).await?; + let bindings = bindings.lock().unwrap(); + for binding in bindings.iter() { + // Snapshot the producer exactly once, then publish it to every + // independent runtime state. Derived consumers sharing one state are + // updated and completed only once. + let snapshot = snapshot_physical_expr(Arc::clone(&binding.predicate))?; + let mut updated = Vec::<&DynamicFilterPhysicalExpr>::new(); + for consumer in &binding.consumers { + let filter = consumer + .downcast_ref::() + .ok_or_else(|| { + DataFusionError::Internal( + "ScalarSubquery dynamic-filter binding has an invalid filter" + .to_string(), + ) + })?; + if !updated + .iter() + .any(|updated_filter| updated_filter.shares_runtime_state(filter)) + { + filter.update(Arc::clone(&snapshot))?; + filter.mark_complete(); + updated.push(filter); + } + } + } Ok(()) } @@ -418,8 +694,10 @@ mod tests { use crate::test::exec::ErrorExec; use arrow::array::{Int32Array, Int64Array}; - use arrow::datatypes::{DataType, Field, Schema}; + use arrow::datatypes::{DataType, Field, IntervalMonthDayNano, Schema, TimeUnit}; use arrow::record_batch::RecordBatch; + use datafusion_expr::Operator; + use datafusion_physical_expr::expressions::DynamicFilterPhysicalExpr; use datafusion_physical_expr::scalar_subquery::ScalarSubqueryExpr; enum ExpectedSubqueryResult { @@ -427,6 +705,90 @@ mod tests { Error(&'static str), } + #[derive(Debug)] + struct ExpressionExec { + input: Arc, + expressions: Arc>>>, + } + + impl ExpressionExec { + fn new( + input: Arc, + expressions: Vec>, + ) -> Self { + Self { + input, + expressions: Arc::new(Mutex::new(expressions)), + } + } + + fn add_expression(&self, expression: Arc) { + self.expressions.lock().unwrap().push(expression); + } + } + + impl DisplayAs for ExpressionExec { + fn fmt_as(&self, t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result { + match t { + DisplayFormatType::Default | DisplayFormatType::Verbose => { + write!(f, "ExpressionExec") + } + DisplayFormatType::TreeRender => write!(f, ""), + } + } + } + + impl ExecutionPlan for ExpressionExec { + fn name(&self) -> &'static str { + "ExpressionExec" + } + + fn properties(&self) -> &Arc { + self.input.properties() + } + + fn children(&self) -> Vec<&Arc> { + vec![&self.input] + } + + fn replace_children( + self: Arc, + mut children: Vec>, + _: ReplaceChildrenOptions, + ) -> Result> { + Ok(Arc::new(Self { + input: children.remove(0), + expressions: Arc::clone(&self.expressions), + })) + } + + fn with_new_children( + self: Arc, + children: Vec>, + ) -> Result> { + self.replace_children( + children, + ReplaceChildrenOptions::new(ChildrenPropertiesMode::Recompute), + ) + } + + fn execute( + &self, + partition: usize, + context: Arc, + ) -> Result { + self.input.execute(partition, context) + } + + fn apply_expressions( + &self, + f: &mut dyn FnMut(&Arc) -> Result, + ) -> Result { + let expressions = self.expressions.lock().unwrap().clone(); + crate::apply_expression_roots(expressions, f) + } + } + #[derive(Debug)] struct CountingExec { inner: Arc, @@ -569,6 +931,294 @@ mod tests { values.value(0) } + fn timestamp_predicate(results: ScalarSubqueryResults) -> Arc { + let scalar = Arc::new(ScalarSubqueryExpr::new( + DataType::Timestamp(TimeUnit::Millisecond, Some("UTC".into())), + false, + SubqueryIndex::new(0), + results, + )); + Arc::new(BinaryExpr::new( + Arc::new(Column::new("timestamp_column", 0)), + Operator::GtEq, + Arc::new(BinaryExpr::new( + scalar, + Operator::Minus, + Arc::new(Literal::new(ScalarValue::IntervalMonthDayNano(Some( + IntervalMonthDayNano::new(0, 0, 60_000_000_000), + )))), + )), + )) + } + + fn dynamic_filter(exec: &ScalarSubqueryExec) -> Arc { + let filter = Arc::clone(&exec.bindings.lock().unwrap()[0].consumers[0]); + Arc::downcast::(filter) + .expect("expected dynamic filter") + } + + #[test] + fn test_scalar_subquery_filter_pushdown_no_binding_retention() -> Result<()> { + for (name, input_contains_filter, expected_bindings) in [ + ("removes_missing_filter", false, 0), + ("retains_filter_in_rewritten_input", true, 1), + ] { + let results = ScalarSubqueryResults::new(1); + let predicate = timestamp_predicate(results.clone()); + let input = Arc::new(ExpressionExec::new( + placeholder_input(), + vec![Arc::clone(&predicate)], + )); + let exec = Arc::new(single_subquery_exec( + input, + make_subquery_plan(vec![int32_batch(vec![1])]), + results, + )); + exec.discover_bindings()?; + let filter = dynamic_filter(&exec); + let mut expressions = + vec![predicate, Arc::clone(&filter) as Arc]; + if !input_contains_filter { + expressions.pop(); + } + let updated_input = + Arc::new(ExpressionExec::new(placeholder_input(), expressions)); + let subquery = Arc::clone(&exec.subqueries()[0].plan); + let updated_exec = exec.replace_children( + vec![updated_input, subquery], + ReplaceChildrenOptions::new(ChildrenPropertiesMode::Recompute), + )?; + + updated_exec.handle_child_pushdown_result( + FilterPushdownPhase::Post, + ChildPushdownResult { + parent_filters: vec![], + self_filters: vec![vec![PushedDown::No + .wrap_expression(Arc::clone(&filter) as Arc)]], + }, + &ConfigOptions::default(), + )?; + assert_eq!( + updated_exec.dynamic_expressions_produced().len(), + expected_bindings, + "{name}" + ); + } + Ok(()) + } + + #[test] + fn test_scalar_subquery_filter_pushdown_retains_accepted_binding() -> Result<()> { + let results = ScalarSubqueryResults::new(1); + let predicate = timestamp_predicate(results.clone()); + let input = Arc::new(ExpressionExec::new( + placeholder_input(), + vec![Arc::clone(&predicate)], + )); + let exec = single_subquery_exec( + input, + make_subquery_plan(vec![int32_batch(vec![1])]), + results, + ); + exec.discover_bindings()?; + let filter = dynamic_filter(&exec); + + exec.handle_child_pushdown_result( + FilterPushdownPhase::Post, + ChildPushdownResult { + parent_filters: vec![], + self_filters: vec![vec![ + PushedDown::Yes + .wrap_expression(Arc::clone(&filter) as Arc), + ]], + }, + &ConfigOptions::default(), + )?; + + assert!(reset_plan_states(Arc::new(exec)).is_err()); + Ok(()) + } + + #[test] + fn test_discover_scalar_subquery_bindings_by_results_scope() -> Result<()> { + let results = ScalarSubqueryResults::new(1); + let predicate = timestamp_predicate(results.clone()); + let input = Arc::new(ExpressionExec::new( + placeholder_input(), + vec![Arc::clone(&predicate)], + )); + let exec = single_subquery_exec( + input.clone(), + make_subquery_plan(vec![int32_batch(vec![1])]), + results, + ); + + exec.discover_bindings()?; + exec.discover_bindings()?; + assert_eq!(exec.bindings.lock().unwrap().len(), 1); + + input.add_expression(timestamp_predicate(exec.results().clone())); + exec.discover_bindings()?; + assert_eq!(exec.bindings.lock().unwrap().len(), 2); + + input.add_expression(timestamp_predicate(ScalarSubqueryResults::new(1))); + exec.discover_bindings()?; + assert_eq!(exec.bindings.lock().unwrap().len(), 2); + Ok(()) + } + + #[tokio::test] + async fn test_execute_subqueries_updates_shared_runtime_state_once() -> Result<()> { + let schema = Arc::new(Schema::new(vec![ + Field::new("a", DataType::Int32, false), + Field::new("b", DataType::Int32, false), + Field::new("c", DataType::Int32, false), + ])); + let original_column = + Arc::new(Column::new_with_schema("a", &schema)?) as Arc; + let filter = Arc::new(DynamicFilterPhysicalExpr::new( + vec![Arc::clone(&original_column)], + lit(true), + )); + let derived_expr_1 = Arc::clone(&filter) + .with_new_children(vec![Arc::new(Column::new_with_schema("b", &schema)?)])?; + let derived_expr_2 = Arc::clone(&filter) + .with_new_children(vec![Arc::new(Column::new_with_schema("c", &schema)?)])?; + let derived_1 = derived_expr_1 + .downcast_ref::() + .ok_or_else(|| { + DataFusionError::Internal("expected dynamic filter".to_string()) + })?; + let derived_2 = derived_expr_2 + .downcast_ref::() + .ok_or_else(|| { + DataFusionError::Internal("expected dynamic filter".to_string()) + })?; + assert!(filter.shares_runtime_state(derived_1)); + assert!(filter.shares_runtime_state(derived_2)); + + let bindings = Arc::new(Mutex::new(vec![ScalarSubqueryBinding { + predicate: Arc::clone(&original_column), + consumers: vec![Arc::clone(&derived_expr_1), Arc::clone(&derived_expr_2)], + }])); + execute_subqueries( + vec![], + ScalarSubqueryResults::new(0), + bindings, + Arc::new(TaskContext::default()), + ) + .await?; + + assert_eq!(filter.snapshot_generation(), 2); + assert_eq!(derived_1.snapshot_generation(), 2); + assert_eq!(derived_2.snapshot_generation(), 2); + assert_eq!( + derived_1 + .current()? + .downcast_ref::() + .unwrap() + .index(), + 1 + ); + assert_eq!( + derived_2 + .current()? + .downcast_ref::() + .unwrap() + .index(), + 2 + ); + for (filter, message) in [ + (filter.as_ref(), "shared filter should be complete"), + (derived_1, "first derived filter should be complete"), + (derived_2, "second derived filter should be complete"), + ] { + tokio::time::timeout( + std::time::Duration::from_secs(1), + filter.wait_complete(), + ) + .await + .expect(message); + } + Ok(()) + } + + #[tokio::test] + async fn test_scalar_subquery_updates_dynamic_filter_with_timestamp() -> Result<()> { + let results = ScalarSubqueryResults::new(1); + let predicate = timestamp_predicate(results.clone()); + let input = Arc::new(ExpressionExec::new(placeholder_input(), vec![predicate])); + let exec = single_subquery_exec( + input, + make_subquery_plan(vec![int32_batch(vec![1])]), + results.clone(), + ); + exec.discover_bindings()?; + results.set( + SubqueryIndex::new(0), + ScalarValue::TimestampMillisecond( + Some(1_672_574_400_000), + Some("UTC".into()), + ), + )?; + + execute_subqueries( + vec![], + results, + Arc::clone(&exec.bindings), + Arc::new(TaskContext::default()), + ) + .await?; + + let filter = dynamic_filter(&exec); + let current = filter.current()?; + let predicate = current.downcast_ref::().unwrap(); + let subtraction = predicate.right().downcast_ref::().unwrap(); + let timestamp = subtraction.left().downcast_ref::().unwrap(); + assert_eq!( + timestamp.value(), + &ScalarValue::TimestampMillisecond( + Some(1_672_574_400_000), + Some("UTC".into()) + ) + ); + tokio::time::timeout(std::time::Duration::from_secs(1), filter.wait_complete()) + .await + .expect("scalar subquery dynamic filter should be complete"); + Ok(()) + } + + #[tokio::test] + async fn test_failed_scalar_subquery_does_not_update_dynamic_filter() -> Result<()> { + let results = ScalarSubqueryResults::new(1); + let predicate = timestamp_predicate(results.clone()); + let input = Arc::new(ExpressionExec::new(placeholder_input(), vec![predicate])); + let exec = single_subquery_exec(input, Arc::new(ErrorExec::new()), results); + exec.discover_bindings()?; + let filter = dynamic_filter(&exec); + let before = filter.current()?; + assert!(before.downcast_ref::().is_some()); + + let result = execute_subqueries( + exec.subqueries().to_vec(), + exec.results().clone(), + Arc::clone(&exec.bindings), + Arc::new(TaskContext::default()), + ) + .await; + assert!(result.is_err()); + assert!(filter.current()?.downcast_ref::().is_some()); + assert!( + tokio::time::timeout( + std::time::Duration::from_millis(50), + filter.wait_complete() + ) + .await + .is_err() + ); + Ok(()) + } + #[tokio::test] async fn test_execute_scalar_subquery_row_count_semantics() -> Result<()> { for (name, plan, expected) in [ diff --git a/datafusion/proto-models/proto/datafusion.proto b/datafusion/proto-models/proto/datafusion.proto index 43a90264c2b1f..2cb3406f7ecf1 100644 --- a/datafusion/proto-models/proto/datafusion.proto +++ b/datafusion/proto-models/proto/datafusion.proto @@ -1689,9 +1689,15 @@ message BufferExecNode { uint64 capacity = 2; } +message ScalarSubqueryDynamicFilterBindingNode { + PhysicalExprNode predicate = 1; + optional uint64 dynamic_filter_id = 2; +} + message ScalarSubqueryExecNode { PhysicalPlanNode input = 1; repeated PhysicalPlanNode subqueries = 2; + repeated ScalarSubqueryDynamicFilterBindingNode dynamic_filter_bindings = 3; } message PhysicalScalarSubqueryExprNode { diff --git a/datafusion/proto-models/src/generated/pbjson.rs b/datafusion/proto-models/src/generated/pbjson.rs index 908f9752b7f18..883aa2759f680 100644 --- a/datafusion/proto-models/src/generated/pbjson.rs +++ b/datafusion/proto-models/src/generated/pbjson.rs @@ -24231,6 +24231,119 @@ impl<'de> serde::Deserialize<'de> for RollupNode { deserializer.deserialize_struct("datafusion.RollupNode", FIELDS, GeneratedVisitor) } } +impl serde::Serialize for ScalarSubqueryDynamicFilterBindingNode { + #[allow(deprecated)] + fn serialize(&self, serializer: S) -> std::result::Result + where + S: serde::Serializer, + { + use serde::ser::SerializeStruct; + let mut len = 0; + if self.predicate.is_some() { + len += 1; + } + if self.dynamic_filter_id.is_some() { + len += 1; + } + let mut struct_ser = serializer.serialize_struct("datafusion.ScalarSubqueryDynamicFilterBindingNode", len)?; + if let Some(v) = self.predicate.as_ref() { + struct_ser.serialize_field("predicate", v)?; + } + if let Some(v) = self.dynamic_filter_id.as_ref() { + #[allow(clippy::needless_borrow)] + #[allow(clippy::needless_borrows_for_generic_args)] + struct_ser.serialize_field("dynamicFilterId", ToString::to_string(&v).as_str())?; + } + struct_ser.end() + } +} +impl<'de> serde::Deserialize<'de> for ScalarSubqueryDynamicFilterBindingNode { + #[allow(deprecated)] + fn deserialize(deserializer: D) -> std::result::Result + where + D: serde::Deserializer<'de>, + { + const FIELDS: &[&str] = &[ + "predicate", + "dynamic_filter_id", + "dynamicFilterId", + ]; + + #[allow(clippy::enum_variant_names)] + enum GeneratedField { + Predicate, + DynamicFilterId, + } + impl<'de> serde::Deserialize<'de> for GeneratedField { + fn deserialize(deserializer: D) -> std::result::Result + where + D: serde::Deserializer<'de>, + { + struct GeneratedVisitor; + + impl serde::de::Visitor<'_> for GeneratedVisitor { + type Value = GeneratedField; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(formatter, "expected one of: {:?}", &FIELDS) + } + + #[allow(unused_variables)] + fn visit_str(self, value: &str) -> std::result::Result + where + E: serde::de::Error, + { + match value { + "predicate" => Ok(GeneratedField::Predicate), + "dynamicFilterId" | "dynamic_filter_id" => Ok(GeneratedField::DynamicFilterId), + _ => Err(serde::de::Error::unknown_field(value, FIELDS)), + } + } + } + deserializer.deserialize_identifier(GeneratedVisitor) + } + } + struct GeneratedVisitor; + impl<'de> serde::de::Visitor<'de> for GeneratedVisitor { + type Value = ScalarSubqueryDynamicFilterBindingNode; + + fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str("struct datafusion.ScalarSubqueryDynamicFilterBindingNode") + } + + fn visit_map(self, mut map_: V) -> std::result::Result + where + V: serde::de::MapAccess<'de>, + { + let mut predicate__ = None; + let mut dynamic_filter_id__ = None; + while let Some(k) = map_.next_key()? { + match k { + GeneratedField::Predicate => { + if predicate__.is_some() { + return Err(serde::de::Error::duplicate_field("predicate")); + } + predicate__ = map_.next_value()?; + } + GeneratedField::DynamicFilterId => { + if dynamic_filter_id__.is_some() { + return Err(serde::de::Error::duplicate_field("dynamicFilterId")); + } + dynamic_filter_id__ = + map_.next_value::<::std::option::Option<::pbjson::private::NumberDeserialize<_>>>()?.map(|x| x.0) + ; + } + } + } + Ok(ScalarSubqueryDynamicFilterBindingNode { + predicate: predicate__, + dynamic_filter_id: dynamic_filter_id__, + }) + } + } + deserializer.deserialize_struct("datafusion.ScalarSubqueryDynamicFilterBindingNode", FIELDS, GeneratedVisitor) + } +} impl serde::Serialize for ScalarSubqueryExecNode { #[allow(deprecated)] fn serialize(&self, serializer: S) -> std::result::Result @@ -24245,6 +24358,9 @@ impl serde::Serialize for ScalarSubqueryExecNode { if !self.subqueries.is_empty() { len += 1; } + if !self.dynamic_filter_bindings.is_empty() { + len += 1; + } let mut struct_ser = serializer.serialize_struct("datafusion.ScalarSubqueryExecNode", len)?; if let Some(v) = self.input.as_ref() { struct_ser.serialize_field("input", v)?; @@ -24252,6 +24368,9 @@ impl serde::Serialize for ScalarSubqueryExecNode { if !self.subqueries.is_empty() { struct_ser.serialize_field("subqueries", &self.subqueries)?; } + if !self.dynamic_filter_bindings.is_empty() { + struct_ser.serialize_field("dynamicFilterBindings", &self.dynamic_filter_bindings)?; + } struct_ser.end() } } @@ -24264,12 +24383,15 @@ impl<'de> serde::Deserialize<'de> for ScalarSubqueryExecNode { const FIELDS: &[&str] = &[ "input", "subqueries", + "dynamic_filter_bindings", + "dynamicFilterBindings", ]; #[allow(clippy::enum_variant_names)] enum GeneratedField { Input, Subqueries, + DynamicFilterBindings, } impl<'de> serde::Deserialize<'de> for GeneratedField { fn deserialize(deserializer: D) -> std::result::Result @@ -24293,6 +24415,7 @@ impl<'de> serde::Deserialize<'de> for ScalarSubqueryExecNode { match value { "input" => Ok(GeneratedField::Input), "subqueries" => Ok(GeneratedField::Subqueries), + "dynamicFilterBindings" | "dynamic_filter_bindings" => Ok(GeneratedField::DynamicFilterBindings), _ => Err(serde::de::Error::unknown_field(value, FIELDS)), } } @@ -24314,6 +24437,7 @@ impl<'de> serde::Deserialize<'de> for ScalarSubqueryExecNode { { let mut input__ = None; let mut subqueries__ = None; + let mut dynamic_filter_bindings__ = None; while let Some(k) = map_.next_key()? { match k { GeneratedField::Input => { @@ -24328,11 +24452,18 @@ impl<'de> serde::Deserialize<'de> for ScalarSubqueryExecNode { } subqueries__ = Some(map_.next_value()?); } + GeneratedField::DynamicFilterBindings => { + if dynamic_filter_bindings__.is_some() { + return Err(serde::de::Error::duplicate_field("dynamicFilterBindings")); + } + dynamic_filter_bindings__ = Some(map_.next_value()?); + } } } Ok(ScalarSubqueryExecNode { input: input__, subqueries: subqueries__.unwrap_or_default(), + dynamic_filter_bindings: dynamic_filter_bindings__.unwrap_or_default(), }) } } diff --git a/datafusion/proto-models/src/generated/prost.rs b/datafusion/proto-models/src/generated/prost.rs index ba00577ab9a1b..280517a41218f 100644 --- a/datafusion/proto-models/src/generated/prost.rs +++ b/datafusion/proto-models/src/generated/prost.rs @@ -2566,11 +2566,22 @@ pub struct BufferExecNode { pub capacity: u64, } #[derive(Clone, PartialEq, ::prost::Message)] +pub struct ScalarSubqueryDynamicFilterBindingNode { + #[prost(message, optional, tag = "1")] + pub predicate: ::core::option::Option, + #[prost(uint64, optional, tag = "2")] + pub dynamic_filter_id: ::core::option::Option, +} +#[derive(Clone, PartialEq, ::prost::Message)] pub struct ScalarSubqueryExecNode { #[prost(message, optional, boxed, tag = "1")] pub input: ::core::option::Option<::prost::alloc::boxed::Box>, #[prost(message, repeated, tag = "2")] pub subqueries: ::prost::alloc::vec::Vec, + #[prost(message, repeated, tag = "3")] + pub dynamic_filter_bindings: ::prost::alloc::vec::Vec< + ScalarSubqueryDynamicFilterBindingNode, + >, } #[derive(Clone, PartialEq, ::prost::Message)] pub struct PhysicalScalarSubqueryExprNode { diff --git a/datafusion/proto/src/physical_plan/mod.rs b/datafusion/proto/src/physical_plan/mod.rs index 222901aff5211..92dd826b5fca8 100644 --- a/datafusion/proto/src/physical_plan/mod.rs +++ b/datafusion/proto/src/physical_plan/mod.rs @@ -2094,6 +2094,17 @@ impl ExecutionPlanDecode for ConverterPlanDecoder<'_, '_> { .proto_to_physical_expr(node, input_schema, self.ctx) } + fn decode_expr_with_scalar_subquery_results( + &self, + node: &protobuf::PhysicalExprNode, + input_schema: &Schema, + results: ScalarSubqueryResults, + ) -> Result> { + let scoped_ctx = self.ctx.with_scalar_subquery_results(results); + self.proto_converter + .proto_to_physical_expr(node, input_schema, &scoped_ctx) + } + fn task_ctx(&self) -> &TaskContext { self.ctx.task_ctx() } diff --git a/datafusion/proto/tests/cases/plans/scalar_subquery.rs b/datafusion/proto/tests/cases/plans/scalar_subquery.rs index 34d30aa03dece..c9b9f2c5a3cc4 100644 --- a/datafusion/proto/tests/cases/plans/scalar_subquery.rs +++ b/datafusion/proto/tests/cases/plans/scalar_subquery.rs @@ -26,17 +26,13 @@ use datafusion::physical_plan::filter::FilterExec; use datafusion::physical_plan::scalar_subquery::{ ScalarSubqueryExec, ScalarSubqueryLink, }; +use datafusion::physical_plan::union::UnionExec; use datafusion::prelude::SessionContext; use datafusion_common::Result; +use datafusion_common::tree_node::TreeNode; use datafusion_expr::physical_planning_context::{ScalarSubqueryResults, SubqueryIndex}; use datafusion_physical_expr::scalar_subquery::ScalarSubqueryExpr; -use datafusion_proto::bytes::{ - physical_plan_from_bytes_with_proto_converter, - physical_plan_to_bytes_with_proto_converter, -}; -use datafusion_proto::physical_plan::{ - DeduplicatingProtoConverter, DefaultPhysicalExtensionCodec, -}; +use datafusion_proto::bytes::{physical_plan_from_bytes, physical_plan_to_bytes}; use std::sync::Arc; use std::vec; @@ -76,23 +72,9 @@ fn roundtrip_scalar_subquery_exec() -> Result<()> { results, )); - // Perform the round-trip using DeduplicatingProtoConverter, which - // creates a DeduplicatingDeserializer that threads scalar subquery - // results through expression deserialization. - let codec = DefaultPhysicalExtensionCodec {}; - let converter = DeduplicatingProtoConverter {}; - let bytes = physical_plan_to_bytes_with_proto_converter( - Arc::clone(&exec), - &codec, - &converter, - )?; + let bytes = physical_plan_to_bytes(Arc::clone(&exec))?; let ctx = SessionContext::new(); - let deserialized = physical_plan_from_bytes_with_proto_converter( - bytes.as_ref(), - ctx.task_ctx().as_ref(), - &codec, - &converter, - )?; + let deserialized = physical_plan_from_bytes(bytes.as_ref(), ctx.task_ctx().as_ref())?; // Verify the deserialized ScalarSubqueryExec's results container is // shared with the ScalarSubqueryExpr in the input plan. @@ -173,12 +155,9 @@ fn roundtrip_nested_scalar_subquery_exec_scopes_results() -> Result<()> { outer_results, )); - let bytes = datafusion_proto::bytes::physical_plan_to_bytes(Arc::clone(&outer_exec))?; + let bytes = physical_plan_to_bytes(Arc::clone(&outer_exec))?; let ctx = SessionContext::new(); - let deserialized = datafusion_proto::bytes::physical_plan_from_bytes( - bytes.as_ref(), - ctx.task_ctx().as_ref(), - )?; + let deserialized = physical_plan_from_bytes(bytes.as_ref(), ctx.task_ctx().as_ref())?; let outer_exec = deserialized .downcast_ref::() @@ -254,12 +233,8 @@ async fn roundtrip_scalar_subquery_exec_with_default_converter_executes() -> Res "expected ScalarSubqueryExec in plan:\n{initial_plan:?}" ); - let bytes = - datafusion_proto::bytes::physical_plan_to_bytes(Arc::clone(&initial_plan))?; - let roundtripped = datafusion_proto::bytes::physical_plan_from_bytes( - bytes.as_ref(), - ctx.task_ctx().as_ref(), - )?; + let bytes = physical_plan_to_bytes(Arc::clone(&initial_plan))?; + let roundtripped = physical_plan_from_bytes(bytes.as_ref(), ctx.task_ctx().as_ref())?; assert!( format!("{roundtripped:?}").contains("ScalarSubqueryExec"), "expected ScalarSubqueryExec after roundtrip:\n{roundtripped:?}" @@ -276,3 +251,116 @@ async fn roundtrip_scalar_subquery_exec_with_default_converter_executes() -> Res Ok(()) } + +/// Verify that built-in protobuf round-tripping preserves scalar-subquery +/// dynamic-filter bindings and that execution updates all decoded consumers. +#[tokio::test] +async fn roundtrip_scalar_subquery_dynamic_filter_binding_executes() -> Result<()> { + use datafusion::physical_plan::PhysicalExpr; + use datafusion_common::tree_node::TreeNodeRecursion; + use datafusion_physical_expr::expressions::DynamicFilterPhysicalExpr; + + let ctx = SessionContext::new(); + let sql = "SELECT x FROM (VALUES (TIMESTAMP '2023-01-01 00:00:00'), \ + (TIMESTAMP '2023-01-03 00:00:00')) AS t(x) \ + WHERE x >= (SELECT max(y) FROM \ + (VALUES (TIMESTAMP '2023-01-02 00:00:00')) AS u(y)) \ + - INTERVAL '0' DAY"; + let initial_plan = ctx.sql(sql).await?.create_physical_plan().await?; + + initial_plan.gather_filters_for_pushdown( + datafusion::physical_plan::filter_pushdown::FilterPushdownPhase::Post, + vec![], + ctx.state().config_options(), + )?; + let scalar_exec = initial_plan + .downcast_ref::() + .expect("expected ScalarSubqueryExec"); + let produced = scalar_exec + .dynamic_expressions_produced() + .into_iter() + .next() + .expect("expected scalar dynamic-filter producer"); + // Two actual input branches consume the same optimized filter instance. + // The default converter decodes these occurrences independently while + // retaining their shared expression_id. + let input_with_consumers = UnionExec::try_new(vec![ + Arc::new(FilterExec::try_new( + Arc::clone(&produced), + Arc::clone(scalar_exec.input()), + )?) as Arc, + Arc::new(FilterExec::try_new( + produced, + Arc::clone(scalar_exec.input()), + )?) as Arc, + ])?; + let mut children = vec![input_with_consumers as Arc]; + children.extend( + scalar_exec + .subqueries() + .iter() + .map(|link| Arc::clone(&link.plan)), + ); + let initial_plan = initial_plan.replace_children( + children, + datafusion::physical_plan::ReplaceChildrenOptions::new( + datafusion::physical_plan::ChildrenPropertiesMode::Recompute, + ), + )?; + + let bytes = physical_plan_to_bytes(Arc::clone(&initial_plan))?; + let roundtripped = physical_plan_from_bytes(bytes.as_ref(), ctx.task_ctx().as_ref())?; + + let scalar_exec = roundtripped + .downcast_ref::() + .expect("expected ScalarSubqueryExec"); + let producer_id = scalar_exec + .dynamic_expressions_produced() + .into_iter() + .next() + .and_then(|expr| expr.expression_id()) + .expect("expected scalar dynamic-filter producer"); + let mut consumers = vec![]; + scalar_exec.input().apply(|node| { + node.apply_expressions(&mut |root| { + root.apply(&mut |expr: &Arc| { + if expr.expression_id() == Some(producer_id) + && expr.downcast_ref::().is_some() + { + consumers.push(Arc::clone(expr)); + } + Ok(TreeNodeRecursion::Continue) + }) + }) + })?; + assert_eq!(consumers.len(), 2); + + let batches = + datafusion::physical_plan::collect_partitioned(roundtripped, ctx.task_ctx()) + .await? + .into_iter() + .flatten() + .collect::>(); + datafusion::assert_batches_eq!( + &[ + "+---------------------+", + "| x |", + "+---------------------+", + "| 2023-01-03T00:00:00 |", + "| 2023-01-03T00:00:00 |", + "+---------------------+" + ], + &batches + ); + + for consumer in consumers { + let filter = consumer + .downcast_ref::() + .expect("consumer should be a dynamic filter"); + tokio::time::timeout(std::time::Duration::from_secs(5), filter.wait_complete()) + .await + .expect("dynamic-filter consumer did not complete"); + assert_ne!(filter.current()?.to_string(), "true"); + } + Ok(()) +}