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.as_any().downcast_ref::<DefaultTableSource>()
44                    && let Some(provider) = source
45                        .table_provider
46                        .as_any()
47                        .downcast_ref::<DfTableProviderAdapter>()
48                    && provider.table().table_type() == TableType::Base
49                {
50                    let info = provider.table().table_info();
51                    self.table_names.insert(TableName::new(
52                        info.catalog_name.clone(),
53                        info.schema_name.clone(),
54                        info.name.clone(),
55                    ));
56                }
57                match &scan.table_name {
58                    TableReference::Full {
59                        catalog,
60                        schema,
61                        table,
62                    } => {
63                        self.table_names.insert(TableName::new(
64                            catalog.to_string(),
65                            schema.to_string(),
66                            table.to_string(),
67                        ));
68                    }
69                    TableReference::Partial { schema, table } => {
70                        self.table_names.insert(TableName::new(
71                            self.query_ctx.current_catalog(),
72                            schema.to_string(),
73                            table.to_string(),
74                        ));
75
76                        scan.table_name = TableReference::Full {
77                            catalog: self.query_ctx.current_catalog().into(),
78                            schema: schema.clone(),
79                            table: table.clone(),
80                        };
81                    }
82                    TableReference::Bare { table } => {
83                        self.table_names.insert(TableName::new(
84                            self.query_ctx.current_catalog(),
85                            self.query_ctx.current_schema(),
86                            table.to_string(),
87                        ));
88
89                        scan.table_name = TableReference::Full {
90                            catalog: self.query_ctx.current_catalog().into(),
91                            schema: self.query_ctx.current_schema().into(),
92                            table: table.clone(),
93                        };
94                    }
95                }
96                Ok(Transformed::yes(LogicalPlan::TableScan(scan)))
97            }
98            node => Ok(Transformed::no(node)),
99        }
100    }
101}
102
103impl TableNamesExtractAndRewriter {
104    fn new(query_ctx: QueryContextRef) -> Self {
105        Self {
106            query_ctx,
107            table_names: HashSet::new(),
108        }
109    }
110}
111
112/// Extracts and rewrites the table names in the plan in the fully qualified style,
113/// return the table names and new plan.
114pub fn extract_and_rewrite_full_table_names(
115    plan: LogicalPlan,
116    query_ctx: QueryContextRef,
117) -> Result<(HashSet<TableName>, LogicalPlan)> {
118    let mut extractor = TableNamesExtractAndRewriter::new(query_ctx);
119    let plan = plan.rewrite_with_subqueries(&mut extractor)?;
120    Ok((extractor.table_names, plan.data))
121}
122
123/// A trait to extract expressions from a logical plan.
124pub trait ExtractExpr {
125    /// Gets expressions from a logical plan.
126    /// It handles [Join] specially so [LogicalPlan::with_new_exprs()] can use the expressions
127    /// this method returns.
128    fn expressions_consider_join(&self) -> Vec<Expr>;
129}
130
131impl ExtractExpr for LogicalPlan {
132    fn expressions_consider_join(&self) -> Vec<Expr> {
133        self.expressions()
134    }
135}
136
137#[cfg(test)]
138pub(crate) mod tests {
139
140    use std::sync::Arc;
141
142    use arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit};
143    use common_catalog::consts::DEFAULT_CATALOG_NAME;
144    use datafusion::logical_expr::builder::LogicalTableSource;
145    use datafusion::logical_expr::{LogicalPlan, LogicalPlanBuilder, col, lit, scalar_subquery};
146    use session::context::QueryContextBuilder;
147
148    use super::*;
149
150    fn mock_table_source() -> Arc<LogicalTableSource> {
151        let schema = Schema::new(vec![
152            Field::new("id", DataType::Int32, true),
153            Field::new("name", DataType::Utf8, true),
154            Field::new("ts", DataType::Timestamp(TimeUnit::Millisecond, None), true),
155        ]);
156        Arc::new(LogicalTableSource::new(SchemaRef::new(schema)))
157    }
158
159    fn mock_plan() -> LogicalPlan {
160        let table_source = mock_table_source();
161
162        let projection = None;
163
164        let builder = LogicalPlanBuilder::scan("devices", table_source, projection).unwrap();
165
166        builder
167            .filter(col("id").gt(lit(500)))
168            .unwrap()
169            .build()
170            .unwrap()
171    }
172
173    fn scalar_subquery_plan(table_name: TableReference) -> LogicalPlan {
174        let subquery = LogicalPlanBuilder::scan(table_name, mock_table_source(), None)
175            .unwrap()
176            .project(vec![col("id")])
177            .unwrap()
178            .build()
179            .unwrap();
180
181        LogicalPlanBuilder::empty(false)
182            .project(vec![scalar_subquery(Arc::new(subquery))])
183            .unwrap()
184            .build()
185            .unwrap()
186    }
187
188    fn assert_dependencies(actual: &HashSet<TableName>, expected: &[(&str, &str, &str)]) {
189        let expected = expected
190            .iter()
191            .map(|(catalog, schema, table)| TableName::new(*catalog, *schema, *table))
192            .collect::<HashSet<_>>();
193        assert_eq!(&expected, actual);
194    }
195
196    fn assert_nested_scalar_subquery_table_name(plan: &LogicalPlan, expected: TableReference) {
197        let LogicalPlan::Projection(projection) = plan else {
198            panic!("expected scalar-subquery projection, got {plan:?}");
199        };
200        let [Expr::ScalarSubquery(subquery)] = projection.expr.as_slice() else {
201            panic!(
202                "expected one scalar-subquery expression, got {:?}",
203                projection.expr
204            );
205        };
206        let LogicalPlan::Projection(projection) = subquery.subquery.as_ref() else {
207            panic!(
208                "expected scalar-subquery projection, got {:?}",
209                subquery.subquery
210            );
211        };
212        let LogicalPlan::TableScan(scan) = projection.input.as_ref() else {
213            panic!(
214                "expected scalar-subquery table scan, got {:?}",
215                projection.input
216            );
217        };
218
219        assert_eq!(expected, scan.table_name);
220    }
221
222    #[test]
223    fn test_extract_full_table_names() {
224        let ctx = QueryContextBuilder::default()
225            .current_schema("test".to_string())
226            .build();
227
228        let (table_names, plan) =
229            extract_and_rewrite_full_table_names(mock_plan(), Arc::new(ctx)).unwrap();
230
231        assert_dependencies(&table_names, &[(DEFAULT_CATALOG_NAME, "test", "devices")]);
232
233        assert_eq!(
234            "Filter: devices.id > Int32(500)\n  TableScan: greptime.test.devices",
235            plan.to_string()
236        );
237    }
238
239    #[test]
240    fn test_extract_full_table_names_from_scalar_subquery_bare_table_scan() {
241        let ctx = QueryContextBuilder::default()
242            .current_catalog("qp031_catalog".to_string())
243            .current_schema("qp031_current_schema".to_string())
244            .build();
245
246        let (table_names, plan) = extract_and_rewrite_full_table_names(
247            scalar_subquery_plan(TableReference::bare("lookup")),
248            Arc::new(ctx),
249        )
250        .unwrap();
251
252        assert_dependencies(
253            &table_names,
254            &[("qp031_catalog", "qp031_current_schema", "lookup")],
255        );
256        assert_nested_scalar_subquery_table_name(
257            &plan,
258            TableReference::full("qp031_catalog", "qp031_current_schema", "lookup"),
259        );
260    }
261
262    #[test]
263    fn test_extract_full_table_names_from_scalar_subquery_partial_table_scan() {
264        let ctx = QueryContextBuilder::default()
265            .current_catalog("qp031_catalog".to_string())
266            .current_schema("qp031_current_schema".to_string())
267            .build();
268
269        let (table_names, plan) = extract_and_rewrite_full_table_names(
270            scalar_subquery_plan(TableReference::partial("qp031_lookup_schema", "lookup")),
271            Arc::new(ctx),
272        )
273        .unwrap();
274
275        assert_dependencies(
276            &table_names,
277            &[("qp031_catalog", "qp031_lookup_schema", "lookup")],
278        );
279        assert_nested_scalar_subquery_table_name(
280            &plan,
281            TableReference::full("qp031_catalog", "qp031_lookup_schema", "lookup"),
282        );
283    }
284
285    #[test]
286    fn test_extract_full_table_names_from_scalar_subquery_full_table_scan() {
287        let ctx = QueryContextBuilder::default()
288            .current_catalog("qp031_catalog".to_string())
289            .current_schema("qp031_current_schema".to_string())
290            .build();
291
292        let (table_names, plan) = extract_and_rewrite_full_table_names(
293            scalar_subquery_plan(TableReference::full(
294                "qp031_external_catalog",
295                "qp031_external_schema",
296                "lookup",
297            )),
298            Arc::new(ctx),
299        )
300        .unwrap();
301
302        assert_dependencies(
303            &table_names,
304            &[("qp031_external_catalog", "qp031_external_schema", "lookup")],
305        );
306        assert_nested_scalar_subquery_table_name(
307            &plan,
308            TableReference::full("qp031_external_catalog", "qp031_external_schema", "lookup"),
309        );
310    }
311}