Skip to main content

query/optimizer/
string_normalization.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 arrow_schema::DataType;
16use datafusion::config::ConfigOptions;
17use datafusion::logical_expr::expr::Cast;
18use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRewriter};
19use datafusion_common::{Result, ScalarValue};
20use datafusion_expr::{Expr, LogicalPlan};
21use datafusion_optimizer::analyzer::AnalyzerRule;
22
23use crate::plan::ExtractExpr;
24
25/// StringNormalizationRule normalizes(trims) string values in logical plan.
26/// Mainly used for timestamp trimming
27#[derive(Debug)]
28pub struct StringNormalizationRule;
29
30impl AnalyzerRule for StringNormalizationRule {
31    fn analyze(&self, plan: LogicalPlan, _config: &ConfigOptions) -> Result<LogicalPlan> {
32        plan.transform(|plan| match plan {
33            LogicalPlan::Projection(_)
34            | LogicalPlan::Filter(_)
35            | LogicalPlan::Window(_)
36            | LogicalPlan::Aggregate(_)
37            | LogicalPlan::Sort(_)
38            | LogicalPlan::Join(_)
39            | LogicalPlan::Repartition(_)
40            | LogicalPlan::Union(_)
41            | LogicalPlan::TableScan(_)
42            | LogicalPlan::EmptyRelation(_)
43            | LogicalPlan::Subquery(_)
44            | LogicalPlan::SubqueryAlias(_)
45            | LogicalPlan::Statement(_)
46            | LogicalPlan::Values(_)
47            | LogicalPlan::Analyze(_)
48            | LogicalPlan::Extension(_)
49            | LogicalPlan::Dml(_)
50            | LogicalPlan::Copy(_)
51            | LogicalPlan::RecursiveQuery(_) => {
52                let mut converter = StringNormalizationConverter;
53                let inputs = plan.inputs().into_iter().cloned().collect::<Vec<_>>();
54                let expr = plan
55                    .expressions_consider_join()
56                    .into_iter()
57                    .map(|e| e.rewrite(&mut converter).map(|x| x.data))
58                    .collect::<Result<Vec<_>>>()?;
59                if expr != plan.expressions_consider_join() {
60                    plan.with_new_exprs(expr, inputs).map(Transformed::yes)
61                } else {
62                    Ok(Transformed::no(plan))
63                }
64            }
65            LogicalPlan::Distinct(_)
66            | LogicalPlan::Limit(_)
67            | LogicalPlan::Explain(_)
68            | LogicalPlan::Unnest(_)
69            | LogicalPlan::Ddl(_)
70            | LogicalPlan::DescribeTable(_) => Ok(Transformed::no(plan)),
71        })
72        .map(|x| x.data)
73    }
74
75    fn name(&self) -> &str {
76        "StringNormalizationRule"
77    }
78}
79
80struct StringNormalizationConverter;
81
82impl TreeNodeRewriter for StringNormalizationConverter {
83    type Node = Expr;
84
85    /// remove extra whitespaces from the String value when
86    /// there is a CAST from a String to Timestamp.
87    /// Otherwise - no modifications applied
88    fn f_up(&mut self, expr: Expr) -> Result<Transformed<Expr>> {
89        let new_expr = match expr {
90            Expr::Cast(Cast { expr, field }) => {
91                let expr = match field.data_type() {
92                    DataType::Timestamp(_, _) => match *expr {
93                        Expr::Literal(value, _) => match value {
94                            ScalarValue::Utf8(Some(s)) => trim_utf_expr(s),
95                            _ => Expr::Literal(value, None),
96                        },
97                        expr => expr,
98                    },
99                    _ => *expr,
100                };
101                Expr::Cast(Cast::new_from_field(Box::new(expr), field))
102            }
103            expr => expr,
104        };
105        Ok(Transformed::yes(new_expr))
106    }
107}
108
109fn trim_utf_expr(s: String) -> Expr {
110    let parts: Vec<_> = s.split_whitespace().collect();
111    let trimmed = parts.join(" ");
112    Expr::Literal(ScalarValue::Utf8(Some(trimmed)), None)
113}
114
115#[cfg(test)]
116mod tests {
117    use std::sync::Arc;
118
119    use arrow::datatypes::TimeUnit::{Microsecond, Millisecond, Nanosecond, Second};
120    use arrow::datatypes::{DataType, SchemaRef};
121    use arrow_schema::{Field, Schema, TimeUnit};
122    use datafusion::datasource::{MemTable, provider_as_source};
123    use datafusion_common::config::ConfigOptions;
124    use datafusion_expr::{Cast, Expr, LogicalPlan, LogicalPlanBuilder, lit};
125    use datafusion_optimizer::analyzer::AnalyzerRule;
126
127    use crate::optimizer::string_normalization::StringNormalizationRule;
128
129    #[test]
130    fn test_normalization_for_string_with_extra_whitespaces_to_timestamp_cast() {
131        let timestamp_str_with_whitespaces = "    2017-07-23    13:10:11   ";
132        let config = &ConfigOptions::default();
133        let projects = vec![
134            create_timestamp_cast_project(Nanosecond, timestamp_str_with_whitespaces),
135            create_timestamp_cast_project(Microsecond, timestamp_str_with_whitespaces),
136            create_timestamp_cast_project(Millisecond, timestamp_str_with_whitespaces),
137            create_timestamp_cast_project(Second, timestamp_str_with_whitespaces),
138        ];
139        for (time_unit, proj) in projects {
140            let plan = create_test_plan_with_project(proj);
141            let result = StringNormalizationRule.analyze(plan, config).unwrap();
142            let expected = format!(
143                "Projection: CAST(Utf8(\"2017-07-23 13:10:11\") AS Timestamp({}))\n  TableScan: t",
144                time_unit
145            );
146            assert_eq!(expected, result.to_string());
147        }
148    }
149
150    #[test]
151    fn test_normalization_for_non_timestamp_casts() {
152        let config = &ConfigOptions::default();
153        let proj_int_to_timestamp = vec![Expr::Cast(Cast::new(
154            Box::new(lit(158412331400600000_i64)),
155            DataType::Timestamp(Nanosecond, None),
156        ))];
157        let int_to_timestamp_plan = create_test_plan_with_project(proj_int_to_timestamp);
158        let result = StringNormalizationRule
159            .analyze(int_to_timestamp_plan, config)
160            .unwrap();
161        let expected = String::from(
162            "Projection: CAST(Int64(158412331400600000) AS Timestamp(ns))\n  TableScan: t",
163        );
164        assert_eq!(expected, result.to_string());
165
166        let proj_string_to_int = vec![Expr::Cast(Cast::new(
167            Box::new(lit("  5   ")),
168            DataType::Int32,
169        ))];
170        let string_to_int_plan = create_test_plan_with_project(proj_string_to_int);
171        let result = StringNormalizationRule
172            .analyze(string_to_int_plan, &ConfigOptions::default())
173            .unwrap();
174        let expected = String::from("Projection: CAST(Utf8(\"  5   \") AS Int32)\n  TableScan: t");
175        assert_eq!(expected, result.to_string());
176    }
177
178    fn create_test_plan_with_project(proj: Vec<Expr>) -> LogicalPlan {
179        prepare_test_plan_builder()
180            .project(proj)
181            .unwrap()
182            .build()
183            .unwrap()
184    }
185
186    fn create_timestamp_cast_project(unit: TimeUnit, timestamp_str: &str) -> (TimeUnit, Vec<Expr>) {
187        let proj = vec![Expr::Cast(Cast::new(
188            Box::new(lit(timestamp_str)),
189            DataType::Timestamp(unit, None),
190        ))];
191        (unit, proj)
192    }
193
194    fn prepare_test_plan_builder() -> LogicalPlanBuilder {
195        let schema = Schema::new(vec![Field::new("f", DataType::Float64, false)]);
196        let table = MemTable::try_new(SchemaRef::from(schema), vec![vec![]]).unwrap();
197        LogicalPlanBuilder::scan("t", provider_as_source(Arc::new(table)), None).unwrap()
198    }
199}