Skip to content
Open
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
30 changes: 14 additions & 16 deletions src/coordinator/distributed.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -265,13 +266,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(())
});

Expand Down
130 changes: 114 additions & 16 deletions src/coordinator/dynamic_filter_registry.rs
Original file line number Diff line number Diff line change
@@ -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};
Expand All @@ -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 {
Expand Down Expand Up @@ -42,13 +47,17 @@ pub(super) struct PlannedDynamicFilter {
pub(super) consumer_tasks: HashSet<TaskKey>,
/// Full dynamic filter containing the merged predicate and its completion state.
pub(super) merged: Option<PhysicalDynamicFilterNode>,
/// Immutable snapshot shared by local and remote delivery, encoded once per change.
merged_bytes: Option<Vec<u8>>,
}

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

/// Query-scoped hub for distributed dynamic filtering.
Expand All @@ -60,14 +69,19 @@ pub(super) struct DynamicFilterRegistryState {
pub(crate) struct DynamicFilterRegistry {
pub(super) state: Mutex<DynamicFilterRegistryState>,
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,
}
}

Expand Down Expand Up @@ -150,14 +164,32 @@ impl DynamicFilterRegistry {
Ok(())
}

pub(crate) fn register_sender(
&self,
task_key: TaskKey,
sender: UnboundedSender<CoordinatorToWorkerMsg>,
) -> 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::<Vec<_>>();
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::<Vec<_>>();
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.
Expand All @@ -166,38 +198,55 @@ 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)
|| previous
.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
Expand Down Expand Up @@ -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.
Expand Down
1 change: 1 addition & 0 deletions src/coordinator/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
7 changes: 4 additions & 3 deletions src/coordinator/prepare_dynamic_plan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -99,16 +99,17 @@ 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_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_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::<NetworkCoalesceExec>() {
(None, Maximum(1))
Expand Down
6 changes: 3 additions & 3 deletions src/coordinator/prepare_static_plan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,12 +42,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,
Expand Down
Loading