Skip to main content

datafusion_federation/sql/
mod.rs

1mod analyzer;
2pub mod ast_analyzer;
3mod executor;
4mod schema;
5mod table;
6mod table_reference;
7
8use std::{any::Any, fmt, sync::Arc, vec};
9
10use analyzer::RewriteTableScanAnalyzer;
11use async_trait::async_trait;
12use datafusion::{
13    arrow::datatypes::{Schema, SchemaRef},
14    catalog::Session,
15    common::{
16        tree_node::{Transformed, TreeNode, TreeNodeRecursion},
17        Statistics,
18    },
19    config::ConfigOptions,
20    error::{DataFusionError, Result},
21    execution::TaskContext,
22    logical_expr::{Extension, LogicalPlan},
23    optimizer::{optimizer::Optimizer, OptimizerConfig, OptimizerRule},
24    physical_expr::EquivalenceProperties,
25    physical_plan::{
26        apply_expression_roots,
27        execution_plan::{Boundedness, EmissionType},
28        filter_pushdown::{
29            ChildPushdownResult, FilterPushdownPhase, FilterPushdownPropagation, PushedDown,
30        },
31        metrics::MetricsSet,
32        DisplayAs, DisplayFormatType, ExecutionPlan, Partitioning, PhysicalExpr, PlanProperties,
33        SendableRecordBatchStream, StatisticsArgs,
34    },
35    sql::{sqlparser::ast::Statement, unparser::Unparser},
36};
37
38pub use executor::{AstAnalyzer, LogicalOptimizer, SQLExecutor, SQLExecutorRef, SqlQueryRewriter};
39pub use schema::{MultiSchemaProvider, SQLSchemaProvider};
40pub use table::{RemoteTable, SQLTable, SQLTableSource};
41pub use table_reference::RemoteTableRef;
42
43use crate::{
44    get_table_source, schema_cast, FederatedPlanNode, FederationPlanner, FederationProvider,
45};
46
47// SQLFederationProvider provides federation to SQL DMBSs.
48#[derive(Debug)]
49pub struct SQLFederationProvider {
50    pub optimizer: Arc<Optimizer>,
51    pub executor: Arc<dyn SQLExecutor>,
52}
53
54impl SQLFederationProvider {
55    pub fn new(executor: Arc<dyn SQLExecutor>) -> Self {
56        Self {
57            optimizer: Arc::new(Optimizer::with_rules(vec![Arc::new(
58                SQLFederationOptimizerRule::new(executor.clone()),
59            )])),
60            executor,
61        }
62    }
63}
64
65impl FederationProvider for SQLFederationProvider {
66    fn name(&self) -> &str {
67        "sql_federation_provider"
68    }
69
70    fn compute_context(&self) -> Option<String> {
71        self.executor.compute_context()
72    }
73
74    fn optimizer(&self) -> Option<Arc<Optimizer>> {
75        Some(self.optimizer.clone())
76    }
77}
78
79#[derive(Debug)]
80struct SQLFederationOptimizerRule {
81    planner: Arc<SQLFederationPlanner>,
82}
83
84impl SQLFederationOptimizerRule {
85    pub fn new(executor: Arc<dyn SQLExecutor>) -> Self {
86        Self {
87            planner: Arc::new(SQLFederationPlanner::new(Arc::clone(&executor))),
88        }
89    }
90}
91
92impl OptimizerRule for SQLFederationOptimizerRule {
93    /// Try to rewrite `plan` to an optimized form, returning `Transformed::yes`
94    /// if the plan was rewritten and `Transformed::no` if it was not.
95    ///
96    /// Note: this function is only called if [`Self::supports_rewrite`] returns
97    /// true. Otherwise the Optimizer calls  [`Self::try_optimize`]
98    fn rewrite(
99        &self,
100        plan: LogicalPlan,
101        _config: &dyn OptimizerConfig,
102    ) -> Result<Transformed<LogicalPlan>> {
103        if let LogicalPlan::Extension(Extension { ref node }) = plan {
104            if node.name() == "Federated" {
105                // Avoid attempting double federation
106                return Ok(Transformed::no(plan));
107            }
108        }
109
110        let fed_plan = FederatedPlanNode::new(plan.clone(), self.planner.clone());
111        let ext_node = Extension {
112            node: Arc::new(fed_plan),
113        };
114
115        let mut plan = LogicalPlan::Extension(ext_node);
116        if let Some(mut rewriter) = self.planner.executor.logical_optimizer() {
117            plan = rewriter(plan)?;
118        }
119
120        Ok(Transformed::yes(plan))
121    }
122
123    /// A human readable name for this analyzer rule
124    fn name(&self) -> &str {
125        "federate_sql"
126    }
127
128    /// Does this rule support rewriting owned plans (rather than by reference)?
129    fn supports_rewrite(&self) -> bool {
130        true
131    }
132}
133
134#[derive(Debug)]
135pub struct SQLFederationPlanner {
136    pub executor: Arc<dyn SQLExecutor>,
137}
138
139impl SQLFederationPlanner {
140    pub fn new(executor: Arc<dyn SQLExecutor>) -> Self {
141        Self { executor }
142    }
143}
144
145#[async_trait]
146impl FederationPlanner for SQLFederationPlanner {
147    async fn plan_federation(
148        &self,
149        node: &FederatedPlanNode,
150        _session_state: &dyn Session,
151    ) -> Result<Arc<dyn ExecutionPlan>> {
152        let schema = Arc::new(node.plan().schema().as_arrow().clone());
153        let plan = node.plan().clone();
154        let statistics = self.executor.statistics(&plan).await?;
155        let input = Arc::new(VirtualExecutionPlan::new(
156            plan,
157            Arc::clone(&self.executor),
158            statistics,
159        ));
160        let schema_cast_exec = schema_cast::SchemaCastScanExec::new(input, schema);
161        Ok(Arc::new(schema_cast_exec))
162    }
163}
164
165#[derive(Debug, Clone)]
166pub struct VirtualExecutionPlan {
167    plan: LogicalPlan,
168    executor: Arc<dyn SQLExecutor>,
169    props: Arc<PlanProperties>,
170    statistics: Statistics,
171    filters: Vec<Arc<dyn PhysicalExpr>>,
172}
173
174impl VirtualExecutionPlan {
175    pub fn new(plan: LogicalPlan, executor: Arc<dyn SQLExecutor>, statistics: Statistics) -> Self {
176        let schema: Schema = plan.schema().as_arrow().clone();
177        let props = Arc::new(PlanProperties::new(
178            EquivalenceProperties::new(Arc::new(schema)),
179            Partitioning::UnknownPartitioning(1),
180            EmissionType::Incremental,
181            Boundedness::Bounded,
182        ));
183        Self {
184            plan,
185            executor,
186            props,
187            statistics,
188            filters: Vec::new(),
189        }
190    }
191
192    pub fn plan(&self) -> &LogicalPlan {
193        &self.plan
194    }
195
196    pub fn executor(&self) -> &Arc<dyn SQLExecutor> {
197        &self.executor
198    }
199
200    pub fn statistics(&self) -> &Statistics {
201        &self.statistics
202    }
203
204    fn schema(&self) -> SchemaRef {
205        let df_schema = self.plan.schema().as_arrow().clone();
206        Arc::new(df_schema)
207    }
208
209    fn final_sql(&self) -> Result<String> {
210        let plan = self.plan.clone();
211        let plan = RewriteTableScanAnalyzer::rewrite(plan)?;
212        let (logical_optimizers, ast_analyzers, sql_query_rewriters) = gather_analyzers(&plan)?;
213        let plan = apply_logical_optimizers(plan, logical_optimizers)?;
214        let ast = self.plan_to_statement(&plan)?;
215        let ast = self.rewrite_with_executor_ast_analyzer(ast)?;
216        let ast = apply_ast_analyzers(ast, ast_analyzers)?;
217        apply_sql_query_rewriters(ast.to_string(), sql_query_rewriters)
218    }
219
220    fn rewrite_with_executor_ast_analyzer(
221        &self,
222        ast: Statement,
223    ) -> Result<Statement, datafusion::error::DataFusionError> {
224        if let Some(mut analyzer) = self.executor.ast_analyzer() {
225            Ok(analyzer(ast)?)
226        } else {
227            Ok(ast)
228        }
229    }
230
231    fn plan_to_statement(&self, plan: &LogicalPlan) -> Result<Statement> {
232        Unparser::new(self.executor.dialect().as_ref()).plan_to_sql(plan)
233    }
234}
235
236fn gather_analyzers(
237    plan: &LogicalPlan,
238) -> Result<(
239    Vec<LogicalOptimizer>,
240    Vec<AstAnalyzer>,
241    Vec<SqlQueryRewriter>,
242)> {
243    let mut logical_optimizers = vec![];
244    let mut ast_analyzers = vec![];
245    let mut sql_query_rewriters = vec![];
246
247    plan.apply(|node| {
248        if let LogicalPlan::TableScan(table) = node {
249            let provider = get_table_source(&table.source)
250                .expect("caller is virtual exec so this is valid")
251                .expect("caller is virtual exec so this is valid");
252            if let Some(source) = (provider.as_ref() as &dyn Any).downcast_ref::<SQLTableSource>() {
253                if let Some(analyzer) = source.table.logical_optimizer() {
254                    logical_optimizers.push(analyzer);
255                }
256                if let Some(analyzer) = source.table.ast_analyzer() {
257                    ast_analyzers.push(analyzer);
258                }
259                if let Some(rewriter) = source.table.sql_query_rewriter() {
260                    sql_query_rewriters.push(rewriter);
261                }
262            }
263        }
264        Ok(datafusion::common::tree_node::TreeNodeRecursion::Continue)
265    })?;
266
267    Ok((logical_optimizers, ast_analyzers, sql_query_rewriters))
268}
269
270fn apply_logical_optimizers(
271    mut plan: LogicalPlan,
272    analyzers: Vec<LogicalOptimizer>,
273) -> Result<LogicalPlan> {
274    for mut analyzer in analyzers {
275        let old_schema = plan.schema().clone();
276        plan = analyzer(plan)?;
277        let new_schema = plan.schema();
278        if &old_schema != new_schema {
279            return Err(DataFusionError::Execution(format!(
280                "Schema altered during logical analysis, expected: {}, found: {}",
281                old_schema, new_schema
282            )));
283        }
284    }
285    Ok(plan)
286}
287
288fn apply_ast_analyzers(mut statement: Statement, analyzers: Vec<AstAnalyzer>) -> Result<Statement> {
289    for mut analyzer in analyzers {
290        statement = analyzer(statement)?;
291    }
292    Ok(statement)
293}
294
295fn apply_sql_query_rewriters(
296    mut query: String,
297    rewriters: Vec<SqlQueryRewriter>,
298) -> Result<String> {
299    for mut rewriter in rewriters {
300        query = rewriter(query)?;
301    }
302    Ok(query)
303}
304
305impl DisplayAs for VirtualExecutionPlan {
306    fn fmt_as(&self, _t: DisplayFormatType, f: &mut fmt::Formatter) -> std::fmt::Result {
307        write!(f, "VirtualExecutionPlan")?;
308        write!(f, " name={}", self.executor.name())?;
309        if let Some(ctx) = self.executor.compute_context() {
310            write!(f, " compute_context={ctx}")?;
311        };
312        let mut plan = match RewriteTableScanAnalyzer::rewrite(self.plan.clone()) {
313            Ok(plan) => plan,
314            Err(_) => self.plan.clone(),
315        };
316        if let Ok(statement) = self.plan_to_statement(&plan) {
317            write!(f, " base_sql={statement}")?;
318        }
319
320        let (logical_optimizers, ast_analyzers, sql_query_rewriters) = match gather_analyzers(&plan)
321        {
322            Ok(analyzers) => analyzers,
323            Err(_) => return Ok(()),
324        };
325
326        let old_plan = plan.clone();
327
328        plan = match apply_logical_optimizers(plan, logical_optimizers) {
329            Ok(plan) => plan,
330            _ => return Ok(()),
331        };
332
333        let statement = match self.plan_to_statement(&plan) {
334            Ok(statement) => statement,
335            _ => return Ok(()),
336        };
337
338        if plan != old_plan {
339            write!(f, " rewritten_logical_sql={statement}")?;
340        }
341
342        let old_statement = statement.clone();
343        let statement = match self.rewrite_with_executor_ast_analyzer(statement) {
344            Ok(statement) => statement,
345            _ => return Ok(()),
346        };
347        if old_statement != statement {
348            write!(f, " rewritten_executor_sql={statement}")?;
349        }
350
351        let old_statement = statement.clone();
352        let statement = match apply_ast_analyzers(statement, ast_analyzers) {
353            Ok(statement) => statement,
354            _ => return Ok(()),
355        };
356        if old_statement != statement {
357            write!(f, " rewritten_ast_analyzer={statement}")?;
358        }
359
360        let sql = statement.to_string();
361        let rewritten_sql = match apply_sql_query_rewriters(sql.clone(), sql_query_rewriters) {
362            Ok(sql) => sql,
363            _ => return Ok(()),
364        };
365        if sql != rewritten_sql {
366            write!(f, " rewritten_sql_query={rewritten_sql}")?;
367        }
368
369        Ok(())
370    }
371}
372
373impl ExecutionPlan for VirtualExecutionPlan {
374    fn name(&self) -> &str {
375        "sql_federation_exec"
376    }
377
378    fn schema(&self) -> SchemaRef {
379        self.schema()
380    }
381
382    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
383        vec![]
384    }
385
386    fn with_new_children(
387        self: Arc<Self>,
388        _: Vec<Arc<dyn ExecutionPlan>>,
389    ) -> Result<Arc<dyn ExecutionPlan>> {
390        Ok(self)
391    }
392
393    fn execute(
394        &self,
395        _partition: usize,
396        _context: Arc<TaskContext>,
397    ) -> Result<SendableRecordBatchStream> {
398        self.executor
399            .execute(&self.final_sql()?, self.schema(), &self.filters)
400    }
401
402    fn properties(&self) -> &Arc<PlanProperties> {
403        &self.props
404    }
405
406    fn statistics_from_inputs(
407        &self,
408        _input_stats: &[Arc<Statistics>],
409        _args: &StatisticsArgs,
410    ) -> Result<Arc<Statistics>> {
411        Ok(Arc::new(self.statistics.clone()))
412    }
413
414    fn metrics(&self) -> Option<MetricsSet> {
415        self.executor.metrics()
416    }
417
418    fn handle_child_pushdown_result(
419        &self,
420        _phase: FilterPushdownPhase,
421        child_pushdown_result: ChildPushdownResult,
422        _config: &ConfigOptions,
423    ) -> Result<FilterPushdownPropagation<Arc<dyn ExecutionPlan>>> {
424        let parent_filters: Vec<_> = child_pushdown_result
425            .clone()
426            .parent_filters
427            .into_iter()
428            .map(|f| f.filter)
429            .collect();
430
431        if parent_filters.is_empty() {
432            return Ok(FilterPushdownPropagation {
433                filters: vec![],
434                updated_node: None,
435            });
436        }
437
438        let filters_pushed_down = vec![PushedDown::Yes; parent_filters.len()];
439        let mut node = self.clone();
440        node.filters = parent_filters;
441
442        Ok(FilterPushdownPropagation {
443            filters: filters_pushed_down,
444            updated_node: Some(Arc::new(node)),
445        })
446    }
447
448    fn apply_expressions(
449        &self,
450        f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> Result<TreeNodeRecursion>,
451    ) -> Result<TreeNodeRecursion> {
452        apply_expression_roots(&self.filters, f)
453    }
454}
455
456#[cfg(test)]
457mod tests {
458    use std::any::Any;
459    use std::collections::HashSet;
460    use std::sync::atomic::{AtomicUsize, Ordering};
461    use std::sync::Arc;
462
463    use crate::sql::{
464        RemoteTableRef, SQLExecutor, SQLFederationProvider, SQLTable, SQLTableSource,
465    };
466    use crate::FederatedTableProviderAdaptor;
467    use async_trait::async_trait;
468    use datafusion::arrow::datatypes::{Schema, SchemaRef};
469    use datafusion::common::tree_node::TreeNodeRecursion;
470    use datafusion::common::TableReference;
471    use datafusion::execution::SendableRecordBatchStream;
472    use datafusion::sql::unparser::dialect::Dialect;
473    use datafusion::sql::unparser::{self};
474    use datafusion::{
475        arrow::datatypes::{DataType, Field},
476        datasource::TableProvider,
477        execution::context::SessionContext,
478    };
479
480    use super::table::RemoteTable;
481    use super::*;
482
483    #[derive(Debug, Clone)]
484    struct TestExecutor {
485        compute_context: String,
486    }
487
488    #[async_trait]
489    impl SQLExecutor for TestExecutor {
490        fn name(&self) -> &str {
491            "TestExecutor"
492        }
493
494        fn compute_context(&self) -> Option<String> {
495            Some(self.compute_context.clone())
496        }
497
498        fn dialect(&self) -> Arc<dyn Dialect> {
499            Arc::new(unparser::dialect::DefaultDialect {})
500        }
501
502        fn execute(
503            &self,
504            _query: &str,
505            _schema: SchemaRef,
506            _filters: &[Arc<dyn PhysicalExpr>],
507        ) -> Result<SendableRecordBatchStream> {
508            unimplemented!()
509        }
510
511        async fn table_names(&self) -> Result<Vec<String>> {
512            unimplemented!()
513        }
514
515        async fn get_table_schema(&self, _table_name: &str) -> Result<SchemaRef> {
516            unimplemented!()
517        }
518    }
519
520    fn get_test_table_provider(name: String, executor: TestExecutor) -> Arc<dyn TableProvider> {
521        let schema = Arc::new(Schema::new(vec![
522            Field::new("a", DataType::Int64, false),
523            Field::new("b", DataType::Utf8, false),
524            Field::new("c", DataType::Date32, false),
525        ]));
526        let table_ref = RemoteTableRef::try_from(name).unwrap();
527        let table = Arc::new(RemoteTable::new(table_ref, schema));
528        let provider = Arc::new(SQLFederationProvider::new(Arc::new(executor)));
529        let table_source = Arc::new(SQLTableSource { provider, table });
530        Arc::new(FederatedTableProviderAdaptor::new(table_source))
531    }
532
533    fn get_test_table_provider_with_table(
534        table: Arc<dyn SQLTable>,
535        executor: TestExecutor,
536    ) -> Arc<dyn TableProvider> {
537        let provider = Arc::new(SQLFederationProvider::new(Arc::new(executor)));
538        let table_source = Arc::new(SQLTableSource::new_with_table(provider, table));
539        Arc::new(FederatedTableProviderAdaptor::new(table_source))
540    }
541
542    #[derive(Debug)]
543    struct SqlRewriteTable {
544        table: RemoteTable,
545        rewrite_calls: Arc<AtomicUsize>,
546        suffix: String,
547    }
548
549    impl SqlRewriteTable {
550        fn new(
551            table_ref: RemoteTableRef,
552            schema: SchemaRef,
553            rewrite_calls: Arc<AtomicUsize>,
554            suffix: impl Into<String>,
555        ) -> Self {
556            Self {
557                table: RemoteTable::new(table_ref, schema),
558                rewrite_calls,
559                suffix: suffix.into(),
560            }
561        }
562    }
563
564    impl SQLTable for SqlRewriteTable {
565        fn as_any(&self) -> &dyn Any {
566            self
567        }
568
569        fn table_reference(&self) -> TableReference {
570            self.table.table_reference().clone()
571        }
572
573        fn schema(&self) -> SchemaRef {
574            Arc::clone(self.table.schema())
575        }
576
577        fn sql_query_rewriter(&self) -> Option<SqlQueryRewriter> {
578            let rewrite_calls = Arc::clone(&self.rewrite_calls);
579            let suffix = self.suffix.clone();
580            Some(Box::new(move |sql| {
581                rewrite_calls.fetch_add(1, Ordering::SeqCst);
582                Ok(format!("{sql} {suffix}"))
583            }))
584        }
585    }
586
587    #[tokio::test]
588    async fn basic_sql_federation_test() -> Result<(), DataFusionError> {
589        let test_executor_a = TestExecutor {
590            compute_context: "a".into(),
591        };
592
593        let test_executor_b = TestExecutor {
594            compute_context: "b".into(),
595        };
596
597        let table_a1_ref = "table_a1".to_string();
598        let table_a1 = get_test_table_provider(table_a1_ref.clone(), test_executor_a.clone());
599
600        let table_a2_ref = "table_a2".to_string();
601        let table_a2 = get_test_table_provider(table_a2_ref.clone(), test_executor_a);
602
603        let table_b1_ref = "table_b1(1)".to_string();
604        let table_b1_df_ref = "table_local_b1".to_string();
605
606        let table_b1 = get_test_table_provider(table_b1_ref.clone(), test_executor_b);
607
608        // Create a new SessionState with the optimizer rule we created above
609        let state = crate::default_session_state();
610        let ctx = SessionContext::new_with_state(state);
611
612        ctx.register_table(table_a1_ref.clone(), table_a1).unwrap();
613        ctx.register_table(table_a2_ref.clone(), table_a2).unwrap();
614        ctx.register_table(table_b1_df_ref.clone(), table_b1)
615            .unwrap();
616
617        let query = r#"
618            SELECT * FROM table_a1
619            UNION ALL
620            SELECT * FROM table_a2
621            UNION ALL
622            SELECT * FROM table_local_b1;
623        "#;
624
625        let df = ctx.sql(query).await?;
626
627        let logical_plan = df.into_optimized_plan()?;
628
629        let mut table_a1_federated = false;
630        let mut table_a2_federated = false;
631        let mut table_b1_federated = false;
632
633        let _ = logical_plan.apply(|node| {
634            if let LogicalPlan::Extension(node) = node {
635                if let Some(node) = node.node.as_any().downcast_ref::<FederatedPlanNode>() {
636                    let _ = node.plan().apply(|node| {
637                        if let LogicalPlan::TableScan(table) = node {
638                            if table.table_name.table() == table_a1_ref {
639                                table_a1_federated = true;
640                            }
641                            if table.table_name.table() == table_a2_ref {
642                                table_a2_federated = true;
643                            }
644                            // assuming table name is rewritten via analyzer
645                            if table.table_name.table() == table_b1_df_ref {
646                                table_b1_federated = true;
647                            }
648                        }
649                        Ok(TreeNodeRecursion::Continue)
650                    });
651                }
652            }
653            Ok(TreeNodeRecursion::Continue)
654        });
655
656        assert!(table_a1_federated);
657        assert!(table_a2_federated);
658        assert!(table_b1_federated);
659
660        let physical_plan = ctx.state().create_physical_plan(&logical_plan).await?;
661
662        let mut final_queries = vec![];
663
664        let _ = physical_plan.apply(|node| {
665            if node.name() == "sql_federation_exec" {
666                let node = (node.as_ref() as &dyn Any)
667                    .downcast_ref::<VirtualExecutionPlan>()
668                    .unwrap();
669
670                final_queries.push(node.final_sql()?);
671            }
672            Ok(TreeNodeRecursion::Continue)
673        });
674
675        let expected = vec![
676            "SELECT table_a1.a, table_a1.b, table_a1.c FROM table_a1",
677            "SELECT table_a2.a, table_a2.b, table_a2.c FROM table_a2",
678            "SELECT table_b1.a, table_b1.b, table_b1.c FROM table_b1(1) AS table_b1",
679        ];
680
681        assert_eq!(
682            HashSet::<&str>::from_iter(final_queries.iter().map(|x| x.as_str())),
683            HashSet::from_iter(expected)
684        );
685
686        Ok(())
687    }
688
689    #[tokio::test]
690    async fn multi_reference_sql_federation_test() -> Result<(), DataFusionError> {
691        let test_executor_a = TestExecutor {
692            compute_context: "test".into(),
693        };
694
695        let lowercase_table_ref = "default.table".to_string();
696        let lowercase_local_table_ref = "dftable".to_string();
697        let lowercase_table =
698            get_test_table_provider(lowercase_table_ref.clone(), test_executor_a.clone());
699
700        let capitalized_table_ref = "default.Table(1)".to_string();
701        let capitalized_local_table_ref = "dfview".to_string();
702        let capitalized_table =
703            get_test_table_provider(capitalized_table_ref.clone(), test_executor_a);
704
705        // Create a new SessionState with the optimizer rule we created above
706        let state = crate::default_session_state();
707        let ctx = SessionContext::new_with_state(state);
708
709        ctx.register_table(lowercase_local_table_ref.clone(), lowercase_table)
710            .unwrap();
711        ctx.register_table(capitalized_local_table_ref.clone(), capitalized_table)
712            .unwrap();
713
714        let query = r#"
715                SELECT * FROM dftable
716                UNION ALL
717                SELECT * FROM dfview;
718            "#;
719
720        let df = ctx.sql(query).await?;
721
722        let logical_plan = df.into_optimized_plan()?;
723
724        let mut lowercase_table = false;
725        let mut capitalized_table = false;
726
727        let _ = logical_plan.apply(|node| {
728            if let LogicalPlan::Extension(node) = node {
729                if let Some(node) = node.node.as_any().downcast_ref::<FederatedPlanNode>() {
730                    let _ = node.plan().apply(|node| {
731                        if let LogicalPlan::TableScan(table) = node {
732                            if table.table_name.table() == lowercase_local_table_ref {
733                                lowercase_table = true;
734                            }
735                            if table.table_name.table() == capitalized_local_table_ref {
736                                capitalized_table = true;
737                            }
738                        }
739                        Ok(TreeNodeRecursion::Continue)
740                    });
741                }
742            }
743            Ok(TreeNodeRecursion::Continue)
744        });
745
746        assert!(lowercase_table);
747        assert!(capitalized_table);
748
749        let physical_plan = ctx.state().create_physical_plan(&logical_plan).await?;
750
751        let mut final_queries = vec![];
752
753        let _ = physical_plan.apply(|node| {
754            if node.name() == "sql_federation_exec" {
755                let node = (node.as_ref() as &dyn Any)
756                    .downcast_ref::<VirtualExecutionPlan>()
757                    .unwrap();
758
759                final_queries.push(node.final_sql()?);
760            }
761            Ok(TreeNodeRecursion::Continue)
762        });
763
764        let expected = vec![
765            r#"SELECT "table".a, "table".b, "table".c FROM "default"."table" UNION ALL SELECT "Table".a, "Table".b, "Table".c FROM "default"."Table"(1) AS Table"#,
766        ];
767
768        assert_eq!(
769            HashSet::<&str>::from_iter(final_queries.iter().map(|x| x.as_str())),
770            HashSet::from_iter(expected)
771        );
772
773        Ok(())
774    }
775
776    /// EXPLAIN ANALYZE must not federate the Analyze wrapper — only the inner
777    /// query should be federated. Otherwise the SQL Unparser fails because it
778    /// cannot convert Analyze to SQL.
779    #[tokio::test]
780    async fn explain_analyze_not_federated() -> Result<(), DataFusionError> {
781        let executor = TestExecutor {
782            compute_context: "a".into(),
783        };
784
785        let table_ref = "test_table".to_string();
786        let table = get_test_table_provider(table_ref.clone(), executor);
787
788        let state = crate::default_session_state();
789        let ctx = SessionContext::new_with_state(state);
790        ctx.register_table(table_ref, table).unwrap();
791
792        let plan = ctx
793            .sql("EXPLAIN ANALYZE SELECT * FROM test_table")
794            .await?
795            .into_optimized_plan()?;
796
797        // The top-level node must be Analyze, not Federated.
798        assert!(
799            matches!(plan, LogicalPlan::Analyze(_)),
800            "Expected Analyze at root, got: {}",
801            plan.display_indent()
802        );
803
804        // The inner plan should contain a Federated extension node.
805        let mut found_federated = false;
806        plan.apply(|node| {
807            if let LogicalPlan::Extension(ext) = node {
808                if ext.node.name() == "Federated" {
809                    found_federated = true;
810                    return Ok(TreeNodeRecursion::Stop);
811                }
812            }
813            Ok(TreeNodeRecursion::Continue)
814        })?;
815        assert!(
816            found_federated,
817            "Expected a Federated node inside the Analyze plan"
818        );
819
820        // Physical planning should succeed (this is where it used to fail).
821        let physical_plan = ctx.state().create_physical_plan(&plan).await?;
822        assert_eq!(physical_plan.name(), "AnalyzeExec");
823
824        Ok(())
825    }
826
827    #[tokio::test]
828    async fn sql_query_rewriter_hook_invoked_and_rewrites_sql() -> Result<(), DataFusionError> {
829        let executor = TestExecutor {
830            compute_context: "rewrite".into(),
831        };
832        let rewrite_calls = Arc::new(AtomicUsize::new(0));
833        let table_ref = "table_with_rewriter".to_string();
834        let table = Arc::new(SqlRewriteTable::new(
835            table_ref.clone().try_into().unwrap(),
836            Arc::new(Schema::new(vec![
837                Field::new("a", DataType::Int64, false),
838                Field::new("b", DataType::Utf8, false),
839                Field::new("c", DataType::Date32, false),
840            ])),
841            Arc::clone(&rewrite_calls),
842            "/* rewritten by sql_query_rewriter */",
843        ));
844        let table_provider = get_test_table_provider_with_table(table, executor);
845
846        let state = crate::default_session_state();
847        let ctx = SessionContext::new_with_state(state);
848        ctx.register_table(table_ref.clone(), table_provider)
849            .unwrap();
850
851        let query = format!("SELECT * FROM {table_ref}");
852        let df = ctx.sql(&query).await?;
853        let logical_plan = df.into_optimized_plan()?;
854        let physical_plan = ctx.state().create_physical_plan(&logical_plan).await?;
855
856        let mut final_queries = vec![];
857        let _ = physical_plan.apply(|node| {
858            if node.name() == "sql_federation_exec" {
859                let node = (node.as_ref() as &dyn Any)
860                    .downcast_ref::<VirtualExecutionPlan>()
861                    .unwrap();
862                final_queries.push(node.final_sql()?);
863            }
864            Ok(TreeNodeRecursion::Continue)
865        });
866
867        let [final_query] = final_queries.as_slice() else {
868            panic!("expected a single federated SQL query");
869        };
870
871        assert!(final_query.ends_with("/* rewritten by sql_query_rewriter */"));
872        assert_eq!(rewrite_calls.load(Ordering::SeqCst), 1);
873
874        Ok(())
875    }
876}