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#[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 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 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 fn name(&self) -> &str {
123 "federate_sql"
124 }
125
126 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 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 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 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 #[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 assert!(
786 matches!(plan, LogicalPlan::Analyze(_)),
787 "Expected Analyze at root, got: {}",
788 plan.display_indent()
789 );
790
791 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 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}