common_function/scalars/json/
json_get_rewriter.rs1#[cfg(test)]
16use std::sync::Arc;
17
18use arrow_schema::{DataType, TimeUnit};
19use datafusion::common::config::ConfigOptions;
20use datafusion::common::tree_node::Transformed;
21use datafusion::common::{DFSchema, Result};
22use datafusion::logical_expr::expr_rewriter::FunctionRewrite;
23use datafusion::scalar::ScalarValue;
24use datafusion_expr::expr::ScalarFunction;
25use datafusion_expr::{Cast, Expr};
26
27use crate::scalars::json::JsonGetWithType;
28
29#[derive(Debug)]
30pub struct JsonGetRewriter;
31
32impl FunctionRewrite for JsonGetRewriter {
33 fn name(&self) -> &'static str {
34 "JsonGetRewriter"
35 }
36
37 fn rewrite(
38 &self,
39 expr: Expr,
40 _schema: &DFSchema,
41 _config: &ConfigOptions,
42 ) -> Result<Transformed<Expr>> {
43 Ok(match expr {
44 Expr::Cast(cast) => inject_type_from_cast_expr(cast)?,
45 Expr::ScalarFunction(cast) => inject_type_from_cast_func(cast)?,
46 expr => Transformed::no(expr),
47 })
48 }
49}
50
51fn inject_type_from_cast_expr(cast: Cast) -> Result<Transformed<Expr>> {
62 let Cast { expr, field } = cast;
63 let mut data_type = field.data_type().clone();
64
65 let mut json_get = match *expr {
66 Expr::ScalarFunction(f)
67 if f.func.name().eq_ignore_ascii_case(JsonGetWithType::NAME) && f.args.len() == 2 =>
68 {
69 f
70 }
71 expr => {
72 return Ok(Transformed::no(Expr::Cast(Cast {
73 expr: Box::new(expr),
74 field,
75 })));
76 }
77 };
78
79 if data_type.is_string() {
80 data_type = DataType::Utf8View;
81 }
82 let with_type = ScalarValue::try_new_null(&data_type).map(|x| Expr::Literal(x, None))?;
83 json_get.args.push(with_type);
84 Ok(Transformed::yes(Expr::ScalarFunction(json_get)))
85}
86
87fn inject_type_from_cast_func(cast: ScalarFunction) -> Result<Transformed<Expr>> {
100 let ScalarFunction { func, args } = cast;
101
102 let func_name = func.name().to_ascii_lowercase();
105 if !func_name.contains("arrow_cast") {
106 let original = Expr::ScalarFunction(ScalarFunction { func, args });
107 return Ok(Transformed::no(original));
108 }
109
110 if args.len() != 2 {
114 let original = Expr::ScalarFunction(ScalarFunction { func, args });
115 return Ok(Transformed::no(original));
116 }
117 let [arg0, arg1] = args.try_into().unwrap_or_else(|_| unreachable!());
118
119 let Some(with_type) = arg1
120 .as_literal()
121 .and_then(|x| x.try_as_str())
122 .flatten()
123 .and_then(parse_data_type_from_string)
124 else {
125 let original = Expr::ScalarFunction(ScalarFunction {
126 func,
127 args: vec![arg0, arg1],
128 });
129 return Ok(Transformed::no(original));
130 };
131
132 let mut json_get = match arg0 {
133 Expr::ScalarFunction(f)
134 if f.func.name().eq_ignore_ascii_case(JsonGetWithType::NAME) && f.args.len() == 2 =>
135 {
136 f
137 }
138 arg0 => {
139 let original = Expr::ScalarFunction(ScalarFunction {
140 func,
141 args: vec![arg0, arg1],
142 });
143 return Ok(Transformed::no(original));
144 }
145 };
146
147 let with_type = ScalarValue::try_new_null(&with_type).map(|x| Expr::Literal(x, None))?;
148 json_get.args.push(with_type);
149
150 let rewritten = Expr::ScalarFunction(json_get);
151 Ok(Transformed::yes(rewritten))
152}
153
154fn parse_data_type_from_string(type_str: &str) -> Option<DataType> {
156 match type_str.to_lowercase().as_str() {
157 "int8" | "tinyint" => Some(DataType::Int8),
158 "int16" | "smallint" => Some(DataType::Int16),
159 "int32" | "integer" => Some(DataType::Int32),
160 "int64" | "bigint" => Some(DataType::Int64),
161 "uint8" => Some(DataType::UInt8),
162 "uint16" => Some(DataType::UInt16),
163 "uint32" => Some(DataType::UInt32),
164 "uint64" => Some(DataType::UInt64),
165 "float32" | "real" => Some(DataType::Float32),
166 "float64" | "double" => Some(DataType::Float64),
167 "boolean" | "bool" => Some(DataType::Boolean),
168 "string" | "text" | "varchar" => Some(DataType::Utf8),
169 "timestamp" => Some(DataType::Timestamp(TimeUnit::Microsecond, None)),
170 "date" => Some(DataType::Date32),
171 _ => None,
172 }
173}
174
175#[cfg(test)]
176mod tests {
177 use arrow_schema::DataType;
178 use datafusion::common::DFSchema;
179 use datafusion::common::config::ConfigOptions;
180 use datafusion::logical_expr::expr::Cast;
181 use datafusion::scalar::ScalarValue;
182 use datafusion_expr::Expr;
183 use datafusion_expr::expr::ScalarFunction;
184
185 use super::*;
186
187 #[test]
188 fn test_rewrite_regular_cast() {
189 let rewriter = JsonGetRewriter;
190 let schema = DFSchema::empty();
191 let config = ConfigOptions::new();
192
193 let json_expr = Expr::ScalarFunction(ScalarFunction {
195 func: Arc::new(crate::scalars::udf::create_udf(Arc::new(
196 crate::scalars::json::JsonGetWithType::default(),
197 ))),
198 args: vec![
199 Expr::Literal(ScalarValue::Utf8(Some("{\"a\":1}".to_string())), None),
200 Expr::Literal(ScalarValue::Utf8(Some("$.a".to_string())), None),
201 ],
202 });
203
204 let cast_expr = Expr::Cast(Cast::new(Box::new(json_expr), DataType::Int8));
206
207 let result = rewriter.rewrite(cast_expr, &schema, &config).unwrap();
209
210 assert!(result.transformed);
212
213 match result.data {
215 Expr::ScalarFunction(func) => {
216 assert_eq!(func.args.len(), 3);
218
219 match &func.args[0] {
221 Expr::Literal(ScalarValue::Utf8(Some(json)), _) => {
222 assert_eq!(json, "{\"a\":1}");
223 }
224 _ => panic!("First argument should be a string literal"),
225 }
226
227 match &func.args[1] {
229 Expr::Literal(ScalarValue::Utf8(Some(path)), _) => {
230 assert_eq!(path, "$.a");
231 }
232 _ => panic!("Second argument should be a string literal"),
233 }
234
235 match &func.args[2] {
237 Expr::Literal(value, _) => {
238 assert_eq!(value.data_type(), DataType::Int8);
239 }
240 _ => panic!("Third argument should be a cast expression"),
241 }
242 }
243 _ => panic!("Result should be a ScalarFunction"),
244 }
245 }
246
247 #[test]
248 fn test_rewrite_arrow_cast_function() {
249 let rewriter = JsonGetRewriter;
250 let schema = DFSchema::empty();
251 let config = ConfigOptions::new();
252
253 let parse_json_expr = Expr::ScalarFunction(ScalarFunction {
255 func: Arc::new(crate::scalars::udf::create_udf(Arc::new(
256 crate::scalars::json::ParseJsonFunction::default(),
257 ))),
258 args: vec![Expr::Literal(
259 ScalarValue::Utf8(Some("{\"a\":1}".to_string())),
260 None,
261 )],
262 });
263
264 let json_get_expr = Expr::ScalarFunction(ScalarFunction {
266 func: Arc::new(crate::scalars::udf::create_udf(Arc::new(
267 crate::scalars::json::JsonGetWithType::default(),
268 ))),
269 args: vec![
270 parse_json_expr,
271 Expr::Literal(ScalarValue::Utf8(Some("a".to_string())), None),
272 ],
273 });
274
275 let arrow_cast_expr = Expr::Cast(Cast::new(Box::new(json_get_expr), DataType::Int64));
278
279 let result = rewriter.rewrite(arrow_cast_expr, &schema, &config).unwrap();
281
282 assert!(result.transformed);
284
285 match result.data {
287 Expr::ScalarFunction(func) => {
288 assert_eq!(func.args.len(), 3);
290
291 match &func.args[0] {
293 Expr::ScalarFunction(parse_func) => {
294 assert!(
296 parse_func
297 .func
298 .name()
299 .to_ascii_lowercase()
300 .contains("parse_json")
301 );
302 assert_eq!(parse_func.args.len(), 1);
303 match &parse_func.args[0] {
304 Expr::Literal(ScalarValue::Utf8(Some(json)), _) => {
305 assert_eq!(json, "{\"a\":1}");
306 }
307 _ => panic!("Parse json argument should be a string literal"),
308 }
309 }
310 _ => panic!("First argument should be a parse_json function"),
311 }
312
313 match &func.args[1] {
315 Expr::Literal(ScalarValue::Utf8(Some(path)), _) => {
316 assert_eq!(path, "a");
317 }
318 _ => panic!("Second argument should be a string literal"),
319 }
320
321 match &func.args[2] {
323 Expr::Literal(value, _) => {
324 assert_eq!(value.data_type(), DataType::Int64);
325 }
326 _ => panic!("Third argument should be a cast expression"),
327 }
328 }
329 _ => panic!("Result should be a ScalarFunction"),
330 }
331 }
332
333 #[test]
334 fn test_no_rewrite_for_other_functions() {
335 let rewriter = JsonGetRewriter;
336 let schema = DFSchema::empty();
337 let config = ConfigOptions::new();
338
339 let other_func = Expr::ScalarFunction(ScalarFunction {
341 func: Arc::new(crate::scalars::udf::create_udf(Arc::new(
342 crate::scalars::test::TestAndFunction::default(),
343 ))),
344 args: vec![Expr::Literal(ScalarValue::Int64(Some(4)), None)],
345 });
346
347 let result = rewriter.rewrite(other_func, &schema, &config).unwrap();
349
350 assert!(!result.transformed);
352 }
353
354 #[test]
355 fn test_no_rewrite_for_non_cast_functions() {
356 let rewriter = JsonGetRewriter;
357 let schema = DFSchema::empty();
358 let config = ConfigOptions::new();
359
360 let other_func = Expr::ScalarFunction(ScalarFunction {
362 func: Arc::new(crate::scalars::udf::create_udf(Arc::new(
363 crate::scalars::test::TestAndFunction::default(),
364 ))),
365 args: vec![
366 Expr::ScalarFunction(ScalarFunction {
367 func: Arc::new(crate::scalars::udf::create_udf(Arc::new(
368 crate::scalars::json::JsonGetWithType::default(),
369 ))),
370 args: vec![
371 Expr::Literal(ScalarValue::Utf8(Some("{\"a\":1}".to_string())), None),
372 Expr::Literal(ScalarValue::Utf8(Some("$.a".to_string())), None),
373 ],
374 }),
375 Expr::Literal(ScalarValue::Utf8(Some("Int64".to_string())), None),
376 ],
377 });
378
379 let result = rewriter.rewrite(other_func, &schema, &config).unwrap();
381
382 assert!(!result.transformed);
384 }
385}