Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
227 changes: 211 additions & 16 deletions src/coordinator/dynamic_filter_registry.rs
Original file line number Diff line number Diff line change
@@ -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<DynamicFilterMergeMode>,
// 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<TaskKey>,
/// Registered producers and their latest accepted snapshots, if any.
pub(super) producers: HashMap<TaskKey, Option<PhysicalDynamicFilterNode>>,
pub(super) consumer_tasks: HashSet<TaskKey>,
/// Full dynamic filter containing the merged predicate and its completion state.
pub(super) merged: Option<PhysicalDynamicFilterNode>,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I see this is working in terms of protobuf messages rather than raw DataFusion expressions.

This means that proto conversion would happen even in a fully in-memory setup. Is it easy to avoid and just work with normal DataFusion types here rather than protobuf messages?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main reason is honestly that is_complete isn't public on DynamicFilterPhysicalExpr. So there's actually no way to know if it's complete or not here.

We could detect it on the worker here and send a is_complete bit in the `ProducedDynamicFilter. Do you think making that change is worth it?

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🤔 I don't think there should be any reason for is_complete() to be private upstream, even the mark_complete() method is public.

Just put a PR for that upstream (apache/datafusion#25738). In the meantime, it's fine to keep what you have here 👍

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nice!

Comment thread
jayshrivastava marked this conversation as resolved.
}

#[derive(Default)]
pub(super) struct DynamicFilterRegistryState {
pub(super) filters: HashMap<u64, PlannedDynamicFilter>,
/// Track which stages have registered all of their tasks.
pub(super) sealed_stages: HashSet<usize>,
}

/// Query-scoped hub for distributed dynamic filtering.
Expand Down Expand Up @@ -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<dyn ExecutionPlan>,
task_specialized_plan: &Arc<dyn ExecutionPlan>,
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::<HashJoinExec>()
.is_some_and(|join| matches!(join.partition_mode(), PartitionMode::CollectLeft))
{
DynamicFilterMergeMode::FirstProducerComplete
} else if node.is::<SortExec>() || node.is::<AggregateExec>() {
DynamicFilterMergeMode::Incremental
} else {
DynamicFilterMergeMode::AllProducersComplete
};
let produced_ids: HashSet<_> = node
.dynamic_expressions_produced()
.into_iter()
.filter_map(|expression| {
expression
.downcast_ref::<DynamicFilterPhysicalExpr>()
.map(|_| expression.expression_id())
})
.map(|id| match id {
Some(id) => Ok(id),
None => {
internal_err!("DynamicFilterPhysicalExpr did not have an expression ID")
}
})
.collect::<Result<_>>()?;
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
Expand All @@ -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::<Vec<_>>();
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<PhysicalExprNode>) -> Option<PhysicalExprNode> {
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,
}))),
}),
}
}
1 change: 1 addition & 0 deletions src/coordinator/prepare_dynamic_plan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::<NetworkCoalesceExec>() {
(None, Maximum(1))
Expand Down
2 changes: 1 addition & 1 deletion src/coordinator/prepare_static_plan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
10 changes: 8 additions & 2 deletions src/coordinator/query_coordinator.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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);

Expand Down Expand Up @@ -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);
}
}
}
Expand Down
31 changes: 12 additions & 19 deletions tests/multi_task_collect_join_repros.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>]]}, projection=[id], limit=50, file_type=parquet
│ t1: DataSourceExec: file_groups={3 groups: [[/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>]]}, projection=[id], limit=50, file_type=parquet
│ t2: DataSourceExec: file_groups={3 groups: [[/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>]]}, projection=[id], limit=50, file_type=parquet
│ t3: DataSourceExec: file_groups={3 groups: [[/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>, /target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>, /target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>]]}, projection=[id], limit=50, file_type=parquet
│ t0: DataSourceExec: file_groups={2 groups: [[/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>]]}, 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-<shard>.parquet:<int>..<int>, /target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>, /target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>]]}, 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-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>]]}, 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-<shard>.parquet:<int>..<int>, /target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>], [/target/multi_task_collect_join_repros/build_side/part-<shard>.parquet:<int>..<int>]]}, projection=[id], file_type=parquet, predicate=DynamicFilter [ empty ], dynamic_rg_pruning=eligible
└──────────────────────────────────────────────────
"));
}
Expand Down
Loading