diff --git a/src/coordinator/distributed.rs b/src/coordinator/distributed.rs index 9b29f37e9..179498c73 100644 --- a/src/coordinator/distributed.rs +++ b/src/coordinator/distributed.rs @@ -222,28 +222,29 @@ impl ExecutionPlan for DistributedExec { let prepared_plan = Arc::clone(&self.prepared_plan); let collect_dynamic_filters = self.completed_dynamic_filter_store.is_some(); + let mut builder = RecordBatchReceiverStreamBuilder::new(self.schema(), 1); + let tx = builder.tx(); let query_coordinator = Arc::new(QueryCoordinator::new( Arc::clone(&context), &self.metrics, self.metrics_store.clone(), self.completed_dynamic_filter_store.clone(), + &mut builder, )); - - let mut builder = RecordBatchReceiverStreamBuilder::new(self.schema(), 1); - let tx = builder.tx(); - + // Capture the guard before spawning so even an unpolled execution future cancels on drop. + let guard = query_coordinator.spawner.end_query_guard(); builder.spawn(async move { // Dropping this `guard` is what signals the coordinator->worker channel to be dropped, // which triggers a chain reaction that ends up also gracefully closing the // worker->coordinator channel. The flow looks like this: - // 1. The query ends normally, as all Arrow RecordBatches are already streamed. - // 2. The `guard` here is dropped. - // 3. In StageCoordinator::send_plan_task(), `end_stream_notifier` fires and the - // coordinator->worker channel is gracefully ended. - // 4. The coordinator->worker channel EOS is received in `impl_coordinator_channel.rs`. - // 5. The metrics are send back in the worker->coordinator channel, and then that - // channel is closed. - let guard = query_coordinator.end_query_guard(); + // 1. Execution ends normally. All Arrow RecordBatches are emitted (note that + // the response stream does not terminate yet). + // 2. The `guard` here is dropped, signalling the coordinator->worker streams to finish. + // 3. The worker observes end-of-stream in `impl_coordinator_channel.rs`. + // 4. The worker sends final metrics etc., if enabled. + // 5. The tasks owned by the response stream finish, such as tasks collecting metrics etc. + // 6. The response stream finishes, ending the query for the user. + let _guard = guard; let d_cfg = DistributedConfig::from_config_options(context.session_config().options())?; let mut prepared = match d_cfg.dynamic_task_count { @@ -266,13 +267,10 @@ impl ExecutionPlan for DistributedExec { })?; let mut stream = head_stage.execute(partition, context)?; while let Some(msg) = stream.next().await { - if tx.send(msg).await.is_err() { + if tx.send(Ok(msg?)).await.is_err() { break; // channel closed } } - drop(guard); - drop(tx); - query_coordinator.drain_pending_tasks().await?; Ok(()) }); diff --git a/src/coordinator/dynamic_filter_registry.rs b/src/coordinator/dynamic_filter_registry.rs index 22aee24f6..4905f5c9c 100644 --- a/src/coordinator/dynamic_filter_registry.rs +++ b/src/coordinator/dynamic_filter_registry.rs @@ -1,7 +1,9 @@ use crate::dynamic_filtering::discover_dynamic_filter_consumers; -use crate::{ProducedDynamicFilter, TaskKey}; +use crate::{ + ApplyDynamicFilter, CoordinatorToWorkerMsg, MaybeEncoded, ProducedDynamicFilter, TaskKey, +}; use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion}; -use datafusion::common::{HashMap, HashSet, Result, internal_err}; +use datafusion::common::{HashMap, HashSet, Result, exec_err, internal_err}; use datafusion::execution::TaskContext; use datafusion::physical_expr::expressions::DynamicFilterPhysicalExpr; use datafusion::physical_expr_common::metrics::{ExecutionPlanMetricsSet, MetricBuilder}; @@ -14,7 +16,10 @@ use datafusion_proto::protobuf::physical_expr_node::ExprType; use datafusion_proto::protobuf::{ PhysicalBinaryExprNode, PhysicalDynamicFilterNode, PhysicalExprNode, }; +use prost::Message; use std::sync::{Arc, Mutex}; +use tokio::sync::mpsc::UnboundedSender; +use tokio_util::sync::CancellationToken; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(super) enum DynamicFilterMergeMode { @@ -42,6 +47,8 @@ pub(super) struct PlannedDynamicFilter { pub(super) consumer_tasks: HashSet, /// Full dynamic filter containing the merged predicate and its completion state. pub(super) merged: Option, + /// Immutable snapshot shared by local and remote delivery, encoded once per change. + merged_bytes: Option>, } #[derive(Default)] @@ -49,6 +56,8 @@ pub(super) struct DynamicFilterRegistryState { pub(super) filters: HashMap, /// Track which stages have registered all of their tasks. pub(super) sealed_stages: HashSet, + task_senders: HashMap>, + delivered: HashSet<(u64, TaskKey)>, } /// Query-scoped hub for distributed dynamic filtering. @@ -60,14 +69,19 @@ pub(super) struct DynamicFilterRegistryState { pub(crate) struct DynamicFilterRegistry { pub(super) state: Mutex, dynamic_filter_updates_received: Count, + query_finished: CancellationToken, } impl DynamicFilterRegistry { - pub(crate) fn new(metrics: &ExecutionPlanMetricsSet) -> Self { + pub(crate) fn new( + metrics: &ExecutionPlanMetricsSet, + query_finished: CancellationToken, + ) -> Self { Self { state: Mutex::new(DynamicFilterRegistryState::default()), dynamic_filter_updates_received: MetricBuilder::new(metrics) .global_counter("dynamic_filter_updates_received"), + query_finished, } } @@ -150,14 +164,32 @@ impl DynamicFilterRegistry { Ok(()) } + pub(crate) fn register_sender( + &self, + task_key: TaskKey, + sender: UnboundedSender, + ) -> Result<()> { + let mut state = self.state.lock().expect("dynamic filter registry poisoned"); + state.task_senders.insert(task_key, sender); + state.delivered.retain(|(_, task)| *task != task_key); + let ids = state.filters.keys().copied().collect::>(); + for id in ids { + self.dispatch(&mut state, id)?; + } + Ok(()) + } + /// Mark that a stage has registered all of its tasks. - pub(crate) fn seal_stage(&self, stage_id: usize) { + pub(crate) fn seal_stage(&self, stage_id: usize) -> Result<()> { 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); + if Self::merge(&mut state, id) { + self.dispatch(&mut state, id)?; + } } + Ok(()) } /// Records a producer's latest dynamic-filter update and recomputes the merged filter. @@ -166,27 +198,41 @@ impl DynamicFilterRegistry { task_key: TaskKey, report: ProducedDynamicFilter, task_ctx: &TaskContext, - ) { + ) -> Result<()> { self.record_update_received(); - let Ok(expression) = report.expression.to_proto(task_ctx) else { - return; - }; + let expression = report.expression.to_proto(task_ctx)?; if expression.expr_id != Some(report.expression_id) { - return; + return exec_err!( + "Dynamic filter {} from task {task_key:?} has mismatched expression ID {:?}", + report.expression_id, + expression.expr_id + ); } let Some(ExprType::DynamicFilter(dynamic_filter)) = expression.expr_type else { - return; + return exec_err!( + "Expected dynamic filter expression for filter {} from task {task_key:?}", + report.expression_id + ); }; if dynamic_filter.inner_expr.is_none() { - return; + return exec_err!( + "Missing inner expression for dynamic filter {} from task {task_key:?}", + report.expression_id + ); } let mut state = self.state.lock().expect("dynamic filter registry poisoned"); let Some(filter) = state.filters.get_mut(&report.expression_id) else { - return; + return exec_err!( + "Received update for unregistered dynamic filter {} from task {task_key:?}", + report.expression_id + ); }; let Some(previous) = filter.producers.get_mut(&task_key) else { - return; + return exec_err!( + "Task {task_key:?} is not a registered producer of dynamic filter {}", + report.expression_id + ); }; if (filter.merge_mode != Some(DynamicFilterMergeMode::Incremental) && !dynamic_filter.is_complete) @@ -194,10 +240,13 @@ impl DynamicFilterRegistry { .as_ref() .is_some_and(|previous| previous.is_complete || previous == dynamic_filter.as_ref()) { - return; + return Ok(()); } *previous = Some(*dynamic_filter); - Self::merge(&mut state, report.expression_id); + if Self::merge(&mut state, report.expression_id) { + self.dispatch(&mut state, report.expression_id)?; + } + Ok(()) } /// Merges partial dynamic filters together for the provided dynamic filter @@ -253,9 +302,58 @@ impl DynamicFilterRegistry { // 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_bytes = Some( + PhysicalExprNode { + expr_id: Some(id), + expr_type: Some(ExprType::DynamicFilter(Box::new(merged.clone()))), + } + .encode_to_vec(), + ); filter.merged = Some(merged); + state + .delivered + .retain(|(expression_id, _)| *expression_id != id); true } + + // Merge and enqueue under the same lock so successive snapshots cannot overtake each other. + // Unbounded channel sends do not wait for the network or the receiving worker. + fn dispatch(&self, state: &mut DynamicFilterRegistryState, id: u64) -> Result<()> { + if self.query_finished.is_cancelled() { + return Ok(()); + } + let Some(filter) = state.filters.get(&id) else { + return Ok(()); + }; + let Some(expression) = &filter.merged_bytes else { + return Ok(()); + }; + for &task_key in &filter.consumer_tasks { + // A task-local consumer is already updated directly by its producer. + if filter.producers.contains_key(&task_key) || state.delivered.contains(&(id, task_key)) + { + continue; + } + let Some(sender) = state.task_senders.get(&task_key) else { + continue; + }; + let update = CoordinatorToWorkerMsg::ApplyDynamicFilter(Box::new(ApplyDynamicFilter { + expression_id: id, + expression: MaybeEncoded::Encoded(expression.clone()), + })); + if sender.send(update).is_err() { + // Closing the channel is expected only once query shutdown has started. + if !self.query_finished.is_cancelled() { + return exec_err!( + "Failed to send dynamic filter {id} to task {task_key:?}: channel closed" + ); + } + return Ok(()); + } + state.delivered.insert((id, task_key)); + } + Ok(()) + } } /// Merges [`PhysicalExprNode`] together by ORing them. diff --git a/src/coordinator/mod.rs b/src/coordinator/mod.rs index 746052a6f..a6990579f 100644 --- a/src/coordinator/mod.rs +++ b/src/coordinator/mod.rs @@ -4,6 +4,7 @@ mod latency_metric; mod prepare_dynamic_plan; mod prepare_static_plan; mod query_coordinator; +mod spawner; mod store; pub use distributed::DistributedExec; diff --git a/src/coordinator/prepare_dynamic_plan.rs b/src/coordinator/prepare_dynamic_plan.rs index 3416458f8..504661faf 100644 --- a/src/coordinator/prepare_dynamic_plan.rs +++ b/src/coordinator/prepare_dynamic_plan.rs @@ -106,16 +106,18 @@ pub(super) async fn prepare_dynamic_plan( let results = futures::future::try_join_all(futures).await?; let mut workers = Vec::with_capacity(input_stage.tasks); - for (task_i, (url, worker_tx, worker_rx)) in results.into_iter().enumerate() { + for (task_i, (url, worker_tx, worker_rx_stream)) in results.into_iter().enumerate() + { workers.push(url); load_info_rxs.push({ - let rx = stage_coordinator.worker_to_coordinator_task(task_i, worker_rx); + let rx = + stage_coordinator.worker_to_coordinator_task(task_i, worker_rx_stream); UnboundedReceiverStream::new(rx) }); let _ = worker_tx.send(CoordinatorToWorkerMsg::KickOffSampling); stage_coordinator.coordinator_to_worker_task(task_i, worker_tx)?; } - stage_coordinator.seal_dynamic_filter_stage(); + 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 0d7b33bf0..c0a5e193f 100644 --- a/src/coordinator/prepare_static_plan.rs +++ b/src/coordinator/prepare_static_plan.rs @@ -49,12 +49,12 @@ pub(super) async fn prepare_static_plan( let results = futures::future::try_join_all(futures).await?; let mut workers = Vec::with_capacity(stage.tasks); - for (task_i, (url, worker_tx, worker_rx)) in results.into_iter().enumerate() { + for (task_i, (url, worker_tx, worker_stream)) in results.into_iter().enumerate() { workers.push(url); - stage_coordinator.worker_to_coordinator_task(task_i, worker_rx); + stage_coordinator.worker_to_coordinator_task(task_i, worker_stream); stage_coordinator.coordinator_to_worker_task(task_i, worker_tx)?; } - stage_coordinator.seal_dynamic_filter_stage(); + 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 afc576d73..3ec35ebff 100644 --- a/src/coordinator/query_coordinator.rs +++ b/src/coordinator/query_coordinator.rs @@ -4,6 +4,7 @@ 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::coordinator::spawner::Spawner; use crate::dynamic_filtering::{ dynamic_filter_remote_producer_ids, is_local_dynamic_filtering_enabled, is_remote_dynamic_filtering_enabled, @@ -20,25 +21,26 @@ use crate::work_unit_feed::{build_work_unit_batch_msg, set_work_unit_send_time}; use crate::{ CoordinatorToWorkerMsg, DISTRIBUTED_DATAFUSION_TASK_ID_LABEL, DistributedGetterExt, DistributedTaskContext, DistributedWorkUnitFeedContext, LoadInfo, LocalWorkerContext, - MaybeEncoded, SetPlanRequest, TaskCompletedDynamicFilters, TaskKey, TaskMetrics, - WorkUnitFeedDeclaration, WorkerToCoordinatorMsg, get_distributed_channel_resolver, + MaybeEncoded, ProducedDynamicFilter, SetPlanRequest, TaskCompletedDynamicFilters, TaskKey, + TaskMetrics, WorkUnitFeedDeclaration, WorkerToCoordinatorMsg, get_distributed_channel_resolver, }; use datafusion::common::Result; use datafusion::common::instant::Instant; -use datafusion::common::runtime::JoinSet; use datafusion::common::tree_node::{Transformed, TreeNodeRecursion}; use datafusion::common::{DataFusionError, internal_err}; use datafusion::execution::TaskContext; use datafusion::physical_expr_common::metrics::{ExecutionPlanMetricsSet, Label, MetricBuilder}; use datafusion::physical_plan::metrics::Count; use datafusion::physical_plan::repartition::RepartitionExec; +use datafusion::physical_plan::stream::RecordBatchReceiverStreamBuilder; use datafusion::physical_plan::{ChildrenPropertiesMode, ExecutionPlan, ReplaceChildrenOptions}; use datafusion::prelude::SessionConfig; -use futures::{Stream, StreamExt, TryStreamExt}; -use std::ops::DerefMut; +use futures::future::try_join_all; +use futures::stream::BoxStream; +use futures::{StreamExt, TryStreamExt}; use std::sync::{Arc, Mutex}; -use tokio::sync::Notify; -use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender}; +use tokio::select; +use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel}; use tokio_stream::wrappers::UnboundedReceiverStream; use url::Url; use uuid::Uuid; @@ -59,8 +61,7 @@ pub(super) struct QueryCoordinator { metrics_store: Option>>, completed_dynamic_filter_store: Option>>, dynamic_filter_registry: Arc, - end_stream_notifier: Arc, - join_set: Mutex>>, + pub(super) spawner: Spawner, } impl QueryCoordinator { @@ -70,16 +71,20 @@ impl QueryCoordinator { metrics_set: &ExecutionPlanMetricsSet, metrics_store: Option>>, completed_dynamic_filter_store: Option>>, + output: &mut RecordBatchReceiverStreamBuilder, ) -> Self { + let spawner = Spawner::new(output); Self { task_ctx, metrics: metrics_set.clone(), metrics_store, completed_dynamic_filter_store, - dynamic_filter_registry: Arc::new(DynamicFilterRegistry::new(metrics_set)), + dynamic_filter_registry: Arc::new(DynamicFilterRegistry::new( + metrics_set, + spawner.query_finished(), + )), coordinator_to_worker_metrics: CoordinatorToWorkerMetrics::new(metrics_set), - end_stream_notifier: Arc::new(Notify::new()), - join_set: Mutex::new(JoinSet::new()), + spawner, } } @@ -97,8 +102,7 @@ impl QueryCoordinator { 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, + spawner: &self.spawner, } } @@ -106,23 +110,6 @@ impl QueryCoordinator { pub(super) fn session_config(&self) -> &SessionConfig { self.task_ctx.session_config() } - - /// returns a guard that, when dropped, it signals all the coordinator->worker connections that - /// the query is finished, ending them, and propagating the EOS to the workers so that they can - /// clean up any remaining state. - pub(super) fn end_query_guard(&self) -> NotifyGuard { - NotifyGuard(Arc::clone(&self.end_stream_notifier)) - } - - /// Blocks until all background tasks have finished (e.g., sending WorkUnit feeds, or collecting - /// metrics) - pub(super) async fn drain_pending_tasks(self: Arc) -> Result<()> { - let join_set = std::mem::take(self.join_set.lock().unwrap().deref_mut()); - for res in join_set.join_all().await { - res?; - } - Ok(()) - } } /// Manages all the coordinator->worker and worker->coordinator comms that happen during the @@ -144,8 +131,7 @@ pub(super) struct StageCoordinator<'a> { metrics_store: &'a Option>>, completed_dynamic_filter_store: &'a Option>>, dynamic_filter_registry: &'a Arc, - end_stream_notifier: &'a Arc, - join_set: &'a Mutex>>, + spawner: &'a Spawner, } impl<'a> StageCoordinator<'a> { @@ -161,7 +147,7 @@ impl<'a> StageCoordinator<'a> { ) -> Result<( Url, UnboundedSender, - UnboundedReceiver, + BoxStream<'static, Result>, )> { let session_config = self.task_ctx.session_config(); @@ -193,8 +179,7 @@ impl<'a> StageCoordinator<'a> { let coordinator_to_worker_tx_slot = Mutex::new(None); let dialer = new_coordinator_to_worker_dialer(|url| { - let (coordinator_to_worker_tx, coordinator_to_worker_rx) = - tokio::sync::mpsc::unbounded_channel(); + let (coordinator_to_worker_tx, coordinator_to_worker_rx) = unbounded_channel(); coordinator_to_worker_tx_slot .lock() .unwrap() @@ -203,17 +188,15 @@ impl<'a> StageCoordinator<'a> { let coordinator_to_worker_stream = UnboundedReceiverStream::new(coordinator_to_worker_rx) .map(set_work_unit_send_time) - // Keep the request side of the channel open until the query ends: this tail emits - // no messages and only completes, once the `Notify` fires. Workers interpret this - // EOS of this stream as a query finished/aborted signal. The flow looks like this: - // 1. The query ends normally, as all Arrow RecordBatches are already streamed. + // Keep the request side of the channel open until the query ends. Workers interpret this + // end-of-stream (EOS) as a query finished/aborted signal. The flow looks like this: + // 1. Execution ends normally. All Arrow RecordBatches are emitted. // 2. The end stream notifier guard is dropped in `DistributedExec::execute()`. - // 3. Here, `end_stream_notifier` fires and the coordinator->worker channel is - // gracefully ended. + // 3. Here, the coordinator->worker channel is gracefully ended. // 4. The coordinator->worker channel EOS is received in `impl_coordinator_channel.rs`. // 5. The metrics and final dynamic filters are sent back in the // worker->coordinator channel, and then that channel is closed. - .chain(keep_stream_alive(Arc::clone(self.end_stream_notifier))) + .take_until(self.spawner.query_finished().cancelled_owned()) .boxed(); let set_plan_request = SetPlanRequest { @@ -277,42 +260,29 @@ impl<'a> StageCoordinator<'a> { }; metrics.plan_send_latency.record(&start); - let (worker_to_coordinator_tx, worker_to_coordinator_rx) = - tokio::sync::mpsc::unbounded_channel(); - - let mut worker_to_coordinator_stream = response.worker_to_coordinator_stream; - self.join_set.lock().unwrap().spawn(async move { - while let Some(msg) = worker_to_coordinator_stream.try_next().await? { - if worker_to_coordinator_tx.send(msg).is_err() { - break; // receiver dropped - } - } - Ok::<_, DataFusionError>(()) - }); - let Some(coordinator_to_worker_tx) = coordinator_to_worker_tx_slot.lock().unwrap().take() else { return internal_err!("Missing coordinator_to_worker_tx"); }; + self.dynamic_filter_registry + .register_sender(task_key, coordinator_to_worker_tx.clone())?; Ok(( response.url, coordinator_to_worker_tx, - worker_to_coordinator_rx, + response.worker_to_coordinator_stream, )) } - pub(super) fn seal_dynamic_filter_stage(&self) { - self.dynamic_filter_registry.seal_stage(self.stage_id); + pub(super) fn seal_dynamic_filter_stage(&self) -> Result<()> { + 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. + /// Spawns background tasks in charge of collecting messages sent by a worker. pub(super) fn worker_to_coordinator_task( &mut self, task_i: usize, - mut worker_to_coordinator_rx: UnboundedReceiver, + mut worker_to_coordinator_stream: BoxStream<'static, Result>, ) -> UnboundedReceiver { let task_key = TaskKey { query_id: self.query_id, @@ -323,37 +293,73 @@ impl<'a> StageCoordinator<'a> { let mut 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 (load_info_tx, load_info_rx) = unbounded_channel(); + let (reports_tx, mut reports_rx) = unbounded_channel::(); + let (runtime_tx, mut runtime_rx) = unbounded_channel::(); let mut load_info_tx_opt = Some(load_info_tx); - // Cannot use self.join_set because that's tied to the lifetime of the query, and the - // metrics collection process might outlive the query's lifetime. - #[allow(clippy::disallowed_methods)] - tokio::spawn(async move { - while let Some(msg) = worker_to_coordinator_rx.recv().await { + // Read from the worker -> coordinator stream and fail the query on error. + self.spawner.spawn(async move { + while let Some(msg) = worker_to_coordinator_stream.try_next().await? { + let receiver_closed = match msg { + WorkerToCoordinatorMsg::TaskMetrics(metrics) => { + reports_tx.send(FinalReport::Metrics(metrics)).is_err() + } + WorkerToCoordinatorMsg::TaskCompletedDynamicFilters(filters) => reports_tx + .send(FinalReport::DynamicFilters(filters)) + .is_err(), + WorkerToCoordinatorMsg::ProducedDynamicFilter(filter) => runtime_tx + .send(RuntimeUpdate::DynamicFilter(filter)) + .is_err(), + WorkerToCoordinatorMsg::LoadInfo(info) => { + runtime_tx.send(RuntimeUpdate::LoadInfo(info)).is_err() + } + WorkerToCoordinatorMsg::LoadInfoEos => { + runtime_tx.send(RuntimeUpdate::LoadInfoEos).is_err() + } + }; + if receiver_closed { + break; + } + } + Ok(()) + }); + + // Handle runtime updates and fail the query on error. + self.spawner.spawn(async move { + while let Some(msg) = runtime_rx.recv().await { match msg { - WorkerToCoordinatorMsg::TaskMetrics(v) => { - if let Some(store) = task_metrics.take() { - store.insert(task_key, v); - } + RuntimeUpdate::DynamicFilter(filter) => { + dynamic_filter_registry + .record_dynamic_filter_update(task_key, *filter, &task_ctx)?; } - WorkerToCoordinatorMsg::LoadInfo(load_info) => { + RuntimeUpdate::LoadInfo(load_info) => { if let Some(tx) = &load_info_tx_opt { let _ = tx.send(load_info); } } - WorkerToCoordinatorMsg::LoadInfoEos => { + RuntimeUpdate::LoadInfoEos => { let _ = load_info_tx_opt.take(); } - WorkerToCoordinatorMsg::TaskCompletedDynamicFilters(filters) => { + } + } + Ok(()) + }); + + // Handle reports regardless of when the query is finished. + self.spawner.spawn_unbounded(async move { + while let Some(msg) = reports_rx.recv().await { + match msg { + FinalReport::Metrics(v) => { + if let Some(task_metrics) = task_metrics.take() { + task_metrics.insert(task_key, v); + } + } + FinalReport::DynamicFilters(filters) => { if let Some(store) = completed_dynamic_filter_store.take() { store.insert(task_key, filters); } } - WorkerToCoordinatorMsg::ProducedDynamicFilter(filter) => { - dynamic_filter_registry - .record_dynamic_filter_update(task_key, *filter, &task_ctx); - } } } // An unexecuted task sends no final reports; still complete its waits. @@ -364,6 +370,7 @@ impl<'a> StageCoordinator<'a> { store.insert(task_key, TaskCompletedDynamicFilters::default()); } }); + load_info_rx } @@ -432,9 +439,14 @@ impl<'a> StageCoordinator<'a> { } } - self.join_set.lock().unwrap().spawn(async move { + // Cancel any feeds if the query finishes or aborts. + let query_finished = self.spawner.query_finished(); + self.spawner.spawn(async move { let _guard = WorkUnitEosOnDrop(tx); - futures::future::try_join_all(futures).await?; + select! { + result = try_join_all(futures) => { result?; } + _ = query_finished.cancelled() => {} + } Ok(()) }); Ok(()) @@ -522,8 +534,17 @@ impl<'a> StageCoordinator<'a> { } } -fn keep_stream_alive(notify: Arc) -> impl Stream + 'static { - futures::stream::once(notify.notified_owned()).filter_map(|()| futures::future::ready(None)) +/// [`WorkerToCoordinatorMsg`] which may arrive after the query has finished executing. +enum FinalReport { + Metrics(TaskMetrics), + DynamicFilters(TaskCompletedDynamicFilters), +} + +/// [`WorkerToCoordinatorMsg`] which occurs during execution. +enum RuntimeUpdate { + DynamicFilter(Box), + LoadInfo(LoadInfo), + LoadInfoEos, } struct TaskSpecializedPlan { @@ -532,14 +553,6 @@ struct TaskSpecializedPlan { dynamic_filter_remote_producer_ids: Vec, } -pub(super) struct NotifyGuard(Arc); - -impl Drop for NotifyGuard { - fn drop(&mut self) { - self.0.notify_waiters(); - } -} - /// Metrics that measure network details about communications between [DistributedExec] and a worker. #[derive(Clone)] pub(super) struct CoordinatorToWorkerMetrics { diff --git a/src/coordinator/spawner.rs b/src/coordinator/spawner.rs new file mode 100644 index 000000000..15e9cd58f --- /dev/null +++ b/src/coordinator/spawner.rs @@ -0,0 +1,83 @@ +use datafusion::common::runtime::JoinSet; +use datafusion::common::{Result, exec_err}; +use datafusion::physical_plan::stream::RecordBatchReceiverStreamBuilder; +use futures::FutureExt; +use futures::future::BoxFuture; +use std::panic::resume_unwind; +use tokio::select; +use tokio::sync::mpsc::{UnboundedSender, unbounded_channel}; +use tokio_util::sync::{CancellationToken, DropGuard}; + +/// Owns query tokio task spawning and error propagation. +pub(super) struct Spawner { + // Cancelled when the main execution future completes or is dropped. + query_finished: CancellationToken, + task_tx: UnboundedSender>>, +} + +impl Spawner { + pub(super) fn new(output: &mut RecordBatchReceiverStreamBuilder) -> Self { + let query_finished = CancellationToken::new(); + let (task_tx, mut task_rx) = unbounded_channel::>>(); + + // Supervisor task responsible for propagating errors from tasks to the main + // query response stream. Terminates when all tasks finished, otherwise + // blocks the response stream from finishing. + output.spawn(async move { + let mut tasks = JoinSet::new(); + loop { + select! { + Some(result) = tasks.join_next() => { + match result { + Ok(result) => result?, + Err(error) if error.is_panic() => { + resume_unwind(error.into_panic()); + } + Err(error) => { + return exec_err!("non panic JoinSet Error: {error}"); + } + } + } + Some(task) = task_rx.recv() => { + tasks.spawn(task); + } + else => break, + } + } + Ok(()) + }); + + Self { + query_finished, + task_tx, + } + } + + /// Returns the completion signal that tasks may use to abort. + pub(super) fn query_finished(&self) -> CancellationToken { + self.query_finished.clone() + } + + /// Returns a guard that signals query completion when dropped. + pub(super) fn end_query_guard(&self) -> DropGuard { + self.query_finished.clone().drop_guard() + } + + /// Registers a query-scoped task whose errors propagate to the output stream. + /// + /// Tasks may choose to abort themselves when the query stream finishes using + /// [`Spawner::end_query_guard`] or [`Spawner::query_finished`]. + /// + /// Otherwise, they may choose to terminate gracefully, blocking the response stream. + pub(super) fn spawn(&self, task: impl Future> + Send + 'static) { + // Once the supervisor stops, sending fails and drops the future without running it. + let _ = self.task_tx.send(task.boxed()); + } + + /// Spawns a task whose lifetime is not tied to the response stream and whose + /// errors do not propagate to the reposne stream. + pub(super) fn spawn_unbounded(&self, task: impl Future + Send + 'static) { + #[allow(clippy::disallowed_methods)] + tokio::spawn(task); + } +} diff --git a/src/lib.rs b/src/lib.rs index 523252982..7ca505d3f 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -56,8 +56,8 @@ pub use common::MaybeEncoded; pub use worker_resolver::{WorkerResolver, get_distributed_worker_resolver}; pub use protocol::{ - ChannelResolver, CoordinatorToWorkerMsg, ExecuteTaskRequest, GetWorkerInfoRequest, - GetWorkerInfoResponse, LoadInfo, ProducedDynamicFilter, SetPlanRequest, + ApplyDynamicFilter, ChannelResolver, CoordinatorToWorkerMsg, ExecuteTaskRequest, + GetWorkerInfoRequest, GetWorkerInfoResponse, LoadInfo, ProducedDynamicFilter, SetPlanRequest, TaskCompletedDynamicFilters, TaskDynamicFilter, TaskKey, TaskMetrics, WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg, WorkerChannel, WorkerToCoordinatorMsg, get_distributed_channel_resolver, diff --git a/src/protocol/grpc/generated/worker.rs b/src/protocol/grpc/generated/worker.rs index 268cf9537..91bcc9d2c 100644 --- a/src/protocol/grpc/generated/worker.rs +++ b/src/protocol/grpc/generated/worker.rs @@ -1,7 +1,7 @@ // This file is @generated by prost-build. #[derive(Clone, PartialEq, ::prost::Message)] pub struct CoordinatorToWorkerMsg { - #[prost(oneof = "coordinator_to_worker_msg::Inner", tags = "1, 2, 3, 4")] + #[prost(oneof = "coordinator_to_worker_msg::Inner", tags = "1, 2, 3, 4, 5")] pub inner: ::core::option::Option, } /// Nested message and enum types in `CoordinatorToWorkerMsg`. @@ -23,10 +23,21 @@ pub mod coordinator_to_worker_msg { /// Signals the worker to begin sampling during adaptive query execution. #[prost(message, tag = "4")] KickOffSampling(super::KickOffSampling), + /// Sends a coordinator-merged dynamic-filter update to one consumer task. + #[prost(message, tag = "5")] + ApplyDynamicFilter(super::ApplyDynamicFilter), } } #[derive(Clone, Copy, PartialEq, Eq, Hash, ::prost::Message)] pub struct KickOffSampling {} +#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] +pub struct ApplyDynamicFilter { + #[prost(uint64, tag = "1")] + pub expression_id: u64, + /// Serialized datafusion.proto.PhysicalExprNode containing a full DynamicFilter expression. + #[prost(bytes = "vec", tag = "2")] + pub expression_proto: ::prost::alloc::vec::Vec, +} #[derive(Clone, PartialEq, ::prost::Message)] pub struct WorkerToCoordinatorMsg { #[prost(oneof = "worker_to_coordinator_msg::Inner", tags = "1, 2, 3, 4, 5")] diff --git a/src/protocol/grpc/worker.proto b/src/protocol/grpc/worker.proto index af4e51798..74ccab378 100644 --- a/src/protocol/grpc/worker.proto +++ b/src/protocol/grpc/worker.proto @@ -25,6 +25,8 @@ message CoordinatorToWorkerMsg { bool work_unit_eos = 3; // Signals the worker to begin sampling during adaptive query execution. KickOffSampling kick_off_sampling = 4; + // Sends a coordinator-merged dynamic-filter update to one consumer task. + ApplyDynamicFilter apply_dynamic_filter = 5; } } @@ -32,6 +34,12 @@ message KickOffSampling { } +message ApplyDynamicFilter { + uint64 expression_id = 1; + // Serialized datafusion.proto.PhysicalExprNode containing a full DynamicFilter expression. + bytes expression_proto = 2; +} + message WorkerToCoordinatorMsg { oneof inner { // Sends the metrics collected during task execution back to the coordinator. diff --git a/src/protocol/grpc/worker_client.rs b/src/protocol/grpc/worker_client.rs index af1fa272c..ea7355db7 100644 --- a/src/protocol/grpc/worker_client.rs +++ b/src/protocol/grpc/worker_client.rs @@ -26,7 +26,8 @@ use datafusion::execution::TaskContext; use datafusion::execution::memory_pool::MemoryConsumer; use datafusion::physical_expr_common::metrics::{Count, Label, MetricBuilder, MetricValue, Time}; use datafusion::physical_plan::metrics::{ExecutionPlanMetricsSet, Gauge}; -use futures::stream::BoxStream; +use futures::future::{Either, pending, ready}; +use futures::stream::{BoxStream, select}; use futures::{FutureExt, Stream, StreamExt, TryStreamExt}; use http::{Extensions, HeaderMap}; use pin_project::{pin_project, pinned_drop}; @@ -38,8 +39,8 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use std::task::{Context, Poll}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use tokio::sync::Notify; -use tokio::sync::mpsc::UnboundedSender; -use tokio_stream::wrappers::UnboundedReceiverStream; +use tokio::sync::mpsc::{UnboundedSender, channel}; +use tokio_stream::wrappers::{ReceiverStream, UnboundedReceiverStream}; use tokio_util::sync::CancellationToken; use tonic::metadata::MetadataMap; use tonic::{Code, Request, Status}; @@ -56,6 +57,10 @@ impl WorkerChannel for pb::worker_service_client::WorkerServiceClient Result>> { let set_plan_request = encode_set_plan_request(set_plan_request, ctx)?; let plan_bytes_sent = set_plan_request.plan_proto.len(); + let task_ctx = Arc::clone(ctx); + // Tonic request streams cannot yield errors, so return encoding failures through + // the response stream instead. Only the first failure matters. + let (error_tx, mut error_rx) = channel(1); let input_stream = futures::stream::once(async move { pb::CoordinatorToWorkerMsg { inner: Some(pb::coordinator_to_worker_msg::Inner::SetPlanRequest( @@ -63,16 +68,28 @@ impl WorkerChannel for pb::worker_service_client::WorkerServiceClient Either::Left(ready(msg)), + Err(error) => { + let _ = error_tx.try_send(error); + // Retain the receiver until the error is observed. + Either::Right(pending()) + } + } + })); - let output_stream = self - .coordinator_channel(Request::from_parts( + let response = tokio::select! { + biased; + Some(error) = error_rx.recv() => return Err(error), + response = self.coordinator_channel(Request::from_parts( MetadataMap::from_headers(headers), Extensions::default(), input_stream, )) - .boxed() - .await + .boxed() => response, + }; + let output_stream = response .map_err(|err| { if let Some(err) = tonic_status_to_datafusion_error(&err) { return err; @@ -118,7 +135,9 @@ impl WorkerChannel for pb::worker_service_client::WorkerServiceClient pb::CoordinatorToWorkerMsg { - pb::CoordinatorToWorkerMsg { +fn encode_coordinator_to_worker_msg( + msg: CoordinatorToWorkerMsg, + task_ctx: &Arc, +) -> Result { + Ok(pb::CoordinatorToWorkerMsg { inner: Some(match msg { CoordinatorToWorkerMsg::KickOffSampling => { pb::coordinator_to_worker_msg::Inner::KickOffSampling(pb::KickOffSampling {}) @@ -472,8 +494,14 @@ fn encode_coordinator_to_worker_msg(msg: CoordinatorToWorkerMsg) -> pb::Coordina CoordinatorToWorkerMsg::WorkUnitEos => { pb::coordinator_to_worker_msg::Inner::WorkUnitEos(true) } + CoordinatorToWorkerMsg::ApplyDynamicFilter(filter) => { + pb::coordinator_to_worker_msg::Inner::ApplyDynamicFilter(pb::ApplyDynamicFilter { + expression_id: filter.expression_id, + expression_proto: filter.expression.encode(task_ctx)?, + }) + } }), - } + }) } fn encode_set_plan_request( diff --git a/src/protocol/grpc/worker_service.rs b/src/protocol/grpc/worker_service.rs index c96e62a8f..7a71db92a 100644 --- a/src/protocol/grpc/worker_service.rs +++ b/src/protocol/grpc/worker_service.rs @@ -6,10 +6,10 @@ use super::spawn_select_all::spawn_select_all; use crate::common::{deserialize_uuid, now_ns}; use crate::protocol::grpc::{ObservabilityServiceImpl, ObservabilityServiceServer}; use crate::{ - CoordinatorToWorkerMsg, DistributedConfig, ExecuteTaskRequest, LoadInfo, MaybeEncoded, - ProducedDynamicFilter, ProducerHead, SetPlanRequest, TaskCompletedDynamicFilters, TaskKey, - TaskMetrics, WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg, Worker, WorkerResolver, - WorkerToCoordinatorMsg, + ApplyDynamicFilter, CoordinatorToWorkerMsg, DistributedConfig, ExecuteTaskRequest, LoadInfo, + MaybeEncoded, ProducedDynamicFilter, ProducerHead, SetPlanRequest, TaskCompletedDynamicFilters, + TaskKey, TaskMetrics, WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg, Worker, + WorkerResolver, WorkerToCoordinatorMsg, }; use crate::worker::CoordinatorChannelResult; @@ -219,6 +219,12 @@ fn decode_coordinator_to_worker_msg( pb::coordinator_to_worker_msg::Inner::KickOffSampling(_) => { CoordinatorToWorkerMsg::KickOffSampling } + pb::coordinator_to_worker_msg::Inner::ApplyDynamicFilter(filter) => { + CoordinatorToWorkerMsg::ApplyDynamicFilter(Box::new(ApplyDynamicFilter { + expression_id: filter.expression_id, + expression: MaybeEncoded::Encoded(filter.expression_proto), + })) + } }, ) } diff --git a/src/protocol/mod.rs b/src/protocol/mod.rs index f3fb839b4..93e2467e3 100644 --- a/src/protocol/mod.rs +++ b/src/protocol/mod.rs @@ -10,8 +10,8 @@ pub(crate) use channel_resolver::set_distributed_channel_resolver; pub use channel_resolver::{ChannelResolver, get_distributed_channel_resolver}; pub use in_process::LocalWorkerContext; pub use worker_channel::{ - CoordinatorToWorkerMsg, ExecuteTaskRequest, GetWorkerInfoRequest, GetWorkerInfoResponse, - LoadInfo, ProducedDynamicFilter, SetPlanRequest, TaskCompletedDynamicFilters, - TaskDynamicFilter, TaskKey, TaskMetrics, WorkUnitBatch, WorkUnitFeedDeclaration, WorkUnitMsg, - WorkerChannel, WorkerToCoordinatorMsg, + ApplyDynamicFilter, CoordinatorToWorkerMsg, ExecuteTaskRequest, GetWorkerInfoRequest, + GetWorkerInfoResponse, 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 78323b45c..e548f69a4 100644 --- a/src/protocol/worker_channel.rs +++ b/src/protocol/worker_channel.rs @@ -55,6 +55,15 @@ pub enum CoordinatorToWorkerMsg { WorkUnitBatch(WorkUnitBatch), /// Signals an EOS for WorkUnits. After this message is received, no more WorkUnits will be sent. WorkUnitEos, + /// Sends a dynamic-filter update to one consumer task which should be applied locally. + ApplyDynamicFilter(Box), +} + +#[derive(Clone, Debug)] +pub struct ApplyDynamicFilter { + pub expression_id: u64, + /// A full dynamic filter containing the merged predicate. + pub expression: MaybeEncoded>, } #[derive(Clone, Copy, Debug, Hash, PartialEq, Eq)] diff --git a/src/worker/impl_coordinator_channel.rs b/src/worker/impl_coordinator_channel.rs index bef27f3b9..e1dcd1305 100644 --- a/src/worker/impl_coordinator_channel.rs +++ b/src/worker/impl_coordinator_channel.rs @@ -153,8 +153,8 @@ impl Worker { // is the following: // 1. The query ends normally, as all Arrow RecordBatches are already streamed. // 2. In DistributedExec::execute(), the end query guard is dropped. - // 3. In StageCoordinator::send_plan_task(), `end_stream_notifier` fires and the - // coordinator->worker channel is gracefully ended. + // 3. The query's cancellation token ends the coordinator->worker stream, regardless + // of any senders retained by the dynamic filter registry. // 4. The coordinator->worker channel EOS is received by this same function, ending the // while loop inside this `tokio::spawn` below. // 5. The metrics and final dynamic filters are sent back in the worker->coordinator @@ -195,6 +195,10 @@ impl Worker { CoordinatorToWorkerMsg::KickOffSampling => { sampler_gate.kick_off(); } + CoordinatorToWorkerMsg::ApplyDynamicFilter(_) => { + // Runtime application is introduced independently from the routing + // protocol. Until then, accepting the message is intentionally a no-op. + } } } diff --git a/tests/error_propagation.rs b/tests/error_propagation.rs index 355a6ace5..8d5d5c39e 100644 --- a/tests/error_propagation.rs +++ b/tests/error_propagation.rs @@ -1,6 +1,8 @@ #[cfg(all(feature = "integration", test))] mod tests { + use async_trait::async_trait; use datafusion::common::tree_node::{Transformed, TreeNode, TreeNodeRecursion}; + use datafusion::common::{assert_contains, exec_err}; use datafusion::error::DataFusionError; use datafusion::execution::{SendableRecordBatchStream, SessionState, TaskContext}; use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr}; @@ -12,16 +14,23 @@ mod tests { }; use datafusion_distributed::test_utils::localhost::start_localhost_context; use datafusion_distributed::test_utils::parquet::register_parquet_tables; - use datafusion_distributed::{DistributedExt, WorkerQueryContext}; + use datafusion_distributed::test_utils::routing::UrlEmitterRouteTaskHandler; + use datafusion_distributed::{ + DefaultSessionBuilder, DistributedExt, RouteTaskEvent, RouteTaskEventResponse, + RouteTaskHandler, WorkerQueryContext, ok_or_some_err, + }; use datafusion_proto::physical_plan::{ PhysicalExtensionCodec, PhysicalProtoConverterExtension, }; use datafusion_proto::protobuf::proto_error; - use futures::{TryStreamExt, stream}; + use futures::{StreamExt, TryStreamExt, stream}; use prost::Message; use std::error::Error; use std::fmt::Formatter; use std::sync::Arc; + use std::time::Duration; + use test_case::test_case; + use tokio::time::timeout; #[tokio::test] async fn test_error_propagation() -> Result<(), Box> { @@ -66,6 +75,42 @@ mod tests { Ok(()) } + #[test_case(false; "static_planner")] + #[test_case(true; "adaptive_planner")] + #[tokio::test] + async fn worker_channel_error_fails_query(adaptive: bool) -> Result<(), Box> { + let (ctx, _guard, _) = start_localhost_context(2, DefaultSessionBuilder).await; + let ctx = ctx + .with_distributed_dynamic_task_count(adaptive)? + .with_distributed_route_task_handler(FailingWorkerChannel); + register_parquet_tables(&ctx).await?; + let query = ctx + .sql(r#"SELECT "MinTemp" FROM weather WHERE "MinTemp" > 20.0"#) + .await?; + let error = timeout(Duration::from_secs(5), query.collect()) + .await? + .expect_err("worker channel error must fail the query"); + assert_contains!(error.to_string(), "injected worker channel error"); + Ok(()) + } + + struct FailingWorkerChannel; + + #[async_trait] + impl RouteTaskHandler for FailingWorkerChannel { + async fn handle( + &self, + event: RouteTaskEvent<'_>, + ) -> Option> { + let mut response = ok_or_some_err!(UrlEmitterRouteTaskHandler.handle(event).await?); + response.worker_to_coordinator_stream = + stream::once(async { exec_err!("injected worker channel error") }) + .chain(response.worker_to_coordinator_stream) + .boxed(); + Some(Ok(response)) + } + } + /// A custom execution plan that wraps a child but always throws an error. /// This tests that errors are properly propagated in distributed execution. #[derive(Debug)] diff --git a/tests/work_unit_feed.rs b/tests/work_unit_feed.rs index 697e9b4a8..f7980791c 100644 --- a/tests/work_unit_feed.rs +++ b/tests/work_unit_feed.rs @@ -15,6 +15,7 @@ mod tests { use futures::TryStreamExt; use std::sync::Arc; use std::time::{Duration, Instant}; + use tokio::time::timeout; #[tokio::test] async fn single_task_no_distribution() -> Result<(), Box> { @@ -714,15 +715,15 @@ mod tests { /// Same as [`err_op_in_single_task_propagates`] but with two tasks, so the /// erroring feed actually goes through the coordinator → worker gRPC path. - /// Guards against errors being silently swallowed as EOF on the worker side. #[tokio::test] async fn err_op_in_distributed_feed_propagates() -> Result<(), Box> { - let res = run_query( - r#" - SELECT * FROM test_work_unit('a', 2, 'rows(1)', 'rows(1), err(boom_distributed)') - "#, - ) - .await; + let query = r#" + SELECT * FROM test_work_unit('a', 2, 'wait(30000), rows(1)', 'rows(1), err(boom_distributed)') + "#; + // Expect the query to fail fast rather than waitinng for the unrelated 30s feed. + let res = timeout(Duration::from_secs(5), run_query(query)) + .await + .expect("feed error waited for an unrelated pending feed"); let err = res.expect_err("distributed query should have failed"); let msg = err.to_string(); assert!(