From 0c7ce6c8bcadb0dd03c36fbe6100d8cd88316f00 Mon Sep 17 00:00:00 2001 From: Geoffrey Claude Date: Tue, 21 Jul 2026 16:44:10 +0200 Subject: [PATCH] worker: add a plan setup wrapper --- docs/source/user-guide/03-worker.md | 6 + src/test_utils/mod.rs | 1 + src/test_utils/worker_plan_setup.rs | 8 + src/worker/impl_coordinator_channel.rs | 34 ++- src/worker/session_builder.rs | 99 +++++++ tests/worker_plan_setup.rs | 377 +++++++++++++++++++++++++ 6 files changed, 515 insertions(+), 10 deletions(-) create mode 100644 src/test_utils/worker_plan_setup.rs create mode 100644 tests/worker_plan_setup.rs diff --git a/docs/source/user-guide/03-worker.md b/docs/source/user-guide/03-worker.md index 3084cda6..1dc18195 100644 --- a/docs/source/user-guide/03-worker.md +++ b/docs/source/user-guide/03-worker.md @@ -72,6 +72,12 @@ It receives a `WorkerQueryContext` with two fields: - `headers` — the HTTP headers from the incoming request, handy for metadata like authentication tokens or per-query configuration. +Implementations that need to trace or account for worker plan setup can also +override `WorkerSessionBuilder::run_plan_setup`. Its callback covers physical +plan decoding, worker plan hooks, and initial sampler startup. A wrapper that +succeeds must call the callback exactly once and return its result. Builders composed with +`MappedWorkerSessionBuilderExt::map` retain this behavior automatically. + ```{note} A worker only *executes* fragments — it never plans queries. So it needs your codecs (to decode any custom nodes) but **not** the distributed planner or diff --git a/src/test_utils/mod.rs b/src/test_utils/mod.rs index dc66f977..ef65d427 100644 --- a/src/test_utils/mod.rs +++ b/src/test_utils/mod.rs @@ -10,3 +10,4 @@ pub mod routing; pub mod session_context; pub mod test_work_unit_feed; pub mod work_unit_file_scan; +pub mod worker_plan_setup; diff --git a/src/test_utils/worker_plan_setup.rs b/src/test_utils/worker_plan_setup.rs new file mode 100644 index 00000000..6ea8031b --- /dev/null +++ b/src/test_utils/worker_plan_setup.rs @@ -0,0 +1,8 @@ +use crate::execution_plans::SamplerExec; +use datafusion::physical_plan::ExecutionPlan; +use std::sync::Arc; + +/// Wraps a plan in the internal sampler node for worker setup integration tests. +pub fn wrap_in_sampler(plan: Arc) -> Arc { + Arc::new(SamplerExec::new(plan)) +} diff --git a/src/worker/impl_coordinator_channel.rs b/src/worker/impl_coordinator_channel.rs index cf16e890..f52e0c1f 100644 --- a/src/worker/impl_coordinator_channel.rs +++ b/src/worker/impl_coordinator_channel.rs @@ -83,16 +83,30 @@ impl Worker { }) .await?; - let codec = DistributedCodec::new_combined_with_user(session_state.config()); - let task_ctx = session_state.task_ctx(); - let proto_node = PhysicalPlanNode::try_decode(request.plan_proto.as_ref())?; - let mut plan = proto_node.try_into_physical_plan(&task_ctx, &codec)?; - - for hook in self.hooks.on_plan.iter() { - plan = hook(plan, session_state.config())?; - } - load_info_rxs = - SamplerExec::kick_off_first_sampler(Arc::clone(&plan), Arc::clone(&task_ctx))?; + let mut plan_setup = None; + self.session_builder + .run_plan_setup(&session_state, &mut || { + let codec = DistributedCodec::new_combined_with_user(session_state.config()); + let task_ctx = session_state.task_ctx(); + let proto_node = PhysicalPlanNode::try_decode(request.plan_proto.as_ref())?; + let mut plan = proto_node.try_into_physical_plan(&task_ctx, &codec)?; + + for hook in self.hooks.on_plan.iter() { + plan = hook(plan, session_state.config())?; + } + let sampler_receivers = SamplerExec::kick_off_first_sampler( + Arc::clone(&plan), + Arc::clone(&task_ctx), + )?; + plan_setup = Some((plan, task_ctx, sampler_receivers)); + Ok(()) + })?; + let Some((plan, task_ctx, sampler_receivers)) = plan_setup else { + return internal_err!( + "WorkerSessionBuilder::run_plan_setup did not run the setup callback" + ); + }; + load_info_rxs = sampler_receivers; // Initialize partition count to the number of partitions in the stage Ok::<_, DataFusionError>(TaskData { diff --git a/src/worker/session_builder.rs b/src/worker/session_builder.rs index c0f0bf41..60e2ccb7 100644 --- a/src/worker/session_builder.rs +++ b/src/worker/session_builder.rs @@ -59,6 +59,20 @@ pub trait WorkerSessionBuilder { &self, ctx: WorkerQueryContext, ) -> Result; + + /// Runs the synchronous setup for a plan received by the worker. + /// + /// The callback decodes the physical plan, applies worker plan hooks, and starts any initial + /// sampling work. Implementations may override this method to wrap that work with tracing or + /// resource accounting. An implementation that returns `Ok(())` must invoke `setup` exactly + /// once and return its result; it may reject setup by returning an error without invoking it. + fn run_plan_setup( + &self, + _session_state: &SessionState, + setup: &mut dyn FnMut() -> Result<(), DataFusionError>, + ) -> Result<(), DataFusionError> { + setup() + } } /// Noop implementation of the [WorkerSessionBuilder]. Used by default if no [WorkerSessionBuilder] @@ -156,4 +170,89 @@ where let builder = SessionStateBuilder::new_from_existing(state); (self.f)(builder) } + + fn run_plan_setup( + &self, + session_state: &SessionState, + setup: &mut dyn FnMut() -> Result<(), DataFusionError>, + ) -> Result<(), DataFusionError> { + self.inner.run_plan_setup(session_state, setup) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use datafusion::common::{assert_contains, internal_err}; + use std::sync::atomic::{AtomicUsize, Ordering}; + + #[test] + fn default_plan_setup_runs_callback_once() { + let state = SessionStateBuilder::new().build(); + let mut callback_calls = 0; + + DefaultSessionBuilder + .run_plan_setup(&state, &mut || { + callback_calls += 1; + Ok(()) + }) + .unwrap(); + + assert_eq!(callback_calls, 1); + } + + #[test] + fn default_plan_setup_propagates_callback_error() { + let state = SessionStateBuilder::new().build(); + + let error = DefaultSessionBuilder + .run_plan_setup(&state, &mut || internal_err!("plan setup failed")) + .expect_err("the callback error should be returned"); + + assert_contains!(error.to_string(), "plan setup failed"); + } + + #[test] + fn mapped_plan_setup_delegates_to_inner_builder() { + let wrapper_calls = Arc::new(AtomicUsize::new(0)); + let builder = RecordingSessionBuilder { + wrapper_calls: Arc::clone(&wrapper_calls), + } + .map(|builder| Ok(builder.build())); + let state = SessionStateBuilder::new().build(); + let mut callback_calls = 0; + + builder + .run_plan_setup(&state, &mut || { + callback_calls += 1; + Ok(()) + }) + .unwrap(); + + assert_eq!(wrapper_calls.load(Ordering::Relaxed), 1); + assert_eq!(callback_calls, 1); + } + + struct RecordingSessionBuilder { + wrapper_calls: Arc, + } + + #[async_trait] + impl WorkerSessionBuilder for RecordingSessionBuilder { + async fn build_session_state( + &self, + ctx: WorkerQueryContext, + ) -> Result { + Ok(ctx.builder.build()) + } + + fn run_plan_setup( + &self, + _session_state: &SessionState, + setup: &mut dyn FnMut() -> Result<(), DataFusionError>, + ) -> Result<(), DataFusionError> { + self.wrapper_calls.fetch_add(1, Ordering::Relaxed); + setup() + } + } } diff --git a/tests/worker_plan_setup.rs b/tests/worker_plan_setup.rs new file mode 100644 index 00000000..bb5ad06f --- /dev/null +++ b/tests/worker_plan_setup.rs @@ -0,0 +1,377 @@ +#[cfg(all(feature = "integration", test))] +mod tests { + use arrow::array::Int32Array; + use arrow::datatypes::{DataType, Field, Schema}; + use arrow::record_batch::RecordBatch; + use async_trait::async_trait; + use datafusion::common::tree_node::{Transformed, TreeNode}; + use datafusion::common::{Result, assert_contains, internal_err}; + use datafusion::error::DataFusionError; + use datafusion::execution::{SendableRecordBatchStream, SessionState, TaskContext}; + use datafusion::physical_plan::{ + DisplayAs, DisplayFormatType, ExecutionPlan, PlanProperties, execute_stream, + }; + use datafusion_distributed::test_utils::in_memory_channel_resolver::start_configured_in_memory_context; + use datafusion_distributed::test_utils::session_context::register_temp_parquet_table; + use datafusion_distributed::test_utils::worker_plan_setup::wrap_in_sampler; + use datafusion_distributed::{ + DistributedExt, Worker, WorkerQueryContext, WorkerSessionBuilder, + }; + use datafusion_proto::physical_plan::PhysicalExtensionCodec; + use datafusion_proto::protobuf::proto_error; + use futures::TryStreamExt; + use prost::Message; + use std::cell::Cell; + use std::fmt::Formatter; + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + + thread_local! { + static PLAN_SETUP_ACTIVE: Cell = const { Cell::new(false) }; + } + + #[tokio::test] + async fn plan_setup_wraps_decode_hooks_and_sampler_kickoff() + -> Result<(), Box> { + let checks = Arc::new(PlanSetupChecks::default()); + let builder = ScopedSessionBuilder { + checks: Arc::clone(&checks), + }; + let mut ctx = start_configured_in_memory_context(3, builder, { + let checks = Arc::clone(&checks); + move |mut worker| { + add_sampler_check_hook(&mut worker, Arc::clone(&checks)); + add_scope_check_hook(&mut worker, Arc::clone(&checks)); + worker + } + }) + .await; + ctx.set_distributed_user_codec(ScopedPassThroughExecCodec { checks: None }); + + let _left_file = register_input_table("plan_setup_left", &ctx).await?; + let plan = ctx + .sql("SELECT id FROM plan_setup_left WHERE id > 1 ORDER BY id") + .await? + .create_physical_plan() + .await?; + let plan = plan + .transform_up(|plan| { + if plan.children().is_empty() { + return Ok(Transformed::yes(Arc::new(ScopedPassThroughExec::new(plan)))); + } + Ok(Transformed::no(plan)) + })? + .data; + + let batches = execute_stream(plan, ctx.task_ctx())? + .try_collect::>() + .await?; + + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 2); + assert!(checks.wrapper_calls.load(Ordering::Relaxed) > 0); + assert!(checks.decode_checks.load(Ordering::Relaxed) > 0); + assert_eq!( + checks.hook_checks.load(Ordering::Relaxed), + checks.wrapper_calls.load(Ordering::Relaxed) * 2, + ); + assert!(checks.sampler_kickoff_checks.load(Ordering::Relaxed) > 0); + + Ok(()) + } + + #[tokio::test] + async fn plan_setup_wrapper_errors_propagate_to_query() -> Result<(), Box> + { + let mut ctx = + start_configured_in_memory_context(3, RejectingSessionBuilder, |worker| worker).await; + ctx.set_distributed_user_codec(ScopedPassThroughExecCodec { checks: None }); + let _left_file = register_input_table("plan_setup_left", &ctx).await?; + let plan = ctx + .sql("SELECT id FROM plan_setup_left WHERE id > 1 ORDER BY id") + .await? + .create_physical_plan() + .await?; + let plan = plan + .transform_up(|plan| { + if plan.children().is_empty() { + return Ok(Transformed::yes(Arc::new(ScopedPassThroughExec::new(plan)))); + } + Ok(Transformed::no(plan)) + })? + .data; + + let error = execute_stream(plan, ctx.task_ctx())? + .try_collect::>() + .await + .expect_err("worker plan setup should fail"); + + assert_contains!(error.to_string(), "worker plan setup rejected"); + Ok(()) + } + + #[derive(Debug, Default)] + struct PlanSetupChecks { + wrapper_calls: AtomicUsize, + decode_checks: AtomicUsize, + hook_checks: AtomicUsize, + sampler_kickoff_checks: AtomicUsize, + } + + #[derive(Clone)] + struct ScopedSessionBuilder { + checks: Arc, + } + + #[async_trait] + impl WorkerSessionBuilder for ScopedSessionBuilder { + async fn build_session_state( + &self, + ctx: WorkerQueryContext, + ) -> Result { + Ok(ctx + .builder + .with_distributed_user_codec(ScopedPassThroughExecCodec { + checks: Some(Arc::clone(&self.checks)), + }) + .build()) + } + + fn run_plan_setup( + &self, + _session_state: &SessionState, + setup: &mut dyn FnMut() -> Result<(), DataFusionError>, + ) -> Result<(), DataFusionError> { + self.checks.wrapper_calls.fetch_add(1, Ordering::Relaxed); + PLAN_SETUP_ACTIVE.with(|active| { + let previously_active = active.replace(true); + let result = setup(); + active.set(previously_active); + result + }) + } + } + + #[derive(Clone, Copy)] + struct RejectingSessionBuilder; + + #[async_trait] + impl WorkerSessionBuilder for RejectingSessionBuilder { + async fn build_session_state( + &self, + ctx: WorkerQueryContext, + ) -> Result { + Ok(ctx + .builder + .with_distributed_user_codec(ScopedPassThroughExecCodec { checks: None }) + .build()) + } + + fn run_plan_setup( + &self, + _session_state: &SessionState, + _setup: &mut dyn FnMut() -> Result<(), DataFusionError>, + ) -> Result<(), DataFusionError> { + internal_err!("worker plan setup rejected") + } + } + + fn add_sampler_check_hook(worker: &mut Worker, checks: Arc) { + worker.add_on_plan_hook(move |plan, _session_config| { + check_plan_setup_active("first plan hook")?; + checks.hook_checks.fetch_add(1, Ordering::Relaxed); + let checked_plan: Arc = + Arc::new(SamplerKickoffCheckExec::new(plan, Arc::clone(&checks))); + Ok(wrap_in_sampler(checked_plan)) + }); + } + + fn add_scope_check_hook(worker: &mut Worker, checks: Arc) { + worker.add_on_plan_hook(move |plan, _session_config| { + check_plan_setup_active("second plan hook")?; + checks.hook_checks.fetch_add(1, Ordering::Relaxed); + Ok(plan) + }); + } + + fn check_plan_setup_active(operation: &str) -> Result<()> { + PLAN_SETUP_ACTIVE.with(|active| { + if active.get() { + Ok(()) + } else { + internal_err!("{operation} ran outside WorkerSessionBuilder::run_plan_setup") + } + }) + } + + async fn register_input_table( + table_name: &str, + ctx: &datafusion::prelude::SessionContext, + ) -> Result { + let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)])); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(Int32Array::from(vec![1, 2, 3]))], + )?; + register_temp_parquet_table(table_name, schema, vec![batch], ctx).await + } + + #[derive(Debug)] + struct ScopedPassThroughExec { + properties: Arc, + child: Arc, + } + + impl ScopedPassThroughExec { + fn new(child: Arc) -> Self { + Self { + properties: Arc::clone(child.properties()), + child, + } + } + } + + impl DisplayAs for ScopedPassThroughExec { + fn fmt_as(&self, _: DisplayFormatType, f: &mut Formatter) -> std::fmt::Result { + write!(f, "ScopedPassThroughExec") + } + } + + impl ExecutionPlan for ScopedPassThroughExec { + fn name(&self) -> &str { + "ScopedPassThroughExec" + } + + fn properties(&self) -> &Arc { + &self.properties + } + + fn children(&self) -> Vec<&Arc> { + vec![&self.child] + } + + fn with_new_children( + self: Arc, + children: Vec>, + ) -> Result> { + let [child] = children.as_slice() else { + return internal_err!("ScopedPassThroughExec should have exactly one child"); + }; + Ok(Arc::new(Self::new(Arc::clone(child)))) + } + + fn execute( + &self, + partition: usize, + context: Arc, + ) -> Result { + self.child.execute(partition, context) + } + } + + #[derive(Debug)] + struct ScopedPassThroughExecCodec { + checks: Option>, + } + + #[derive(Clone, PartialEq, Message)] + struct ScopedPassThroughExecProto {} + + impl PhysicalExtensionCodec for ScopedPassThroughExecCodec { + fn try_decode( + &self, + buf: &[u8], + inputs: &[Arc], + _ctx: &TaskContext, + ) -> Result> { + let _node = ScopedPassThroughExecProto::decode(buf) + .map_err(|error| proto_error(format!("{error}")))?; + let [input] = inputs else { + return Err(proto_error(format!( + "ScopedPassThroughExec expects one child, got {}", + inputs.len() + ))); + }; + + if let Some(checks) = &self.checks { + check_plan_setup_active("physical extension decode")?; + checks.decode_checks.fetch_add(1, Ordering::Relaxed); + } + + Ok(Arc::new(ScopedPassThroughExec::new(Arc::clone(input)))) + } + + fn try_encode(&self, node: Arc, buf: &mut Vec) -> Result<()> { + if node.downcast_ref::().is_none() { + return Err(proto_error(format!( + "expected ScopedPassThroughExec, got {}", + node.name() + ))); + } + ScopedPassThroughExecProto {} + .encode(buf) + .map_err(|error| proto_error(format!("{error}"))) + } + } + + #[derive(Debug)] + struct SamplerKickoffCheckExec { + properties: Arc, + child: Arc, + checks: Arc, + } + + impl SamplerKickoffCheckExec { + fn new(child: Arc, checks: Arc) -> Self { + Self { + properties: Arc::clone(child.properties()), + child, + checks, + } + } + } + + impl DisplayAs for SamplerKickoffCheckExec { + fn fmt_as(&self, _: DisplayFormatType, f: &mut Formatter) -> std::fmt::Result { + write!(f, "SamplerKickoffCheckExec") + } + } + + impl ExecutionPlan for SamplerKickoffCheckExec { + fn name(&self) -> &str { + "SamplerKickoffCheckExec" + } + + fn properties(&self) -> &Arc { + &self.properties + } + + fn children(&self) -> Vec<&Arc> { + vec![&self.child] + } + + fn with_new_children( + self: Arc, + children: Vec>, + ) -> Result> { + let [child] = children.as_slice() else { + return internal_err!("SamplerKickoffCheckExec should have exactly one child"); + }; + Ok(Arc::new(Self::new( + Arc::clone(child), + Arc::clone(&self.checks), + ))) + } + + fn execute( + &self, + partition: usize, + context: Arc, + ) -> Result { + check_plan_setup_active("sampler kickoff")?; + self.checks + .sampler_kickoff_checks + .fetch_add(1, Ordering::Relaxed); + self.child.execute(partition, context) + } + } +}