diff --git a/src/coordinator/dynamic_filter_registry.rs b/src/coordinator/dynamic_filter_registry.rs index c293f8265..22aee24f6 100644 --- a/src/coordinator/dynamic_filter_registry.rs +++ b/src/coordinator/dynamic_filter_registry.rs @@ -1,27 +1,54 @@ -use crate::TaskKey; -use crate::dynamic_filtering::{ - discover_dynamic_filter_consumers, discover_dynamic_filter_producers, -}; -use datafusion::common::{HashMap, HashSet, Result}; +use crate::dynamic_filtering::discover_dynamic_filter_consumers; +use crate::{ProducedDynamicFilter, TaskKey}; +use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion}; +use datafusion::common::{HashMap, HashSet, Result, internal_err}; +use datafusion::execution::TaskContext; +use datafusion::physical_expr::expressions::DynamicFilterPhysicalExpr; use datafusion::physical_expr_common::metrics::{ExecutionPlanMetricsSet, MetricBuilder}; use datafusion::physical_plan::ExecutionPlan; +use datafusion::physical_plan::aggregates::AggregateExec; +use datafusion::physical_plan::joins::{HashJoinExec, PartitionMode}; use datafusion::physical_plan::metrics::Count; +use datafusion::physical_plan::sorts::sort::SortExec; +use datafusion_proto::protobuf::physical_expr_node::ExprType; +use datafusion_proto::protobuf::{ + PhysicalBinaryExprNode, PhysicalDynamicFilterNode, PhysicalExprNode, +}; use std::sync::{Arc, Mutex}; +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(super) enum DynamicFilterMergeMode { + /// Wait for every planned producer task to report a complete dynamic filter, then + /// merge and forward the filter. Used for partitioned joins. + AllProducersComplete, + /// Wait for any producer to report a complete dynamic filter and forward it. + /// Used for collect left joins. + FirstProducerComplete, + /// Forward any update received from any producer, merging updates from + /// different producers. Used for TopK dynamic filters in aggregates and sorts. + Incremental, +} + #[derive(Default)] pub(super) struct PlannedDynamicFilter { + pub(super) merge_mode: Option, // 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, + /// Registered producers and their latest accepted snapshots, if any. + pub(super) producers: HashMap>, pub(super) consumer_tasks: HashSet, + /// Full dynamic filter containing the merged predicate and its completion state. + pub(super) merged: Option, } #[derive(Default)] pub(super) struct DynamicFilterRegistryState { pub(super) filters: HashMap, + /// Track which stages have registered all of their tasks. + pub(super) sealed_stages: HashSet, } /// Query-scoped hub for distributed dynamic filtering. @@ -51,22 +78,66 @@ impl DynamicFilterRegistry { /// Adds any dynamic filter producers and consumers found in `plan` to the registry. pub(crate) fn register_task( &self, - plan: &Arc, + task_specialized_plan: &Arc, task_key: TaskKey, ) -> Result<()> { - let producers = discover_dynamic_filter_producers(plan)?; + let mut producers = vec![]; + + task_specialized_plan.apply(|node| { + // `CollectLeft` joins broadcast an equivalent build side to every producer task, + // so we can forward the first completed dynamic filter. + // + // TopK and MIN/MAX bounds are independently useful and incremental, so we can use + // all updates. + // + // Any remaining dynamic filter should wait for completion and should be merged + // at the coordinator before being applied. + let merge_mode = if node + .downcast_ref::() + .is_some_and(|join| matches!(join.partition_mode(), PartitionMode::CollectLeft)) + { + DynamicFilterMergeMode::FirstProducerComplete + } else if node.is::() || node.is::() { + DynamicFilterMergeMode::Incremental + } else { + DynamicFilterMergeMode::AllProducersComplete + }; + let produced_ids: HashSet<_> = node + .dynamic_expressions_produced() + .into_iter() + .filter_map(|expression| { + expression + .downcast_ref::() + .map(|_| expression.expression_id()) + }) + .map(|id| match id { + Some(id) => Ok(id), + None => { + internal_err!("DynamicFilterPhysicalExpr did not have an expression ID") + } + }) + .collect::>()?; + producers.extend(produced_ids.iter().map(|id| (*id, merge_mode))); + Ok(TreeNodeRecursion::Continue) + })?; // 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 consumers = discover_dynamic_filter_consumers(task_specialized_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 (id, merge_mode) in producers { + let filter = state.filters.entry(id).or_default(); + filter.merge_mode = Some(match filter.merge_mode { + Some(existing) if existing != merge_mode => { + return internal_err!( + "Dynamic filter {id} has conflicting merge modes: \ + {existing:?} and {merge_mode:?}" + ); + } + Some(existing) => existing, + None => merge_mode, + }); + filter.producers.entry(task_key).or_insert(None); } for consumer in consumers { state @@ -78,4 +149,128 @@ impl DynamicFilterRegistry { } Ok(()) } + + /// Mark that a stage has registered all of its tasks. + pub(crate) fn seal_stage(&self, stage_id: usize) { + let mut state = self.state.lock().expect("dynamic filter registry poisoned"); + state.sealed_stages.insert(stage_id); + let ids = state.filters.keys().copied().collect::>(); + for id in ids { + Self::merge(&mut state, id); + } + } + + /// Records a producer's latest dynamic-filter update and recomputes the merged filter. + pub(crate) fn record_dynamic_filter_update( + &self, + task_key: TaskKey, + report: ProducedDynamicFilter, + task_ctx: &TaskContext, + ) { + self.record_update_received(); + let Ok(expression) = report.expression.to_proto(task_ctx) else { + return; + }; + if expression.expr_id != Some(report.expression_id) { + return; + } + let Some(ExprType::DynamicFilter(dynamic_filter)) = expression.expr_type else { + return; + }; + if dynamic_filter.inner_expr.is_none() { + return; + } + + let mut state = self.state.lock().expect("dynamic filter registry poisoned"); + let Some(filter) = state.filters.get_mut(&report.expression_id) else { + return; + }; + let Some(previous) = filter.producers.get_mut(&task_key) else { + return; + }; + if (filter.merge_mode != Some(DynamicFilterMergeMode::Incremental) + && !dynamic_filter.is_complete) + || previous + .as_ref() + .is_some_and(|previous| previous.is_complete || previous == dynamic_filter.as_ref()) + { + return; + } + *previous = Some(*dynamic_filter); + Self::merge(&mut state, report.expression_id); + } + + /// Merges partial dynamic filters together for the provided dynamic filter + /// id only if there are enough updates present. + fn merge(state: &mut DynamicFilterRegistryState, id: u64) -> bool { + let Some(filter) = state.filters.get_mut(&id) else { + return false; + }; + let previous = filter.merged.as_ref(); + if previous.is_some_and(|filter| filter.is_complete) { + return false; + } + let Some(mode) = filter.merge_mode else { + return false; + }; + let all_complete = !filter.producers.is_empty() + && filter.producers.iter().all(|(task, report)| { + state.sealed_stages.contains(&task.stage_id) + && report.as_ref().is_some_and(|report| report.is_complete) + }); + if mode == DynamicFilterMergeMode::AllProducersComplete && !all_complete { + return false; + } + + let mut reports: Vec<_> = filter + .producers + .iter() + .filter_map(|(task, report)| report.as_ref().map(|report| (task, report))) + .collect(); + reports.sort_unstable_by_key(|(key, _)| (key.stage_id, key.task_number)); + if mode == DynamicFilterMergeMode::FirstProducerComplete { + reports.truncate(1); + } + let Some((_, template)) = reports.first() else { + return false; + }; + let inner_expr = merge_predicates( + reports + .iter() + .filter_map(|(_, report)| report.inner_expr.as_deref().cloned()) + .collect(), + ) + .map(Box::new); + let is_complete = mode == DynamicFilterMergeMode::FirstProducerComplete || all_complete; + if previous.is_some_and(|previous| { + previous.inner_expr == inner_expr && previous.is_complete == is_complete + }) { + return false; + } + let mut merged = (*template).clone(); + merged.inner_expr = inner_expr; + merged.is_complete = is_complete; + // Use a synthetic generation number for the merged filter. Each partial update has it's own generation + // is not useful here. + merged.generation = previous.map_or(0, |previous| previous.generation + 1); + filter.merged = Some(merged); + true + } +} + +/// Merges [`PhysicalExprNode`] together by ORing them. +fn merge_predicates(mut predicates: Vec) -> Option { + match predicates.len() { + 0 => None, + 1 => predicates.pop(), + _ => Some(PhysicalExprNode { + expr_id: None, + expr_type: Some(ExprType::BinaryExpr(Box::new(PhysicalBinaryExprNode { + l: None, + r: None, + op: "Or".to_owned(), + operands: predicates, + }))), + }), + } } diff --git a/src/coordinator/prepare_dynamic_plan.rs b/src/coordinator/prepare_dynamic_plan.rs index 6b7f95857..e0a0458f3 100644 --- a/src/coordinator/prepare_dynamic_plan.rs +++ b/src/coordinator/prepare_dynamic_plan.rs @@ -108,6 +108,7 @@ pub(super) async fn prepare_dynamic_plan( let _ = worker_tx.send(CoordinatorToWorkerMsg::KickOffSampling); stage_coordinator.coordinator_to_worker_task(task_i, worker_tx)?; } + stage_coordinator.seal_dynamic_filter_stage(); let (stats, consumer_tc) = if nb_type == TypeId::of::() { (None, Maximum(1)) diff --git a/src/coordinator/prepare_static_plan.rs b/src/coordinator/prepare_static_plan.rs index f83838d5c..c1fb7288a 100644 --- a/src/coordinator/prepare_static_plan.rs +++ b/src/coordinator/prepare_static_plan.rs @@ -47,7 +47,7 @@ pub(super) async fn prepare_static_plan( stage_coordinator.worker_to_coordinator_task(task_i, worker_rx); stage_coordinator.coordinator_to_worker_task(task_i, worker_tx)?; } - + stage_coordinator.seal_dynamic_filter_stage(); Ok(Transformed::yes(plan.with_input_stage(Stage::Remote( RemoteStage { query_id: stage.query_id, diff --git a/src/coordinator/query_coordinator.rs b/src/coordinator/query_coordinator.rs index 4889f888f..258525995 100644 --- a/src/coordinator/query_coordinator.rs +++ b/src/coordinator/query_coordinator.rs @@ -299,6 +299,10 @@ impl<'a> StageCoordinator<'a> { )) } + pub(super) fn seal_dynamic_filter_stage(&self) { + self.dynamic_filter_registry.seal_stage(self.stage_id); + } + /// Spawns a background task in charge of collecting messages sent by a worker. Some things that /// are collected from workers are: /// - Execution metrics information, sent once the worker has finished executing the task. @@ -315,6 +319,7 @@ impl<'a> StageCoordinator<'a> { let task_metrics = self.metrics_store.clone(); let completed_dynamic_filter_store = self.completed_dynamic_filter_store.clone(); let dynamic_filter_registry = Arc::clone(self.dynamic_filter_registry); + let task_ctx = Arc::clone(self.task_ctx); let (load_info_tx, load_info_rx) = tokio::sync::mpsc::unbounded_channel(); let mut load_info_tx_opt = Some(load_info_tx); @@ -342,8 +347,9 @@ impl<'a> StageCoordinator<'a> { store.insert(task_key, filters); } } - WorkerToCoordinatorMsg::ProducedDynamicFilter(_) => { - dynamic_filter_registry.record_update_received(); + WorkerToCoordinatorMsg::ProducedDynamicFilter(filter) => { + dynamic_filter_registry + .record_dynamic_filter_update(task_key, *filter, &task_ctx); } } } diff --git a/tests/multi_task_collect_join_repros.rs b/tests/multi_task_collect_join_repros.rs index 03724f8e5..39f59eb39 100644 --- a/tests/multi_task_collect_join_repros.rs +++ b/tests/multi_task_collect_join_repros.rs @@ -377,20 +377,13 @@ mod tests { .unwrap(); } - /// A build-side `LIMIT` is carried by the `CoalescePartitionsExec` that a - /// broadcast rewrite replaces. The replacement must preserve that fetch or - /// the join observes every build row instead of the requested 50. - /// - /// Only 3 of the 4 build-side files survive the limit: DataFusion stops listing files - /// once their row counts exceed the fetch and sorts the survivors by path only - /// afterwards, so which files remain depends on the directory listing order and on - /// `buffer_unordered` statistics fetching. That differs between filesystems and even - /// between CI runs on the same image. Redact the build-side file index; the probe-side - /// files are unaffected and stay asserted. + /// The build-side limit must be applied before its rows are broadcast to probe tasks. + /// Without that limit, the join observes every build row instead of the requested 50. + /// ORDER BY keeps the limited rows and scan plan deterministic. #[tokio::test] async fn build_side_fetch_is_preserved_by_broadcast() { let plan = assert_distributed_matches_single_node( - "SELECT count(*) FROM (SELECT id FROM build_side LIMIT 50) b \ + "SELECT count(*) FROM (SELECT id FROM build_side ORDER BY id LIMIT 50) b \ JOIN probe_side p ON b.id = p.id", true, ) @@ -419,16 +412,16 @@ mod tests { └────────────────────────────────────────────────── ┌───── Stage 2 ── tasks=1, partitions=4 │ BroadcastExec: input_partitions=1, consumer_tasks=4, output_partitions=4 - │ CoalescePartitionsExec: fetch=50 - │ [Stage 1] => NetworkCoalesceExec: output_partitions=12, input_tasks=4 + │ SortPreservingMergeExec: [id@0 ASC NULLS LAST], fetch=50 + │ [Stage 1] => NetworkCoalesceExec: output_partitions=8, input_tasks=4 └────────────────────────────────────────────────── - ┌───── Stage 1 ── tasks=4, partitions=12 - │ LocalLimitExec: fetch=50 + ┌───── Stage 1 ── tasks=4, partitions=8 + │ SortExec: TopK(fetch=50), expr=[id@0 ASC NULLS LAST], preserve_partitioning=[true] │ DistributedLeafExec: - │ t0: DataSourceExec: file_groups={3 groups: [[/target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..]]}, projection=[id], limit=50, file_type=parquet - │ t1: DataSourceExec: file_groups={3 groups: [[/target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..]]}, projection=[id], limit=50, file_type=parquet - │ t2: DataSourceExec: file_groups={3 groups: [[/target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..]]}, projection=[id], limit=50, file_type=parquet - │ t3: DataSourceExec: file_groups={3 groups: [[/target/multi_task_collect_join_repros/build_side/part-.parquet:.., /target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:.., /target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..]]}, projection=[id], limit=50, file_type=parquet + │ t0: DataSourceExec: file_groups={2 groups: [[/target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..]]}, projection=[id], file_type=parquet, predicate=DynamicFilter [ empty ], dynamic_rg_pruning=eligible + │ t1: DataSourceExec: file_groups={2 groups: [[/target/multi_task_collect_join_repros/build_side/part-.parquet:.., /target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:.., /target/multi_task_collect_join_repros/build_side/part-.parquet:..]]}, projection=[id], file_type=parquet, predicate=DynamicFilter [ empty ], dynamic_rg_pruning=eligible + │ t2: DataSourceExec: file_groups={2 groups: [[/target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..]]}, projection=[id], file_type=parquet, predicate=DynamicFilter [ empty ], dynamic_rg_pruning=eligible + │ t3: DataSourceExec: file_groups={2 groups: [[/target/multi_task_collect_join_repros/build_side/part-.parquet:.., /target/multi_task_collect_join_repros/build_side/part-.parquet:..], [/target/multi_task_collect_join_repros/build_side/part-.parquet:..]]}, projection=[id], file_type=parquet, predicate=DynamicFilter [ empty ], dynamic_rg_pruning=eligible └────────────────────────────────────────────────── ")); }