Skip to main content

common_function/scalars/json/
json_get_rewriter.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
15#[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
51// Expr::Cast(
52//   Expr::ScalarFunction(
53//     json_get(column, path),
54//     <data_type>
55//   )
56// )
57// =>
58// Expr::ScalarFunction(
59//   json_get(column, path, <data_type>)
60// )
61fn 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
87// Expr::ScalarFunction(
88//   arrow_cast(
89//     Expr::ScalarFunction(
90//       json_get(column, path),
91//     ),
92//     <data_type>
93//   )
94// )
95// =>
96// Expr::ScalarFunction(
97//   json_get(column, path, <data_type>)
98// )
99fn inject_type_from_cast_func(cast: ScalarFunction) -> Result<Transformed<Expr>> {
100    let ScalarFunction { func, args } = cast;
101
102    // Check if this is an Arrow cast function
103    // The function name might be "arrow_cast" or similar
104    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    // Arrow cast function should have exactly 2 arguments:
111    // 1. The expression to cast (could be json_get)
112    // 2. The target type as a string literal
113    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
154// Parse a data type from a string representation
155fn 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        // Create a json_get function
194        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        // Create a cast expression: json_get(...)::int8
205        let cast_expr = Expr::Cast(Cast::new(Box::new(json_expr), DataType::Int8));
206
207        // Apply the rewriter
208        let result = rewriter.rewrite(cast_expr, &schema, &config).unwrap();
209
210        // Verify the result is transformed
211        assert!(result.transformed);
212
213        // Verify the result is a ScalarFunction
214        match result.data {
215            Expr::ScalarFunction(func) => {
216                // Should have 3 arguments now (original 2 + null cast)
217                assert_eq!(func.args.len(), 3);
218
219                // First argument should be the original json
220                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                // Second argument should be the path
228                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                // Third argument should be a null cast to Int8
236                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        // Create a parse_json function
254        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        // Create a json_get function
265        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        // Create an arrow cast function: cast(json_get(...), 'Int64')
276        // Note: ArrowCastFunc doesn't exist in this codebase, so this test uses a simple cast instead
277        let arrow_cast_expr = Expr::Cast(Cast::new(Box::new(json_get_expr), DataType::Int64));
278
279        // Apply the rewriter
280        let result = rewriter.rewrite(arrow_cast_expr, &schema, &config).unwrap();
281
282        // Verify the result is transformed
283        assert!(result.transformed);
284
285        // Verify the result is a ScalarFunction (json_get_with_type)
286        match result.data {
287            Expr::ScalarFunction(func) => {
288                // Should have 3 arguments now (original 2 + null cast)
289                assert_eq!(func.args.len(), 3);
290
291                // First argument should be the original parse_json function
292                match &func.args[0] {
293                    Expr::ScalarFunction(parse_func) => {
294                        // Verify it's a parse_json function with the right argument
295                        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                // Second argument should be the path
314                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                // Third argument should be a null cast to Int64
322                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        // Create a non-json function
340        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        // Apply the rewriter
348        let result = rewriter.rewrite(other_func, &schema, &config).unwrap();
349
350        // Verify the result is not transformed
351        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        // Create a scalar function that doesn't contain "cast"
361        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        // Apply the rewriter
380        let result = rewriter.rewrite(other_func, &schema, &config).unwrap();
381
382        // Verify the result is not transformed
383        assert!(!result.transformed);
384    }
385}