datafusion_federation/sql/
ast_analyzer.rs1use std::ops::ControlFlow;
2
3use datafusion::common::TableReference;
4use datafusion::sql::sqlparser::ast::{
5 FunctionArg, Ident, ObjectName, Statement, TableAlias, TableFactor, TableFunctionArgs,
6 VisitMut, VisitorMut,
7};
8
9use super::AstAnalyzer;
10
11pub fn replace_table_args_analyzer(mut visitor: TableArgReplace) -> AstAnalyzer {
12 let x = move |mut statement: Statement| {
13 let _ = VisitMut::visit(&mut statement, &mut visitor);
14 Ok(statement)
15 };
16 Box::new(x)
17}
18
19#[derive(Debug, Clone, PartialEq, Eq, Default)]
38pub struct TableArgReplace {
39 pub tables: Vec<(TableReference, TableFunctionArgs)>,
40}
41
42impl TableArgReplace {
43 pub fn new(tables: Vec<(TableReference, Vec<FunctionArg>)>) -> Self {
45 Self {
46 tables: tables
47 .into_iter()
48 .map(|(table, args)| {
49 (
50 table,
51 TableFunctionArgs {
52 args,
53 settings: None,
54 },
55 )
56 })
57 .collect(),
58 }
59 }
60
61 pub fn with(mut self, table: TableReference, args: Vec<FunctionArg>) -> Self {
63 self.tables.push((
64 table,
65 TableFunctionArgs {
66 args,
67 settings: None,
68 },
69 ));
70 self
71 }
72
73 pub fn into_analyzer(self) -> AstAnalyzer {
75 replace_table_args_analyzer(self)
76 }
77}
78
79impl VisitorMut for TableArgReplace {
80 type Break = ();
81 fn pre_visit_table_factor(
82 &mut self,
83 table_factor: &mut TableFactor,
84 ) -> ControlFlow<Self::Break> {
85 if let TableFactor::Table {
86 name, args, alias, ..
87 } = table_factor
88 {
89 let name_as_tableref = name_to_table_reference(name);
90 if let Some((table, arg)) = self
91 .tables
92 .iter()
93 .find(|(t, _)| t.resolved_eq(&name_as_tableref))
94 {
95 *args = Some(arg.clone());
96 if alias.is_none() {
97 *alias = Some(TableAlias {
98 explicit: true,
99 name: Ident::new(table.table()),
100 columns: vec![],
101 at: None,
102 })
103 }
104 }
105 }
106 ControlFlow::Continue(())
107 }
108}
109
110fn name_to_table_reference(name: &ObjectName) -> TableReference {
111 let first = name
112 .0
113 .first()
114 .map(|n| n.as_ident().expect("expected Ident").value.to_string());
115 let second = name
116 .0
117 .get(1)
118 .map(|n| n.as_ident().expect("expected Ident").value.to_string());
119 let third = name
120 .0
121 .get(2)
122 .map(|n| n.as_ident().expect("expected Ident").value.to_string());
123
124 match (first, second, third) {
125 (Some(first), Some(second), Some(third)) => TableReference::full(first, second, third),
126 (Some(first), Some(second), None) => TableReference::partial(first, second),
127 (Some(first), None, None) => TableReference::bare(first),
128 _ => panic!("Invalid table name"),
129 }
130}