Skip to main content

datafusion_federation/sql/
ast_analyzer.rs

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