Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 49 additions & 2 deletions src/rewrite.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ use datafusion::common::tree_node::Transformed;
use datafusion::common::Column;
use datafusion::common::DFSchema;
use datafusion::common::Result;
use datafusion::logical_expr::expr::{Alias, Cast, Expr, ScalarFunction};
use datafusion::logical_expr::expr::{Alias, Cast, Expr, ScalarFunction, TryCast};
use datafusion::logical_expr::expr_rewriter::FunctionRewrite;
use datafusion::logical_expr::planner::{ExprPlanner, PlannerResult, RawBinaryExpr};
use datafusion::logical_expr::sqlparser::ast::BinaryOperator;
Expand All @@ -32,7 +32,18 @@ impl FunctionRewrite for JsonFunctionRewriter {
}),
})
}
Expr::ScalarFunction(func) => unnest_json_calls(func),
Expr::TryCast(try_cast) => {
optimise_json_get(try_cast.field.data_type(), &try_cast.expr, schema).map(|folded| match folded {
Folded::Exact(accessor) => accessor,
Folded::Narrowing(accessor) => Expr::TryCast(TryCast {
expr: Box::new(accessor),
field: try_cast.field.clone(),
}),
})
}
Expr::ScalarFunction(func) => {
optimise_json_get_arrow_cast(func, schema).or_else(|| unnest_json_calls(func))
}
_ => None,
};
Ok(transform.map_or_else(|| Transformed::no(expr), Transformed::yes))
Expand All @@ -57,6 +68,9 @@ enum Folded {
/// always been. Where it is not, only the cast's *input* is replaced: dropping the cast there would
/// silently change the type of the expression, and for a narrowing cast its value too, while
/// keeping it costs one primitive-to-primitive cast and still never materializes the union.
///
/// `TRY_CAST` gets the same treatment, so a value outside the target type's range still becomes
/// NULL rather than being handed back as the accessor's wider type.
fn optimise_json_get(cast_to: &DataType, cast_expr: &Expr, schema: &DFSchema) -> Option<Folded> {
let scalar_func = extract_scalar_function(cast_expr)?;
if !is_json_get(scalar_func) {
Expand Down Expand Up @@ -91,6 +105,39 @@ fn typed_accessor(cast_to: &DataType) -> Option<Arc<ScalarUDF>> {
})
}

/// The same rewrite for `arrow_cast(json_get(foo, bar), 'Int64')` and its `arrow_try_cast` sibling.
///
/// These two are still scalar function calls when this rewriter runs. `DataFusion` lowers them to
/// `Expr::Cast` / `Expr::TryCast` in `SimplifyExpressions`, an optimizer rule, whereas function
/// rewrites are applied by `ApplyFunctionRewrites` at the start of the analyzer. The analyzer never
/// runs again afterwards, so without this the JSON union is materialized only to be cast away.
///
/// Only the call's first argument is replaced, so the named type is still what comes out. That
/// lowering then drops the cast by itself when the accessor already returns the named type.
fn optimise_json_get_arrow_cast(func: &ScalarFunction, schema: &DFSchema) -> Option<Expr> {
if !matches!(func.func.name(), "arrow_cast" | "arrow_try_cast") {
return None;
}
let [cast_expr, type_arg] = func.args.as_slice() else {
return None;
};
let Expr::Literal(ScalarValue::Utf8(Some(type_name)), _) = type_arg else {
return None;
};
// `arrow_cast` names its target as an Arrow type string, which is how DataFusion itself reads
// it back in `ArrowCastFunc::return_field_from_args`.
let cast_to = type_name.parse::<DataType>().ok()?;
// the call keeps its type argument either way, and `ArrowCastFunc::simplify` drops the cast
// itself when the accessor already returns that type
let accessor = match optimise_json_get(&cast_to, cast_expr, schema)? {
Folded::Exact(accessor) | Folded::Narrowing(accessor) => accessor,
};
Some(Expr::ScalarFunction(ScalarFunction {
func: func.func.clone(),
args: vec![accessor, type_arg.clone()],
}))
}

// Replace nested JSON functions e.g. `json_get(json_get(col, 'foo'), 'bar')` with `json_get(col, 'foo', 'bar')`
fn unnest_json_calls(func: &ScalarFunction) -> Option<Expr> {
if !matches!(
Expand Down
80 changes: 80 additions & 0 deletions tests/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1558,6 +1558,86 @@ async fn test_plan_arrow_cast_int() {
assert_eq!(lines, expected);
}

#[tokio::test]
async fn test_plan_try_cast_int() {
let lines = logical_plan(r"explain select try_cast((json_data->'foo') as bigint) from test").await;

let expected = [
"Projection: json_get_int(test.json_data, Utf8(\"foo\")) AS json_data -> 'foo'",
" TableScan: test projection=[json_data]",
];

assert_eq!(lines, expected);
}

#[tokio::test]
async fn test_try_cast_int() {
let sql = r#"select try_cast(('{"foo": 42}'->'foo') as bigint)"#;
let batches = run_query(sql).await.unwrap();
assert_eq!(display_val(batches).await, (DataType::Int64, "42".to_string()));
}

#[tokio::test]
async fn test_try_cast_wrong_type_is_null() {
// json_get_int yields NULL for a value that is not an integer, which is what TRY_CAST asks for
let sql = r#"select try_cast(('{"foo": "not an int"}'->'foo') as bigint)"#;
let batches = run_query(sql).await.unwrap();
assert_eq!(display_val(batches).await, (DataType::Int64, String::new()));
}

#[tokio::test]
async fn test_plan_arrow_cast_fn_int() {
let lines = logical_plan(r"explain select arrow_cast((json_data->'foo'), 'Int64') from test").await;

let expected = [
"Projection: json_get_int(test.json_data, Utf8(\"foo\")) AS arrow_cast(json_data -> 'foo',Utf8(\"Int64\"))",
" TableScan: test projection=[json_data]",
];

assert_eq!(lines, expected);
}

#[tokio::test]
async fn test_plan_arrow_try_cast_fn_int() {
let lines = logical_plan(r"explain select arrow_try_cast((json_data->'foo'), 'Int64') from test").await;

let expected = [
"Projection: json_get_int(test.json_data, Utf8(\"foo\")) AS arrow_try_cast(json_data -> 'foo',Utf8(\"Int64\"))",
" TableScan: test projection=[json_data]",
];

assert_eq!(lines, expected);
}

#[tokio::test]
async fn test_arrow_cast_fn_int() {
let sql = r#"select arrow_cast(('{"foo": 42}'->'foo'), 'Int64')"#;
let batches = run_query(sql).await.unwrap();
assert_eq!(display_val(batches).await, (DataType::Int64, "42".to_string()));
}

#[tokio::test]
async fn test_arrow_try_cast_fn_str() {
let sql = r#"select arrow_try_cast(('{"foo": "bar"}'->'foo'), 'Utf8')"#;
let batches = run_query(sql).await.unwrap();
assert_eq!(display_val(batches).await, (DataType::Utf8, "bar".to_string()));
}

/// `arrow_cast` names an exact Arrow type. A type no accessor returns exactly is still folded, so
/// the union is not materialized, but the cast to that type is kept so the output type is the one
/// that was named.
#[tokio::test]
async fn test_plan_arrow_cast_fn_narrowing_type_keeps_cast() {
let lines = logical_plan(r"explain select arrow_cast((json_data->'foo'), 'Int32') from test").await;

let expected = [
"Projection: CAST(json_get_int(test.json_data, Utf8(\"foo\")) AS Int32) AS arrow_cast(json_data -> 'foo',Utf8(\"Int32\"))",
" TableScan: test projection=[json_data]",
];

assert_eq!(lines, expected);
}

#[tokio::test]
async fn test_arrow_double_nested() {
let sql = "select name, json_data->'foo'->0 from test";
Expand Down
84 changes: 64 additions & 20 deletions tests/rewrite_differential.rs
Original file line number Diff line number Diff line change
Expand Up @@ -63,59 +63,78 @@ enum Fold {
Narrowing,
}

/// A cast target: how it is spelled in SQL, the typed accessor the rewriter folds it to, and
/// what the fold does with it.
/// A cast target: how it is spelled in SQL, the exact Arrow type `arrow_cast` names for it,
/// the typed accessor the rewriter folds it to, and what each of the two folding paths does
/// with it.
#[derive(Debug, Clone, Copy)]
struct Target {
sql: &'static str,
arrow: &'static str,
accessor: &'static str,
sql_cast: Fold,
arrow_cast: Fold,
}

const TARGETS: &[Target] = &[
Target {
sql: "bigint",
arrow: "Int64",
accessor: "json_get_int",
sql_cast: Fold::Exact,
arrow_cast: Fold::Exact,
},
Target {
sql: "double",
arrow: "Float64",
accessor: "json_get_float",
sql_cast: Fold::Exact,
arrow_cast: Fold::Exact,
},
Target {
sql: "boolean",
arrow: "Boolean",
accessor: "json_get_bool",
sql_cast: Fold::Exact,
arrow_cast: Fold::Exact,
},
// SQL `VARCHAR` is `Utf8View` in DataFusion, but `json_as_text` returns `Utf8`.
Target {
sql: "varchar",
arrow: "Utf8",
accessor: "json_as_text",
sql_cast: Fold::Narrowing,
arrow_cast: Fold::Exact,
},
// Narrowing targets: the type asked for is narrower than what the accessor returns, so the
// fold keeps the cast on top of the accessor.
Target {
sql: "int",
arrow: "Int32",
accessor: "json_get_int",
sql_cast: Fold::Narrowing,
arrow_cast: Fold::Narrowing,
},
Target {
sql: "real",
arrow: "Float32",
accessor: "json_get_float",
sql_cast: Fold::Narrowing,
arrow_cast: Fold::Narrowing,
},
Target {
sql: "decimal(10,2)",
arrow: "Decimal128(10, 2)",
accessor: "json_get_float",
sql_cast: Fold::Narrowing,
arrow_cast: Fold::Narrowing,
},
// A type with no accessor, so it is not folded.
// A type neither path has an accessor for, so neither folds it.
Target {
sql: "smallint",
arrow: "Int16",
accessor: "json_get_int",
sql_cast: Fold::None,
arrow_cast: Fold::None,
},
];

Expand All @@ -124,24 +143,42 @@ const TARGETS: &[Target] = &[
enum Spelling {
Cast,
DoubleColon,
TryCast,
ArrowCast,
ArrowTryCast,
}

const SPELLINGS: &[Spelling] = &[Spelling::Cast, Spelling::DoubleColon];
const SPELLINGS: &[Spelling] = &[
Spelling::Cast,
Spelling::DoubleColon,
Spelling::TryCast,
Spelling::ArrowCast,
Spelling::ArrowTryCast,
];

impl Spelling {
/// Write `inner` cast to `target`.
fn apply(self, inner: &str, target: Target) -> String {
match self {
Spelling::Cast => format!("cast({inner} as {})", target.sql),
Spelling::DoubleColon => format!("({inner})::{}", target.sql),
Spelling::TryCast => format!("try_cast({inner} as {})", target.sql),
Spelling::ArrowCast => format!("arrow_cast({inner}, '{}')", target.arrow),
Spelling::ArrowTryCast => format!("arrow_try_cast({inner}, '{}')", target.arrow),
}
}

/// What the rewriter does with this spelling of this target. Every spelling here is a SQL
/// cast, so they all fold alike; the argument stays for the call sites that pair the two.
fn names_an_arrow_type(self) -> bool {
matches!(self, Spelling::ArrowCast | Spelling::ArrowTryCast)
}

/// What the rewriter does with this spelling of this target.
fn fold(self, target: Target) -> Fold {
let _ = self;
target.sql_cast
if self.names_an_arrow_type() {
target.arrow_cast
} else {
target.sql_cast
}
}
}

Expand Down Expand Up @@ -483,28 +520,35 @@ fn narrowing_cast_preserves_type_and_narrows() {
let rt = runtime();
let ctx = create_context().unwrap();

// (json value, sql type, CAST outcome)
// (json value, sql type, CAST outcome, TRY_CAST outcome)
let cases = [
("42", "int", "Int32=42"),
("42", "real", "Float32=42.0"),
("42", "decimal(10,2)", "Decimal128(10, 2)=42.00"),
(r#""abc""#, "varchar", "Utf8View=abc"),
// out of the target type's range: CAST fails
("3000000000", "int", "ERROR"),
("9223372036854775807", "int", "ERROR"),
("3000000000", "decimal(10,2)", "ERROR"),
("42", "int", "Int32=42", "Int32=42"),
("42", "real", "Float32=42.0", "Float32=42.0"),
(
"42",
"decimal(10,2)",
"Decimal128(10, 2)=42.00",
"Decimal128(10, 2)=42.00",
),
(r#""abc""#, "varchar", "Utf8View=abc", "Utf8View=abc"),
// out of the target type's range: CAST fails, TRY_CAST yields NULL
("3000000000", "int", "ERROR", "Int32=NULL"),
("9223372036854775807", "int", "ERROR", "Int32=NULL"),
("3000000000", "decimal(10,2)", "ERROR", "Decimal128(10, 2)=NULL"),
];

for (value, sql_type, want_cast) in cases {
for (value, sql_type, want_cast, want_try_cast) in cases {
let target = TARGETS.iter().find(|t| t.sql == sql_type).unwrap();
let doc = Doc {
json: format!(r#"{{"a": {value}}}"#),
path: vec!["a".to_string()],
};
set_doc(&ctx, &doc.json);

let sql = select(&Spelling::Cast.apply(&doc.json_get(), *target));
assert_eq!(rt.block_on(outcome(&ctx, &sql)).to_string(), want_cast, "{sql}");
for (spelling, want) in [(Spelling::Cast, want_cast), (Spelling::TryCast, want_try_cast)] {
let sql = select(&spelling.apply(&doc.json_get(), *target));
assert_eq!(rt.block_on(outcome(&ctx, &sql)).to_string(), want, "{sql}");
}
}
}

Expand Down
Loading