Skip to main content

datafusion_federation/sql/
ast_analyzer.rs

1use 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/// Used to construct a AstAnalyzer that can replace table arguments.
20///
21/// ```rust
22/// use datafusion::sql::sqlparser::ast::{FunctionArg, Expr, Value};
23/// use datafusion::common::TableReference;
24/// use datafusion_federation::sql::ast_analyzer::TableArgReplace;
25///
26/// let mut analyzer = TableArgReplace::default().with(
27///     TableReference::parse_str("table1"),
28///     vec![FunctionArg::Unnamed(
29///         Expr::value(
30///             Value::Number("1".to_string(), false),
31///         )
32///         .into(),
33///     )],
34/// );
35/// let analyzer = analyzer.into_analyzer();
36/// ```
37#[derive(Debug, Clone, PartialEq, Eq, Default)]
38pub struct TableArgReplace {
39    pub tables: Vec<(TableReference, TableFunctionArgs)>,
40}
41
42impl TableArgReplace {
43    /// Constructs a new `TableArgReplace` instance.
44    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    /// Adds a new table argument replacement.
62    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    /// Converts the `TableArgReplace` instance into an `AstAnalyzer`.
74    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}