Skip to main content

query/optimizer/
json_type_concretize.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::any::Any;
16use std::collections::HashMap;
17
18use common_function::scalars::json::json_get::{JsonGetWithType, parse_json_get_path};
19use datafusion::datasource::{DefaultTableSource, TableProvider};
20use datafusion_common::tree_node::{Transformed, TreeNode, TreeNodeRecursion};
21use datafusion_common::{Result, plan_datafusion_err, plan_err};
22use datafusion_expr::{Expr, LogicalPlan};
23use datafusion_optimizer::{OptimizerConfig, OptimizerRule};
24use datatypes::extension::json::is_json2_extension_type;
25use datatypes::types::json_type::{JsonNativeType, JsonObjectType};
26use jsonb::jsonpath::Path;
27use table::table::adapter::DfTableProviderAdapter;
28
29use crate::dummy_catalog::DummyTableProvider;
30
31/// Concretize (deduce) the expected JSON type from query.
32///
33/// For example, we can concretize a JSON type of `{ a: { b: Number } }` from
34/// `select j.a.b::Int64`. The JSON type will be later set into the scan request,
35/// for converting the JSON arrays.
36#[derive(Debug)]
37pub(crate) struct JsonTypeConcretizeRule;
38
39impl OptimizerRule for JsonTypeConcretizeRule {
40    fn name(&self) -> &str {
41        "JsonTypeConcretizeRule"
42    }
43
44    fn rewrite(
45        &self,
46        plan: LogicalPlan,
47        _config: &dyn OptimizerConfig,
48    ) -> Result<Transformed<LogicalPlan>> {
49        let json_types = deduce_json_types(&plan)?;
50        if json_types.is_empty() {
51            return Ok(Transformed::no(plan));
52        }
53
54        plan.transform_down(|plan| match &plan {
55            LogicalPlan::TableScan(table_scan) => {
56                let Some(source) = table_scan.source.downcast_ref::<DefaultTableSource>() else {
57                    return Ok(Transformed::no(plan));
58                };
59
60                if apply_json_type_hint(source.table_provider.as_ref(), &json_types) {
61                    Ok(Transformed::yes(plan))
62                } else {
63                    Ok(Transformed::no(plan))
64                }
65            }
66            _ => Ok(Transformed::no(plan)),
67        })
68    }
69}
70
71// FIXME: `json_types` is keyed only by unqualified column name. In joins with
72// same-named JSON2 columns, a hint deduced from one scan can be applied to
73// another scan. Carry the originating relation/scan when deducing hints.
74/// Applies JSON type hints to providers that can carry scan request hints.
75///
76/// Returns `true` if at least one JSON2 hint is retained and written to the provider.
77fn apply_json_type_hint(
78    provider: &dyn TableProvider,
79    json_types: &HashMap<String, JsonNativeType>,
80) -> bool {
81    let schema = provider.schema();
82    let json_types = json_types
83        .iter()
84        .filter(|(column, _)| {
85            schema
86                .fields()
87                .iter()
88                .any(|field| field.name() == *column && is_json2_extension_type(field))
89        })
90        .map(|(column, json_type)| (column.clone(), json_type.clone()))
91        .collect::<HashMap<_, _>>();
92
93    if json_types.is_empty() {
94        return false;
95    }
96
97    if let Some(adapter) = (provider as &dyn Any).downcast_ref::<DummyTableProvider>() {
98        adapter.with_json_type_hint(json_types);
99        return true;
100    }
101
102    if let Some(adapter) = (provider as &dyn Any).downcast_ref::<DfTableProviderAdapter>() {
103        adapter.with_json_type_hint(json_types);
104        return true;
105    }
106
107    false
108}
109
110pub(crate) fn deduce_json_types(plan: &LogicalPlan) -> Result<HashMap<String, JsonNativeType>> {
111    let mut json_types = HashMap::<String, JsonNativeType>::new();
112    // JSON2 columns in the final output must retain their complete values even when
113    // predicates or other expressions access only specific paths.
114    // For example, `SELECT j FROM t WHERE json_get(j, 'a') = 1`.
115    plan.schema()
116        .fields()
117        .iter()
118        .filter(|field| is_json2_extension_type(field))
119        .for_each(|field| {
120            json_types.insert(field.name().clone(), JsonNativeType::Variant);
121        });
122
123    plan.apply(|plan| {
124        for expr in plan.expressions() {
125            // Optimizer-generated projections may keep the JSON root only so later json_get
126            // expressions can access another path. A same-name pass-through does not require the
127            // complete root by itself; any real whole-column consumer above it is visited
128            // separately, and a whole root in the final output is captured from the plan schema.
129            if matches!(plan, LogicalPlan::Projection(_)) && is_same_name_column_projection(&expr) {
130                continue;
131            }
132            expr.apply(|expr| {
133                if let Some((column, json_type)) = deduce_json_type(expr)? {
134                    json_types.entry(column).or_default().merge(&json_type);
135                    Ok(TreeNodeRecursion::Jump)
136                } else {
137                    Ok(TreeNodeRecursion::Continue)
138                }
139            })?;
140        }
141        Ok(TreeNodeRecursion::Continue)
142    })?;
143    Ok(json_types)
144}
145
146fn is_same_name_column_projection(expr: &Expr) -> bool {
147    match expr {
148        Expr::Column(_) => true,
149        Expr::Alias(alias) => {
150            matches!(alias.expr.as_ref(), Expr::Column(column) if column.name == alias.name)
151        }
152        _ => false,
153    }
154}
155
156fn deduce_json_type(expr: &Expr) -> Result<Option<(String, JsonNativeType)>> {
157    let f = match expr {
158        Expr::ScalarFunction(f) if f.name().eq_ignore_ascii_case(JsonGetWithType::NAME) => f,
159        Expr::Column(c) => return Ok(Some((c.name.clone(), JsonNativeType::Variant))),
160        _ => return Ok(None),
161    };
162
163    let Some(Expr::Column(column)) = f.args.first() else {
164        return plan_err!(
165            "First argument of {} is expected to be a column expr, actual: {:?}",
166            JsonGetWithType::NAME,
167            f.args.first()
168        );
169    };
170
171    let Some(path) = json_get_path(f) else {
172        return plan_err!(
173            "Second argument of {} is expected to be a string literal, actual: {:?}",
174            JsonGetWithType::NAME,
175            f.args.get(1)
176        );
177    };
178
179    let json_path = parse_json_get_path(path)
180        .map_err(|e| plan_datafusion_err!("Invalid JSONPath {path:?}: {e}"))?;
181
182    if json_path
183        .paths
184        .iter()
185        .all(|segment| matches!(segment, Path::Root))
186    {
187        return Ok(Some((column.name.clone(), JsonNativeType::String)));
188    }
189
190    let with_type = f
191        .args
192        .get(2)
193        .and_then(|expr| expr.as_literal())
194        .map(|x| x.data_type())
195        .map(|with_type| {
196            JsonNativeType::try_from(&with_type).map_err(|e| plan_datafusion_err!("{e:?}"))
197        })
198        .transpose()?
199        .unwrap_or(JsonNativeType::String);
200
201    let mut root = with_type;
202    for segment in json_path.paths.into_iter().rev() {
203        let name = match segment {
204            Path::Root => continue,
205            Path::DotField(name) | Path::ColonField(name) | Path::ObjectField(name) => name,
206            // A full JSONPath expression can select arrays or use filters/wildcards.
207            // Keep the entire value when an object projection cannot represent it.
208            _ => return Ok(Some((column.name.clone(), JsonNativeType::Variant))),
209        };
210        let mut object = JsonObjectType::new();
211        object.insert(name.into_owned(), root);
212        root = JsonNativeType::Object(object);
213    }
214
215    Ok(Some((column.name.clone(), root)))
216}
217
218/// Returns the literal JSON path argument of a `json_get` call.
219fn json_get_path(function: &datafusion_expr::expr::ScalarFunction) -> Option<&str> {
220    function
221        .args
222        .get(1)
223        .and_then(|expr| expr.as_literal())
224        .and_then(|value| value.try_as_str())
225        .flatten()
226}
227
228#[cfg(test)]
229mod tests {
230    use std::sync::Arc;
231
232    use api::v1::SemanticType;
233    use arrow_schema::DataType;
234    use common_function::scalars::udf::create_udf;
235    use datafusion::datasource::provider_as_source;
236    use datafusion::functions_aggregate::expr_fn::count;
237    use datafusion_common::{Column, ScalarValue};
238    use datafusion_expr::expr::ScalarFunction;
239    use datafusion_expr::{LogicalPlanBuilder, col, lit};
240    use datafusion_optimizer::OptimizerContext;
241    use datatypes::extension::json::{Json2ExtensionType, JsonMetadata};
242    use datatypes::json::JsonSettings;
243    use datatypes::schema::ColumnSchema;
244    use store_api::metadata::{ColumnMetadata, RegionMetadataBuilder};
245    use store_api::storage::{ConcreteDataType, RegionId};
246
247    use super::*;
248    use crate::optimizer::test_util::{MetaRegionEngine, mock_table_provider};
249
250    fn json_get_expr(base: Expr, path: Expr, with_type: Option<DataType>) -> Result<Expr> {
251        let json_get = Arc::new(create_udf(Arc::new(JsonGetWithType::default())));
252        let mut args = vec![base, path];
253        if let Some(with_type) = with_type {
254            let with_type = ScalarValue::try_new_null(&with_type)?;
255            args.push(Expr::Literal(with_type, None));
256        }
257        Ok(Expr::ScalarFunction(ScalarFunction::new_udf(
258            json_get, args,
259        )))
260    }
261
262    fn path_expr(path: &str) -> Expr {
263        Expr::Literal(ScalarValue::Utf8(Some(path.to_string())), None)
264    }
265
266    fn build_plan(exprs: Vec<Expr>) -> Result<(Arc<DummyTableProvider>, LogicalPlan)> {
267        let provider = Arc::new(mock_table_provider(RegionId::new(1024, 1)));
268        let plan = LogicalPlanBuilder::scan("t", provider_as_source(provider.clone()), None)?
269            .project(exprs)?
270            .build()?;
271        Ok((provider, plan))
272    }
273
274    fn build_json2_scan() -> Result<(Arc<DummyTableProvider>, LogicalPlanBuilder)> {
275        build_json2_scan_with_settings(JsonSettings::default())
276    }
277
278    fn build_json2_scan_with_settings(
279        settings: JsonSettings,
280    ) -> Result<(Arc<DummyTableProvider>, LogicalPlanBuilder)> {
281        let region_id = RegionId::new(1024, 2);
282        let mut builder = RegionMetadataBuilder::new(region_id);
283        let mut json_column = ColumnSchema::new(
284            "j",
285            ConcreteDataType::json2(JsonNativeType::Object(JsonObjectType::new())),
286            true,
287        );
288        json_column.with_extension_type(&Json2ExtensionType::new(Arc::new(JsonMetadata::new(
289            settings,
290        ))));
291        builder
292            .push_column_metadata(ColumnMetadata {
293                column_schema: json_column,
294                semantic_type: SemanticType::Field,
295                column_id: 1,
296            })
297            .push_column_metadata(ColumnMetadata {
298                column_schema: ColumnSchema::new(
299                    "ts",
300                    ConcreteDataType::timestamp_millisecond_datatype(),
301                    false,
302                ),
303                semantic_type: SemanticType::Timestamp,
304                column_id: 2,
305            });
306        let metadata = Arc::new(builder.build().unwrap());
307        let engine = Arc::new(MetaRegionEngine::with_metadata(metadata.clone()));
308        let provider = Arc::new(DummyTableProvider::new(region_id, engine, metadata));
309        let plan = LogicalPlanBuilder::scan("t", provider_as_source(provider.clone()), None)?;
310        Ok((provider, plan))
311    }
312
313    fn build_json2_plan(exprs: Vec<Expr>) -> Result<(Arc<DummyTableProvider>, LogicalPlan)> {
314        let (provider, plan) = build_json2_scan()?;
315        let plan = plan.project(exprs)?.build()?;
316        Ok((provider, plan))
317    }
318
319    #[test]
320    fn test_json_type_concretize_rule_rewrite() -> Result<()> {
321        let exprs = vec![
322            json_get_expr(col("j"), path_expr("a.b"), Some(DataType::Int64))?.alias("ab"),
323            json_get_expr(col("j"), path_expr("a.c"), None)?.alias("ac"),
324            json_get_expr(col("j"), path_expr("d"), Some(DataType::Boolean))?.alias("d"),
325        ];
326        let (provider, plan) = build_json2_plan(exprs)?;
327
328        assert!(
329            JsonTypeConcretizeRule
330                .rewrite(plan, &OptimizerContext::default())?
331                .transformed
332        );
333
334        let expected = JsonNativeType::Object(JsonObjectType::from([
335            (
336                "a".to_string(),
337                JsonNativeType::Object(JsonObjectType::from([
338                    ("b".to_string(), JsonNativeType::i64()),
339                    ("c".to_string(), JsonNativeType::String),
340                ])),
341            ),
342            ("d".to_string(), JsonNativeType::Bool),
343        ]));
344
345        let request = provider.scan_request();
346        assert_eq!(1, request.json_type_hint.len());
347        assert_eq!(Some(&expected), request.json_type_hint.get("j"));
348        Ok(())
349    }
350
351    #[test]
352    fn test_deduce_json_type_object_paths() -> Result<()> {
353        for paths in [
354            ["a.b", r#"$."a"."b""#, r#"["a"]["b"]"#],
355            [r#"$."a.b"."c.d""#, r#"["a.b"]["c.d"]"#, r#"$."a.b"["c.d"]"#],
356        ] {
357            let deduce = |path| {
358                deduce_json_type(&json_get_expr(
359                    col("j"),
360                    path_expr(path),
361                    Some(DataType::Int64),
362                )?)
363            };
364            let expected = deduce(paths[0])?;
365            assert!(!matches!(expected, Some((_, JsonNativeType::Variant))));
366            for path in &paths[1..] {
367                assert_eq!(deduce(path)?, expected);
368            }
369        }
370        Ok(())
371    }
372
373    #[test]
374    fn test_deduce_json_type_invalid_path() -> Result<()> {
375        let expr = json_get_expr(col("j"), path_expr("$.a["), Some(DataType::Int64))?;
376        let err = deduce_json_type(&expr).unwrap_err();
377        assert!(err.to_string().contains("Invalid JSONPath"), "{err}");
378        Ok(())
379    }
380
381    #[test]
382    fn test_deduce_json_type_with_list_index() -> Result<()> {
383        for path in [
384            "l[0]",
385            "$.l[0]",
386            "$.l[*]",
387            "$.o.*",
388            "$.l[0 to 2]",
389            "$.l ? (@.a == 1)",
390        ] {
391            let expr = json_get_expr(col("j"), path_expr(path), Some(DataType::Int64))?;
392            assert_eq!(
393                Some(("j".to_string(), JsonNativeType::Variant)),
394                deduce_json_type(&expr)?,
395                "{path}"
396            );
397        }
398        Ok(())
399    }
400
401    #[test]
402    fn test_json_type_concretize_rule_conflict_to_variant() -> Result<()> {
403        let exprs = vec![
404            json_get_expr(col("j"), path_expr("a"), Some(DataType::Int64))?.alias("a_num"),
405            json_get_expr(col("j"), path_expr("a.b"), Some(DataType::Boolean))?.alias("a_obj"),
406        ];
407        let (provider, plan) = build_json2_plan(exprs)?;
408
409        assert!(
410            JsonTypeConcretizeRule
411                .rewrite(plan, &OptimizerContext::default())?
412                .transformed
413        );
414
415        let expected = JsonNativeType::Object(JsonObjectType::from([(
416            "a".to_string(),
417            JsonNativeType::Variant,
418        )]));
419        assert_eq!(
420            Some(&expected),
421            provider.scan_request().json_type_hint.get("j")
422        );
423        Ok(())
424    }
425
426    #[test]
427    fn test_json_type_concretize_rule_ignores_non_json2_columns() -> Result<()> {
428        let exprs =
429            vec![json_get_expr(col("k0"), path_expr("a.b"), Some(DataType::Int64))?.alias("ab")];
430        let (provider, plan) = build_plan(exprs)?;
431
432        assert!(
433            !JsonTypeConcretizeRule
434                .rewrite(plan, &OptimizerContext::default())?
435                .transformed
436        );
437        assert!(provider.scan_request().json_type_hint.is_empty());
438        Ok(())
439    }
440
441    #[test]
442    fn test_json_type_concretize_rule_no_json_get() -> Result<()> {
443        let (provider, plan) = build_plan(vec![col("k0"), col("v0")])?;
444
445        assert!(
446            !JsonTypeConcretizeRule
447                .rewrite(plan, &OptimizerContext::default())?
448                .transformed
449        );
450        assert!(provider.scan_request().json_type_hint.is_empty());
451        Ok(())
452    }
453
454    #[test]
455    fn test_allow_json2_path_use_in_intermediate_plan() -> Result<()> {
456        let json_get = json_get_expr(col("j"), path_expr("a"), Some(DataType::Int64))?;
457        let (provider, plan) = build_json2_scan()?;
458        let plan = plan
459            .aggregate(vec![json_get], Vec::<Expr>::new())?
460            .aggregate(Vec::<Expr>::new(), vec![count(lit(1))])?
461            .build()?;
462
463        assert!(
464            JsonTypeConcretizeRule
465                .rewrite(plan, &OptimizerContext::default())?
466                .transformed
467        );
468        assert_eq!(
469            Some(&JsonNativeType::Object(JsonObjectType::from([(
470                "a".to_string(),
471                JsonNativeType::i64(),
472            )]))),
473            provider.scan_request().json_type_hint.get("j")
474        );
475        Ok(())
476    }
477
478    #[test]
479    fn test_allow_json2_projection_by_path() -> Result<()> {
480        let expr = json_get_expr(col("j"), path_expr("a"), Some(DataType::Int64))?;
481        let (provider, plan) = build_json2_plan(vec![expr])?;
482
483        assert!(
484            JsonTypeConcretizeRule
485                .rewrite(plan, &OptimizerContext::default())?
486                .transformed
487        );
488        assert_eq!(
489            Some(&JsonNativeType::Object(JsonObjectType::from([(
490                "a".to_string(),
491                JsonNativeType::i64(),
492            )]))),
493            provider.scan_request().json_type_hint.get("j")
494        );
495        Ok(())
496    }
497
498    #[test]
499    fn test_allow_json2_filter_with_root_projection() -> Result<()> {
500        let predicate =
501            json_get_expr(col("j"), path_expr("a"), Some(DataType::Int64))?.eq(lit(1_i64));
502        let (provider, plan) = build_json2_scan()?;
503        let plan = plan.filter(predicate)?.build()?;
504
505        assert!(
506            JsonTypeConcretizeRule
507                .rewrite(plan, &OptimizerContext::default())?
508                .transformed
509        );
510        assert_eq!(
511            Some(&JsonNativeType::Variant),
512            provider.scan_request().json_type_hint.get("j")
513        );
514        Ok(())
515    }
516
517    #[test]
518    fn test_deduce_json_type_with_non_column_base() -> Result<()> {
519        let expr = json_get_expr(
520            Expr::Literal(ScalarValue::Utf8(Some("{}".to_string())), None),
521            path_expr("a"),
522            Some(DataType::Int64),
523        )?;
524
525        let err = deduce_json_type(&expr).unwrap_err();
526        assert!(
527            err.to_string()
528                .contains("First argument of json_get is expected to be a column expr")
529        );
530        Ok(())
531    }
532
533    #[test]
534    fn test_deduce_json_type_with_non_literal_path() -> Result<()> {
535        let expr = json_get_expr(
536            Expr::Column(Column::new_unqualified("k0")),
537            Expr::Column(Column::new_unqualified("path_col")),
538            Some(DataType::Int64),
539        )?;
540
541        let err = deduce_json_type(&expr).unwrap_err();
542        assert!(
543            err.to_string()
544                .contains("Second argument of json_get is expected to be a string literal")
545        );
546        Ok(())
547    }
548
549    #[test]
550    fn test_deduce_json_type_default_string() -> Result<()> {
551        let expr = json_get_expr(
552            Expr::Column(Column::new_unqualified("k0")),
553            path_expr("a.b"),
554            None,
555        )?;
556
557        let deduced = deduce_json_type(&expr)?;
558        let expected = JsonNativeType::Object(JsonObjectType::from([(
559            "a".to_string(),
560            JsonNativeType::Object(JsonObjectType::from([(
561                "b".to_string(),
562                JsonNativeType::String,
563            )])),
564        )]));
565
566        assert_eq!(Some(("k0".to_string(), expected)), deduced);
567        Ok(())
568    }
569}