Skip to main content

datafusion_distributed/events/
worker_plan_rewrite.rs

1use super::common::EventHandlerChain;
2use datafusion::error::Result;
3use datafusion::execution::config::SessionConfig;
4use datafusion::physical_plan::ExecutionPlan;
5use std::sync::Arc;
6
7/// Information supplied while rewriting a decoded worker stage plan before registration.
8pub struct WorkerPlanRewriteEvent<'a> {
9    /// The worker-local plan. Each handler receives the plan returned by the previous handler.
10    pub plan: Arc<dyn ExecutionPlan>,
11    /// The configuration of the worker session that will execute the plan.
12    pub session_config: &'a SessionConfig,
13}
14
15/// The worker-local plan produced by a [`WorkerPlanRewriteHandler`].
16pub struct WorkerPlanRewriteEventResponse {
17    /// The original or transformed plan. If transformed, the plan needs to maintain the same
18    /// topology.
19    pub plan: Arc<dyn ExecutionPlan>,
20}
21
22impl WorkerPlanRewriteEventResponse {
23    /// Returns a response containing the rewritten worker-local plan.
24    pub fn new(plan: Arc<dyn ExecutionPlan>) -> Self {
25        Self { plan }
26    }
27}
28
29/// Rewrites a decoded worker-local plan before it is registered for execution.
30///
31/// Every registered handler runs in registration order and receives the plan returned by the
32/// previous handler. Returning an error aborts plan registration.
33pub trait WorkerPlanRewriteHandler: Send + Sync + 'static {
34    /// Returns the plan to pass to the next handler.
35    fn rewrite_worker_plan(
36        &self,
37        ev: WorkerPlanRewriteEvent,
38    ) -> Result<WorkerPlanRewriteEventResponse>;
39}
40
41impl<F> WorkerPlanRewriteHandler for F
42where
43    F: Send + Sync + 'static,
44    F: for<'a> Fn(WorkerPlanRewriteEvent<'a>) -> Result<WorkerPlanRewriteEventResponse>,
45{
46    fn rewrite_worker_plan(
47        &self,
48        ev: WorkerPlanRewriteEvent,
49    ) -> Result<WorkerPlanRewriteEventResponse> {
50        self(ev)
51    }
52}
53
54impl WorkerPlanRewriteHandler for Arc<dyn WorkerPlanRewriteHandler> {
55    fn rewrite_worker_plan(
56        &self,
57        ev: WorkerPlanRewriteEvent,
58    ) -> Result<WorkerPlanRewriteEventResponse> {
59        self.as_ref().rewrite_worker_plan(ev)
60    }
61}
62
63pub(crate) type WorkerPlanRewriteHandlers = EventHandlerChain<dyn WorkerPlanRewriteHandler>;
64
65impl WorkerPlanRewriteHandlers {
66    pub(crate) fn handle(ev: WorkerPlanRewriteEvent) -> Result<WorkerPlanRewriteEventResponse> {
67        let WorkerPlanRewriteEvent {
68            plan,
69            session_config,
70        } = ev;
71        let plan = match session_config.get_extension::<WorkerPlanRewriteHandlers>() {
72            Some(handlers) => handlers.try_fold(plan, |plan, handler| {
73                handler
74                    .rewrite_worker_plan(WorkerPlanRewriteEvent {
75                        plan,
76                        session_config,
77                    })
78                    .map(|response| response.plan)
79            })?,
80            None => plan,
81        };
82        Ok(WorkerPlanRewriteEventResponse::new(plan))
83    }
84}