1use 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#[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
71fn 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 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 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 _ => 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
218fn 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}