diff --git a/datafusion-tracing/src/planner.rs b/datafusion-tracing/src/planner.rs index f7d30ab..c6e49f7 100644 --- a/datafusion-tracing/src/planner.rs +++ b/datafusion-tracing/src/planner.rs @@ -20,8 +20,7 @@ use async_trait::async_trait; use datafusion::catalog::Session; use datafusion::common::Result; -use datafusion::execution::SessionStateBuilder; -use datafusion::execution::context::{QueryPlanner, SessionState}; +use datafusion::execution::context::QueryPlanner; use datafusion::logical_expr::LogicalPlan; use datafusion::physical_plan::{ExecutionPlan, displayable}; use std::sync::Arc; @@ -54,26 +53,10 @@ pub(crate) struct TracingQueryPlanner { } impl TracingQueryPlanner { - /// Create a new `TracingQueryPlanner` that wraps the provided inner planner at a specific level. - fn new_with_level(inner: Arc, level: Level) -> Self { + /// Wrap the session's query planner with tracing at the specified level. + pub(crate) fn new(inner: Arc, level: Level) -> Self { Self { inner, level } } - - /// Wraps the query planner of an existing `SessionState` with tracing instrumentation at a specific level. - /// - /// This preserves any custom `QueryPlanner` that may already be configured in the state, - /// ensuring that tracing is added as a layer on top of existing functionality. - pub(crate) fn instrument_state_with_level( - state: SessionState, - level: Level, - ) -> SessionState { - let current_planner = state.query_planner().clone(); - let wrapped_planner = Arc::new(Self::new_with_level(current_planner, level)); - - SessionStateBuilder::from(state) - .with_query_planner(wrapped_planner) - .build() - } } #[async_trait] diff --git a/datafusion-tracing/src/rule_instrumentation.rs b/datafusion-tracing/src/rule_instrumentation.rs index ce652f6..36e2ade 100644 --- a/datafusion-tracing/src/rule_instrumentation.rs +++ b/datafusion-tracing/src/rule_instrumentation.rs @@ -977,18 +977,20 @@ pub fn instrument_session_state( ); // Rebuild SessionState with instrumented rules - let state = SessionStateBuilder::from(state) - .with_analyzer_rules(analyzers) + let planner = Arc::clone(state.query_planner()); + let mut builder = SessionStateBuilder::from(state) .with_optimizer_rules(optimizers) - .with_physical_optimizer_rules(physical_optimizers) - .build(); + .with_physical_optimizer_rules(physical_optimizers); + // Keep the existing analyzer's function rewrites. + builder.analyzer().get_or_insert_default().rules = analyzers; // Automatically instrument the query planner when physical optimizer is enabled if options.physical_optimizer.phase_span_enabled() { - TracingQueryPlanner::instrument_state_with_level(state, span_level) - } else { - state + builder = builder + .with_query_planner(Arc::new(TracingQueryPlanner::new(planner, span_level))); } + + builder.build() } /// Instruments analyzer rules with phase sentinel and optional rule-level spans. @@ -1126,8 +1128,10 @@ fn instrument_physical_optimizer_rules( #[cfg(test)] mod tests { use super::*; - use datafusion::common::DataFusionError; + use datafusion::common::{DFSchema, DataFusionError}; use datafusion::execution::SessionStateBuilder; + use datafusion::logical_expr::Expr; + use datafusion::logical_expr::expr_rewriter::FunctionRewrite; use datafusion::prelude::{SessionConfig, SessionContext}; use std::fmt; use std::sync::atomic::{AtomicBool, Ordering}; @@ -1400,6 +1404,40 @@ mod tests { Ok(()) } + #[derive(Debug)] + struct NoOpFunctionRewrite; + + impl FunctionRewrite for NoOpFunctionRewrite { + fn name(&self) -> &str { + "retained_function_rewrite" + } + + fn rewrite( + &self, + expr: Expr, + _: &DFSchema, + _: &ConfigOptions, + ) -> Result> { + Ok(Transformed::no(expr)) + } + } + + #[test] + fn instrumentation_preserves_analyzer_function_rewrites() { + let mut builder = SessionStateBuilder::new(); + builder + .analyzer() + .get_or_insert_default() + .add_function_rewrite(Arc::new(NoOpFunctionRewrite)); + + let state = crate::instrument_rules_with_info_spans!( + options: RuleInstrumentationOptions::full(), + state: builder.build() + ); + + assert_eq!(state.analyzer().function_rewrites.len(), 1); + } + // ----------------------------------------------------------------------- // Tests // -----------------------------------------------------------------------