From 84f69fe21d979bad2fcf06cb54d9ad10cafc4419 Mon Sep 17 00:00:00 2001 From: discord9 <55937128+discord9@users.noreply.github.com> Date: Tue, 25 Aug 2026 20:01:36 +0800 Subject: [PATCH 1/3] fix: snapshot scalar subqueries for pruning Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> --- .../physical-expr/src/scalar_subquery.rs | 78 +++++++++++++++ datafusion/pruning/src/pruning_predicate.rs | 96 ++++++++++++++++++- 2 files changed, 172 insertions(+), 2 deletions(-) diff --git a/datafusion/physical-expr/src/scalar_subquery.rs b/datafusion/physical-expr/src/scalar_subquery.rs index 473b52a5cb45c..83465f4293023 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,73 @@ 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_snapshot_reset_returns_to_pending() -> Result<()> { + let results = ScalarSubqueryResults::new(1); + let expr = ScalarSubqueryExpr::new( + DataType::Int64, + true, + SubqueryIndex::new(0), + results.clone(), + ); + results.set(SubqueryIndex::new(0), ScalarValue::Int64(Some(7)))?; + assert!(expr.snapshot()?.is_some()); + + results.clear(); + assert!(expr.snapshot().is_err()); + Ok(()) + } + #[test] fn test_identity_equality() { let results = make_results(vec![None, None]); diff --git a/datafusion/pruning/src/pruning_predicate.rs b/datafusion/pruning/src/pruning_predicate.rs index 3a63451495e4c..37d750ad3c328 100644 --- a/datafusion/pruning/src/pruning_predicate.rs +++ b/datafusion/pruning/src/pruning_predicate.rs @@ -2187,10 +2187,16 @@ mod tests { use arrow::array::Decimal128Array; use arrow::{ - array::{BinaryArray, Int32Array, Int64Array, StringArray, UInt64Array}, - datatypes::TimeUnit, + array::{ + BinaryArray, Int32Array, Int64Array, StringArray, TimestampNanosecondArray, + UInt64Array, + }, + datatypes::{IntervalMonthDayNano, TimeUnit}, }; use datafusion_expr::expr::InList; + use datafusion_expr::physical_planning_context::{ + ScalarSubqueryResults, SubqueryIndex, + }; use datafusion_expr::{BinaryExpr, Expr, cast, is_null, try_cast}; use datafusion_functions_nested::expr_fn::{array_has, make_array}; use datafusion_physical_expr::expressions::{ @@ -2611,6 +2617,92 @@ mod tests { } } + #[test] + fn scalar_subquery_timestamp_snapshot_builds_pruning_predicate() -> Result<()> { + let schema = Arc::new(Schema::new(vec![Field::new( + "ts", + DataType::Timestamp(TimeUnit::Nanosecond, None), + true, + )])); + let results = ScalarSubqueryResults::new(1); + let scalar = datafusion_physical_expr::scalar_subquery::ScalarSubqueryExpr::new( + DataType::Timestamp(TimeUnit::Nanosecond, None), + true, + SubqueryIndex::new(0), + results.clone(), + ); + let timestamp = 1_700_000_000_000_000_000; + let expected_bound = timestamp - 60_000_000_000; + results.set( + SubqueryIndex::new(0), + ScalarValue::TimestampNanosecond(Some(timestamp), None), + )?; + + let predicate: Arc = Arc::new(phys_expr::BinaryExpr::new( + phys_expr::col("ts", &schema)?, + Operator::GtEq, + Arc::new(phys_expr::BinaryExpr::new( + Arc::new(scalar), + Operator::Minus, + Arc::new(phys_expr::Literal::new(ScalarValue::IntervalMonthDayNano( + Some(IntervalMonthDayNano { + months: 0, + days: 0, + nanoseconds: 60_000_000_000, + }), + ))), + )), + )); + + let snapshot = snapshot_physical_expr_opt(Arc::clone(&predicate))?; + assert!(snapshot.transformed); + let simplified = PhysicalExprSimplifier::new(&schema).simplify(snapshot.data)?; + let bound = simplified + .downcast_ref::() + .expect("simplified predicate should remain a binary comparison") + .right() + .downcast_ref::() + .expect("timestamp bound should be folded to a literal"); + assert_eq!( + bound.value(), + &ScalarValue::TimestampNanosecond(Some(expected_bound), None) + ); + + let pruning = PruningPredicateBuilder::new() + .with_file_schema(Arc::clone(&schema)) + .try_build(predicate)?; + assert_eq!( + pruning.orig_expr().to_string(), + format!( + "ts@0 >= {}", + ScalarValue::TimestampNanosecond(Some(expected_bound), None) + ) + ); + let below = TestStatistics::new().with( + "ts", + ContainerStats::new() + .with_min(Arc::new(TimestampNanosecondArray::from(vec![Some( + expected_bound - 1, + )]))) + .with_max(Arc::new(TimestampNanosecondArray::from(vec![Some( + expected_bound - 1, + )]))), + ); + let matching = TestStatistics::new().with( + "ts", + ContainerStats::new() + .with_min(Arc::new(TimestampNanosecondArray::from(vec![Some( + expected_bound, + )]))) + .with_max(Arc::new(TimestampNanosecondArray::from(vec![Some( + expected_bound, + )]))), + ); + assert_eq!(pruning.prune(&below)?, vec![false]); + assert_eq!(pruning.prune(&matching)?, vec![true]); + Ok(()) + } + #[test] fn prune_all_rows_null_counts() { // if null_count = row_count then we should prune the container for i = 0 From f6fc83dcacc35e71150f37f6d5e81bbd4632345c Mon Sep 17 00:00:00 2001 From: discord9 <55937128+discord9@users.noreply.github.com> Date: Wed, 26 Aug 2026 00:19:10 +0800 Subject: [PATCH 2/3] feat: push scalar subquery filters into scans Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> --- .../physical_optimizer/filter_pushdown.rs | 206 ++++- .../src/expressions/dynamic_filters/mod.rs | 9 + datafusion/physical-plan/src/proto.rs | 21 + .../physical-plan/src/scalar_subquery.rs | 711 +++++++++++++++++- .../proto-models/proto/datafusion.proto | 6 + .../proto-models/src/generated/pbjson.rs | 131 ++++ .../proto-models/src/generated/prost.rs | 11 + datafusion/proto/src/physical_plan/mod.rs | 11 + .../tests/cases/plans/scalar_subquery.rs | 136 ++++ 9 files changed, 1218 insertions(+), 24 deletions(-) diff --git a/datafusion/core/tests/physical_optimizer/filter_pushdown.rs b/datafusion/core/tests/physical_optimizer/filter_pushdown.rs index a98c1b7bcf98b..a2d18e105b287 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,14 +38,17 @@ use datafusion::{ use datafusion_catalog::memory::DataSourceExec; use datafusion_common::{ JoinType, + arrow::datatypes::IntervalMonthDayNano, config::ConfigOptions, tree_node::{TreeNode, TreeNodeRecursion}, }; use datafusion_datasource::{ - PartitionedFile, file_groups::FileGroup, file_scan_config::FileScanConfigBuilder, + PartitionedFile, file::FileSource, file_groups::FileGroup, + file_scan_config::FileScanConfigBuilder, }; 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 +57,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 +67,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 +79,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 +3073,197 @@ fn test_hashjoin_dynamic_filter_pushdown_is_used() { } } +/// Regression test for a scalar-subquery dynamic filter on a real Parquet source. +/// +/// With Parquet row-filter pushdown disabled, the source still retains the +/// predicate for statistics pruning and reports it as `PushedDown::No`. The +/// scalar-subquery producer must therefore retain its binding until execution +/// updates and completes the dynamic filter. +#[tokio::test] +async fn scalar_subquery_dynamic_filter_parquet_pruning_with_pushdown_disabled() { + let schema = Arc::new(Schema::new(vec![Field::new( + "ts", + DataType::Timestamp(TimeUnit::Nanosecond, None), + false, + )])); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(TimestampNanosecondArray::from(vec![ + 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::new(Schema::new(vec![Field::new( + "value", + DataType::Timestamp(TimeUnit::Nanosecond, None), + false, + )])), + vec![Arc::new(TimestampNanosecondArray::from(vec![ + 1_500_000_000, + ]))], + ) + .unwrap(); + let producer = datafusion::datasource::memory::MemorySourceConfig::try_new_exec( + &[vec![producer_batch]], + Arc::new(Schema::new(vec![Field::new( + "value", + DataType::Timestamp(TimeUnit::Nanosecond, None), + false, + )])), + None, + ) + .unwrap(); + + let scalar = Arc::new(ScalarSubqueryExpr::new( + DataType::Timestamp(TimeUnit::Nanosecond, None), + false, + SubqueryIndex::new(0), + results.clone(), + )); + let scalar_minus_interval = Arc::new(BinaryExpr::new( + scalar, + Operator::Minus, + Arc::new(Literal::new(ScalarValue::IntervalMonthDayNano(Some( + IntervalMonthDayNano { + months: 0, + days: 0, + nanoseconds: 500_000_000, + }, + )))), + )); + let predicate = Arc::new(BinaryExpr::new( + col("ts", &schema).unwrap(), + Operator::GtEq, + scalar_minus_interval, + )) as Arc; + let main = + Arc::new(FilterExec::try_new(predicate, scan).unwrap()) as Arc; + let plan = Arc::new(ScalarSubqueryExec::new( + main, + 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 (binding_predicate, consumer) = scalar_exec + .dynamic_filter_bindings() + .into_iter() + .next() + .expect("scalar subquery producer should retain its dynamic filter binding"); + let expression_id = consumer + .expression_id() + .expect("dynamic filter should have an expression ID"); + + let mut scan_predicate = None; + optimized + .apply(|node| { + if let Some(scan) = node.downcast_ref::() + && let Some((_, parquet)) = + scan.downcast_to_file_source::() + { + scan_predicate = parquet.filter(); + } + Ok(TreeNodeRecursion::Continue) + }) + .unwrap(); + let scan_predicate = + scan_predicate.expect("Parquet scan should retain pruning predicate"); + let mut found_consumer = false; + scan_predicate + .apply(|expr| { + if expr.expression_id() == Some(expression_id) { + found_consumer = true; + Ok(TreeNodeRecursion::Stop) + } else { + Ok(TreeNodeRecursion::Continue) + } + }) + .unwrap(); + assert!( + found_consumer, + "scan predicate should retain the dynamic filter ID" + ); + assert!(format_plan_for_test(&optimized).contains("dynamic_rg_pruning=eligible")); + + 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, context.task_ctx()).await.unwrap(); + assert_batches_eq!( + &[ + "+---------------------+", + "| ts |", + "+---------------------+", + "| 1970-01-01T00:00:02 |", + "+---------------------+", + ], + &batches + ); + + let current = consumer + .downcast_ref::() + .expect("binding consumer should be a dynamic filter") + .current() + .unwrap(); + assert_ne!(current.to_string(), "true"); + consumer + .downcast_ref::() + .unwrap() + .wait_complete() + .await; + assert!( + binding_predicate + .to_string() + .contains("scalar_subquery(1500000000)") + ); +} + /// 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..1bdd0003466ba 100644 --- a/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs +++ b/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs @@ -183,6 +183,15 @@ 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. + 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 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..5f78c4324c67c 100644 --- a/datafusion/physical-plan/src/scalar_subquery.rs +++ b/datafusion/physical-plan/src/scalar_subquery.rs @@ -25,15 +25,31 @@ //! [`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::utils::collect_columns; +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 +108,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 +139,7 @@ impl ScalarSubqueryExec { subqueries, subquery_future: Arc::default(), results, + bindings: Arc::default(), cache, } } @@ -124,6 +156,23 @@ impl ScalarSubqueryExec { &self.results } + /// Returns the dynamic-filter bindings associated with this execution plan. + pub fn dynamic_filter_bindings( + &self, + ) -> Vec<(Arc, Arc)> { + self.bindings + .lock() + .unwrap() + .iter() + .map(|binding| { + ( + Arc::clone(&binding.predicate), + Arc::clone(&binding.consumers[0]), + ) + }) + .collect() + } + /// Returns a per-child bool vec that is `true` for the main input /// (child 0) and `false` for every subquery child. fn true_for_input_only(&self) -> Vec { @@ -131,6 +180,63 @@ 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 children = collect_columns(predicate) + .into_iter() + .map(|column| Arc::new(column) as Arc) + .collect(); + let filter = Arc::new(DynamicFilterPhysicalExpr::new(children, lit(true))); + let _ = left; + 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 +289,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 +306,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 +390,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 +417,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 +478,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 +538,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 +604,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 +619,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 +636,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 +714,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 +725,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 +951,293 @@ 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_removes_binding_when_input_lacks_filter() + -> 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 = Arc::new(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::No + .wrap_expression(Arc::clone(&filter) as Arc), + ]], + }, + &ConfigOptions::default(), + )?; + assert!(exec.dynamic_expressions_produced().is_empty()); + reset_plan_states(exec)?; + Ok(()) + } + + #[test] + fn test_scalar_subquery_filter_pushdown_no_retains_binding_when_input_contains_filter() + -> 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 subquery_plan = make_subquery_plan(vec![int32_batch(vec![1])]); + let exec = Arc::new(single_subquery_exec(input, subquery_plan, results)); + exec.discover_bindings()?; + let filter = dynamic_filter(&exec); + + let updated_input = Arc::new(ExpressionExec::new( + placeholder_input(), + vec![predicate, Arc::clone(&filter) as Arc], + )); + let subquery_plan = Arc::clone(&exec.subqueries()[0].plan); + let updated_exec = exec.replace_children( + vec![updated_input, subquery_plan], + 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(), 1); + 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(), + )?; + + let produced = exec.dynamic_expressions_produced(); + assert_eq!(produced.len(), 1); + assert!(Arc::ptr_eq( + &produced[0], + &(Arc::clone(&filter) as Arc) + )); + + let mut roots = vec![]; + exec.apply_expressions(&mut |root| { + roots.push(Arc::clone(root)); + Ok(TreeNodeRecursion::Continue) + })?; + assert_eq!(roots.len(), 2); + assert!(roots.iter().any(|root| Arc::ptr_eq(root, &predicate))); + assert!(roots.iter().any(|root| Arc::ptr_eq( + root, + &(Arc::clone(&filter) as Arc) + ))); + 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 filter = Arc::new(DynamicFilterPhysicalExpr::new(vec![], lit(true))); + let derived_expr = Arc::clone(&filter).with_new_children(vec![])?; + let derived = derived_expr + .downcast_ref::() + .ok_or_else(|| { + DataFusionError::Internal("expected dynamic filter".to_string()) + })?; + assert!(filter.shares_runtime_state(&derived)); + + let bindings = Arc::new(Mutex::new(vec![ScalarSubqueryBinding { + predicate: lit(false), + consumers: vec![ + Arc::clone(&filter) as Arc, + Arc::clone(&derived_expr), + ], + }])); + execute_subqueries( + vec![], + ScalarSubqueryResults::new(0), + bindings, + Arc::new(TaskContext::default()), + ) + .await?; + + assert_eq!(filter.snapshot_generation(), 2); + assert_eq!(derived.snapshot_generation(), 2); + tokio::time::timeout(std::time::Duration::from_secs(1), filter.wait_complete()) + .await + .expect("shared filter should be complete"); + tokio::time::timeout(std::time::Duration::from_secs(1), derived.wait_complete()) + .await + .expect("derived filter should be complete"); + 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..abd0a2a46a41f 100644 --- a/datafusion/proto/tests/cases/plans/scalar_subquery.rs +++ b/datafusion/proto/tests/cases/plans/scalar_subquery.rs @@ -26,8 +26,10 @@ 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::{ @@ -276,3 +278,137 @@ 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, including the shared filter instance used by the +/// input plan, and that execution updates that filter. +#[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 binding_filter = scalar_exec.dynamic_filter_bindings()[0].1.clone(); + // 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(&binding_filter), + Arc::clone(scalar_exec.input()), + )?) as Arc, + Arc::new(FilterExec::try_new( + binding_filter, + 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 mut initial_bindings = 0; + initial_plan.apply(|node| { + if let Some(exec) = node.downcast_ref::() { + let produced = exec.dynamic_expressions_produced(); + if !produced.is_empty() { + initial_bindings += produced.len(); + } + } + Ok(TreeNodeRecursion::Continue) + })?; + assert_eq!( + initial_bindings, 1, + "expected one scalar dynamic-filter binding" + ); + + 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 mut roundtripped_bindings = 0; + roundtripped.apply(|node| { + if let Some(exec) = node.downcast_ref::() { + let produced = exec.dynamic_expressions_produced(); + if produced.is_empty() { + return Ok(TreeNodeRecursion::Continue); + } + assert_eq!(produced.len(), 1); + let binding = exec.dynamic_filter_bindings(); + assert_eq!(binding.len(), 1); + let bound_filter = &binding[0].1; + let bound_filter_expr = bound_filter + .downcast_ref::() + .expect("binding should contain a dynamic filter"); + assert_eq!( + Some(bound_filter_expr.expression_id().unwrap()), + produced[0].expression_id() + ); + + let mut consumers = vec![]; + exec.input().apply(|child| { + child.apply_expressions(&mut |expr| { + if expr.expression_id() == produced[0].expression_id() + && expr.downcast_ref::().is_some() + { + consumers.push(Arc::clone(expr)); + } + Ok(TreeNodeRecursion::Continue) + })?; + Ok(TreeNodeRecursion::Continue) + })?; + assert_eq!(consumers.len(), 2); + assert!(!Arc::ptr_eq(&consumers[0], &consumers[1])); + assert!(Arc::ptr_eq(bound_filter, &consumers[0])); + roundtripped_bindings += 1; + } + Ok(TreeNodeRecursion::Continue) + })?; + assert_eq!(roundtripped_bindings, 1); + + 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 + ); + Ok(()) +} From 42c870c6172d77f6295cea0d059ccab5168df614 Mon Sep 17 00:00:00 2001 From: discord9 <55937128+discord9@users.noreply.github.com> Date: Wed, 26 Aug 2026 17:27:38 +0800 Subject: [PATCH 3/3] refactor: trim scalar filter pushdown coverage Signed-off-by: discord9 <55937128+discord9@users.noreply.github.com> --- .../physical_optimizer/filter_pushdown.rs | 146 +++++------ .../src/expressions/dynamic_filters/mod.rs | 11 +- .../physical-expr/src/scalar_subquery.rs | 17 -- .../physical-plan/src/scalar_subquery.rs | 227 ++++++++---------- .../tests/cases/plans/scalar_subquery.rs | 144 ++++------- datafusion/pruning/src/pruning_predicate.rs | 96 +------- 6 files changed, 217 insertions(+), 424 deletions(-) diff --git a/datafusion/core/tests/physical_optimizer/filter_pushdown.rs b/datafusion/core/tests/physical_optimizer/filter_pushdown.rs index a2d18e105b287..340e78d161f45 100644 --- a/datafusion/core/tests/physical_optimizer/filter_pushdown.rs +++ b/datafusion/core/tests/physical_optimizer/filter_pushdown.rs @@ -43,8 +43,7 @@ use datafusion_common::{ tree_node::{TreeNode, TreeNodeRecursion}, }; use datafusion_datasource::{ - PartitionedFile, file::FileSource, file_groups::FileGroup, - file_scan_config::FileScanConfigBuilder, + PartitionedFile, file_groups::FileGroup, file_scan_config::FileScanConfigBuilder, }; use datafusion_execution::object_store::ObjectStoreUrl; use datafusion_expr::ScalarUDF; @@ -3073,25 +3072,19 @@ fn test_hashjoin_dynamic_filter_pushdown_is_used() { } } -/// Regression test for a scalar-subquery dynamic filter on a real Parquet source. -/// -/// With Parquet row-filter pushdown disabled, the source still retains the -/// predicate for statistics pruning and reports it as `PushedDown::No`. The -/// scalar-subquery producer must therefore retain its binding until execution -/// updates and completes the dynamic filter. +/// 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", - DataType::Timestamp(TimeUnit::Nanosecond, None), + timestamp.clone(), false, )])); let batch = RecordBatch::try_new( Arc::clone(&schema), - vec![Arc::new(TimestampNanosecondArray::from(vec![ - 0, - 2_000_000_000, - ]))], + vec![Arc::new(TimestampNanosecondArray::from([0, 2_000_000_000]))], ) .unwrap(); @@ -3123,53 +3116,39 @@ async fn scalar_subquery_dynamic_filter_parquet_pruning_with_pushdown_disabled() let results = ScalarSubqueryResults::new(1); let producer_batch = RecordBatch::try_new( - Arc::new(Schema::new(vec![Field::new( - "value", - DataType::Timestamp(TimeUnit::Nanosecond, None), - false, - )])), - vec![Arc::new(TimestampNanosecondArray::from(vec![ - 1_500_000_000, - ]))], + 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::new(Schema::new(vec![Field::new( - "value", - DataType::Timestamp(TimeUnit::Nanosecond, None), - false, - )])), + Arc::clone(&schema), None, ) .unwrap(); - let scalar = Arc::new(ScalarSubqueryExpr::new( - DataType::Timestamp(TimeUnit::Nanosecond, None), + timestamp, false, SubqueryIndex::new(0), results.clone(), )); - let scalar_minus_interval = Arc::new(BinaryExpr::new( - scalar, - Operator::Minus, - Arc::new(Literal::new(ScalarValue::IntervalMonthDayNano(Some( - IntervalMonthDayNano { - months: 0, - days: 0, - nanoseconds: 500_000_000, - }, - )))), - )); let predicate = Arc::new(BinaryExpr::new( col("ts", &schema).unwrap(), Operator::GtEq, - scalar_minus_interval, + 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 main = - Arc::new(FilterExec::try_new(predicate, scan).unwrap()) as Arc; let plan = Arc::new(ScalarSubqueryExec::new( - main, + Arc::new(FilterExec::try_new(predicate, scan).unwrap()), vec![ScalarSubqueryLink { plan: producer, index: SubqueryIndex::new(0), @@ -3184,57 +3163,50 @@ async fn scalar_subquery_dynamic_filter_parquet_pruning_with_pushdown_disabled() 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 (binding_predicate, consumer) = scalar_exec - .dynamic_filter_bindings() - .into_iter() - .next() - .expect("scalar subquery producer should retain its dynamic filter binding"); - let expression_id = consumer + 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_predicate = None; - optimized + 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() { - scan_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_predicate = - scan_predicate.expect("Parquet scan should retain pruning predicate"); - let mut found_consumer = false; - scan_predicate - .apply(|expr| { - if expr.expression_id() == Some(expression_id) { - found_consumer = true; - Ok(TreeNodeRecursion::Stop) - } else { - Ok(TreeNodeRecursion::Continue) - } - }) - .unwrap(); - assert!( - found_consumer, - "scan predicate should retain the dynamic filter ID" - ); - assert!(format_plan_for_test(&optimized).contains("dynamic_rg_pruning=eligible")); + 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, context.task_ctx()).await.unwrap(); + let batches = collect(optimized.clone(), context.task_ctx()) + .await + .unwrap(); assert_batches_eq!( &[ "+---------------------+", @@ -3246,21 +3218,25 @@ async fn scalar_subquery_dynamic_filter_parquet_pruning_with_pushdown_disabled() &batches ); - let current = consumer + let consumer = scan_consumer .downcast_ref::() - .expect("binding consumer should be a dynamic filter") - .current() - .unwrap(); - assert_ne!(current.to_string(), "true"); - consumer - .downcast_ref::() - .unwrap() - .wait_complete() - .await; + .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!( - binding_predicate - .to_string() - .contains("scalar_subquery(1500000000)") + pruned >= 1, + "expected statistics pruning, observed {pruned}" ); } diff --git a/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs b/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs index 1bdd0003466ba..a712677c4c0b9 100644 --- a/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs +++ b/datafusion/physical-expr/src/expressions/dynamic_filters/mod.rs @@ -188,6 +188,7 @@ impl DynamicFilterPhysicalExpr { /// 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) } @@ -323,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 83465f4293023..4f1a837300cce 100644 --- a/datafusion/physical-expr/src/scalar_subquery.rs +++ b/datafusion/physical-expr/src/scalar_subquery.rs @@ -335,23 +335,6 @@ mod tests { Ok(()) } - #[test] - fn test_snapshot_reset_returns_to_pending() -> Result<()> { - let results = ScalarSubqueryResults::new(1); - let expr = ScalarSubqueryExpr::new( - DataType::Int64, - true, - SubqueryIndex::new(0), - results.clone(), - ); - results.set(SubqueryIndex::new(0), ScalarValue::Int64(Some(7)))?; - assert!(expr.snapshot()?.is_some()); - - results.clear(); - assert!(expr.snapshot().is_err()); - Ok(()) - } - #[test] fn test_identity_equality() { let results = make_results(vec![None, None]); diff --git a/datafusion/physical-plan/src/scalar_subquery.rs b/datafusion/physical-plan/src/scalar_subquery.rs index 5f78c4324c67c..5ed0eb6608982 100644 --- a/datafusion/physical-plan/src/scalar_subquery.rs +++ b/datafusion/physical-plan/src/scalar_subquery.rs @@ -38,7 +38,6 @@ use datafusion_physical_expr::expressions::{ BinaryExpr, Column, DynamicFilterPhysicalExpr, Literal, lit, }; use datafusion_physical_expr::scalar_subquery::ScalarSubqueryExpr; -use datafusion_physical_expr::utils::collect_columns; use datafusion_physical_expr_common::physical_expr::{ PhysicalExpr, snapshot_physical_expr, }; @@ -156,23 +155,6 @@ impl ScalarSubqueryExec { &self.results } - /// Returns the dynamic-filter bindings associated with this execution plan. - pub fn dynamic_filter_bindings( - &self, - ) -> Vec<(Arc, Arc)> { - self.bindings - .lock() - .unwrap() - .iter() - .map(|binding| { - ( - Arc::clone(&binding.predicate), - Arc::clone(&binding.consumers[0]), - ) - }) - .collect() - } - /// Returns a per-child bool vec that is `true` for the main input /// (child 0) and `false` for every subquery child. fn true_for_input_only(&self) -> Vec { @@ -200,12 +182,10 @@ impl ScalarSubqueryExec { if !ScalarSubqueryResults::ptr_eq(scalar.results(), &self.results) { return None; } - let children = collect_columns(predicate) - .into_iter() - .map(|column| Arc::new(column) as Arc) - .collect(); - let filter = Arc::new(DynamicFilterPhysicalExpr::new(children, lit(true))); - let _ = left; + 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], @@ -978,73 +958,52 @@ mod tests { } #[test] - fn test_scalar_subquery_filter_pushdown_no_removes_binding_when_input_lacks_filter() - -> 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 = Arc::new(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::No - .wrap_expression(Arc::clone(&filter) as Arc), - ]], - }, - &ConfigOptions::default(), - )?; - assert!(exec.dynamic_expressions_produced().is_empty()); - reset_plan_states(exec)?; - Ok(()) - } - - #[test] - fn test_scalar_subquery_filter_pushdown_no_retains_binding_when_input_contains_filter() - -> 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 subquery_plan = make_subquery_plan(vec![int32_batch(vec![1])]); - let exec = Arc::new(single_subquery_exec(input, subquery_plan, results)); - exec.discover_bindings()?; - let filter = dynamic_filter(&exec); + 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), + )?; - let updated_input = Arc::new(ExpressionExec::new( - placeholder_input(), - vec![predicate, Arc::clone(&filter) as Arc], - )); - let subquery_plan = Arc::clone(&exec.subqueries()[0].plan); - let updated_exec = exec.replace_children( - vec![updated_input, subquery_plan], - 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(), 1); + 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(()) } @@ -1076,24 +1035,7 @@ mod tests { &ConfigOptions::default(), )?; - let produced = exec.dynamic_expressions_produced(); - assert_eq!(produced.len(), 1); - assert!(Arc::ptr_eq( - &produced[0], - &(Arc::clone(&filter) as Arc) - )); - - let mut roots = vec![]; - exec.apply_expressions(&mut |root| { - roots.push(Arc::clone(root)); - Ok(TreeNodeRecursion::Continue) - })?; - assert_eq!(roots.len(), 2); - assert!(roots.iter().any(|root| Arc::ptr_eq(root, &predicate))); - assert!(roots.iter().any(|root| Arc::ptr_eq( - root, - &(Arc::clone(&filter) as Arc) - ))); + assert!(reset_plan_states(Arc::new(exec)).is_err()); Ok(()) } @@ -1127,21 +1069,37 @@ mod tests { #[tokio::test] async fn test_execute_subqueries_updates_shared_runtime_state_once() -> Result<()> { - let filter = Arc::new(DynamicFilterPhysicalExpr::new(vec![], lit(true))); - let derived_expr = Arc::clone(&filter).with_new_children(vec![])?; - let derived = derived_expr + 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)); + assert!(filter.shares_runtime_state(derived_1)); + assert!(filter.shares_runtime_state(derived_2)); let bindings = Arc::new(Mutex::new(vec![ScalarSubqueryBinding { - predicate: lit(false), - consumers: vec![ - Arc::clone(&filter) as Arc, - Arc::clone(&derived_expr), - ], + predicate: Arc::clone(&original_column), + consumers: vec![Arc::clone(&derived_expr_1), Arc::clone(&derived_expr_2)], }])); execute_subqueries( vec![], @@ -1152,13 +1110,36 @@ mod tests { .await?; assert_eq!(filter.snapshot_generation(), 2); - assert_eq!(derived.snapshot_generation(), 2); - tokio::time::timeout(std::time::Duration::from_secs(1), filter.wait_complete()) - .await - .expect("shared filter should be complete"); - tokio::time::timeout(std::time::Duration::from_secs(1), derived.wait_complete()) + 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("derived filter should be complete"); + .expect(message); + } Ok(()) } diff --git a/datafusion/proto/tests/cases/plans/scalar_subquery.rs b/datafusion/proto/tests/cases/plans/scalar_subquery.rs index abd0a2a46a41f..c9b9f2c5a3cc4 100644 --- a/datafusion/proto/tests/cases/plans/scalar_subquery.rs +++ b/datafusion/proto/tests/cases/plans/scalar_subquery.rs @@ -32,13 +32,7 @@ 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; @@ -78,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. @@ -175,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::() @@ -256,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:?}" @@ -280,8 +253,7 @@ async fn roundtrip_scalar_subquery_exec_with_default_converter_executes() -> Res } /// Verify that built-in protobuf round-tripping preserves scalar-subquery -/// dynamic-filter bindings, including the shared filter instance used by the -/// input plan, and that execution updates that filter. +/// 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; @@ -304,17 +276,21 @@ async fn roundtrip_scalar_subquery_dynamic_filter_binding_executes() -> Result<( let scalar_exec = initial_plan .downcast_ref::() .expect("expected ScalarSubqueryExec"); - let binding_filter = scalar_exec.dynamic_filter_bindings()[0].1.clone(); + 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(&binding_filter), + Arc::clone(&produced), Arc::clone(scalar_exec.input()), )?) as Arc, Arc::new(FilterExec::try_new( - binding_filter, + produced, Arc::clone(scalar_exec.input()), )?) as Arc, ])?; @@ -331,67 +307,33 @@ async fn roundtrip_scalar_subquery_dynamic_filter_binding_executes() -> Result<( datafusion::physical_plan::ChildrenPropertiesMode::Recompute, ), )?; - let mut initial_bindings = 0; - initial_plan.apply(|node| { - if let Some(exec) = node.downcast_ref::() { - let produced = exec.dynamic_expressions_produced(); - if !produced.is_empty() { - initial_bindings += produced.len(); - } - } - Ok(TreeNodeRecursion::Continue) - })?; - assert_eq!( - initial_bindings, 1, - "expected one scalar dynamic-filter binding" - ); - - 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 mut roundtripped_bindings = 0; - roundtripped.apply(|node| { - if let Some(exec) = node.downcast_ref::() { - let produced = exec.dynamic_expressions_produced(); - if produced.is_empty() { - return Ok(TreeNodeRecursion::Continue); - } - assert_eq!(produced.len(), 1); - let binding = exec.dynamic_filter_bindings(); - assert_eq!(binding.len(), 1); - let bound_filter = &binding[0].1; - let bound_filter_expr = bound_filter - .downcast_ref::() - .expect("binding should contain a dynamic filter"); - assert_eq!( - Some(bound_filter_expr.expression_id().unwrap()), - produced[0].expression_id() - ); + 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 mut consumers = vec![]; - exec.input().apply(|child| { - child.apply_expressions(&mut |expr| { - if expr.expression_id() == produced[0].expression_id() - && expr.downcast_ref::().is_some() - { - consumers.push(Arc::clone(expr)); - } - Ok(TreeNodeRecursion::Continue) - })?; + 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); - assert!(!Arc::ptr_eq(&consumers[0], &consumers[1])); - assert!(Arc::ptr_eq(bound_filter, &consumers[0])); - roundtripped_bindings += 1; - } - Ok(TreeNodeRecursion::Continue) + }) + }) })?; - assert_eq!(roundtripped_bindings, 1); + assert_eq!(consumers.len(), 2); let batches = datafusion::physical_plan::collect_partitioned(roundtripped, ctx.task_ctx()) @@ -410,5 +352,15 @@ async fn roundtrip_scalar_subquery_dynamic_filter_binding_executes() -> Result<( ], &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(()) } diff --git a/datafusion/pruning/src/pruning_predicate.rs b/datafusion/pruning/src/pruning_predicate.rs index 37d750ad3c328..3a63451495e4c 100644 --- a/datafusion/pruning/src/pruning_predicate.rs +++ b/datafusion/pruning/src/pruning_predicate.rs @@ -2187,16 +2187,10 @@ mod tests { use arrow::array::Decimal128Array; use arrow::{ - array::{ - BinaryArray, Int32Array, Int64Array, StringArray, TimestampNanosecondArray, - UInt64Array, - }, - datatypes::{IntervalMonthDayNano, TimeUnit}, + array::{BinaryArray, Int32Array, Int64Array, StringArray, UInt64Array}, + datatypes::TimeUnit, }; use datafusion_expr::expr::InList; - use datafusion_expr::physical_planning_context::{ - ScalarSubqueryResults, SubqueryIndex, - }; use datafusion_expr::{BinaryExpr, Expr, cast, is_null, try_cast}; use datafusion_functions_nested::expr_fn::{array_has, make_array}; use datafusion_physical_expr::expressions::{ @@ -2617,92 +2611,6 @@ mod tests { } } - #[test] - fn scalar_subquery_timestamp_snapshot_builds_pruning_predicate() -> Result<()> { - let schema = Arc::new(Schema::new(vec![Field::new( - "ts", - DataType::Timestamp(TimeUnit::Nanosecond, None), - true, - )])); - let results = ScalarSubqueryResults::new(1); - let scalar = datafusion_physical_expr::scalar_subquery::ScalarSubqueryExpr::new( - DataType::Timestamp(TimeUnit::Nanosecond, None), - true, - SubqueryIndex::new(0), - results.clone(), - ); - let timestamp = 1_700_000_000_000_000_000; - let expected_bound = timestamp - 60_000_000_000; - results.set( - SubqueryIndex::new(0), - ScalarValue::TimestampNanosecond(Some(timestamp), None), - )?; - - let predicate: Arc = Arc::new(phys_expr::BinaryExpr::new( - phys_expr::col("ts", &schema)?, - Operator::GtEq, - Arc::new(phys_expr::BinaryExpr::new( - Arc::new(scalar), - Operator::Minus, - Arc::new(phys_expr::Literal::new(ScalarValue::IntervalMonthDayNano( - Some(IntervalMonthDayNano { - months: 0, - days: 0, - nanoseconds: 60_000_000_000, - }), - ))), - )), - )); - - let snapshot = snapshot_physical_expr_opt(Arc::clone(&predicate))?; - assert!(snapshot.transformed); - let simplified = PhysicalExprSimplifier::new(&schema).simplify(snapshot.data)?; - let bound = simplified - .downcast_ref::() - .expect("simplified predicate should remain a binary comparison") - .right() - .downcast_ref::() - .expect("timestamp bound should be folded to a literal"); - assert_eq!( - bound.value(), - &ScalarValue::TimestampNanosecond(Some(expected_bound), None) - ); - - let pruning = PruningPredicateBuilder::new() - .with_file_schema(Arc::clone(&schema)) - .try_build(predicate)?; - assert_eq!( - pruning.orig_expr().to_string(), - format!( - "ts@0 >= {}", - ScalarValue::TimestampNanosecond(Some(expected_bound), None) - ) - ); - let below = TestStatistics::new().with( - "ts", - ContainerStats::new() - .with_min(Arc::new(TimestampNanosecondArray::from(vec![Some( - expected_bound - 1, - )]))) - .with_max(Arc::new(TimestampNanosecondArray::from(vec![Some( - expected_bound - 1, - )]))), - ); - let matching = TestStatistics::new().with( - "ts", - ContainerStats::new() - .with_min(Arc::new(TimestampNanosecondArray::from(vec![Some( - expected_bound, - )]))) - .with_max(Arc::new(TimestampNanosecondArray::from(vec![Some( - expected_bound, - )]))), - ); - assert_eq!(pruning.prune(&below)?, vec![false]); - assert_eq!(pruning.prune(&matching)?, vec![true]); - Ok(()) - } - #[test] fn prune_all_rows_null_counts() { // if null_count = row_count then we should prune the container for i = 0