Skip to main content

query/
plan.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::HashSet;
16
17use datafusion::datasource::DefaultTableSource;
18use datafusion_common::TableReference;
19use datafusion_common::tree_node::{Transformed, TreeNodeRewriter};
20use datafusion_expr::{Expr, LogicalPlan};
21use session::context::QueryContextRef;
22pub use table::metadata::TableType;
23use table::table::adapter::DfTableProviderAdapter;
24use table::table_name::TableName;
25
26use crate::error::Result;
27
28struct TableNamesExtractAndRewriter {
29    pub(crate) table_names: HashSet<TableName>,
30    query_ctx: QueryContextRef,
31}
32
33impl TreeNodeRewriter for TableNamesExtractAndRewriter {
34    type Node = LogicalPlan;
35
36    /// descend
37    fn f_down<'a>(
38        &mut self,
39        node: Self::Node,
40    ) -> datafusion::error::Result<Transformed<Self::Node>> {
41        match node {
42            LogicalPlan::TableScan(mut scan) => {
43                if let Some(source) = scan.source.downcast_ref::<DefaultTableSource>()
44                    && let Some(provider) = source
45                        .table_provider
46                        .downcast_ref::<DfTableProviderAdapter>()
47                    && provider.table().table_type() == TableType::Base
48                {
49                    let info = provider.table().table_info();
50                    self.table_names.insert(TableName::new(
51                        info.catalog_name.clone(),
52                        info.schema_name.clone(),
53                        info.name.clone(),
54                    ));
55                }
56                match &scan.table_name {
57                    TableReference::Full {
58                        catalog,
59                        schema,
60                        table,
61                    } => {
62                        self.table_names.insert(TableName::new(
63                            catalog.to_string(),
64                            schema.to_string(),
65                            table.to_string(),
66                        ));
67                    }
68                    TableReference::Partial { schema, table } => {
69                        self.table_names.insert(TableName::new(
70                            self.query_ctx.current_catalog(),
71                            schema.to_string(),
72                            table.to_string(),
73                        ));
74
75                        scan.table_name = TableReference::Full {
76                            catalog: self.query_ctx.current_catalog().into(),
77                            schema: schema.clone(),
78                            table: table.clone(),
79                        };
80                    }
81                    TableReference::Bare { table } => {
82                        self.table_names.insert(TableName::new(
83                            self.query_ctx.current_catalog(),
84                            self.query_ctx.current_schema(),
85                            table.to_string(),
86                        ));
87
88                        scan.table_name = TableReference::Full {
89                            catalog: self.query_ctx.current_catalog().into(),
90                            schema: self.query_ctx.current_schema().into(),
91                            table: table.clone(),
92                        };
93                    }
94                }
95                Ok(Transformed::yes(LogicalPlan::TableScan(scan)))
96            }
97            node => Ok(Transformed::no(node)),
98        }
99    }
100}
101
102impl TableNamesExtractAndRewriter {
103    fn new(query_ctx: QueryContextRef) -> Self {
104        Self {
105            query_ctx,
106            table_names: HashSet::new(),
107        }
108    }
109}
110
111/// Extracts and rewrites the table names in the plan in the fully qualified style,
112/// return the table names and new plan.
113pub fn extract_and_rewrite_full_table_names(
114    plan: LogicalPlan,
115    query_ctx: QueryContextRef,
116) -> Result<(HashSet<TableName>, LogicalPlan)> {
117    let mut extractor = TableNamesExtractAndRewriter::new(query_ctx);
118    let plan = plan.rewrite_with_subqueries(&mut extractor)?;
119    Ok((extractor.table_names, plan.data))
120}
121
122/// A trait to extract expressions from a logical plan.
123pub trait ExtractExpr {
124    /// Gets expressions from a logical plan.
125    /// It handles [Join] specially so [LogicalPlan::with_new_exprs()] can use the expressions
126    /// this method returns.
127    fn expressions_consider_join(&self) -> Vec<Expr>;
128}
129
130impl ExtractExpr for LogicalPlan {
131    fn expressions_consider_join(&self) -> Vec<Expr> {
132        self.expressions()
133    }
134}
135
136#[cfg(test)]
137pub(crate) mod tests {
138
139    use std::sync::Arc;
140
141    use arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit};
142    use common_catalog::consts::DEFAULT_CATALOG_NAME;
143    use datafusion::logical_expr::builder::LogicalTableSource;
144    use datafusion::logical_expr::{LogicalPlan, LogicalPlanBuilder, col, lit, scalar_subquery};
145    use session::context::QueryContextBuilder;
146
147    use super::*;
148
149    fn mock_table_source() -> Arc<LogicalTableSource> {
150        let schema = Schema::new(vec![
151            Field::new("id", DataType::Int32, true),
152            Field::new("name", DataType::Utf8, true),
153            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
154        ]);
155        Arc::new(LogicalTableSource::new(SchemaRef::new(schema)))
156    }
157
158    fn mock_plan() -> LogicalPlan {
159        let table_source = mock_table_source();
160
161        let projection = None;
162
163        let builder = LogicalPlanBuilder::scan("devices", table_source, projection).unwrap();
164
165        builder
166            .filter(col("id").gt(lit(500)))
167            .unwrap()
168            .build()
169            .unwrap()
170    }
171
172    fn scalar_subquery_plan(table_name: TableReference) -> LogicalPlan {
173        let subquery = LogicalPlanBuilder::scan(table_name, mock_table_source(), None)
174            .unwrap()
175            .project(vec![col("id")])
176            .unwrap()
177            .build()
178            .unwrap();
179
180        LogicalPlanBuilder::empty(false)
181            .project(vec![scalar_subquery(Arc::new(subquery))])
182            .unwrap()
183            .build()
184            .unwrap()
185    }
186
187    fn assert_dependencies(actual: &HashSet<TableName>, expected: &[(&str, &str, &str)]) {
188        let expected = expected
189            .iter()
190            .map(|(catalog, schema, table)| TableName::new(*catalog, *schema, *table))
191            .collect::<HashSet<_>>();
192        assert_eq!(&expected, actual);
193    }
194
195    fn assert_nested_scalar_subquery_table_name(plan: &LogicalPlan, expected: TableReference) {
196        let LogicalPlan::Projection(projection) = plan else {
197            panic!("expected scalar-subquery projection, got {plan:?}");
198        };
199        let [Expr::ScalarSubquery(subquery)] = projection.expr.as_slice() else {
200            panic!(
201                "expected one scalar-subquery expression, got {:?}",
202                projection.expr
203            );
204        };
205        let LogicalPlan::Projection(projection) = subquery.subquery.as_ref() else {
206            panic!(
207                "expected scalar-subquery projection, got {:?}",
208                subquery.subquery
209            );
210        };
211        let LogicalPlan::TableScan(scan) = projection.input.as_ref() else {
212            panic!(
213                "expected scalar-subquery table scan, got {:?}",
214                projection.input
215            );
216        };
217
218        assert_eq!(expected, scan.table_name);
219    }
220
221    #[test]
222    fn test_extract_full_table_names() {
223        let ctx = QueryContextBuilder::default()
224            .current_schema("test".to_string())
225            .build();
226
227        let (table_names, plan) =
228            extract_and_rewrite_full_table_names(mock_plan(), Arc::new(ctx)).unwrap();
229
230        assert_dependencies(&table_names, &[(DEFAULT_CATALOG_NAME, "test", "devices")]);
231
232        assert_eq!(
233            "Filter: devices.id > Int32(500)\n  TableScan: greptime.test.devices",
234            plan.to_string()
235        );
236    }
237
238    #[test]
239    fn test_extract_full_table_names_from_scalar_subquery_bare_table_scan() {
240        let ctx = QueryContextBuilder::default()
241            .current_catalog("qp031_catalog".to_string())
242            .current_schema("qp031_current_schema".to_string())
243            .build();
244
245        let (table_names, plan) = extract_and_rewrite_full_table_names(
246            scalar_subquery_plan(TableReference::bare("lookup")),
247            Arc::new(ctx),
248        )
249        .unwrap();
250
251        assert_dependencies(
252            &table_names,
253            &[("qp031_catalog", "qp031_current_schema", "lookup")],
254        );
255        assert_nested_scalar_subquery_table_name(
256            &plan,
257            TableReference::full("qp031_catalog", "qp031_current_schema", "lookup"),
258        );
259    }
260
261    #[test]
262    fn test_extract_full_table_names_from_scalar_subquery_partial_table_scan() {
263        let ctx = QueryContextBuilder::default()
264            .current_catalog("qp031_catalog".to_string())
265            .current_schema("qp031_current_schema".to_string())
266            .build();
267
268        let (table_names, plan) = extract_and_rewrite_full_table_names(
269            scalar_subquery_plan(TableReference::partial("qp031_lookup_schema", "lookup")),
270            Arc::new(ctx),
271        )
272        .unwrap();
273
274        assert_dependencies(
275            &table_names,
276            &[("qp031_catalog", "qp031_lookup_schema", "lookup")],
277        );
278        assert_nested_scalar_subquery_table_name(
279            &plan,
280            TableReference::full("qp031_catalog", "qp031_lookup_schema", "lookup"),
281        );
282    }
283
284    #[test]
285    fn test_extract_full_table_names_from_scalar_subquery_full_table_scan() {
286        let ctx = QueryContextBuilder::default()
287            .current_catalog("qp031_catalog".to_string())
288            .current_schema("qp031_current_schema".to_string())
289            .build();
290
291        let (table_names, plan) = extract_and_rewrite_full_table_names(
292            scalar_subquery_plan(TableReference::full(
293                "qp031_external_catalog",
294                "qp031_external_schema",
295                "lookup",
296            )),
297            Arc::new(ctx),
298        )
299        .unwrap();
300
301        assert_dependencies(
302            &table_names,
303            &[("qp031_external_catalog", "qp031_external_schema", "lookup")],
304        );
305        assert_nested_scalar_subquery_table_name(
306            &plan,
307            TableReference::full("qp031_external_catalog", "qp031_external_schema", "lookup"),
308        );
309    }
310}