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#[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 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 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 fn name(&self) -> &str {
125 "federate_sql"
126 }
127
128 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 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 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 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 #[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 assert!(
799 matches!(plan, LogicalPlan::Analyze(_)),
800 "Expected Analyze at root, got: {}",
801 plan.display_indent()
802 );
803
804 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 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}