Skip to content
Closed
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
72 changes: 44 additions & 28 deletions datafusion/functions/src/core/nvl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ use datafusion_common::Result;
use datafusion_expr::simplify::{ExprSimplifyResult, SimplifyContext};
use datafusion_expr::{
ColumnarValue, Documentation, Expr, ReturnFieldArgs, ScalarFunctionArgs,
ScalarUDFImpl, Signature, Volatility,
ScalarUDFImpl, Signature,
};
use datafusion_macros::user_doc;

Expand Down Expand Up @@ -59,26 +59,6 @@ pub struct NVLFunc {
aliases: Vec<String>,
}

/// Currently supported types by the nvl/ifnull function.
/// The order of these types correspond to the order on which coercion applies
/// This should thus be from least informative to most informative
static SUPPORTED_NVL_TYPES: &[DataType] = &[
DataType::Boolean,
DataType::UInt8,
DataType::UInt16,
DataType::UInt32,
DataType::UInt64,
DataType::Int8,
DataType::Int16,
DataType::Int32,
DataType::Int64,
DataType::Float32,
DataType::Float64,
DataType::Utf8View,
DataType::Utf8,
DataType::LargeUtf8,
];

impl Default for NVLFunc {
fn default() -> Self {
Self::new()
Expand All @@ -88,13 +68,7 @@ impl Default for NVLFunc {
impl NVLFunc {
pub fn new() -> Self {
Self {
coalesce: CoalesceFunc {
signature: Signature::uniform(
2,
SUPPORTED_NVL_TYPES.to_vec(),
Volatility::Immutable,
),
},
coalesce: CoalesceFunc::new(),
aliases: vec![String::from("ifnull")],
}
}
Expand All @@ -113,6 +87,32 @@ impl ScalarUDFImpl for NVLFunc {
self.coalesce.return_type(arg_types)
}

fn coerce_types(&self, arg_types: &[DataType]) -> Result<Vec<DataType>> {
// Preserve the existing type for an all-NULL call. The uniform
// signature previously selected the first supported type (Boolean).
if arg_types.iter().all(DataType::is_null) {
return Ok(vec![DataType::Boolean; arg_types.len()]);
}

match self.coalesce.coerce_types(arg_types) {
Ok(coerced_types) => Ok(coerced_types),
Err(err) => {
// NVL historically allowed a boolean fallback for numeric
// inputs (for example IFNULL(int_col, false)). Keep that
// behavior while using COALESCE's broader type resolution.
if let [left, right] = arg_types {
if *left == DataType::Boolean && is_legacy_numeric_type(right) {
return Ok(vec![right.clone(), right.clone()]);
}
if *right == DataType::Boolean && is_legacy_numeric_type(left) {
return Ok(vec![left.clone(), left.clone()]);
}
}
Err(err)
}
}
}

fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result<FieldRef> {
self.coalesce.return_field_from_args(args)
}
Expand Down Expand Up @@ -148,3 +148,19 @@ impl ScalarUDFImpl for NVLFunc {
self.doc()
}
}

fn is_legacy_numeric_type(data_type: &DataType) -> bool {
matches!(
data_type,
DataType::UInt8
| DataType::UInt16
| DataType::UInt32
| DataType::UInt64
| DataType::Int8
| DataType::Int16
| DataType::Int32
| DataType::Int64
| DataType::Float32
| DataType::Float64
)
}
33 changes: 33 additions & 0 deletions datafusion/sqllogictest/test_files/nvl.slt
Original file line number Diff line number Diff line change
Expand Up @@ -149,6 +149,39 @@ SELECT NVL(arrow_cast('a', 'Utf8View'), NULL);
----
a

statement ok
CREATE TABLE nvl_type_test AS SELECT
CAST('12345678901234567890123456789012345678' AS DECIMAL(38,0)) AS big,
CAST(1.23 AS DECIMAL(5,2)) AS d,
TIMESTAMP '2024-01-01 10:00:00' AS ts,
DATE '2024-01-01' AS dt;

# NVL and IFNULL should use the same common type resolution as COALESCE.
query TTTTTTTTT
SELECT
arrow_typeof(nvl(big, 0)),
arrow_typeof(nvl(d, d)),
arrow_typeof(ifnull(d, d)),
arrow_typeof(nvl(ts, ts)),
arrow_typeof(ifnull(ts, ts)),
arrow_typeof(nvl(dt, dt)),
arrow_typeof(ifnull(dt, dt)),
arrow_typeof(coalesce(ts, ts)),
arrow_typeof(coalesce(dt, dt))
FROM nvl_type_test;
----
Decimal128(38, 0) Decimal128(5, 2) Decimal128(5, 2) Timestamp(ns) Timestamp(ns) Date32 Date32 Timestamp(ns) Date32

query T
SELECT CAST(nvl(big, 0) AS VARCHAR) FROM nvl_type_test;
----
12345678901234567890123456789012345678

query T
SELECT CAST(nvl(ts, ts) + INTERVAL '1 hour' AS VARCHAR) FROM nvl_type_test;
----
2024-01-01T11:00:00

# nvl is implemented as a case, and short-circuits evaluation
# so the following query should not error
query I
Expand Down