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
16 changes: 13 additions & 3 deletions src/coordinator/dynamic_filter_registry.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,9 @@ use crate::dynamic_filtering::{
discover_dynamic_filter_consumers, discover_dynamic_filter_producers,
};
use datafusion::common::{HashMap, HashSet, Result};
use datafusion::physical_expr_common::metrics::{ExecutionPlanMetricsSet, MetricBuilder};
use datafusion::physical_plan::ExecutionPlan;
use datafusion::physical_plan::metrics::Count;
use std::sync::{Arc, Mutex};

#[derive(Default)]
Expand All @@ -28,14 +30,22 @@ pub(super) struct DynamicFilterRegistryState {
/// - 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<DynamicFilterRegistryState>,
dynamic_filter_updates_received: Count,
}

impl DynamicFilterRegistry {
pub(crate) fn new() -> Self {
Self::default()
pub(crate) fn new(metrics: &ExecutionPlanMetricsSet) -> Self {
Self {
state: Mutex::new(DynamicFilterRegistryState::default()),
dynamic_filter_updates_received: MetricBuilder::new(metrics)
.global_counter("dynamic_filter_updates_received"),
}
}

pub(crate) fn record_update_received(&self) {
self.dynamic_filter_updates_received.add(1);
}

/// Adds any dynamic filter producers and consumers found in `plan` to the registry.
Expand Down
43 changes: 32 additions & 11 deletions src/coordinator/query_coordinator.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ use crate::coordinator::DynamicFilterRegistry;
use crate::coordinator::Store;
use crate::coordinator::latency_metric::LatencyMetric;
use crate::dynamic_filtering::{
is_dynamic_filtering_enabled,
dynamic_filter_remote_producer_ids, is_dynamic_filtering_enabled,
maybe_roundtrip_plan_to_sever_in_memory_dynamic_filter_relationships,
};
use crate::events::{
Expand Down Expand Up @@ -74,7 +74,7 @@ impl QueryCoordinator {
metrics: metrics_set.clone(),
metrics_store,
completed_dynamic_filter_store,
dynamic_filter_registry: Arc::new(DynamicFilterRegistry::new()),
dynamic_filter_registry: Arc::new(DynamicFilterRegistry::new(metrics_set)),
coordinator_to_worker_metrics: CoordinatorToWorkerMetrics::new(metrics_set),
end_stream_notifier: Arc::new(Notify::new()),
join_set: Mutex::new(JoinSet::new()),
Expand Down Expand Up @@ -163,7 +163,11 @@ impl<'a> StageCoordinator<'a> {
)> {
let session_config = self.task_ctx.session_config();

let (specialized, work_unit_feed_declarations) = self.task_specialized_plan(task_i)?;
let TaskSpecializedPlan {
plan,
work_unit_feed_declarations,
dynamic_filter_remote_producer_ids,
} = self.task_specialized_plan(task_i)?;

let task_key = TaskKey {
query_id: self.query_id,
Expand All @@ -172,7 +176,7 @@ impl<'a> StageCoordinator<'a> {
};

self.dynamic_filter_registry
.register_task(&specialized, task_key)?;
.register_task(&plan, task_key)?;

let mut headers = get_config_extension_propagation_headers(session_config)?;
headers.extend(get_passthrough_headers(session_config));
Expand Down Expand Up @@ -211,7 +215,8 @@ impl<'a> StageCoordinator<'a> {
let set_plan_request = SetPlanRequest {
task_key,
task_count: self.task_count,
plan: MaybeEncoded::Decoded(Arc::clone(&specialized)),
plan: MaybeEncoded::Decoded(Arc::clone(&plan)),
dynamic_filter_remote_producer_ids: dynamic_filter_remote_producer_ids.clone(),
work_unit_feed_declarations: work_unit_feed_declarations.clone(),
target_worker_url: url.clone(),
query_start_time_ns: self.metrics.instantiation_time,
Expand Down Expand Up @@ -256,7 +261,7 @@ impl<'a> StageCoordinator<'a> {
task_ctx: self.task_ctx,
metrics: self.metrics_set,
worker_resolver: worker_resolver.as_ref(),
task_specialized_plan: &specialized,
task_specialized_plan: &plan,
task_key,
task_count: self.task_count,
dialer: &dialer,
Expand Down Expand Up @@ -308,6 +313,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 (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 @@ -335,6 +341,9 @@ impl<'a> StageCoordinator<'a> {
store.insert(task_key, filters);
}
}
WorkerToCoordinatorMsg::ProducedDynamicFilter(_) => {
dynamic_filter_registry.record_update_received();
}
}
}
});
Expand Down Expand Up @@ -418,10 +427,7 @@ impl<'a> StageCoordinator<'a> {
/// trimming down any unnecessary information that the specific `task_i` task is not going to
/// need, like unexecuted branches in [ChildrenIsolatorUnionExec], or unexecuted variants of
/// [DistributedLeafExec].
fn task_specialized_plan(
&self,
task_i: usize,
) -> Result<(Arc<dyn ExecutionPlan>, Vec<WorkUnitFeedDeclaration>)> {
fn task_specialized_plan(&self, task_i: usize) -> Result<TaskSpecializedPlan> {
let session_config = self.task_ctx.session_config();
let wuf_registry = session_config
.get_extension::<WorkUnitFeedRegistry>()
Expand Down Expand Up @@ -472,14 +478,29 @@ impl<'a> StageCoordinator<'a> {
} else {
transformed.data
};
Ok((plan, work_unit_feed_declarations))
let dynamic_filter_remote_producer_ids = if dynamic_filtering_enabled {
dynamic_filter_remote_producer_ids(&plan)?
} else {
vec![]
};
Ok(TaskSpecializedPlan {
plan,
work_unit_feed_declarations,
dynamic_filter_remote_producer_ids,
})
}
}

fn keep_stream_alive<T: 'static>(notify: Arc<Notify>) -> impl Stream<Item = T> + 'static {
futures::stream::once(notify.notified_owned()).filter_map(|()| futures::future::ready(None))
}

struct TaskSpecializedPlan {
plan: Arc<dyn ExecutionPlan>,
work_unit_feed_declarations: Vec<WorkUnitFeedDeclaration>,
dynamic_filter_remote_producer_ids: Vec<u64>,
}

pub(super) struct NotifyGuard(Arc<Notify>);

impl Drop for NotifyGuard {
Expand Down
43 changes: 38 additions & 5 deletions src/dynamic_filtering/discovery.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ use std::sync::Arc;
#[derive(Clone)]
pub(crate) struct DiscoveredDynamicFilterProducer {
pub(crate) id: u64,
pub(crate) expression: Arc<dyn PhysicalExpr>,
}

/// A dynamic-filter consumer discovered in an execution plan along with the schema it is evaluated
Expand Down Expand Up @@ -130,7 +131,7 @@ pub(crate) fn discover_dynamic_filter_producers(
};
producers
.entry(id)
.or_insert(DiscoveredDynamicFilterProducer { id });
.or_insert_with(|| DiscoveredDynamicFilterProducer { id, expression });
}
Ok(TreeNodeRecursion::Continue)
})?;
Expand All @@ -140,6 +141,31 @@ pub(crate) fn discover_dynamic_filter_producers(
Ok(producers)
}

/// Returns producer IDs with at least one remote consumer.
///
/// If a producer ID is present in the dynamic-filter anchors of any [`NetworkBoundary`], the plan
/// contains at least one remote consumer and the producer's updates must be forwarded to the
/// coordinator.
///
/// [`NetworkBoundary`]: crate::NetworkBoundary
pub(crate) fn dynamic_filter_remote_producer_ids(
plan: &Arc<dyn ExecutionPlan>,
) -> Result<Vec<u64>> {
let producer_ids: HashSet<_> = discover_dynamic_filter_producers(plan)?
.into_iter()
.map(|producer| producer.id)
.collect();
let anchor_ids: HashSet<_> = discover_dynamic_filter_consumers(plan)?
.anchors
.into_iter()
.map(|anchor| anchor.id)
.collect();

let mut remote_producer_ids: Vec<_> = producer_ids.intersection(&anchor_ids).copied().collect();
remote_producer_ids.sort_unstable();
Ok(remote_producer_ids)
}

/// 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.
Expand Down Expand Up @@ -212,7 +238,7 @@ mod tests {
)
.await?;
assert_snapshot!(display, @r"
Stage 5
Stage 5 remote_producers=[1]
AggregateExec
HashJoinExec producers=[1]
NetworkShuffleExec
Expand All @@ -222,7 +248,7 @@ mod tests {
RepartitionExec
AggregateExec
DataSourceExec consumers=[1]
Stage 3
Stage 3 remote_producers=[2]
RepartitionExec
HashJoinExec producers=[2]
NetworkShuffleExec
Expand Down Expand Up @@ -262,7 +288,7 @@ mod tests {
)
.await?;
assert_snapshot!(display, @r"
Stage 4
Stage 4 remote_producers=[1]
AggregateExec
HashJoinExec producers=[1]
AggregateExec
Expand Down Expand Up @@ -432,8 +458,15 @@ mod tests {
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 remote_producers = dynamic_filter_remote_producer_ids(plan)?
.into_iter()
.collect();
let remote_producers = normalizer
.annotation("remote_producers", remote_producers)
.map_or_else(String::new, |annotation| format!(" {annotation}"));
writeln!(output, "Stage {stage_id}{remote_producers}")
.expect("writing to String cannot fail");
let consumers = discover_dynamic_filter_consumers(plan)?;
let discovered = DynamicFilterIds {
consumers: consumers
Expand Down
7 changes: 4 additions & 3 deletions src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -57,9 +57,10 @@ pub use worker_resolver::{WorkerResolver, get_distributed_worker_resolver};

pub use protocol::{
ChannelResolver, CoordinatorToWorkerMsg, ExecuteTaskRequest, GetWorkerInfoRequest,
GetWorkerInfoResponse, LoadInfo, SetPlanRequest, TaskCompletedDynamicFilters,
TaskDynamicFilter, TaskKey, TaskMetrics, WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg,
WorkerChannel, WorkerToCoordinatorMsg, get_distributed_channel_resolver,
GetWorkerInfoResponse, LoadInfo, ProducedDynamicFilter, SetPlanRequest,
TaskCompletedDynamicFilters, TaskDynamicFilter, TaskKey, TaskMetrics, WorkUnitBatch,
WorkUnitFeedDeclaration, WorkUnitMsg, WorkerChannel, WorkerToCoordinatorMsg,
get_distributed_channel_resolver,
};
pub use stage::{
DistributedTaskContext, Stage, display_plan_ascii, display_plan_graphviz, explain_analyze,
Expand Down
16 changes: 15 additions & 1 deletion src/protocol/grpc/generated/worker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ pub mod coordinator_to_worker_msg {
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct WorkerToCoordinatorMsg {
#[prost(oneof = "worker_to_coordinator_msg::Inner", tags = "1, 2, 3, 4")]
#[prost(oneof = "worker_to_coordinator_msg::Inner", tags = "1, 2, 3, 4, 5")]
pub inner: ::core::option::Option<worker_to_coordinator_msg::Inner>,
}
/// Nested message and enum types in `WorkerToCoordinatorMsg`.
Expand Down Expand Up @@ -57,6 +57,9 @@ pub mod worker_to_coordinator_msg {
/// Another task in the same stage may report a different value for expression_id=10.
#[prost(message, tag = "4")]
TaskCompletedDynamicFilters(super::TaskCompletedDynamicFilters),
/// An observed dynamic filter state produced by this task.
#[prost(message, tag = "5")]
ProducedDynamicFilter(super::ProducedDynamicFilter),
}
}
#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)]
Expand All @@ -67,6 +70,14 @@ pub struct DynamicFilter {
#[prost(bytes = "vec", tag = "2")]
pub expression_proto: ::prost::alloc::vec::Vec<u8>,
}
#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)]
pub struct ProducedDynamicFilter {
#[prost(uint64, tag = "1")]
pub expression_id: u64,
/// Serialized datafusion.proto.PhysicalExprNode.
#[prost(bytes = "vec", tag = "2")]
pub expression_proto: ::prost::alloc::vec::Vec<u8>,
}
#[derive(Clone, PartialEq, ::prost::Message)]
pub struct TaskCompletedDynamicFilters {
#[prost(message, repeated, tag = "1")]
Expand Down Expand Up @@ -146,6 +157,9 @@ pub struct SetPlanRequest {
/// relative to when the query was fired in the coordinator.
#[prost(uint64, tag = "6")]
pub query_start_time_ns: u64,
/// Producer IDs whose updates must be reported to the coordinator for remote consumers.
#[prost(uint64, repeated, tag = "8")]
pub dynamic_filter_remote_producer_ids: ::prost::alloc::vec::Vec<u64>,
}
/// Nested message and enum types in `SetPlanRequest`.
pub mod set_plan_request {
Expand Down
1 change: 1 addition & 0 deletions src/protocol/grpc/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ mod worker_service;

// TODO: this should not be exposed.
pub(crate) use channel_resolver::DEFAULT_CHANNEL_RESOLVER_PER_RUNTIME;
pub(crate) use on_drop_stream::on_drop_stream;

pub use channel_resolver::{BoxCloneSyncChannel, DefaultChannelResolver};
pub use observability::{
Expand Down
11 changes: 11 additions & 0 deletions src/protocol/grpc/worker.proto
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,9 @@ message WorkerToCoordinatorMsg {
//
// Another task in the same stage may report a different value for expression_id=10.
TaskCompletedDynamicFilters task_completed_dynamic_filters = 4;

// An observed dynamic filter state produced by this task.
ProducedDynamicFilter produced_dynamic_filter = 5;
}
}

Expand All @@ -62,6 +65,12 @@ message DynamicFilter {
bytes expression_proto = 2;
}

message ProducedDynamicFilter {
uint64 expression_id = 1;
// Serialized datafusion.proto.PhysicalExprNode.
bytes expression_proto = 2;
}

message TaskCompletedDynamicFilters {
repeated DynamicFilter filters = 1;
}
Expand Down Expand Up @@ -131,6 +140,8 @@ message SetPlanRequest {
// Unix nanos when the query started as reported by the coordinator. Used for collecting temporal metrics
// relative to when the query was fired in the coordinator.
uint64 query_start_time_ns = 6;
// Producer IDs whose updates must be reported to the coordinator for remote consumers.
repeated uint64 dynamic_filter_remote_producer_ids = 8;
}

message WorkUnitBatch {
Expand Down
Loading
Loading