From c071afb2404236feff500270e316a0fb30af0143 Mon Sep 17 00:00:00 2001 From: Jayant Shrivastava Date: Mon, 17 Aug 2026 19:22:39 +0000 Subject: [PATCH] feat: forward remote dynamic filter updates to coordinator MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Task-local consumer producer F ────────────────> consumer F shared expression; no RPC Cross-task consumer producer F ── boundary anchor F ──> SetPlanRequest.report_ids=[F] │ ├── wait_update() -----+ ├── wait_complete() ---+--> latest observed F --RPC--> coordinator └── query cancellation +--> stop Derive the report allowlist from producer IDs intersected with network-boundary anchor IDs. Workers observe only allowlisted producers, while DataFusion continues updating task-local consumers directly in memory. Watch-based updates may naturally coalesce, and completion remains observed separately because it does not advance the generation. --- src/coordinator/dynamic_filter_registry.rs | 16 ++- src/coordinator/query_coordinator.rs | 43 ++++++-- src/dynamic_filtering/discovery.rs | 43 +++++++- src/lib.rs | 7 +- src/protocol/grpc/generated/worker.rs | 16 ++- src/protocol/grpc/mod.rs | 1 + src/protocol/grpc/worker.proto | 11 ++ src/protocol/grpc/worker_client.rs | 45 +++++++- src/protocol/grpc/worker_service.rs | 56 +++++++++- src/protocol/mod.rs | 5 +- src/protocol/worker_channel.rs | 17 +++ src/worker/impl_coordinator_channel.rs | 110 ++++++++++++++++--- src/worker/worker_service.rs | 12 ++ tests/dynamic_filtering/aggregates.rs | 1 + tests/dynamic_filtering/collect_left_join.rs | 1 + tests/dynamic_filtering/common.rs | 33 +++++- tests/dynamic_filtering/partitioned_join.rs | 1 + tests/dynamic_filtering/sorts.rs | 1 + tests/stateful_data_cleanup.rs | 71 ++++++++++-- 19 files changed, 431 insertions(+), 59 deletions(-) diff --git a/src/coordinator/dynamic_filter_registry.rs b/src/coordinator/dynamic_filter_registry.rs index 6642fdf65..c293f8265 100644 --- a/src/coordinator/dynamic_filter_registry.rs +++ b/src/coordinator/dynamic_filter_registry.rs @@ -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)] @@ -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, + 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. diff --git a/src/coordinator/query_coordinator.rs b/src/coordinator/query_coordinator.rs index 499dab7ad..d5bea6990 100644 --- a/src/coordinator/query_coordinator.rs +++ b/src/coordinator/query_coordinator.rs @@ -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::{ @@ -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()), @@ -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, @@ -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)); @@ -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, @@ -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, @@ -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); @@ -335,6 +341,9 @@ impl<'a> StageCoordinator<'a> { store.insert(task_key, filters); } } + WorkerToCoordinatorMsg::ProducedDynamicFilter(_) => { + dynamic_filter_registry.record_update_received(); + } } } }); @@ -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, Vec)> { + fn task_specialized_plan(&self, task_i: usize) -> Result { let session_config = self.task_ctx.session_config(); let wuf_registry = session_config .get_extension::() @@ -472,7 +478,16 @@ 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, + }) } } @@ -480,6 +495,12 @@ fn keep_stream_alive(notify: Arc) -> impl Stream + futures::stream::once(notify.notified_owned()).filter_map(|()| futures::future::ready(None)) } +struct TaskSpecializedPlan { + plan: Arc, + work_unit_feed_declarations: Vec, + dynamic_filter_remote_producer_ids: Vec, +} + pub(super) struct NotifyGuard(Arc); impl Drop for NotifyGuard { diff --git a/src/dynamic_filtering/discovery.rs b/src/dynamic_filtering/discovery.rs index a3b011c33..9b833a01a 100644 --- a/src/dynamic_filtering/discovery.rs +++ b/src/dynamic_filtering/discovery.rs @@ -11,6 +11,7 @@ use std::sync::Arc; #[derive(Clone)] pub(crate) struct DiscoveredDynamicFilterProducer { pub(crate) id: u64, + pub(crate) expression: Arc, } /// A dynamic-filter consumer discovered in an execution plan along with the schema it is evaluated @@ -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) })?; @@ -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, +) -> Result> { + 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. @@ -212,7 +238,7 @@ mod tests { ) .await?; assert_snapshot!(display, @r" - Stage 5 + Stage 5 remote_producers=[1] AggregateExec HashJoinExec producers=[1] NetworkShuffleExec @@ -222,7 +248,7 @@ mod tests { RepartitionExec AggregateExec DataSourceExec consumers=[1] - Stage 3 + Stage 3 remote_producers=[2] RepartitionExec HashJoinExec producers=[2] NetworkShuffleExec @@ -262,7 +288,7 @@ mod tests { ) .await?; assert_snapshot!(display, @r" - Stage 4 + Stage 4 remote_producers=[1] AggregateExec HashJoinExec producers=[1] AggregateExec @@ -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 diff --git a/src/lib.rs b/src/lib.rs index d8c9a7896..523252982 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -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, diff --git a/src/protocol/grpc/generated/worker.rs b/src/protocol/grpc/generated/worker.rs index 78bbd1058..760780e31 100644 --- a/src/protocol/grpc/generated/worker.rs +++ b/src/protocol/grpc/generated/worker.rs @@ -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, } /// Nested message and enum types in `WorkerToCoordinatorMsg`. @@ -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)] @@ -67,6 +70,14 @@ pub struct DynamicFilter { #[prost(bytes = "vec", tag = "2")] pub expression_proto: ::prost::alloc::vec::Vec, } +#[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, +} #[derive(Clone, PartialEq, ::prost::Message)] pub struct TaskCompletedDynamicFilters { #[prost(message, repeated, tag = "1")] @@ -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, } /// Nested message and enum types in `SetPlanRequest`. pub mod set_plan_request { diff --git a/src/protocol/grpc/mod.rs b/src/protocol/grpc/mod.rs index e432b607b..59dc730d5 100644 --- a/src/protocol/grpc/mod.rs +++ b/src/protocol/grpc/mod.rs @@ -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::{ diff --git a/src/protocol/grpc/worker.proto b/src/protocol/grpc/worker.proto index 61319f5f8..3e0ab81c4 100644 --- a/src/protocol/grpc/worker.proto +++ b/src/protocol/grpc/worker.proto @@ -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; } } @@ -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; } @@ -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 { diff --git a/src/protocol/grpc/worker_client.rs b/src/protocol/grpc/worker_client.rs index de5cce49a..140ae88db 100644 --- a/src/protocol/grpc/worker_client.rs +++ b/src/protocol/grpc/worker_client.rs @@ -10,9 +10,9 @@ use crate::{ BytesMetricExt, CoordinatorToWorkerMsg, DISTRIBUTED_DATAFUSION_TASK_ID_LABEL, DistributedConfig, ExecuteTaskRequest, FirstLatencyMetric, GetWorkerInfoRequest, GetWorkerInfoResponse, LatencyMetricExt, LoadInfo, MaxLatencyMetric, MaybeEncoded, - MinLatencyMetric, P50LatencyMetric, P95LatencyMetric, ProducerHead, SetPlanRequest, - TaskCompletedDynamicFilters, TaskDynamicFilter, TaskKey, TaskMetrics, WorkUnitBatch, - WorkUnitFeedDeclaration, WorkUnitMsg, WorkerChannel, WorkerToCoordinatorMsg, + MinLatencyMetric, P50LatencyMetric, P95LatencyMetric, ProducedDynamicFilter, ProducerHead, + SetPlanRequest, TaskCompletedDynamicFilters, TaskDynamicFilter, TaskKey, TaskMetrics, + WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg, WorkerChannel, WorkerToCoordinatorMsg, }; use arrow_flight::FlightData; use arrow_flight::decode::FlightRecordBatchStream; @@ -489,6 +489,7 @@ fn encode_set_plan_request( .collect(), target_worker_url: request.target_worker_url.to_string(), query_start_time_ns: request.query_start_time_ns as u64, + dynamic_filter_remote_producer_ids: request.dynamic_filter_remote_producer_ids, }) } @@ -552,10 +553,24 @@ fn decode_worker_to_coordinator_msg( decode_task_completed_dynamic_filters(filters)?, ) } + pb::worker_to_coordinator_msg::Inner::ProducedDynamicFilter(filter) => { + WorkerToCoordinatorMsg::ProducedDynamicFilter(Box::new( + decode_produced_dynamic_filter(filter)?, + )) + } }, ) } +fn decode_produced_dynamic_filter( + filter: pb::ProducedDynamicFilter, +) -> Result { + Ok(ProducedDynamicFilter { + expression_id: filter.expression_id, + expression: MaybeEncoded::Encoded(filter.expression_proto), + }) +} + fn decode_task_completed_dynamic_filters( filters: pb::TaskCompletedDynamicFilters, ) -> Result { @@ -764,6 +779,7 @@ impl> Future for ElapsedComputeFuture { #[cfg(test)] mod tests { use super::*; + use datafusion_proto::protobuf::PhysicalExprNode; use futures::StreamExt; use futures::stream::unfold; @@ -839,4 +855,27 @@ mod tests { assert!(expensive_time.value() > cheap_time.value()); } + + #[test] + fn decode_produced_dynamic_filter() -> Result<()> { + let expression = PhysicalExprNode::default(); + let decoded = decode_worker_to_coordinator_msg(pb::WorkerToCoordinatorMsg { + inner: Some(pb::worker_to_coordinator_msg::Inner::ProducedDynamicFilter( + pb::ProducedDynamicFilter { + expression_id: 42, + expression_proto: expression.encode_to_vec(), + }, + )), + })?; + + let WorkerToCoordinatorMsg::ProducedDynamicFilter(decoded) = decoded else { + panic!("expected produced dynamic filter"); + }; + assert_eq!(decoded.expression_id, 42); + let MaybeEncoded::Encoded(decoded_expression) = decoded.expression else { + panic!("expected encoded dynamic filter"); + }; + assert_eq!(decoded_expression, expression.encode_to_vec()); + Ok(()) + } } diff --git a/src/protocol/grpc/worker_service.rs b/src/protocol/grpc/worker_service.rs index 534c949b5..9dabf0b55 100644 --- a/src/protocol/grpc/worker_service.rs +++ b/src/protocol/grpc/worker_service.rs @@ -7,8 +7,9 @@ use crate::common::{deserialize_uuid, now_ns}; use crate::protocol::grpc::{ObservabilityServiceImpl, ObservabilityServiceServer}; use crate::{ CoordinatorToWorkerMsg, DistributedConfig, ExecuteTaskRequest, LoadInfo, MaybeEncoded, - ProducerHead, SetPlanRequest, TaskCompletedDynamicFilters, TaskKey, TaskMetrics, WorkUnitBatch, - WorkUnitFeedDeclaration, WorkUnitMsg, Worker, WorkerResolver, WorkerToCoordinatorMsg, + ProducedDynamicFilter, ProducerHead, SetPlanRequest, TaskCompletedDynamicFilters, TaskKey, + TaskMetrics, WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg, Worker, WorkerResolver, + WorkerToCoordinatorMsg, }; use crate::worker::CoordinatorChannelResult; @@ -231,6 +232,7 @@ fn decode_set_plan_request(request: pb::SetPlanRequest) -> Result>()?, target_worker_url: parse_url(&request.target_worker_url, "target_worker_url")?, query_start_time_ns: request.query_start_time_ns as usize, + dynamic_filter_remote_producer_ids: request.dynamic_filter_remote_producer_ids, }) } @@ -281,10 +283,28 @@ fn encode_worker_to_coordinator_msg( encode_task_completed_dynamic_filters(filters, task_ctx)?, ) } + WorkerToCoordinatorMsg::ProducedDynamicFilter(filter) => { + pb::worker_to_coordinator_msg::Inner::ProducedDynamicFilter( + encode_produced_dynamic_filter(*filter, task_ctx)?, + ) + } }), }) } +fn encode_produced_dynamic_filter( + filter: ProducedDynamicFilter, + task_ctx: &Arc, +) -> Result { + Ok(pb::ProducedDynamicFilter { + expression_id: filter.expression_id, + expression_proto: filter + .expression + .encode(task_ctx) + .map_err(datafusion_error_to_tonic_status)?, + }) +} + fn encode_task_completed_dynamic_filters( filters: TaskCompletedDynamicFilters, task_ctx: &Arc, @@ -461,3 +481,35 @@ fn garbage_collect_arrays( &RecordBatchOptions::new().with_row_count(Some(row_count)), )?) } + +#[cfg(test)] +mod tests { + use super::*; + use datafusion::physical_expr::expressions::lit; + use datafusion::prelude::SessionContext; + + #[test] + fn encode_produced_dynamic_filter() { + let expression = lit(true); + let task_ctx = SessionContext::new().task_ctx(); + let expected = MaybeEncoded::Decoded(Arc::clone(&expression)) + .encode(&task_ctx) + .unwrap(); + let encoded = encode_worker_to_coordinator_msg( + WorkerToCoordinatorMsg::ProducedDynamicFilter(Box::new(ProducedDynamicFilter { + expression_id: 42, + expression: MaybeEncoded::Decoded(expression), + })), + &task_ctx, + ) + .unwrap(); + + let Some(pb::worker_to_coordinator_msg::Inner::ProducedDynamicFilter(encoded)) = + encoded.inner + else { + panic!("expected produced dynamic filter"); + }; + assert_eq!(encoded.expression_id, 42); + assert_eq!(encoded.expression_proto, expected); + } +} diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index 1ee655e3a..f3fb839b4 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -11,6 +11,7 @@ pub use channel_resolver::{ChannelResolver, get_distributed_channel_resolver}; pub use in_process::LocalWorkerContext; pub use worker_channel::{ CoordinatorToWorkerMsg, ExecuteTaskRequest, GetWorkerInfoRequest, GetWorkerInfoResponse, - LoadInfo, SetPlanRequest, TaskCompletedDynamicFilters, TaskDynamicFilter, TaskKey, TaskMetrics, - WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg, WorkerChannel, WorkerToCoordinatorMsg, + LoadInfo, ProducedDynamicFilter, SetPlanRequest, TaskCompletedDynamicFilters, + TaskDynamicFilter, TaskKey, TaskMetrics, WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg, + WorkerChannel, WorkerToCoordinatorMsg, }; diff --git a/src/protocol/worker_channel.rs b/src/protocol/worker_channel.rs index a2650ea19..f3cb8f9e3 100644 --- a/src/protocol/worker_channel.rs +++ b/src/protocol/worker_channel.rs @@ -80,6 +80,9 @@ pub struct SetPlanRequest { pub task_count: usize, /// The subplan the worker is expected to execute. pub plan: MaybeEncoded>, + /// Producer expression IDs whose consumers cross a network boundary. Workers observe and + /// report updates to the coordinator. + pub dynamic_filter_remote_producer_ids: Vec, /// Information about all the work unit feeds that will be streamed from coordinator to worker. /// This information is needed here because at the moment of setting the plan, all the appropriate /// channels for the incoming work unit feeds need to be constructed. @@ -125,12 +128,26 @@ pub enum WorkerToCoordinatorMsg { /// Sends the final dynamic filters used by dynamic filter consumers back to the coorindator /// for displaying. TaskCompletedDynamicFilters(TaskCompletedDynamicFilters), + /// Sends an observed producer dynamic-filter state to the coordinator. This update + /// is to be used for runtime dynamic filtering. + ProducedDynamicFilter(Box), /// Load information reported by a task. This information is used for dynamically /// sizing the number of workers involved in a query. LoadInfo(LoadInfo), LoadInfoEos, } +/// A dynamic filter state update produced by a plan node in a task to be sent to the coordinator. +#[derive(Clone, Debug)] +pub struct ProducedDynamicFilter { + pub expression_id: u64, + /// Note that sending an update via a live pointer could mean that the dynamic filter updates during + /// transport. This means that observations at the coordinator may repeat or skip generations, but never + /// regress. Since the worker monitors updates and completion, it's guaranteed that the completed + /// filter state will not be missed. + pub expression: MaybeEncoded>, +} + #[derive(Clone, Debug, Default)] pub struct TaskCompletedDynamicFilters { /// Final expressions keyed by their DataFusion physical-expression ID. The TaskKey is diff --git a/src/worker/impl_coordinator_channel.rs b/src/worker/impl_coordinator_channel.rs index a6477f6b6..bd1898c38 100644 --- a/src/worker/impl_coordinator_channel.rs +++ b/src/worker/impl_coordinator_channel.rs @@ -1,26 +1,34 @@ use crate::common::TreeNodeExt; -use crate::dynamic_filtering::discover_dynamic_filter_consumers; +use crate::dynamic_filtering::{ + discover_dynamic_filter_consumers, discover_dynamic_filter_producers, +}; use crate::events::{WorkerPlanRewriteEvent, WorkerPlanRewriteHandlers}; use crate::execution_plans::SamplerExec; use crate::protocol::LocalWorkerContext; +#[cfg(feature = "integration")] +use crate::protocol::grpc::on_drop_stream; use crate::work_unit_feed::{RemoteWorkUnitFeedRegistry, set_work_unit_received_time}; use crate::worker::task_data::TaskDataMetrics; use crate::{ CoordinatorToWorkerMsg, DistributedConfig, DistributedExt, DistributedTaskContext, - MaybeEncoded, SetPlanRequest, TaskCompletedDynamicFilters, TaskData, TaskDynamicFilter, - TaskMetrics, Worker, WorkerQueryContext, WorkerToCoordinatorMsg, + MaybeEncoded, ProducedDynamicFilter, SetPlanRequest, TaskCompletedDynamicFilters, TaskData, + TaskDynamicFilter, TaskMetrics, Worker, WorkerQueryContext, WorkerToCoordinatorMsg, }; use datafusion::common::tree_node::TreeNodeRecursion; -use datafusion::common::{DataFusionError, Result, exec_datafusion_err}; +use datafusion::common::{DataFusionError, HashSet, Result, exec_datafusion_err}; use datafusion::execution::{SessionStateBuilder, TaskContext}; +use datafusion::physical_expr::PhysicalExpr; +use datafusion::physical_expr::expressions::DynamicFilterPhysicalExpr; use datafusion::physical_plan::ExecutionPlan; use datafusion::prelude::SessionConfig; use futures::stream::{BoxStream, FuturesUnordered, select_all}; use futures::{FutureExt, StreamExt, TryStreamExt}; use http::HeaderMap; +#[cfg(feature = "integration")] +use std::sync::atomic::Ordering; use std::sync::{Arc, OnceLock}; -use tokio::sync::oneshot; use tokio::sync::oneshot::Sender; +use tokio::sync::{oneshot, watch}; /// Return value of the [Worker::coordinator_channel] method. pub struct CoordinatorChannelResult { @@ -123,6 +131,16 @@ impl Worker { let task_data = task_data_result.map_err(DataFusionError::Shared)?; + let dynamic_filter_remote_producer_ids: HashSet<_> = request + .dynamic_filter_remote_producer_ids + .iter() + .copied() + .collect(); + let producer_filters = discover_dynamic_filter_producers(&task_data.base_plan)? + .into_iter() + .filter(|producer| dynamic_filter_remote_producer_ids.contains(&producer.id)); + let (producer_cancel_tx, producer_cancel_rx) = watch::channel(false); + // Continue reading remaining messages (work unit feed data) in the background. let mut work_unit_senders = Some(remote_work_unit_feed_registry.senders); let task_data_entries = Arc::clone(&self.task_data_entries); @@ -175,8 +193,13 @@ impl Worker { } } + // Cancel any dynamic filter producce streams that did not complete for any reason. + // It's expected that dynamic filters should complete and send their updates before + // task execution ends. + producer_cancel_tx.send_replace(true); + // Send metrics and completed dynamic filters if enabled. - // TODO(#686): handle errors + let metrics_tx = task_data.metrics_tx.lock().unwrap().take(); let dynamic_filters_tx = task_data .completed_dynamic_filters_tx @@ -231,17 +254,78 @@ impl Worker { }, ); + let produced_dynamic_filters_stream = + select_all(producer_filters.into_iter().map(|producer| { + produced_dynamic_filter_stream( + producer.id, + producer.expression, + producer_cancel_rx.clone(), + ) + })); + + let stream = select_all([ + produced_dynamic_filters_stream.boxed(), + load_info_stream.boxed(), + metrics_stream.boxed(), + dynamic_filters_stream.boxed(), + ]) + .map(Ok) + .boxed(); + + #[cfg(feature = "integration")] + let stream = self.track_worker_to_coordinator_stream(stream); + Ok(CoordinatorChannelResult { task_ctx: Arc::clone(&task_data.task_ctx), - stream: select_all([ - load_info_stream.boxed(), - metrics_stream.boxed(), - dynamic_filters_stream.boxed(), - ]) - .map(Ok) - .boxed(), + stream, }) } + + #[cfg(feature = "integration")] + fn track_worker_to_coordinator_stream( + &self, + stream: BoxStream<'static, Result>, + ) -> BoxStream<'static, Result> { + self.coordinator_channels_running + .fetch_add(1, Ordering::SeqCst); + let channels_running = Arc::clone(&self.coordinator_channels_running); + on_drop_stream(stream, move || { + channels_running.fetch_sub(1, Ordering::SeqCst); + }) + .boxed() + } +} + +/// Streams updates from one dynamic-filter producer until it completes or cancellation is detected. +fn produced_dynamic_filter_stream( + expression_id: u64, + expression: Arc, + cancel_rx: watch::Receiver, +) -> BoxStream<'static, WorkerToCoordinatorMsg> { + futures::stream::unfold(Some((expression, cancel_rx)), move |state| async move { + let (expression, mut cancel_rx) = state?; + let dynamic_filter = expression + .downcast_ref::() + .expect("producer discovery returns DynamicFilterPhysicalExpr"); + + // `wait_update()` uses a Tokio watch channel, so multiple generations can + // are natively "deduped" into just one update. This helps avoid too many + // update messages. If this becomes an issue, we can introduce an artificial + // backoff. + let completed = tokio::select! { + _ = dynamic_filter.wait_update() => false, + _ = dynamic_filter.wait_complete() => true, + _ = cancel_rx.wait_for(|cancelled| *cancelled) => return None, + }; + let message = + WorkerToCoordinatorMsg::ProducedDynamicFilter(Box::new(ProducedDynamicFilter { + expression_id, + expression: MaybeEncoded::Decoded(Arc::clone(&expression)), + })); + let next = (!completed).then_some((expression, cancel_rx)); + Some((message, next)) + }) + .boxed() } /// Finds all consumed dynamic filters for the completed task report. diff --git a/src/worker/worker_service.rs b/src/worker/worker_service.rs index c0eb71cce..5e32bd8d6 100644 --- a/src/worker/worker_service.rs +++ b/src/worker/worker_service.rs @@ -6,6 +6,8 @@ use datafusion::execution::runtime_env::RuntimeEnv; use moka::future::Cache; use std::borrow::Cow; use std::sync::Arc; +#[cfg(feature = "integration")] +use std::sync::atomic::{AtomicUsize, Ordering}; use std::time::Duration; use url::Url; @@ -24,6 +26,8 @@ pub struct Worker { pub(super) session_builder: Arc, pub(crate) max_message_size: Option, pub(super) version: Cow<'static, str>, + #[cfg(feature = "integration")] + pub(super) coordinator_channels_running: Arc, } impl Default for Worker { @@ -35,6 +39,8 @@ impl Default for Worker { session_builder: Arc::new(DefaultSessionBuilder), max_message_size: Some(usize::MAX), version: Cow::Borrowed(""), + #[cfg(feature = "integration")] + coordinator_channels_running: Arc::new(AtomicUsize::new(0)), } } } @@ -103,4 +109,10 @@ impl Worker { self.task_data_entries.run_pending_tasks().await; self.task_data_entries.entry_count() as usize } + + /// Returns the number of live worker-to-coordinator streams. + #[cfg(feature = "integration")] + pub fn coordinator_channels_running(&self) -> usize { + self.coordinator_channels_running.load(Ordering::SeqCst) + } } diff --git a/tests/dynamic_filtering/aggregates.rs b/tests/dynamic_filtering/aggregates.rs index 41c3cc4f6..c650dbcd8 100644 --- a/tests/dynamic_filtering/aggregates.rs +++ b/tests/dynamic_filtering/aggregates.rs @@ -45,6 +45,7 @@ mod tests { ) "#, ) + .expect_dynamic_filter_updates() .execute() .await?; assert_snapshot!(display, @" diff --git a/tests/dynamic_filtering/collect_left_join.rs b/tests/dynamic_filtering/collect_left_join.rs index 00105d363..49cd54e70 100644 --- a/tests/dynamic_filtering/collect_left_join.rs +++ b/tests/dynamic_filtering/collect_left_join.rs @@ -69,6 +69,7 @@ mod tests { "#, ) .with_broadcast_joins() + .expect_dynamic_filter_updates() .execute() .await?; assert_snapshot!(display, @r" diff --git a/tests/dynamic_filtering/common.rs b/tests/dynamic_filtering/common.rs index 1833da626..e468ecf8d 100644 --- a/tests/dynamic_filtering/common.rs +++ b/tests/dynamic_filtering/common.rs @@ -24,6 +24,7 @@ pub(crate) struct TestQuery<'a> { broadcast_joins: bool, one_task_per_leaf: bool, collect_dynamic_filters: bool, + expect_dynamic_filter_updates: bool, } impl<'a> TestQuery<'a> { @@ -34,6 +35,7 @@ impl<'a> TestQuery<'a> { broadcast_joins: false, one_task_per_leaf: false, collect_dynamic_filters: true, + expect_dynamic_filter_updates: false, } } @@ -61,6 +63,11 @@ impl<'a> TestQuery<'a> { self } + pub(crate) fn expect_dynamic_filter_updates(mut self) -> Self { + self.expect_dynamic_filter_updates = true; + self + } + pub(crate) async fn execute(self) -> Result { let (ctx, _guard, _) = start_localhost_context(2, DefaultSessionBuilder).await; let mut ctx = ctx @@ -69,13 +76,15 @@ impl<'a> TestQuery<'a> { if self.one_task_per_leaf { ctx = ctx.with_distributed_desired_task_count_handler(1usize); } - if !self.broadcast_joins { - // Force partitioned hash joins. + { 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; + if !self.broadcast_joins { + // Force partitioned hash joins. + optimizer.hash_join_single_partition_threshold = 0; + optimizer.hash_join_single_partition_threshold_rows = 0; + } } register_parquet_tables(&ctx).await?; execute_query_and_display( @@ -83,6 +92,7 @@ impl<'a> TestQuery<'a> { self.sql, self.expected_rows, self.collect_dynamic_filters, + self.expect_dynamic_filter_updates, ) .await } @@ -114,7 +124,7 @@ pub(crate) async fn execute_range_partitioned_query( register_range_partitioned_table(&ctx, "dim", "testdata/join/parquet/dim", "d_dkey").await?; register_range_partitioned_table(&ctx, "fact", "testdata/join/parquet/fact", "f_dkey").await?; - execute_query_and_display(&ctx, sql, expected_rows, true).await + execute_query_and_display(&ctx, sql, expected_rows, true, false).await } async fn register_range_partitioned_table( @@ -146,6 +156,7 @@ async fn execute_query_and_display( sql: &str, expected_rows: usize, collect_dynamic_filters: bool, + expect_dynamic_filter_updates: bool, ) -> Result { let plan = ctx.sql(sql).await?.create_physical_plan().await?; let task_ctx = ctx.task_ctx(); @@ -165,6 +176,18 @@ async fn execute_query_and_display( ); assert_eq!(display_plan_ascii(plan.as_ref(), false), original_display); + if expect_dynamic_filter_updates { + let updates = plan + .metrics() + .expect("DistributedExec has metrics") + .sum(|metric| metric.value().name() == "dynamic_filter_updates_received") + .map_or(0, |metric| metric.as_usize()); + assert!( + updates > 0, + "expected dynamic_filter_updates_received > 0, got {updates}" + ); + } + Ok(display_plan_ascii( plan_with_dynamic_filters.as_ref(), false, diff --git a/tests/dynamic_filtering/partitioned_join.rs b/tests/dynamic_filtering/partitioned_join.rs index a291db9f7..1743c5da3 100644 --- a/tests/dynamic_filtering/partitioned_join.rs +++ b/tests/dynamic_filtering/partitioned_join.rs @@ -70,6 +70,7 @@ mod tests { JOIN weather probe ON build.key = probe."RainToday" "#, ) + .expect_dynamic_filter_updates() .execute() .await?; assert_snapshot!(display, @" diff --git a/tests/dynamic_filtering/sorts.rs b/tests/dynamic_filtering/sorts.rs index 48ad97772..3522b3444 100644 --- a/tests/dynamic_filtering/sorts.rs +++ b/tests/dynamic_filtering/sorts.rs @@ -54,6 +54,7 @@ mod tests { "#, ) .with_expected_rows(10) + .expect_dynamic_filter_updates() .execute() .await?; assert_snapshot!(display, @" diff --git a/tests/stateful_data_cleanup.rs b/tests/stateful_data_cleanup.rs index 0208a240c..409546946 100644 --- a/tests/stateful_data_cleanup.rs +++ b/tests/stateful_data_cleanup.rs @@ -1,7 +1,6 @@ #[cfg(all(feature = "integration", feature = "tpch", test))] mod tests { - use datafusion::common::instant::Instant; - use datafusion::error::Result; + use datafusion::common::{Result, instant::Instant}; use datafusion::physical_plan::execute_stream; use datafusion::prelude::SessionContext; use datafusion_distributed::test_utils::localhost::start_localhost_context; @@ -14,8 +13,11 @@ mod tests { use std::path::Path; use std::time::Duration; use test_case::test_case; - use tokio::sync::OnceCell; - use tokio::time::timeout; + use tokio::{ + spawn, + sync::OnceCell, + time::{sleep, timeout}, + }; const NUM_WORKERS: usize = 4; const TPCH_SCALE_FACTOR: f64 = 1.0; @@ -62,10 +64,48 @@ mod tests { Ok(()) } - /// Polls until every worker reports 0 running tasks, or fails after 5s. Task entries are - /// torn down asynchronously once the coordinator->worker channel disconnects (shortly after - /// the query's output stream is dropped), so cleanup is not observable synchronously the - /// instant the query future resolves — hence the poll rather than an immediate assert. + #[test_case((false, false); "metrics_disabled_static_planner")] + #[test_case((true, false); "metrics_enabled_static_planner")] + #[test_case((false, true); "metrics_disabled_dynamic_planner")] + #[test_case((true, true); "metrics_enabled_dynamic_planner")] + #[tokio::test(flavor = "multi_thread")] + async fn cancellation_closes_coordinator_channels( + (collect_metrics, adaptive): (bool, bool), + ) -> Result<()> { + let (mut d_ctx, _guard, workers) = + start_localhost_context(NUM_WORKERS, DefaultSessionBuilder).await; + d_ctx.set_distributed_metrics_collection(collect_metrics)?; + d_ctx.set_distributed_dynamic_task_count(adaptive)?; + + #[allow(clippy::disallowed_methods)] + let execution = spawn(run_tpch_query(d_ctx, "q2")); + + timeout(Duration::from_secs(10), async { + while coordinator_channels_running(&workers) == 0 { + assert!( + !execution.is_finished(), + "query completed before opening a coordinator channel" + ); + sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("query did not open a coordinator channel within 10 seconds"); + assert!(coordinator_channels_running(&workers) > 0); + execution.abort(); + let error = timeout(Duration::from_secs(1), execution) + .await + .expect("cancelled query did not stop within one second") + .expect_err("query completed before it was cancelled"); + assert!(error.is_cancelled()); + assert_no_tasks_running_eventually(&workers).await; + + Ok(()) + } + + /// Polls until every worker reports 0 running tasks and worker-to-coordinator streams, or fails + /// after 5s. Cleanup is asynchronous after the query output is dropped, so it is not observable + /// synchronously when the query future resolves. async fn assert_no_tasks_running_eventually(workers: &[Worker]) { let start = Instant::now(); loop { @@ -73,17 +113,26 @@ mod tests { for worker in workers { tasks_running += worker.tasks_running().await; } - if tasks_running == 0 { + let channels_running = coordinator_channels_running(workers); + if tasks_running == 0 && channels_running == 0 { return; } assert!( start.elapsed() < Duration::from_secs(5), - "Expected 0 tasks running across workers, but still had {tasks_running} after 5s" + "Expected no running tasks or coordinator channels, but still had \ + {tasks_running} tasks and {channels_running} channels after 5s" ); - tokio::time::sleep(Duration::from_millis(50)).await; + sleep(Duration::from_millis(50)).await; } } + fn coordinator_channels_running(workers: &[Worker]) -> usize { + workers + .iter() + .map(Worker::coordinator_channels_running) + .sum() + } + async fn run_tpch_query(d_ctx: SessionContext, query_id: &str) -> Result<()> { let data_dir = ensure_tpch_data(TPCH_SCALE_FACTOR, TPCH_DATA_PARTS).await;