From 8dba6b2d6f432aedbfa8528f63d73ffd4c9109bb Mon Sep 17 00:00:00 2001 From: Jayant Shrivastava Date: Mon, 17 Aug 2026 19:17:21 +0000 Subject: [PATCH 1/2] feat: plan distributed dynamic filters MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Before stage split Producer worker plan HashJoin producer F HashJoin producer F └── NetworkShuffle └── NetworkShuffle └── Local stage ├── Remote stage └── consumer F └── anchor F | apply_expressions() finds F Register task-specialized producer and consumer topology independently from completed-consumer display state. Preserve consumers crossing a stage boundary as non-evaluating network-boundary anchors so both static and dynamic planning retain the information needed to identify remote filters. --- src/codec/distributed_codec.rs | 149 ++++- src/codec/physical_plan.rs | 9 +- src/common/recursion.rs | 1 + src/coordinator/distributed.rs | 7 +- src/coordinator/dynamic_filter_registry.rs | 71 +++ src/coordinator/mod.rs | 2 + src/coordinator/prepare_dynamic_plan.rs | 3 + src/coordinator/prepare_static_plan.rs | 3 + src/coordinator/query_coordinator.rs | 26 +- src/dynamic_filtering/discovery.rs | 512 ++++++++++++------ src/dynamic_filtering/display.rs | 12 +- src/dynamic_filtering/mod.rs | 178 +++++- .../benchmarks/shuffle_bench.rs | 1 + .../benchmarks/transport_bench.rs | 1 + src/execution_plans/network_broadcast.rs | 6 +- src/execution_plans/network_coalesce.rs | 7 +- src/execution_plans/network_shuffle.rs | 6 +- src/stage.rs | 10 + src/test_utils/routing.rs | 34 +- src/worker/impl_coordinator_channel.rs | 2 +- tests/dynamic_filtering/common.rs | 15 +- tests/dynamic_filtering/partitioned_join.rs | 39 +- 22 files changed, 850 insertions(+), 244 deletions(-) create mode 100644 src/coordinator/dynamic_filter_registry.rs diff --git a/src/codec/distributed_codec.rs b/src/codec/distributed_codec.rs index cb22e0fdd..95f2b507e 100644 --- a/src/codec/distributed_codec.rs +++ b/src/codec/distributed_codec.rs @@ -14,8 +14,8 @@ use datafusion::arrow::datatypes::SchemaRef; use datafusion::common::Result; use datafusion::error::DataFusionError; use datafusion::execution::TaskContext; -use datafusion::physical_expr::EquivalenceProperties; use datafusion::physical_expr::equivalence::{EquivalenceClass, EquivalenceGroup}; +use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr}; use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType}; use datafusion::physical_plan::union::UnionExec; use datafusion::physical_plan::{ExecutionPlan, Partitioning, PlanProperties}; @@ -70,11 +70,17 @@ impl PhysicalExtensionCodec for DistributedCodec { fn parse_stage_proto( proto: Option, inputs: &[Arc], + dynamic_filter_anchors: Vec>, ) -> Result { let Some(proto) = proto else { return Err(proto_error("Empty StageProto")); }; if let Some(input) = inputs.first().cloned() { + if !dynamic_filter_anchors.is_empty() { + return Err(proto_error( + "Dynamic filter anchors require a remote input stage", + )); + } Ok(Stage::Local(LocalStage { query_id: deserialize_uuid(proto.query_id.as_ref())?, num: proto.num as usize, @@ -98,6 +104,7 @@ impl PhysicalExtensionCodec for DistributedCodec { num: proto.num as usize, workers: worker_urls, runtime_stats: None, + dynamic_filter_anchors, })) } } @@ -123,8 +130,19 @@ impl PhysicalExtensionCodec for DistributedCodec { proto_converter, )? .ok_or(proto_error("NetworkShuffleExec is missing partitioning"))?; + let sort_exprs = parse_physical_sort_exprs(&ordering, &decode_ctx, &schema, proto_converter)?; + + let dynamic_filter_anchors = input_stage + .as_ref() + .into_iter() + .flat_map(|stage| stage.dynamic_filter_anchors.iter()) + .map(|expression| { + proto_converter.proto_to_physical_expr(expression, &schema, &decode_ctx) + }) + .collect::>>()?; + let schema = Arc::new(schema); let mut equivalence_properties = parse_equivalence_properties( equivalence_classes, @@ -140,7 +158,7 @@ impl PhysicalExtensionCodec for DistributedCodec { Ok(Arc::new(new_network_hash_shuffle_exec( partitioning, equivalence_properties, - parse_stage_proto(input_stage, inputs)?, + parse_stage_proto(input_stage, inputs, dynamic_filter_anchors)?, ))) } DistributedExecNode::NetworkCoalesceTasks(NetworkCoalesceExecProto { @@ -162,6 +180,14 @@ impl PhysicalExtensionCodec for DistributedCodec { proto_converter, )? .ok_or(proto_error("NetworkCoalesceExec is missing partitioning"))?; + let dynamic_filter_anchors = input_stage + .as_ref() + .into_iter() + .flat_map(|stage| stage.dynamic_filter_anchors.iter()) + .map(|expression| { + proto_converter.proto_to_physical_expr(expression, &schema, &decode_ctx) + }) + .collect::>>()?; let schema = Arc::new(schema); let equivalence_properties = parse_equivalence_properties( equivalence_classes, @@ -173,7 +199,7 @@ impl PhysicalExtensionCodec for DistributedCodec { Ok(Arc::new(new_network_coalesce_tasks_exec( partitioning, equivalence_properties, - parse_stage_proto(input_stage, inputs)?, + parse_stage_proto(input_stage, inputs, dynamic_filter_anchors)?, ))) } DistributedExecNode::NetworkBroadcast(NetworkBroadcastExecProto { @@ -195,6 +221,14 @@ impl PhysicalExtensionCodec for DistributedCodec { proto_converter, )? .ok_or(proto_error("NetworkBroadcastExec is missing partitioning"))?; + let dynamic_filter_anchors = input_stage + .as_ref() + .into_iter() + .flat_map(|stage| stage.dynamic_filter_anchors.iter()) + .map(|expression| { + proto_converter.proto_to_physical_expr(expression, &schema, &decode_ctx) + }) + .collect::>>()?; let schema = Arc::new(schema); let equivalence_properties = parse_equivalence_properties( equivalence_classes, @@ -206,7 +240,7 @@ impl PhysicalExtensionCodec for DistributedCodec { Ok(Arc::new(new_network_broadcast_exec( partitioning, equivalence_properties, - parse_stage_proto(input_stage, inputs)?, + parse_stage_proto(input_stage, inputs, dynamic_filter_anchors)?, ))) } DistributedExecNode::Broadcast(BroadcastExecProto { @@ -284,12 +318,22 @@ impl PhysicalExtensionCodec for DistributedCodec { buf: &mut Vec, proto_converter: &dyn PhysicalProtoConverterExtension, ) -> Result<()> { - fn encode_stage_proto(stage: &Stage) -> Result { + fn encode_stage_proto( + stage: &Stage, + codec: &DistributedCodec, + proto_converter: &dyn PhysicalProtoConverterExtension, + ) -> Result { + let dynamic_filter_anchors = stage + .dynamic_filter_anchors() + .iter() + .map(|expression| proto_converter.physical_expr_to_proto(expression, codec)) + .collect::>>()?; Ok(match stage { Stage::Local(local) => StageProto { query_id: serialize_uuid(&local.query_id).into(), num: local.num as u64, tasks: vec![ExecutionTaskProto::default(); local.tasks], + dynamic_filter_anchors, }, Stage::Remote(remote) => { let mut tasks = Vec::with_capacity(remote.workers.len()); @@ -302,6 +346,7 @@ impl PhysicalExtensionCodec for DistributedCodec { query_id: serialize_uuid(&remote.query_id).into(), num: remote.num as u64, tasks, + dynamic_filter_anchors, } } }) @@ -325,7 +370,11 @@ impl PhysicalExtensionCodec for DistributedCodec { self, proto_converter, )?), - input_stage: Some(encode_stage_proto(node.input_stage())?), + input_stage: Some(encode_stage_proto( + node.input_stage(), + self, + proto_converter, + )?), equivalence_classes: serialize_equivalence_group( node.properties().equivalence_properties(), self, @@ -347,7 +396,11 @@ impl PhysicalExtensionCodec for DistributedCodec { self, proto_converter, )?), - input_stage: Some(encode_stage_proto(node.input_stage())?), + input_stage: Some(encode_stage_proto( + node.input_stage(), + self, + proto_converter, + )?), equivalence_classes: serialize_equivalence_group( node.properties().equivalence_properties(), self, @@ -368,7 +421,11 @@ impl PhysicalExtensionCodec for DistributedCodec { self, proto_converter, )?), - input_stage: Some(encode_stage_proto(node.input_stage())?), + input_stage: Some(encode_stage_proto( + node.input_stage(), + self, + proto_converter, + )?), equivalence_classes: serialize_equivalence_group( node.properties().equivalence_properties(), self, @@ -492,6 +549,9 @@ pub struct StageProto { /// the plan #[prost(message, repeated, tag = "3")] pub tasks: Vec, + /// Dynamic-filter consumers retained after a remote stage's plan has moved to its workers. + #[prost(message, repeated, tag = "4")] + pub dynamic_filter_anchors: Vec, } #[derive(Clone, PartialEq, ::prost::Message)] @@ -673,14 +733,20 @@ fn new_network_broadcast_exec( #[cfg(test)] mod tests { - use super::super::physical_plan::new_proto_converter as default_proto_converter; + use super::super::physical_plan::{ + new_proto_converter as default_proto_converter, roundtrip_pb, + }; use super::*; use datafusion::arrow::datatypes::{DataType, Field}; use datafusion::physical_expr::{LexOrdering, PhysicalExpr}; use datafusion::physical_plan::empty::EmptyExec; + use datafusion::physical_plan::filter::FilterExec; use datafusion::prelude::SessionContext; use datafusion::{ - physical_expr::{Partitioning, PhysicalSortExpr, expressions::Column, expressions::col}, + physical_expr::{ + Partitioning, PhysicalSortExpr, + expressions::{Column, DynamicFilterPhysicalExpr, col, lit}, + }, physical_plan::{ExecutionPlan, displayable, sorts::sort::SortExec, union::UnionExec}, }; @@ -694,6 +760,7 @@ mod tests { num: 0, workers: vec![], runtime_stats: None, + dynamic_filter_anchors: vec![], }) } @@ -741,6 +808,68 @@ mod tests { Ok(()) } + #[test] + fn test_roundtrip_network_dynamic_filter_anchor() -> datafusion::common::Result<()> { + let ctx = create_context(); + let schema = schema_i32("a"); + let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new( + vec![Arc::new(Column::new("a", 0))], + lit(true), + )) as Arc; + let expected_id = dynamic_filter.expression_id(); + let stage = Stage::Remote(RemoteStage { + query_id: Default::default(), + num: 0, + workers: vec![], + runtime_stats: None, + dynamic_filter_anchors: vec![Arc::clone(&dynamic_filter)], + }); + let network: Arc = Arc::new(new_network_hash_shuffle_exec( + Partitioning::Hash(vec![Arc::new(Column::new("a", 0))], 4), + EquivalenceProperties::new(schema), + stage, + )); + + let mut buf = vec![]; + DistributedCodec.try_encode(Arc::clone(&network), &mut buf, &default_proto_converter())?; + let encoded = DistributedExecProto::decode(buf.as_slice()) + .map_err(|error| proto_error(format!("{error}")))?; + let Some(DistributedExecNode::NetworkHashShuffle(encoded)) = encoded.node else { + panic!("expected a network shuffle") + }; + assert_eq!( + encoded + .input_stage + .expect("network shuffle should contain its input stage") + .dynamic_filter_anchors + .len(), + 1, + ); + + let plan: Arc = Arc::new(FilterExec::try_new(dynamic_filter, network)?); + + let decoded = roundtrip_pb(plan, &ctx)?; + let filter = decoded.downcast_ref::().unwrap(); + let predicate = filter + .predicate() + .downcast_ref::() + .unwrap(); + let network = filter.input().downcast_ref::().unwrap(); + let anchor = network.input_stage().dynamic_filter_anchors()[0] + .downcast_ref::() + .unwrap(); + + assert_eq!(predicate.expression_id(), expected_id); + assert_eq!(anchor.expression_id(), expected_id); + predicate.update(lit(false))?; + assert_eq!( + anchor.current()?.to_string(), + "false", + "the filter predicate and network anchor should share state", + ); + Ok(()) + } + #[test] fn test_roundtrip_union() -> datafusion::common::Result<()> { let codec = DistributedCodec; diff --git a/src/codec/physical_plan.rs b/src/codec/physical_plan.rs index 632178cc3..f1a970ace 100644 --- a/src/codec/physical_plan.rs +++ b/src/codec/physical_plan.rs @@ -49,8 +49,13 @@ pub(crate) fn roundtrip_pb( plan: Arc, task_ctx: &TaskContext, ) -> Result> { - let encoded = encode_execution_plan(plan, task_ctx)?; - decode_execution_plan(&encoded, task_ctx) + let encode_codec = DistributedCodec::new_combined_with_user(task_ctx.session_config()); + let encode_converter = new_proto_converter(); + let proto = encode_converter.execution_plan_to_proto(&plan, &encode_codec)?; + let decode_codec = DistributedCodec::new_combined_with_user(task_ctx.session_config()); + let decode_ctx = PhysicalPlanDecodeContext::new(task_ctx, &decode_codec); + let decode_converter = new_proto_converter(); + decode_converter.proto_to_execution_plan(&proto, &decode_ctx) } pub(crate) fn encode_physical_expr( diff --git a/src/common/recursion.rs b/src/common/recursion.rs index b44df3462..b7f2ec6bd 100644 --- a/src/common/recursion.rs +++ b/src/common/recursion.rs @@ -1012,6 +1012,7 @@ mod tests { num: 0, workers: vec![], runtime_stats: None, + dynamic_filter_anchors: vec![], })) .unwrap() } diff --git a/src/coordinator/distributed.rs b/src/coordinator/distributed.rs index 3d9481f90..e2b2508b7 100644 --- a/src/coordinator/distributed.rs +++ b/src/coordinator/distributed.rs @@ -3,7 +3,9 @@ use crate::coordinator::prepare_dynamic_plan::prepare_dynamic_plan; use crate::coordinator::prepare_static_plan::prepare_static_plan; use crate::coordinator::query_coordinator::QueryCoordinator; use crate::coordinator::store::{Store, task_keys_for_plan}; -use crate::dynamic_filtering::sever_dynamic_filter_relationships_in_plan_for_display; +use crate::dynamic_filtering::{ + is_dynamic_filtering_enabled, sever_dynamic_filter_relationships_in_plan_for_display, +}; use crate::{DistributedConfig, TaskCompletedDynamicFilters, TaskKey, TaskMetrics}; use datafusion::common::internal_datafusion_err; use datafusion::common::tree_node::TreeNodeRecursion; @@ -249,7 +251,8 @@ impl ExecutionPlan for DistributedExec { false => prepare_static_plan(&query_coordinator, &base_plan).await?, }; - prepared.plan_for_viz = match collect_dynamic_filters { + let dynamic_filtering_enabled = is_dynamic_filtering_enabled(context.session_config()); + prepared.plan_for_viz = match dynamic_filtering_enabled && collect_dynamic_filters { true => sever_dynamic_filter_relationships_in_plan_for_display( prepared.plan_for_viz, &context, diff --git a/src/coordinator/dynamic_filter_registry.rs b/src/coordinator/dynamic_filter_registry.rs new file mode 100644 index 000000000..6642fdf65 --- /dev/null +++ b/src/coordinator/dynamic_filter_registry.rs @@ -0,0 +1,71 @@ +use crate::TaskKey; +use crate::dynamic_filtering::{ + discover_dynamic_filter_consumers, discover_dynamic_filter_producers, +}; +use datafusion::common::{HashMap, HashSet, Result}; +use datafusion::physical_plan::ExecutionPlan; +use std::sync::{Arc, Mutex}; + +#[derive(Default)] +pub(super) struct PlannedDynamicFilter { + // Producer and consumer tasks for a dynamic filter. + // + // Note that it is not guaranteed that every task within a stage produces / consumes dynamic filters. For + // example, a distributed union may prevent a dynamic filter from appearing in all tasks. So, we + // store task keys rather than stage ids. + pub(super) producer_tasks: HashSet, + pub(super) consumer_tasks: HashSet, +} + +#[derive(Default)] +pub(super) struct DynamicFilterRegistryState { + pub(super) filters: HashMap, +} + +/// Query-scoped hub for distributed dynamic filtering. +/// +/// It stores the locations of dynamic filters and their runtime state. Informs the coordinator +/// - where dynamic filter updates are coming from +/// - how/if dynamic filter updates should be merged +/// - where dynamic filter updates should be forwarded +#[derive(Default)] +pub(crate) struct DynamicFilterRegistry { + pub(super) state: Mutex, +} + +impl DynamicFilterRegistry { + pub(crate) fn new() -> Self { + Self::default() + } + + /// Adds any dynamic filter producers and consumers found in `plan` to the registry. + pub(crate) fn register_task( + &self, + plan: &Arc, + task_key: TaskKey, + ) -> Result<()> { + let producers = discover_dynamic_filter_producers(plan)?; + // We can safely ignore anchors because they are not evaluated by network boundaries. This + // means they do not need updates forwarded to them. + let consumers = discover_dynamic_filter_consumers(plan)?.consumers; + + let mut state = self.state.lock().expect("dynamic filter registry poisoned"); + for producer in producers { + state + .filters + .entry(producer.id) + .or_default() + .producer_tasks + .insert(task_key); + } + for consumer in consumers { + state + .filters + .entry(consumer.id) + .or_default() + .consumer_tasks + .insert(task_key); + } + Ok(()) + } +} diff --git a/src/coordinator/mod.rs b/src/coordinator/mod.rs index fd7706fda..746052a6f 100644 --- a/src/coordinator/mod.rs +++ b/src/coordinator/mod.rs @@ -1,4 +1,5 @@ mod distributed; +mod dynamic_filter_registry; mod latency_metric; mod prepare_dynamic_plan; mod prepare_static_plan; @@ -6,4 +7,5 @@ mod query_coordinator; mod store; pub use distributed::DistributedExec; +pub(crate) use dynamic_filter_registry::DynamicFilterRegistry; pub(crate) use store::Store; diff --git a/src/coordinator/prepare_dynamic_plan.rs b/src/coordinator/prepare_dynamic_plan.rs index 2ffb482db..bee8b2f70 100644 --- a/src/coordinator/prepare_dynamic_plan.rs +++ b/src/coordinator/prepare_dynamic_plan.rs @@ -5,6 +5,7 @@ use crate::distributed_planner::{ InjectNetworkBoundaryContext, NetworkBoundaryBuilderResult, ProducerHead, calculate_cost, inject_network_boundaries, }; +use crate::dynamic_filtering::orphan_dynamic_filter_consumers; use crate::events::TaskCountAnnotation::{Desired, Maximum}; use crate::execution_plans::SamplerExec; use crate::stage::{LocalStage, RemoteStage}; @@ -80,6 +81,7 @@ pub(super) async fn prepare_dynamic_plan( // In order to infer the compute the cost of the stage above this one, here a sampler // is injected to gather runtime statistics. input_stage.plan = ProducerHead::insert_sampler(input_stage.plan)?; + let dynamic_filter_anchors = orphan_dynamic_filter_consumers(&input_stage.plan)?; let mut load_info_rxs = Vec::with_capacity(input_stage.tasks); @@ -132,6 +134,7 @@ pub(super) async fn prepare_dynamic_plan( num: input_stage.num, workers, runtime_stats: stats, + dynamic_filter_anchors, }), input_properties, }) diff --git a/src/coordinator/prepare_static_plan.rs b/src/coordinator/prepare_static_plan.rs index 5bb35b63e..f83838d5c 100644 --- a/src/coordinator/prepare_static_plan.rs +++ b/src/coordinator/prepare_static_plan.rs @@ -1,6 +1,7 @@ use crate::common::TreeNodeExt; use crate::coordinator::distributed::PreparedPlan; use crate::coordinator::query_coordinator::QueryCoordinator; +use crate::dynamic_filtering::orphan_dynamic_filter_consumers; use crate::stage::RemoteStage; use crate::{NetworkBoundaryExt, Stage}; use datafusion::common::tree_node::Transformed; @@ -31,6 +32,7 @@ pub(super) async fn prepare_static_plan( let Stage::Local(stage) = plan.input_stage() else { return exec_err!("Input stage from network boundary was not in Local state"); }; + let dynamic_filter_anchors = orphan_dynamic_filter_consumers(&stage.plan)?; let mut stage_coordinator = query_coordinator.stage_coordinator(stage); let mut futures = Vec::with_capacity(stage.tasks); @@ -52,6 +54,7 @@ pub(super) async fn prepare_static_plan( num: stage.num, workers, runtime_stats: None, + dynamic_filter_anchors, }, ))?)) }); diff --git a/src/coordinator/query_coordinator.rs b/src/coordinator/query_coordinator.rs index 6bb6faddb..499dab7ad 100644 --- a/src/coordinator/query_coordinator.rs +++ b/src/coordinator/query_coordinator.rs @@ -1,9 +1,13 @@ use crate::codec::roundtrip_pb; use crate::common::{TreeNodeExt, now_ns, task_ctx_with_extension}; use crate::config_extension_ext::get_config_extension_propagation_headers; +use crate::coordinator::DynamicFilterRegistry; use crate::coordinator::Store; use crate::coordinator::latency_metric::LatencyMetric; -use crate::dynamic_filtering::maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships; +use crate::dynamic_filtering::{ + is_dynamic_filtering_enabled, + maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships, +}; use crate::events::{ RouteTaskEvent, RouteTaskEventResponse, RouteTaskHandlers, new_coordinator_to_worker_dialer, }; @@ -52,6 +56,7 @@ pub(super) struct QueryCoordinator { coordinator_to_worker_metrics: CoordinatorToWorkerMetrics, metrics_store: Option>>, completed_dynamic_filter_store: Option>>, + dynamic_filter_registry: Arc, end_stream_notifier: Arc, join_set: Mutex>>, } @@ -69,6 +74,7 @@ impl QueryCoordinator { metrics: metrics_set.clone(), metrics_store, completed_dynamic_filter_store, + dynamic_filter_registry: Arc::new(DynamicFilterRegistry::new()), coordinator_to_worker_metrics: CoordinatorToWorkerMetrics::new(metrics_set), end_stream_notifier: Arc::new(Notify::new()), join_set: Mutex::new(JoinSet::new()), @@ -88,6 +94,7 @@ impl QueryCoordinator { metrics: &self.coordinator_to_worker_metrics, metrics_store: &self.metrics_store, completed_dynamic_filter_store: &self.completed_dynamic_filter_store, + dynamic_filter_registry: &self.dynamic_filter_registry, end_stream_notifier: &self.end_stream_notifier, join_set: &self.join_set, } @@ -134,6 +141,7 @@ pub(super) struct StageCoordinator<'a> { metrics: &'a CoordinatorToWorkerMetrics, metrics_store: &'a Option>>, completed_dynamic_filter_store: &'a Option>>, + dynamic_filter_registry: &'a Arc, end_stream_notifier: &'a Arc, join_set: &'a Mutex>>, } @@ -163,6 +171,9 @@ impl<'a> StageCoordinator<'a> { task_number: task_i, }; + self.dynamic_filter_registry + .register_task(&specialized, task_key)?; + let mut headers = get_config_extension_propagation_headers(session_config)?; headers.extend(get_passthrough_headers(session_config)); @@ -415,6 +426,7 @@ impl<'a> StageCoordinator<'a> { let wuf_registry = session_config .get_extension::() .unwrap_or_default(); + let dynamic_filtering_enabled = is_dynamic_filtering_enabled(session_config); let mut work_unit_feed_declarations = vec![]; let d_ctx = DistributedTaskContext { @@ -452,10 +464,14 @@ impl<'a> StageCoordinator<'a> { Ok(Transformed::no(plan)) })?; - let plan = maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships( - Arc::clone(&transformed.data), - self.task_ctx, - )?; + let plan = if dynamic_filtering_enabled { + maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships( + Arc::clone(&transformed.data), + self.task_ctx, + )? + } else { + transformed.data + }; Ok((plan, work_unit_feed_declarations)) } } diff --git a/src/dynamic_filtering/discovery.rs b/src/dynamic_filtering/discovery.rs index e07b052ad..a3b011c33 100644 --- a/src/dynamic_filtering/discovery.rs +++ b/src/dynamic_filtering/discovery.rs @@ -1,3 +1,4 @@ +use crate::NetworkBoundaryExt; use datafusion::arrow::datatypes::SchemaRef; use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion}; use datafusion::common::{HashMap, HashSet, Result, internal_err}; @@ -6,6 +7,12 @@ use datafusion::physical_expr::expressions::DynamicFilterPhysicalExpr; use datafusion::physical_plan::ExecutionPlan; use std::sync::Arc; +/// A dynamic filter produced by an [`ExecutionPlan`]. +#[derive(Clone)] +pub(crate) struct DiscoveredDynamicFilterProducer { + pub(crate) id: u64, +} + /// A dynamic-filter consumer discovered in an execution plan along with the schema it is evaluated /// against. #[derive(Clone)] @@ -15,11 +22,30 @@ pub(crate) struct DiscoveredDynamicFilter { pub(crate) input_schema: SchemaRef, } -/// Finds dynamic-filter consumers in `plan`, deduplicated by expression ID. +/// An anchor is an artificial dynamic filter consumer injected into network boundaries +/// to keep consumer references alive when they are moved across network boundaries. +/// +/// TODO(#697): remove anchors in df-56. +#[derive(Clone)] +pub(crate) struct DiscoveredDynamicFilterAnchor { + pub(crate) id: u64, + pub(crate) expression: Arc, +} + +pub(crate) struct DiscoveredDynamicFilterConsumers { + // Real consumers, ordered by expression id. + pub(crate) consumers: Vec, + // Artificial consumers. Dynamic filters in network boundaries. Also ordered by expression id. + pub(crate) anchors: Vec, +} + +/// Finds dynamic-filter consumers and network-boundary anchors in `plan`, deduplicated by +/// expression ID within each category. pub(crate) fn discover_dynamic_filter_consumers( plan: &Arc, -) -> Result> { +) -> Result { let mut consumers = HashMap::new(); + let mut anchors = HashMap::new(); plan.apply(|node| { let produced_ids: HashSet<_> = node @@ -40,6 +66,7 @@ pub(crate) fn discover_dynamic_filter_consumers( .first() .map(|child| child.schema()) .unwrap_or_else(|| node.schema()); + let is_network_boundary = node.is_network_boundary(); node.apply_expressions(&mut |root| { root.apply(|expression| { @@ -53,8 +80,16 @@ pub(crate) fn discover_dynamic_filter_consumers( "DynamicFilterPhysicalExpr did not have an expression ID" ); }; - let is_producer_occurrence = produced_ids.contains(&id); - if !is_producer_occurrence { + if is_network_boundary { + // Network-boundary expressions are metadata-only dependencies, not expressions + // evaluated by the node. + anchors + .entry(id) + .or_insert_with(|| DiscoveredDynamicFilterAnchor { + id, + expression: expression.clone(), + }); + } else if !produced_ids.contains(&id) { consumers .entry(id) .or_insert_with(|| DiscoveredDynamicFilter { @@ -72,207 +107,352 @@ pub(crate) fn discover_dynamic_filter_consumers( let mut consumers: Vec<_> = consumers.into_values().collect(); consumers.sort_unstable_by_key(|consumer| consumer.id); - Ok(consumers) + let mut anchors: Vec<_> = anchors.into_values().collect(); + anchors.sort_unstable_by_key(|anchor| anchor.id); + Ok(DiscoveredDynamicFilterConsumers { consumers, anchors }) } -/// Returns whether `plan` contains only the consumer side of a dynamic filter -/// relationship. -pub(crate) fn has_nonlocal_dynamic_filter_relationships( +/// Finds dynamic-filter producers in `plan`, deduplicated and ordered by expression ID. +pub(crate) fn discover_dynamic_filter_producers( plan: &Arc, -) -> Result { - let consumer_ids: HashSet<_> = discover_dynamic_filter_consumers(plan)? - .into_iter() - .map(|consumer| consumer.id) - .collect(); - - let mut producer_ids = HashSet::new(); +) -> Result> { + let mut producers = HashMap::new(); plan.apply(|node| { - for produced in node.dynamic_expressions_produced() { - let Some(id) = produced.expression_id() else { - return internal_err!( - "{}::dynamic_expressions_produced returned an expression without an expression ID", - node.name() - ); + for expression in node.dynamic_expressions_produced() { + if expression + .downcast_ref::() + .is_none() + { + continue; + } + let Some(id) = expression.expression_id() else { + return internal_err!("DynamicFilterPhysicalExpr did not have an expression ID"); }; - producer_ids.insert(id); + producers + .entry(id) + .or_insert(DiscoveredDynamicFilterProducer { id }); } Ok(TreeNodeRecursion::Continue) })?; - Ok(consumer_ids != producer_ids) + let mut producers: Vec<_> = producers.into_values().collect(); + producers.sort_unstable_by_key(|producer| producer.id); + Ok(producers) +} + +/// Finds consumers whose producer does not occur in `plan`. These consumers become orphaned +/// from their producer when the producer is moved behind a remote network boundary. These +/// orphans become network boundary anchors, artificially keeping the producers alive. +/// +/// TODO(697): remove anchors in df-56 +pub(crate) fn orphan_dynamic_filter_consumers( + plan: &Arc, +) -> Result>> { + let produced_here: HashSet<_> = discover_dynamic_filter_producers(plan)? + .into_iter() + .map(|producer| producer.id) + .collect(); + let discovered = discover_dynamic_filter_consumers(plan)?; + // Include anchors here because we want anchors to work recursively. For example, + // if a producer is in stage 4 and its consumer is in stage 1, an + // anchor should exist in stage 4. The easiest way to guarantee that is to ensure + // the anchor exists in stages 2, 3, and 4 recursively via this function. + let orphaned: HashMap<_, _> = discovered + .consumers + .into_iter() + .map(|consumer| (consumer.id, consumer.expression as Arc)) + .chain( + discovered + .anchors + .into_iter() + .map(|anchor| (anchor.id, anchor.expression)), + ) + .filter(|(id, _)| !produced_here.contains(id)) + .collect(); + let mut orphaned: Vec<_> = orphaned.into_iter().collect(); + orphaned.sort_unstable_by_key(|(id, _)| *id); + Ok(orphaned + .into_iter() + .map(|(_, expression)| expression) + .collect()) } #[cfg(test)] mod tests { use super::*; - use datafusion::arrow::datatypes::{DataType, Field, Schema}; - use datafusion::common::Result; - use datafusion::execution::{SendableRecordBatchStream, TaskContext}; - use datafusion::logical_expr::Operator; - use datafusion::physical_expr::expressions::{BinaryExpr, Column, lit}; - use datafusion::physical_plan::empty::EmptyExec; - use datafusion::physical_plan::union::UnionExec; - use datafusion::physical_plan::{ - DisplayAs, DisplayFormatType, PlanProperties, apply_expression_roots, + use crate::test_utils::localhost::start_localhost_context; + use crate::test_utils::parquet::register_parquet_tables; + use crate::{ + DefaultSessionBuilder, DistributedExt, RouteTaskEvent, RouteTaskEventResponse, + RouteTaskHandler, assert_snapshot, }; - use std::fmt::Formatter; + use async_trait::async_trait; + use datafusion::physical_plan::collect; + use itertools::Itertools; + use std::collections::BTreeSet; + use std::fmt::Write; + use tokio::sync::Mutex; #[tokio::test] - async fn discovers_nested_consumer_but_not_its_producer_occurrence() -> Result<()> { - let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])); - let input = Arc::new(EmptyExec::new(Arc::clone(&schema))) as Arc; - let column = Arc::new(Column::new("a", 0)) as Arc; - let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new( - vec![Arc::clone(&column)], - lit(true), - )) as Arc; - let nested = Arc::new(BinaryExpr::new( - Arc::clone(&dynamic_filter), - Operator::And, - lit(true), - )) as Arc; - - let consumer = - Arc::new(ExpressionExec::new(input, nested, false)) as Arc; - let plan = Arc::new(ExpressionExec::new( - consumer, - Arc::clone(&dynamic_filter), - true, - )) as Arc; - - let discovered = discover_dynamic_filter_consumers(&plan)?; - assert_eq!(discovered.len(), 1); - assert_eq!(discovered[0].id, dynamic_filter.expression_id().unwrap()); - assert!(!has_nonlocal_dynamic_filter_relationships(&plan)?); - - dynamic_filter - .downcast_ref::() - .unwrap() - .update(Arc::new(BinaryExpr::new(column, Operator::Gt, lit(10_i32))))?; - dynamic_filter - .downcast_ref::() - .unwrap() - .mark_complete(); - - let current = discovered[0].expression.current()?; - assert_eq!(current.to_string(), "a@0 > 10"); + async fn discovers_dynamic_filters_in_sql_plan() -> Result<()> { + let display = display_query( + r#" + SELECT COUNT(*) + FROM ( + SELECT DISTINCT "RainToday" AS key + FROM weather + ) build + JOIN weather probe ON build.key = probe."RainToday" + JOIN ( + SELECT DISTINCT "RainTomorrow" AS key + FROM weather + ) other_build ON other_build.key = probe."RainTomorrow" + WHERE probe."MinTemp" > 0 + "#, + ) + .await?; + assert_snapshot!(display, @r" + Stage 5 + AggregateExec + HashJoinExec producers=[1] + NetworkShuffleExec + AggregateExec + NetworkShuffleExec anchors=[1] + Stage 4 + RepartitionExec + AggregateExec + DataSourceExec consumers=[1] + Stage 3 + RepartitionExec + HashJoinExec producers=[2] + NetworkShuffleExec + AggregateExec + NetworkShuffleExec anchors=[2] + Stage 2 + RepartitionExec + AggregateExec + DataSourceExec consumers=[2] + Stage 1 + RepartitionExec + FilterExec + DataSourceExec + "); Ok(()) } - #[test] - fn deduplicates_consumers_with_the_same_expression_id() -> Result<()> { - let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])); - let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new( - vec![Arc::new(Column::new("a", 0))], - lit(true), - )) as Arc; - let consumers = (0..2) - .map(|_| { - Arc::new(ExpressionExec::new( - Arc::new(EmptyExec::new(Arc::clone(&schema))), - Arc::clone(&dynamic_filter), - false, - )) as Arc - }) - .collect(); - let plan = UnionExec::try_new(consumers)?; - - let discovered = discover_dynamic_filter_consumers(&plan)?; - - assert_eq!(discovered.len(), 1); - assert_eq!(discovered[0].id, dynamic_filter.expression_id().unwrap()); - assert!(has_nonlocal_dynamic_filter_relationships(&plan)?); + #[tokio::test] + async fn passes_anchor_through_two_shuffles() -> Result<()> { + let display = display_query( + r#" + SELECT COUNT(*) + FROM ( + SELECT DISTINCT "RainToday" AS key + FROM weather + ) build + JOIN ( + SELECT "RainTomorrow" AS key, SUM(n) AS total + FROM ( + SELECT "RainTomorrow", "RainToday", COUNT(*) AS n + FROM weather + GROUP BY "RainTomorrow", "RainToday" + ) grouped + GROUP BY "RainTomorrow" + ) probe ON build.key = probe.key + "#, + ) + .await?; + assert_snapshot!(display, @r" + Stage 4 + AggregateExec + HashJoinExec producers=[1] + AggregateExec + NetworkShuffleExec + ProjectionExec + AggregateExec + NetworkShuffleExec anchors=[1] + Stage 3 + RepartitionExec + AggregateExec + ProjectionExec + AggregateExec + NetworkShuffleExec anchors=[1] + Stage 2 + RepartitionExec + AggregateExec + DataSourceExec consumers=[1] + Stage 1 + RepartitionExec + AggregateExec + DataSourceExec + "); Ok(()) } - #[test] - fn identifies_a_producer_without_a_local_consumer() -> Result<()> { - let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])); - let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new( - vec![Arc::new(Column::new("a", 0))], - lit(true), - )) as Arc; - let plan = Arc::new(ExpressionExec::new( - Arc::new(EmptyExec::new(schema)), - dynamic_filter, - true, - )) as Arc; - - assert!(has_nonlocal_dynamic_filter_relationships(&plan)?); - Ok(()) + async fn display_query(sql: &str) -> Result { + let captured_plans = CapturePlans::default(); + let (ctx, _guard, _) = start_localhost_context(2, DefaultSessionBuilder).await; + let ctx = ctx + .with_distributed_broadcast_joins(false)? + .with_distributed_route_task_handler(captured_plans.clone()); + { + let state = ctx.state_ref(); + let mut state = state.write(); + let optimizer = &mut state.config_mut().options_mut().optimizer; + optimizer.hash_join_single_partition_threshold = 0; + optimizer.hash_join_single_partition_threshold_rows = 0; + } + register_parquet_tables(&ctx).await?; + let plan = ctx.sql(sql).await?.create_physical_plan().await?; + collect(plan, ctx.task_ctx()).await?; + let captured_plans = captured_plans.0.lock().await; + display_dynamic_filter_discovery(&captured_plans) } - #[derive(Debug)] - struct ExpressionExec { - input: Arc, - expression: Arc, - produces_expression: bool, - } + /// Captures the first task of each stage for displaying purposes. + #[derive(Clone, Default)] + struct CapturePlans(Arc>>>); - impl ExpressionExec { - fn new( - input: Arc, - expression: Arc, - produces_expression: bool, - ) -> Self { - Self { - input, - expression, - produces_expression, + #[async_trait] + impl RouteTaskHandler for CapturePlans { + async fn handle( + &self, + event: RouteTaskEvent<'_>, + ) -> Option> { + if event.task_key.task_number == 0 { + self.0.lock().await.insert( + event.task_key.stage_id, + Arc::clone(event.task_specialized_plan), + ); } + None } } - impl DisplayAs for ExpressionExec { - fn fmt_as(&self, _: DisplayFormatType, f: &mut Formatter) -> std::fmt::Result { - write!(f, "ExpressionExec") - } - } + /// Map random dynamic filter expression ids to monotonic numbers 1, 2, 3... + /// for stable snapshots. + #[derive(Default)] + struct IdNormalizer(HashMap); - impl ExecutionPlan for ExpressionExec { - fn name(&self) -> &str { - "ExpressionExec" + impl IdNormalizer { + fn annotation(&mut self, name: &str, ids: BTreeSet) -> Option { + (!ids.is_empty()).then(|| { + let ids = ids + .into_iter() + .map(|id| { + let next = self.0.len() + 1; + self.0.entry(id).or_insert(next).to_string() + }) + .join(", "); + format!("{name}=[{ids}]") + }) } + } - fn properties(&self) -> &Arc { - self.input.properties() - } + struct DynamicFilterIds { + consumers: BTreeSet, + anchors: BTreeSet, + producers: BTreeSet, + } - fn children(&self) -> Vec<&Arc> { - vec![&self.input] - } + fn dynamic_filter_annotations( + node: &dyn ExecutionPlan, + discovered: &DynamicFilterIds, + normalizer: &mut IdNormalizer, + ) -> Result { + let producers = node + .dynamic_expressions_produced() + .iter() + .filter_map(dynamic_filter_id) + .filter(|id| discovered.producers.contains(id)) + .collect::>(); + let is_network_boundary = node.is_network_boundary(); + let mut anchors = BTreeSet::new(); + let mut consumers = BTreeSet::new(); + node.apply_expressions(&mut |root| { + root.apply(|expression| { + if let Some(id) = dynamic_filter_id(expression) { + if is_network_boundary && discovered.anchors.contains(&id) { + anchors.insert(id); + } else if discovered.consumers.contains(&id) && !producers.contains(&id) { + consumers.insert(id); + } + } + Ok(TreeNodeRecursion::Continue) + }) + })?; - fn dynamic_expressions_produced(&self) -> Vec> { - self.produces_expression - .then(|| Arc::clone(&self.expression)) - .into_iter() - .collect() - } + let annotations = [ + ("anchors", anchors), + ("consumers", consumers), + ("producers", producers), + ] + .into_iter() + .filter_map(|(name, ids)| normalizer.annotation(name, ids)) + .join(" "); + Ok(if annotations.is_empty() { + String::new() + } else { + format!(" {annotations}") + }) + } - fn apply_expressions( - &self, - f: &mut dyn FnMut(&Arc) -> Result, - ) -> Result { - apply_expression_roots([&self.expression], f) - } + fn dynamic_filter_id(expression: &Arc) -> Option { + expression.downcast_ref::()?; + Some( + expression + .expression_id() + .expect("dynamic filters always have an expression ID"), + ) + } - fn with_new_children( - self: Arc, - mut children: Vec>, - ) -> Result> { - Ok(Arc::new(Self::new( - children.remove(0), - Arc::clone(&self.expression), - self.produces_expression, - ))) + fn display_dynamic_filter_discovery( + plans: &HashMap>, + ) -> Result { + fn render( + node: &dyn ExecutionPlan, + depth: usize, + discovered: &DynamicFilterIds, + normalizer: &mut IdNormalizer, + output: &mut String, + ) -> Result<()> { + writeln!( + output, + "{}{}{}", + " ".repeat(depth), + node.name(), + dynamic_filter_annotations(node, discovered, normalizer)?, + ) + .expect("writing to String cannot fail"); + for child in node.children() { + render(child.as_ref(), depth + 1, discovered, normalizer, output)?; + } + Ok(()) } - fn execute( - &self, - partition: usize, - context: Arc, - ) -> Result { - self.input.execute(partition, context) + let mut output = String::new(); + let mut normalizer = IdNormalizer::default(); + for stage_id in plans.keys().sorted().rev() { + writeln!(output, "Stage {stage_id}").expect("writing to String cannot fail"); + let plan = &plans[stage_id]; + let consumers = discover_dynamic_filter_consumers(plan)?; + let discovered = DynamicFilterIds { + consumers: consumers + .consumers + .into_iter() + .map(|consumer| consumer.id) + .collect(), + anchors: consumers + .anchors + .into_iter() + .map(|anchor| anchor.id) + .collect(), + producers: discover_dynamic_filter_producers(plan)? + .into_iter() + .map(|producer| producer.id) + .collect(), + }; + render(plan.as_ref(), 1, &discovered, &mut normalizer, &mut output)?; } + Ok(output) } } diff --git a/src/dynamic_filtering/display.rs b/src/dynamic_filtering/display.rs index 8075020d6..abcfa5537 100644 --- a/src/dynamic_filtering/display.rs +++ b/src/dynamic_filtering/display.rs @@ -51,6 +51,9 @@ pub async fn rewrite_distributed_plan_with_dynamic_filters( /// Severs dynamic filter connections so we can update filter values for /// display purposes without having an update in one node propagate to another. /// +/// Unlike [`maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships()`], this +/// isolates every dynamic filter found in the plan. +/// /// For example, in this plan, we would like to be able to [`update()`] every variant independently /// without mutating the producer or other variants /// @@ -62,14 +65,15 @@ pub async fn rewrite_distributed_plan_with_dynamic_filters( /// t0: DataSourceExec: ... /// t1: DataSourceExec: ... /// DistributedLeafExec: -/// t0: DataSourceExec: predicate=DynamicFilter [ f_dkey@2 >= A AND f_dkey@2 <= A AND f_dkey@2 IN (SET) ([]) ] <- unique filter -/// t1: DataSourceExec: predicate=DynamicFilter [ f_dkey@2 >= B AND f_dkey@2 <= B AND f_dkey@2 IN (SET) ([]) ] <- unique filter +/// t0: DataSourceExec: predicate=DynamicFilter [ f_dkey@2 > A ] <- unique filter +/// t1: DataSourceExec: predicate=DynamicFilter [ f_dkey@2 < B ] <- unique filter /// ``` /// /// This is done by deep-copying every leaf variant so we don't have to /// worry about any shared state. /// /// [`update()`]: DynamicFilterPhysicalExpr::update() +/// [`maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships()`]: super::maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships() pub(crate) fn sever_dynamic_filter_relationships_in_plan_for_display( plan: Arc, task_ctx: &Arc, @@ -159,10 +163,10 @@ fn apply_reports_to_distributed_leaves( .iter() .map(|filter| (filter.expression_id, &filter.expression)) .collect(); - let Ok(consumers) = discover_dynamic_filter_consumers(variant) else { + let Ok(discovered) = discover_dynamic_filter_consumers(variant) else { continue; }; - for consumer in consumers { + for consumer in discovered.consumers { let Some(expression) = updates.get(&consumer.id).copied() else { continue; }; diff --git a/src/dynamic_filtering/mod.rs b/src/dynamic_filtering/mod.rs index 5fa38c2e9..61832a114 100644 --- a/src/dynamic_filtering/mod.rs +++ b/src/dynamic_filtering/mod.rs @@ -4,6 +4,7 @@ mod display; use crate::codec::roundtrip_pb; use datafusion::common::Result; use datafusion::execution::TaskContext; +use datafusion::execution::config::SessionConfig; use datafusion::physical_plan::ExecutionPlan; use std::sync::Arc; @@ -11,37 +12,160 @@ pub(crate) use discovery::*; pub use display::rewrite_distributed_plan_with_dynamic_filters; pub(crate) use display::sever_dynamic_filter_relationships_in_plan_for_display; -// We must take care to avoid partial dynamic filter updates when sending an -// in-memory plan. -// -// Consider this partitioned hash join topology where the consumer task is -// collocated with one producer on worker A: -// ```text -// Worker A -// -// Stage 2 Task 0 -// HashJoinExec <- Dynamic Filter Produced: (foo > 100) -// -// Stage 1 Task 0 -// DataSourceExec <- consumer -// -// Worker B -// Stage 2 Task 1 -// HashJoinExec <- Dynamic Filter Produced: (foo != 150) -// ``` -// -// The in-process transport allows the Worker A join to propagate its filter to -// the consumer and mark it as completed, so the consumer incorrectly applies -// (foo > 100) instead of (foo > 100 OR foo != 150). -// -// In this situation, we roundtrip Stage 1 Task 0 to sever the in-memory -// relationship. The dynamic filter update from the producer must reach the -// coordinator for merging prior to being forwarded to the consumer. +pub(crate) fn is_dynamic_filtering_enabled(session_config: &SessionConfig) -> bool { + session_config + .options() + .optimizer + .enable_dynamic_filter_pushdown +} + +/// Deepcopies the plan if it contains any dynamic filter producers or consumers. This isolates +/// any dynamic filters in this plan from dynamic filters *outside* the plan. Plan nodes +/// *within* this plan will share in-memory dynamic filter state with eachother, even after copying. +/// +/// # Why Task-Local Dynamic Filters are Safe +/// +/// ## TopK Dynamic Filters +/// +/// Stage 2 Tasks: M +/// └── SortPreservingMergeExec: fetch=10 (global TopK) +/// └── NetworkCoalesceExec +/// +/// Stage 1 Tasks: N +/// └── SortExec: fetch=10 (local TopK and dynamic-filter producer) +/// └── DataSourceExec: dynamic-filter consumer +/// +/// Each `SortExec` may push its task-local TopK bound into its own input. It still emits the local +/// TopK candidates, which the `SortPreservingMergeExec` reduces to the global TopK in the parent +/// stage. +/// +/// ## Min/Max Dynamic Filters in Partial Aggregates with No Group +/// +/// Stage 2 Tasks: M +/// └── AggregateExec: mode=FinalPartitioned, gby=[], aggr=[max(foo)] (global max) +/// └── NetworkShuffleExec +/// +/// Stage 1 Tasks: N +/// └── RepartitionExec +/// └── AggregateExec: mode=Partial, gby=[], aggr=[max(foo)] (local max and dynamic filter producer) +/// └── DataSourceExec: dynamic-filter consumer +/// +/// The partial aggregate may push its task-local min/max bound into its own input. It still emits +/// the local min/max, which the final aggregate reduces to the global min/max in the parent stage. +/// +/// ## CollectLeft Joins +/// +/// Stage Y Tasks: M +/// +/// HashJoinExec: mode=CollectLeft +/// ├── build: CoalescePartitionsExec +/// │ └── all build partitions (the complete build side) +/// └── probe: +/// └── DataSourceExec: dynamic-filter consumer +/// +/// A `CollectLeft` hash join collects the complete build side in every task (by broadcasting or +/// otherwise) and pushes it down to every probe partition. Since every probe partition sees +/// the entire build-side filter, it will only filter rows which the join would filter out. +/// +/// ## Partitioned Joins +/// +/// A partitioned hash join builds per-partition predicates from the task-local hash table. +/// +/// Stage Y Task i +/// +/// HashJoinExec: mode=Partitioned +/// ├── build partition i → producer predicate P(i) +/// └── probe partition i → consumer of P(i) +/// +/// The build and probe execute corresponding partitions in the same task, so the probe +/// gets the correct/complete filter from it's corresponding build side. +/// +/// If the probe's partitioning is different than the join, then there must be a +/// `RepartitionExec` on the probe side. This will either become a network shuffle, falling +/// outside the local case, or it's a local repartition, meaning the join is not distributed +/// and contains all of the partitions and the entire build side. +/// +/// In all cases, a probe-side partition sees the correct filter. +/// +/// # Cases this Function Avoids +/// +/// Outside of the cases above, we assume that dynamic filter propagation between two tasks requires +/// remote communication and coordination through the coordinator stage. However, there's +/// some edge cases that break this assumption. +/// +/// It's possible for any two tasks to be collocated and share memory because +/// - a user to implement a custom transport layer and skip all proto serialization +/// - the coordinator may send plans to it's local worker via an in-memory channel without serializing +/// +/// The examples below show how this may cause incorrect results. +/// +/// ## Example 1: Producer-Consumer +/// +/// Consider this partitioned hash join topology where the consumer task is +/// collocated with one producer on worker A: +/// ```text +/// Worker A +/// +/// Stage 2 Task 0 +/// HashJoinExec <- Dynamic Filter Produced: (foo > 100) +/// +/// Stage 1 Task 0 +/// DataSourceExec <- consumer +/// +/// Worker B +/// Stage 2 Task 1 +/// HashJoinExec <- Dynamic Filter Produced: (foo != 150) +/// ``` +/// +/// The in-process transport allows the Worker A join to propagate its filter to +/// the consumer and mark it as completed, so the consumer incorrectly applies +/// (foo > 100) instead of (foo > 100 OR foo != 150). +/// +/// ## Example 2: Producer-Producer +/// +/// ```text +/// Worker A +/// +/// Stage 2 Task 0 +/// HashJoinExec <- Dynamic Filter Produced: (foo > 100) +/// +/// Stage 2 Task 1 +/// HashJoinExec <- Dynamic Filter Produced: (foo != 150) +/// +/// Stage 1 Task 0 +/// DataSourceExec <- consumer +/// ``` +/// +/// Both producers are collocated on worker A. In this situation, they both race to +/// update the dynamic filter, meaning the final expression will either be foo > 100 +/// or foo != 150. The correct expression is (foo > 100 OR foo != 150). +/// +/// Example 3: Local Producer-Consumer +/// +/// ```text +/// Worker A +/// +/// Stage 2 Task 0 +/// HashJoinExec <- Dynamic Filter Produced: (foo > 100) +/// DataSourceExec <- consumer +/// +/// Stage 2 Task 1 +/// HashJoinExec <- Dynamic Filter Produced: (foo != 150) +/// DataSourceExec <- consumer +/// ``` +/// +/// Since both producers and both consumers are located on the same worker, they all share +/// one in-memory dynamic filter. This ends up being a race between two writers and two readers. pub(crate) fn maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships( plan: Arc, task_ctx: &Arc, ) -> Result> { - if has_nonlocal_dynamic_filter_relationships(&plan)? { + let has_producers = !discover_dynamic_filter_producers(&plan)?.is_empty(); + let has_consumers = !discover_dynamic_filter_consumers(&plan)? + .consumers + .is_empty(); + + if has_producers || has_consumers { roundtrip_pb(plan, task_ctx) } else { Ok(plan) diff --git a/src/execution_plans/benchmarks/shuffle_bench.rs b/src/execution_plans/benchmarks/shuffle_bench.rs index 4f46f1be7..0656f79d9 100644 --- a/src/execution_plans/benchmarks/shuffle_bench.rs +++ b/src/execution_plans/benchmarks/shuffle_bench.rs @@ -217,6 +217,7 @@ impl ShuffleFixture { num: 0, workers: self.input_stage_workers.clone(), runtime_stats: None, + dynamic_filter_anchors: vec![], }); let mut join_set = JoinSet::default(); diff --git a/src/execution_plans/benchmarks/transport_bench.rs b/src/execution_plans/benchmarks/transport_bench.rs index ad159cd59..9cbaa15e4 100644 --- a/src/execution_plans/benchmarks/transport_bench.rs +++ b/src/execution_plans/benchmarks/transport_bench.rs @@ -265,6 +265,7 @@ impl TransportFixture { num: 0, workers: self.input_stage_tasks.clone(), runtime_stats: None, + dynamic_filter_anchors: vec![], }); let mut join_set = JoinSet::default(); diff --git a/src/execution_plans/network_broadcast.rs b/src/execution_plans/network_broadcast.rs index 94bca6de3..10581506a 100644 --- a/src/execution_plans/network_broadcast.rs +++ b/src/execution_plans/network_broadcast.rs @@ -12,7 +12,7 @@ use datafusion::physical_expr_common::metrics::MetricsSet; use datafusion::physical_plan::stream::RecordBatchStreamAdapter; use datafusion::physical_plan::{ DisplayAs, DisplayFormatType, ExecutionPlan, Partitioning, PlanProperties, Statistics, - StatisticsArgs, + StatisticsArgs, apply_expression_roots, }; use std::fmt::Formatter; use std::sync::Arc; @@ -219,9 +219,9 @@ impl ExecutionPlan for NetworkBroadcastExec { fn apply_expressions( &self, - _f: &mut dyn FnMut(&Arc) -> Result, + f: &mut dyn FnMut(&Arc) -> Result, ) -> Result { - Ok(TreeNodeRecursion::Continue) + apply_expression_roots(self.input_stage.dynamic_filter_anchors().iter(), f) } fn with_new_children( diff --git a/src/execution_plans/network_coalesce.rs b/src/execution_plans/network_coalesce.rs index 4cbaec0ce..db9b964a8 100644 --- a/src/execution_plans/network_coalesce.rs +++ b/src/execution_plans/network_coalesce.rs @@ -15,7 +15,8 @@ use datafusion::physical_plan::projection::ProjectionExec; use datafusion::physical_plan::stream::RecordBatchStreamAdapter; use datafusion::physical_plan::{ ChildrenPropertiesMode, DisplayAs, DisplayFormatType, EmptyRecordBatchStream, ExecutionPlan, - PlanProperties, ReplaceChildrenOptions, Statistics, StatisticsArgs, internal_err, + PlanProperties, ReplaceChildrenOptions, Statistics, StatisticsArgs, apply_expression_roots, + internal_err, }; use std::fmt::{Debug, Formatter}; use std::sync::Arc; @@ -241,9 +242,9 @@ impl ExecutionPlan for NetworkCoalesceExec { fn apply_expressions( &self, - _f: &mut dyn FnMut(&Arc) -> Result, + f: &mut dyn FnMut(&Arc) -> Result, ) -> Result { - Ok(TreeNodeRecursion::Continue) + apply_expression_roots(self.input_stage.dynamic_filter_anchors().iter(), f) } fn with_new_children( diff --git a/src/execution_plans/network_shuffle.rs b/src/execution_plans/network_shuffle.rs index 8fcbc980d..94fbcb08d 100644 --- a/src/execution_plans/network_shuffle.rs +++ b/src/execution_plans/network_shuffle.rs @@ -16,7 +16,7 @@ use datafusion::physical_plan::sorts::streaming_merge::StreamingMergeBuilder; use datafusion::physical_plan::stream::RecordBatchStreamAdapter; use datafusion::physical_plan::{ DisplayAs, DisplayFormatType, EmptyRecordBatchStream, ExecutionPlan, PlanProperties, - Statistics, StatisticsArgs, + Statistics, StatisticsArgs, apply_expression_roots, }; use std::fmt::Formatter; use std::sync::Arc; @@ -238,9 +238,9 @@ impl ExecutionPlan for NetworkShuffleExec { fn apply_expressions( &self, - _f: &mut dyn FnMut(&Arc) -> Result, + f: &mut dyn FnMut(&Arc) -> Result, ) -> Result { - Ok(TreeNodeRecursion::Continue) + apply_expression_roots(self.input_stage.dynamic_filter_anchors().iter(), f) } fn with_new_children( diff --git a/src/stage.rs b/src/stage.rs index 6e1c520cc..edd7af909 100644 --- a/src/stage.rs +++ b/src/stage.rs @@ -5,6 +5,7 @@ use datafusion::common::{HashMap, Statistics, config_err}; use datafusion::common::{exec_err, plan_err}; use datafusion::error::Result; use datafusion::execution::{SendableRecordBatchStream, TaskContext}; +use datafusion::physical_expr::PhysicalExpr; use datafusion::physical_plan::display::DisplayableExecutionPlan; use datafusion::physical_plan::metrics::{Label, Metric, MetricsSet}; use datafusion::physical_plan::{ @@ -114,6 +115,8 @@ pub struct RemoteStage { pub workers: Vec, /// Statistics collected at runtime, if any. pub runtime_stats: Option>, + /// Dynamic-filter consumers retained after the stage's plan is moved to its workers. + pub dynamic_filter_anchors: Vec>, } impl Stage { @@ -145,6 +148,13 @@ impl Stage { } } + pub(crate) fn dynamic_filter_anchors(&self) -> &[Arc] { + match self { + Self::Local(_) => &[], + Self::Remote(remote) => &remote.dynamic_filter_anchors, + } + } + pub fn metrics(&self) -> MetricsSet { match &self { Self::Local(v) => v.metrics_set.clone(), diff --git a/src/test_utils/routing.rs b/src/test_utils/routing.rs index 0c726e65c..5832378ca 100644 --- a/src/test_utils/routing.rs +++ b/src/test_utils/routing.rs @@ -5,7 +5,8 @@ use arrow::{ use datafusion::{ catalog::{Session, TableFunctionImpl, TableProvider}, common::{ - Result, ScalarValue, Statistics, internal_err, plan_err, tree_node::TreeNodeRecursion, + Result, ScalarValue, Statistics, exec_err, internal_err, plan_err, + tree_node::TreeNodeRecursion, }, datasource::TableType, execution::TaskContext, @@ -23,7 +24,9 @@ use datafusion_proto::{ use futures::stream; use prost::Message; use std::{fmt::Formatter, sync::Arc}; +use tokio::sync::Mutex; use tonic::async_trait; +use url::Url; use crate::{ DesiredTaskCountEvent, DesiredTaskCountEventResponse, DistributedLeafExec, @@ -386,3 +389,32 @@ impl PhysicalExtensionCodec for URLEmitterExtensionCodec { .map_err(|e| proto_error(format!("Failed to encode URLEmitterExec: {e}"))) } } + +/// Colocates all tasks on the same worker by choosing a URL once and caching it. +#[derive(Default)] +pub struct ColocateAllTasksHandler { + cached: Mutex>, +} + +#[async_trait] +impl RouteTaskHandler for ColocateAllTasksHandler { + async fn handle(&self, ev: RouteTaskEvent<'_>) -> Option> { + let url = { + let mut cached = self.cached.lock().await; + if let Some(url) = cached.as_ref() { + url.clone() + } else { + let Some(url) = ok_or_some_err!(ev.worker_resolver.get_urls()) + .into_iter() + .next() + else { + return Some(exec_err!("expected at least one worker URL")); + }; + *cached = Some(url.clone()); + url + } + }; + + Some(ev.dialer.dial(url).await) + } +} diff --git a/src/worker/impl_coordinator_channel.rs b/src/worker/impl_coordinator_channel.rs index 001519b19..a6477f6b6 100644 --- a/src/worker/impl_coordinator_channel.rs +++ b/src/worker/impl_coordinator_channel.rs @@ -256,7 +256,7 @@ fn build_task_completed_dynamic_filters( plan: &Arc, ) -> Result { let mut filters = vec![]; - for consumer in discover_dynamic_filter_consumers(plan)? { + for consumer in discover_dynamic_filter_consumers(plan)?.consumers { filters.push(TaskDynamicFilter { expression_id: consumer.id, expression: MaybeEncoded::Decoded(consumer.expression), diff --git a/tests/dynamic_filtering/common.rs b/tests/dynamic_filtering/common.rs index 87966952f..1833da626 100644 --- a/tests/dynamic_filtering/common.rs +++ b/tests/dynamic_filtering/common.rs @@ -9,7 +9,9 @@ use datafusion::physical_plan::collect; use datafusion::prelude::{SessionContext, col}; use datafusion_distributed::test_utils::localhost::start_localhost_context; use datafusion_distributed::test_utils::parquet::register_parquet_tables; -use datafusion_distributed::test_utils::routing::UrlEmitterRouteTaskHandler; +use datafusion_distributed::test_utils::routing::{ + ColocateAllTasksHandler, UrlEmitterRouteTaskHandler, +}; use datafusion_distributed::{ DefaultSessionBuilder, DistributedExt, display_plan_ascii, rewrite_distributed_plan_with_dynamic_filters, @@ -89,12 +91,17 @@ impl<'a> TestQuery<'a> { pub(crate) async fn execute_range_partitioned_query( sql: &str, expected_rows: usize, + colocate_tasks: bool, ) -> Result { let (ctx, _guard, _) = start_localhost_context(3, DefaultSessionBuilder).await; - let ctx = ctx + let mut ctx = ctx .with_distributed_broadcast_joins(false)? - .with_distributed_desired_task_count_handler(2usize) - .with_distributed_route_task_handler(UrlEmitterRouteTaskHandler); + .with_distributed_desired_task_count_handler(2usize); + ctx = if colocate_tasks { + ctx.with_distributed_route_task_handler(ColocateAllTasksHandler::default()) + } else { + ctx.with_distributed_route_task_handler(UrlEmitterRouteTaskHandler) + }; { let state = ctx.state_ref(); let mut state = state.write(); diff --git a/tests/dynamic_filtering/partitioned_join.rs b/tests/dynamic_filtering/partitioned_join.rs index b0c530c4c..a291db9f7 100644 --- a/tests/dynamic_filtering/partitioned_join.rs +++ b/tests/dynamic_filtering/partitioned_join.rs @@ -3,22 +3,35 @@ mod tests { use crate::common::{TestQuery, execute_range_partitioned_query}; use datafusion::common::Result; use datafusion_distributed::assert_snapshot; + use datafusion_distributed::test_utils::insta::insta::allow_duplicates; /// A Partitioned HashJoinExec propagates dynamic filters to local consumers. #[tokio::test] async fn local_dynamic_filters() -> Result<()> { - let display = execute_range_partitioned_query( - r#" - SELECT d.env, COUNT(*) AS n - FROM dim d - JOIN fact f ON d.d_dkey = f.f_dkey - WHERE d.service = 'log' - GROUP BY d.env - "#, - 2, - ) - .await?; - assert_snapshot!(display, @r" + let display = execute_range_partitioned_query(LOCAL_DYNAMIC_FILTER_QUERY, 2, false).await?; + assert_local_dynamic_filters_plan(display); + Ok(()) + } + + /// Colocated tasks must not share their task-local dynamic filters. + #[tokio::test] + async fn colocated_local_dynamic_filters() -> Result<()> { + let display = execute_range_partitioned_query(LOCAL_DYNAMIC_FILTER_QUERY, 2, true).await?; + assert_local_dynamic_filters_plan(display); + Ok(()) + } + + const LOCAL_DYNAMIC_FILTER_QUERY: &str = r#" + SELECT d.env, COUNT(*) AS n + FROM dim d + JOIN fact f ON d.d_dkey = f.f_dkey + WHERE d.service = 'log' + GROUP BY d.env + "#; + + fn assert_local_dynamic_filters_plan(display: String) { + allow_duplicates! { + assert_snapshot!(display, @r" ┌───── DistributedExec │ CoalescePartitionsExec │ [Stage 2] => NetworkCoalesceExec: output_partitions=4, input_tasks=2 @@ -41,7 +54,7 @@ mod tests { │ t1: DataSourceExec: file_groups={2 groups: [[/testdata/join/parquet/fact/f_dkey=B/data0.parquet], [/testdata/join/parquet/fact/f_dkey=D/data0.parquet]]}, projection=[f_dkey], output_partitioning=Range([f_dkey@0 ASC NULLS LAST], [(C)], 2), file_type=parquet, predicate=DynamicFilter [ f_dkey@2 >= B AND f_dkey@2 <= B AND f_dkey@2 IN (SET) ([]) ], dynamic_rg_pruning=eligible, pruning_predicate=f_dkey_null_count@1 != row_count@2 AND f_dkey_max@0 >= B AND f_dkey_null_count@1 != row_count@2 AND f_dkey_min@3 <= B AND f_dkey_null_count@1 != row_count@2 AND f_dkey_min@3 <= B AND B <= f_dkey_max@0, required_guarantees=[f_dkey in (B)] └────────────────────────────────────────────────── "); - Ok(()) + } } /// A Partitioned HashJoinExec does not propagate dynamic filters to a remote consumer. From 3582c103c5679462a91aa856af480acd2f84955c Mon Sep 17 00:00:00 2001 From: Jayant Shrivastava Date: Mon, 21 Sep 2026 19:29:07 +0000 Subject: [PATCH 2/2] wrap comments in code blocks --- src/dynamic_filtering/mod.rs | 58 ++++++++++++++++++++---------------- 1 file changed, 33 insertions(+), 25 deletions(-) diff --git a/src/dynamic_filtering/mod.rs b/src/dynamic_filtering/mod.rs index 61832a114..3752ba82b 100644 --- a/src/dynamic_filtering/mod.rs +++ b/src/dynamic_filtering/mod.rs @@ -27,13 +27,15 @@ pub(crate) fn is_dynamic_filtering_enabled(session_config: &SessionConfig) -> bo /// /// ## TopK Dynamic Filters /// -/// Stage 2 Tasks: M -/// └── SortPreservingMergeExec: fetch=10 (global TopK) -/// └── NetworkCoalesceExec -/// -/// Stage 1 Tasks: N -/// └── SortExec: fetch=10 (local TopK and dynamic-filter producer) -/// └── DataSourceExec: dynamic-filter consumer +/// ```text +/// Stage 2 Tasks: M +/// └── SortPreservingMergeExec: fetch=10 (global TopK) +/// └── NetworkCoalesceExec +/// +/// Stage 1 Tasks: N +/// └── SortExec: fetch=10 (local TopK and dynamic-filter producer) +/// └── DataSourceExec: dynamic-filter consumer +/// ``` /// /// Each `SortExec` may push its task-local TopK bound into its own input. It still emits the local /// TopK candidates, which the `SortPreservingMergeExec` reduces to the global TopK in the parent @@ -41,27 +43,31 @@ pub(crate) fn is_dynamic_filtering_enabled(session_config: &SessionConfig) -> bo /// /// ## Min/Max Dynamic Filters in Partial Aggregates with No Group /// -/// Stage 2 Tasks: M -/// └── AggregateExec: mode=FinalPartitioned, gby=[], aggr=[max(foo)] (global max) -/// └── NetworkShuffleExec -/// -/// Stage 1 Tasks: N -/// └── RepartitionExec -/// └── AggregateExec: mode=Partial, gby=[], aggr=[max(foo)] (local max and dynamic filter producer) -/// └── DataSourceExec: dynamic-filter consumer +/// ```text +/// Stage 2 Tasks: M +/// └── AggregateExec: mode=FinalPartitioned, gby=[], aggr=[max(foo)] (global max) +/// └── NetworkShuffleExec +/// +/// Stage 1 Tasks: N +/// └── RepartitionExec +/// └── AggregateExec: mode=Partial, gby=[], aggr=[max(foo)] (local max and dynamic filter producer) +/// └── DataSourceExec: dynamic-filter consumer +/// ``` /// /// The partial aggregate may push its task-local min/max bound into its own input. It still emits /// the local min/max, which the final aggregate reduces to the global min/max in the parent stage. /// /// ## CollectLeft Joins /// -/// Stage Y Tasks: M +/// ```text +/// Stage Y Tasks: M /// -/// HashJoinExec: mode=CollectLeft -/// ├── build: CoalescePartitionsExec -/// │ └── all build partitions (the complete build side) -/// └── probe: -/// └── DataSourceExec: dynamic-filter consumer +/// HashJoinExec: mode=CollectLeft +/// ├── build: CoalescePartitionsExec +/// │ └── all build partitions (the complete build side) +/// └── probe: +/// └── DataSourceExec: dynamic-filter consumer +/// ``` /// /// A `CollectLeft` hash join collects the complete build side in every task (by broadcasting or /// otherwise) and pushes it down to every probe partition. Since every probe partition sees @@ -71,11 +77,13 @@ pub(crate) fn is_dynamic_filtering_enabled(session_config: &SessionConfig) -> bo /// /// A partitioned hash join builds per-partition predicates from the task-local hash table. /// -/// Stage Y Task i +/// ```text +/// Stage Y Task i /// -/// HashJoinExec: mode=Partitioned -/// ├── build partition i → producer predicate P(i) -/// └── probe partition i → consumer of P(i) +/// HashJoinExec: mode=Partitioned +/// ├── build partition i → producer predicate P(i) +/// └── probe partition i → consumer of P(i) +/// ``` /// /// The build and probe execute corresponding partitions in the same task, so the probe /// gets the correct/complete filter from it's corresponding build side.