diff --git a/native/spark-expr/benches/checked_arithmetic.rs b/native/spark-expr/benches/checked_arithmetic.rs index 280852171c..ed03f3c616 100644 --- a/native/spark-expr/benches/checked_arithmetic.rs +++ b/native/spark-expr/benches/checked_arithmetic.rs @@ -18,6 +18,7 @@ use arrow::array::{ArrayRef, Float64Array, Int32Array, Int64Array}; use arrow::datatypes::DataType; use criterion::{criterion_group, criterion_main, Criterion}; +use datafusion::common::ScalarValue; use datafusion::physical_plan::ColumnarValue; use datafusion_comet_spark_expr::{checked_add, checked_div, checked_mul, checked_sub, EvalMode}; use std::hint::black_box; @@ -95,6 +96,21 @@ fn criterion_benchmark(c: &mut Criterion) { ColumnarValue::Array(f64_array(0x9e37_79b9, nulls)), ]; + let i32_scalar_right_args = [ + ColumnarValue::Array(i32_array(0x1234_5678, nulls)), + ColumnarValue::Scalar(ScalarValue::Int32(Some(42))), + ]; + + let i32_scalar_left_args = [ + ColumnarValue::Scalar(ScalarValue::Int32(Some(42))), + ColumnarValue::Array(i32_array(0x1234_5678, nulls)), + ]; + + let i64_scalar_args = [ + ColumnarValue::Array(i64_array(0x1234_5678, nulls)), + ColumnarValue::Scalar(ScalarValue::Int64(Some(100))), + ]; + group.bench_function(format!("checked_add_i32_ansi_{label}"), |b| { b.iter(|| { black_box(checked_add( @@ -144,6 +160,36 @@ fn criterion_benchmark(c: &mut Criterion) { )) }) }); + + group.bench_function(format!("checked_add_i32_scalar_right_ansi_{label}"), |b| { + b.iter(|| { + black_box(checked_add( + black_box(&i32_scalar_right_args), + &DataType::Int32, + EvalMode::Ansi, + )) + }) + }); + + group.bench_function(format!("checked_add_i32_scalar_left_ansi_{label}"), |b| { + b.iter(|| { + black_box(checked_add( + black_box(&i32_scalar_left_args), + &DataType::Int32, + EvalMode::Ansi, + )) + }) + }); + + group.bench_function(format!("checked_mul_i64_scalar_ansi_{label}"), |b| { + b.iter(|| { + black_box(checked_mul( + black_box(&i64_scalar_args), + &DataType::Int64, + EvalMode::Ansi, + )) + }) + }); } group.finish(); diff --git a/native/spark-expr/src/math_funcs/checked_arithmetic.rs b/native/spark-expr/src/math_funcs/checked_arithmetic.rs index a76574e16d..8c0dc0fc33 100644 --- a/native/spark-expr/src/math_funcs/checked_arithmetic.rs +++ b/native/spark-expr/src/math_funcs/checked_arithmetic.rs @@ -19,59 +19,115 @@ use arrow::array::{Array, ArrowNativeTypeOp, BooleanBufferBuilder, PrimitiveArra use arrow::array::{ArrayRef, AsArray}; use crate::{divide_by_zero_error, EvalMode, SparkError}; +use array::{Datum, Scalar}; +use arrow::array; use arrow::buffer::NullBuffer; +use arrow::compute::kernels::{arity, numeric}; use arrow::datatypes::{ ArrowPrimitiveType, DataType, Float16Type, Float32Type, Float64Type, Int16Type, Int32Type, Int64Type, Int8Type, }; use arrow::error::ArrowError; use datafusion::common::DataFusionError; +use datafusion::common::ScalarValue; use datafusion::physical_plan::ColumnarValue; use std::sync::Arc; -pub fn try_arithmetic_kernel( +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum MathOp { + Add, + Sub, + Mul, + Div, +} + +fn try_arithmetic_kernel( left: &PrimitiveArray, right: &PrimitiveArray, - op: &str, - is_ansi_mode: bool, + op: MathOp, ) -> Result where T: ArrowPrimitiveType, { match op { - "checked_add" => checked_binary(left, right, is_ansi_mode, false, |l, r| l.add_checked(r)), - "checked_sub" => checked_binary(left, right, is_ansi_mode, false, |l, r| l.sub_checked(r)), - "checked_mul" => checked_binary(left, right, is_ansi_mode, false, |l, r| l.mul_checked(r)), - "checked_div" => checked_binary(left, right, is_ansi_mode, true, |l, r| l.div_checked(r)), - _ => Err(DataFusionError::Internal(format!( - "Unsupported operation: {:?}", - op - ))), + MathOp::Add => checked_binary(left, right, |l, r| l.add_checked(r)), + MathOp::Sub => checked_binary(left, right, |l, r| l.sub_checked(r)), + MathOp::Mul => checked_binary(left, right, |l, r| l.mul_checked(r)), + MathOp::Div => checked_binary(left, right, |l, r| l.div_checked(r)), } } +fn ansi_arithmetic_kernel( + left: &ColumnarValue, + right: &ColumnarValue, + op: MathOp, +) -> Result { + let run_kernel = |l: &dyn Datum, r: &dyn Datum| match op { + MathOp::Add => numeric::add(l, r), + MathOp::Sub => numeric::sub(l, r), + MathOp::Mul => numeric::mul(l, r), + MathOp::Div => numeric::div(l, r), + }; + + let result_array = match (left, right) { + (ColumnarValue::Array(l), ColumnarValue::Array(r)) => run_kernel(&l, &r), + (ColumnarValue::Scalar(l), ColumnarValue::Array(r)) => { + let l_arr = l.to_array()?; + let l_scalar = Scalar::new(l_arr); + run_kernel(&l_scalar, &r) + } + (ColumnarValue::Array(l), ColumnarValue::Scalar(r)) => { + let r_arr = r.to_array()?; + let r_scalar = Scalar::new(r_arr); + run_kernel(&l, &r_scalar) + } + (ColumnarValue::Scalar(l), ColumnarValue::Scalar(r)) => { + let l_arr = l.to_array()?; + let r_arr = r.to_array()?; + let l_scalar = Scalar::new(l_arr); + let r_scalar = Scalar::new(r_arr); + let res = run_kernel(&l_scalar, &r_scalar)?; + let scalar_val = ScalarValue::try_from_array(res.as_ref(), 0)?; + return Ok(ColumnarValue::Scalar(scalar_val)); + } + }; + + let array = result_array.map_err(|e| match e { + ArrowError::DivideByZero => divide_by_zero_error().into(), + _ => DataFusionError::from(SparkError::ArithmeticOverflow { + from_type: String::from("integer"), + }), + })?; + + Ok(ColumnarValue::Array(array)) +} + +fn ansi_float_div( + left: &PrimitiveArray, + right: &PrimitiveArray, +) -> Result +where + T: ArrowPrimitiveType, +{ + arity::try_binary::<_, _, _, T>(left, right, |l, r| l.div_checked(r)) + .map(|array| Arc::new(array) as ArrayRef) + .map_err(|e| match e { + ArrowError::DivideByZero => divide_by_zero_error().into(), + _ => DataFusionError::from(SparkError::ArithmeticOverflow { + from_type: String::from("integer"), + }), + }) +} + fn checked_binary( left: &PrimitiveArray, right: &PrimitiveArray, - is_ansi_mode: bool, - is_div: bool, op: F, ) -> Result where T: ArrowPrimitiveType, F: Fn(T::Native, T::Native) -> Result, { - if is_ansi_mode { - return arrow::compute::kernels::arity::try_binary::<_, _, _, T>(left, right, op) - .map(|array| Arc::new(array) as ArrayRef) - .map_err(|e| match e { - ArrowError::DivideByZero => divide_by_zero_error().into(), - _ => DataFusionError::from(SparkError::ArithmeticOverflow { - from_type: String::from("integer"), - }), - }); - } - let len = left.len(); let lhs = &left.values()[..len]; let rhs = &right.values()[..len]; @@ -84,18 +140,7 @@ where match op(l, r) { Ok(v) => *out = v, Err(_) => { - if !is_ansi_mode { - overflowed.push(i); - } else if nulls.as_ref().is_none_or(|n| n.is_valid(i)) { - return if is_div && r.is_zero() { - Err(divide_by_zero_error().into()) - } else { - Err(SparkError::ArithmeticOverflow { - from_type: String::from("integer"), - } - .into()) - }; - } + overflowed.push(i); } } } @@ -126,12 +171,13 @@ where Ok(Arc::new(PrimitiveArray::::new(values.into(), nulls)) as ArrayRef) } + pub fn checked_add( args: &[ColumnarValue], data_type: &DataType, eval_mode: EvalMode, ) -> Result { - checked_arithmetic_internal(args, data_type, "checked_add", eval_mode) + checked_arithmetic_internal(args, data_type, MathOp::Add, eval_mode) } pub fn checked_sub( @@ -139,7 +185,7 @@ pub fn checked_sub( data_type: &DataType, eval_mode: EvalMode, ) -> Result { - checked_arithmetic_internal(args, data_type, "checked_sub", eval_mode) + checked_arithmetic_internal(args, data_type, MathOp::Sub, eval_mode) } pub fn checked_mul( @@ -147,7 +193,7 @@ pub fn checked_mul( data_type: &DataType, eval_mode: EvalMode, ) -> Result { - checked_arithmetic_internal(args, data_type, "checked_mul", eval_mode) + checked_arithmetic_internal(args, data_type, MathOp::Mul, eval_mode) } pub fn checked_div( @@ -155,13 +201,21 @@ pub fn checked_div( data_type: &DataType, eval_mode: EvalMode, ) -> Result { - checked_arithmetic_internal(args, data_type, "checked_div", eval_mode) + checked_arithmetic_internal(args, data_type, MathOp::Div, eval_mode) +} + +#[inline] +fn is_integer_type(data_type: &DataType) -> bool { + matches!( + data_type, + DataType::Int8 | DataType::Int16 | DataType::Int32 | DataType::Int64 + ) } fn checked_arithmetic_internal( args: &[ColumnarValue], data_type: &DataType, - op: &str, + op: MathOp, eval_mode: EvalMode, ) -> Result { let left = &args[0]; @@ -178,6 +232,12 @@ fn checked_arithmetic_internal( } }; + // Early return for integer types in ANSI mode using the fast Datum/Scalar path + if is_ansi_mode && is_integer_type(data_type) { + return ansi_arithmetic_kernel(left, right, op); + } + + // Materialize operands for Try-mode and float division let (left_arr, right_arr): (ArrayRef, ArrayRef) = match (left, right) { (ColumnarValue::Array(l), ColumnarValue::Array(r)) => (Arc::clone(l), Arc::clone(r)), (ColumnarValue::Scalar(l), ColumnarValue::Array(r)) => { @@ -189,51 +249,69 @@ fn checked_arithmetic_internal( (ColumnarValue::Scalar(l), ColumnarValue::Scalar(r)) => (l.to_array()?, r.to_array()?), }; - // Rust only supports checked_arithmetic on numeric types let result_array = match data_type { - DataType::Int8 => try_arithmetic_kernel::( + DataType::Int8 => try_arithmetic_kernel( left_arr.as_primitive::(), right_arr.as_primitive::(), op, - is_ansi_mode, ), - DataType::Int16 => try_arithmetic_kernel::( + DataType::Int16 => try_arithmetic_kernel( left_arr.as_primitive::(), right_arr.as_primitive::(), op, - is_ansi_mode, ), - DataType::Int32 => try_arithmetic_kernel::( + DataType::Int32 => try_arithmetic_kernel( left_arr.as_primitive::(), right_arr.as_primitive::(), op, - is_ansi_mode, ), - DataType::Int64 => try_arithmetic_kernel::( + DataType::Int64 => try_arithmetic_kernel( left_arr.as_primitive::(), right_arr.as_primitive::(), op, - is_ansi_mode, - ), - // Spark always casts division operands to floats - DataType::Float16 if (op == "checked_div") => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - is_ansi_mode, - ), - DataType::Float32 if (op == "checked_div") => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - is_ansi_mode, - ), - DataType::Float64 if (op == "checked_div") => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - is_ansi_mode, ), + DataType::Float16 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ) + } + } + DataType::Float32 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ) + } + } + DataType::Float64 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ) + } + } _ => Err(DataFusionError::Internal(format!( "Unsupported data type: {:?}", data_type @@ -332,4 +410,59 @@ mod tests { assert_eq!(result, Int32Array::from(vec![None, Some(2)])); assert_eq!(result.values()[0], 0); } + + #[test] + fn test_ansi_integer_div_by_zero() { + let args = int32_args(vec![Some(10)], vec![Some(0)]); + let result = checked_div(&args, &DataType::Int32, EvalMode::Ansi); + assert!(result.is_err()); + } + + #[test] + fn test_ansi_scalar_operand() { + let args = vec![ + ColumnarValue::Array(Arc::new(Int32Array::from(vec![Some(10), Some(20)]))), + ColumnarValue::Scalar(ScalarValue::Int32(Some(5))), + ]; + let result = as_int32(checked_add(&args, &DataType::Int32, EvalMode::Ansi).unwrap()); + assert_eq!(result, Int32Array::from(vec![Some(15), Some(25)])); + } + + #[test] + fn test_ansi_int64_overflow() { + let args = vec![ + ColumnarValue::Array(Arc::new(array::Int64Array::from(vec![Some(i64::MAX)]))), + ColumnarValue::Array(Arc::new(array::Int64Array::from(vec![Some(1)]))), + ]; + let result = checked_add(&args, &DataType::Int64, EvalMode::Ansi); + assert!(result.is_err()); + } + + #[test] + fn test_ansi_scalar_operands() { + let args_right = vec![ + ColumnarValue::Array(Arc::new(Int32Array::from(vec![Some(10), Some(20)]))), + ColumnarValue::Scalar(ScalarValue::Int32(Some(5))), + ]; + let res_right = + as_int32(checked_add(&args_right, &DataType::Int32, EvalMode::Ansi).unwrap()); + assert_eq!(res_right, Int32Array::from(vec![Some(15), Some(25)])); + + let args_left = vec![ + ColumnarValue::Scalar(ScalarValue::Int32(Some(5))), + ColumnarValue::Array(Arc::new(Int32Array::from(vec![Some(10), Some(20)]))), + ]; + let res_left = as_int32(checked_add(&args_left, &DataType::Int32, EvalMode::Ansi).unwrap()); + assert_eq!(res_left, Int32Array::from(vec![Some(15), Some(25)])); + + let args_scalar = vec![ + ColumnarValue::Scalar(ScalarValue::Int32(Some(10))), + ColumnarValue::Scalar(ScalarValue::Int32(Some(20))), + ]; + let res_scalar = checked_add(&args_scalar, &DataType::Int32, EvalMode::Ansi).unwrap(); + match res_scalar { + ColumnarValue::Scalar(ScalarValue::Int32(v)) => assert_eq!(v, Some(30)), + _ => panic!("Expected scalar result"), + } + } }