Skip to content
Open
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
57 changes: 32 additions & 25 deletions datafusion/spark/src/function/math/abs.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,11 +24,16 @@ use datafusion_expr::{
Volatility,
};
use datafusion_functions::{
downcast_named_arg, make_abs_function, make_try_abs_function,
make_wrapping_abs_function,
downcast_named_arg, make_abs_function, make_wrapping_abs_function,
};
use std::sync::Arc;

const ARITHMETIC_OVERFLOW_ERROR: &str = r#"[ARITHMETIC_OVERFLOW] overflow. If necessary set "spark.sql.ansi.enabled" to "false" to bypass this error. SQLSTATE: 22003"#;

fn arithmetic_overflow_error() -> ArrowError {
ArrowError::ComputeError(ARITHMETIC_OVERFLOW_ERROR.to_string())
}

/// Spark-compatible `abs` expression
/// <https://spark.apache.org/docs/latest/api/sql/index.html#abs>
///
Expand Down Expand Up @@ -91,13 +96,7 @@ impl ScalarUDFImpl for SparkAbs {
macro_rules! scalar_compute_op {
($ENABLE_ANSI_MODE:expr, $INPUT:ident, $SCALAR_TYPE:ident) => {{
let result = if $ENABLE_ANSI_MODE {
$INPUT.checked_abs().ok_or_else(|| {
ArrowError::ComputeError(format!(
"{} overflow on abs({:?})",
stringify!($SCALAR_TYPE),
$INPUT
))
})?
$INPUT.checked_abs().ok_or_else(arithmetic_overflow_error)?
} else {
$INPUT.wrapping_abs()
};
Expand All @@ -107,13 +106,7 @@ macro_rules! scalar_compute_op {
}};
($ENABLE_ANSI_MODE:expr, $INPUT:ident, $PRECISION:expr, $SCALE:expr, $SCALAR_TYPE:ident) => {{
let result = if $ENABLE_ANSI_MODE {
$INPUT.checked_abs().ok_or_else(|| {
ArrowError::ComputeError(format!(
"{} overflow on abs({:?})",
stringify!($SCALAR_TYPE),
$INPUT
))
})?
$INPUT.checked_abs().ok_or_else(arithmetic_overflow_error)?
} else {
$INPUT.wrapping_abs()
};
Expand All @@ -125,6 +118,18 @@ macro_rules! scalar_compute_op {
}};
}

macro_rules! make_try_spark_abs_function {
($ARRAY_TYPE:ident) => {{
|input: &ArrayRef| {
let array = downcast_named_arg!(&input, "abs arg", $ARRAY_TYPE);
let res: $ARRAY_TYPE = array
.try_unary(|x| x.checked_abs().ok_or_else(arithmetic_overflow_error))
.and_then(|v| Ok(v.with_data_type(input.data_type().clone())))?;
Ok(Arc::new(res) as ArrayRef)
}
}};
}

pub fn spark_abs(
args: &[ColumnarValue],
enable_ansi_mode: bool,
Expand All @@ -142,31 +147,31 @@ pub fn spark_abs(
| DataType::UInt64 => Ok(args[0].clone()),
DataType::Int8 => {
let abs_fun = if enable_ansi_mode {
make_try_abs_function!(Int8Array)
make_try_spark_abs_function!(Int8Array)
} else {
make_wrapping_abs_function!(Int8Array)
};
abs_fun(array).map(ColumnarValue::Array)
}
DataType::Int16 => {
let abs_fun = if enable_ansi_mode {
make_try_abs_function!(Int16Array)
make_try_spark_abs_function!(Int16Array)
} else {
make_wrapping_abs_function!(Int16Array)
};
abs_fun(array).map(ColumnarValue::Array)
}
DataType::Int32 => {
let abs_fun = if enable_ansi_mode {
make_try_abs_function!(Int32Array)
make_try_spark_abs_function!(Int32Array)
} else {
make_wrapping_abs_function!(Int32Array)
};
abs_fun(array).map(ColumnarValue::Array)
}
DataType::Int64 => {
let abs_fun = if enable_ansi_mode {
make_try_abs_function!(Int64Array)
make_try_spark_abs_function!(Int64Array)
} else {
make_wrapping_abs_function!(Int64Array)
};
Expand All @@ -182,15 +187,15 @@ pub fn spark_abs(
}
DataType::Decimal128(_, _) => {
let abs_fun = if enable_ansi_mode {
make_try_abs_function!(Decimal128Array)
make_try_spark_abs_function!(Decimal128Array)
} else {
make_wrapping_abs_function!(Decimal128Array)
};
abs_fun(array).map(ColumnarValue::Array)
}
DataType::Decimal256(_, _) => {
let abs_fun = if enable_ansi_mode {
make_try_abs_function!(Decimal256Array)
make_try_spark_abs_function!(Decimal256Array)
} else {
make_wrapping_abs_function!(Decimal256Array)
};
Expand Down Expand Up @@ -361,9 +366,11 @@ mod tests {
let args = ColumnarValue::Array(Arc::new(input));
match spark_abs(&[args], true) {
Err(e) => {
assert!(
e.to_string().contains("overflow on abs"),
"Error message did not match. Actual message: {e}"
assert_eq!(
e.to_string(),
format!(
"Arrow error: Compute error: {ARITHMETIC_OVERFLOW_ERROR}"
)
);
}
_ => unreachable!(),
Expand Down
16 changes: 8 additions & 8 deletions datafusion/sqllogictest/test_files/spark/math/abs.slt
Original file line number Diff line number Diff line change
Expand Up @@ -42,16 +42,16 @@ select abs((-128)::TINYINT), abs((-32768)::SMALLINT), abs((-2147483648)::INT), a
statement ok
set datafusion.execution.enable_ansi_mode = true;

query error DataFusion error: Arrow error: Compute error: Int8 overflow on abs\(\-128\)
query error DataFusion error: Arrow error: Compute error: \[ARITHMETIC_OVERFLOW\] overflow\. If necessary set "spark\.sql\.ansi\.enabled" to "false" to bypass this error\. SQLSTATE: 22003
select abs((-128)::TINYINT);

query error DataFusion error: Arrow error: Compute error: Int16 overflow on abs\(\-32768\)
query error DataFusion error: Arrow error: Compute error: \[ARITHMETIC_OVERFLOW\] overflow\. If necessary set "spark\.sql\.ansi\.enabled" to "false" to bypass this error\. SQLSTATE: 22003
select abs((-32768)::SMALLINT);

query error DataFusion error: Arrow error: Compute error: Int32 overflow on abs\(\-2147483648\)
query error DataFusion error: Arrow error: Compute error: \[ARITHMETIC_OVERFLOW\] overflow\. If necessary set "spark\.sql\.ansi\.enabled" to "false" to bypass this error\. SQLSTATE: 22003
select abs((-2147483648)::INT);

query error DataFusion error: Arrow error: Compute error: Int64 overflow on abs\(\-9223372036854775808\)
query error DataFusion error: Arrow error: Compute error: \[ARITHMETIC_OVERFLOW\] overflow\. If necessary set "spark\.sql\.ansi\.enabled" to "false" to bypass this error\. SQLSTATE: 22003
select abs((-9223372036854775808)::BIGINT);

statement ok
Expand Down Expand Up @@ -126,16 +126,16 @@ NULL
statement ok
set datafusion.execution.enable_ansi_mode = true;

query error DataFusion error: Arrow error: Compute error: Int8Array overflow on abs\(\-128\)
query error DataFusion error: Arrow error: Compute error: \[ARITHMETIC_OVERFLOW\] overflow\. If necessary set "spark\.sql\.ansi\.enabled" to "false" to bypass this error\. SQLSTATE: 22003
SELECT abs(a) FROM (VALUES (-127::TINYINT), ((-128)::TINYINT)) AS t(a);

query error DataFusion error: Arrow error: Compute error: Int16Array overflow on abs\(\-32768\)
query error DataFusion error: Arrow error: Compute error: \[ARITHMETIC_OVERFLOW\] overflow\. If necessary set "spark\.sql\.ansi\.enabled" to "false" to bypass this error\. SQLSTATE: 22003
select abs(a) FROM (VALUES (-32767::SMALLINT), ((-32768)::SMALLINT)) AS t(a);

query error DataFusion error: Arrow error: Compute error: Int32Array overflow on abs\(\-2147483648\)
query error DataFusion error: Arrow error: Compute error: \[ARITHMETIC_OVERFLOW\] overflow\. If necessary set "spark\.sql\.ansi\.enabled" to "false" to bypass this error\. SQLSTATE: 22003
select abs(a) FROM (VALUES (-2147483647::INT), ((-2147483648)::INT)) AS t(a);

query error DataFusion error: Arrow error: Compute error: Int64Array overflow on abs\(\-9223372036854775808\)
query error DataFusion error: Arrow error: Compute error: \[ARITHMETIC_OVERFLOW\] overflow\. If necessary set "spark\.sql\.ansi\.enabled" to "false" to bypass this error\. SQLSTATE: 22003
select abs(a) FROM (VALUES (-9223372036854775807::BIGINT), ((-9223372036854775808)::BIGINT)) AS t(a);

statement ok
Expand Down
Loading