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