query/optimizer/
string_normalization.rs1use 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#[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 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}