Skip to content
Closed
Show file tree
Hide file tree
Changes from 5 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
24 changes: 24 additions & 0 deletions crates/sail-execution/proto/sail/plan/physical.proto
Original file line number Diff line number Diff line change
Expand Up @@ -182,6 +182,10 @@ message ExtendedScalarUdf {
SparkBinUdf spark_bin = 30;
SparkNegativeUdf spark_negative = 31;
SparkFromXmlUdf spark_from_xml = 32;
SparkAddUdf spark_add = 33;
SparkSubtractUdf spark_subtract = 34;
SparkMultiplyUdf spark_multiply = 35;
SparkDivideUdf spark_divide = 36;
}
}

Expand Down Expand Up @@ -316,6 +320,26 @@ message SparkNegativeUdf {
bool ansi_mode = 1;
}

message SparkAddUdf {
bool ansi_mode = 1;
bool safe = 2;
}

message SparkSubtractUdf {
bool ansi_mode = 1;
bool safe = 2;
}

message SparkMultiplyUdf {
bool ansi_mode = 1;
bool safe = 2;
}

message SparkDivideUdf {
bool ansi_mode = 1;
bool safe = 2;
}

message SparkFromCsvUdf {
string session_timezone = 1;
}
Expand Down
56 changes: 40 additions & 16 deletions crates/sail-execution/src/proto/codec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -176,19 +176,19 @@ use sail_function::scalar::math::rand_poisson::RandPoisson;
use sail_function::scalar::math::randn::Randn;
use sail_function::scalar::math::random::Random;
use sail_function::scalar::math::spark_abs::SparkAbs;
use sail_function::scalar::math::spark_add::SparkAdd;
use sail_function::scalar::math::spark_bin::SparkBin;
use sail_function::scalar::math::spark_bround::SparkBRound;
use sail_function::scalar::math::spark_ceil_floor::{SparkCeil, SparkFloor};
use sail_function::scalar::math::spark_conv::SparkConv;
use sail_function::scalar::math::spark_div::SparkIntervalDiv;
use sail_function::scalar::math::spark_divide::SparkDivide;
use sail_function::scalar::math::spark_multiply::SparkMultiply;
use sail_function::scalar::math::spark_negative::SparkNegative;
use sail_function::scalar::math::spark_pmod::SparkPmod;
use sail_function::scalar::math::spark_signum::SparkSignum;
use sail_function::scalar::math::spark_try_add::SparkTryAdd;
use sail_function::scalar::math::spark_try_div::SparkTryDiv;
use sail_function::scalar::math::spark_subtract::SparkSubtract;
use sail_function::scalar::math::spark_try_mod::SparkTryMod;
use sail_function::scalar::math::spark_try_mult::SparkTryMult;
use sail_function::scalar::math::spark_try_subtract::SparkTrySubtract;
use sail_function::scalar::math::spark_unhex::SparkUnHex;
use sail_function::scalar::math::spark_uniform::SparkUniform;
use sail_function::scalar::misc::hll_sketch::{HllSketchEstimateFunction, HllUnionFunction};
Expand Down Expand Up @@ -2418,6 +2418,22 @@ impl PhysicalExtensionCodec for RemoteExecutionCodec {
UdfKind::SparkNegative(r#gen::SparkNegativeUdf { ansi_mode }) => {
return Ok(Arc::new(ScalarUDF::from(SparkNegative::new(ansi_mode))));
}
UdfKind::SparkAdd(r#gen::SparkAddUdf { ansi_mode, safe }) => {
return Ok(Arc::new(ScalarUDF::from(SparkAdd::new(ansi_mode, safe))));
}
UdfKind::SparkSubtract(r#gen::SparkSubtractUdf { ansi_mode, safe }) => {
return Ok(Arc::new(ScalarUDF::from(SparkSubtract::new(
ansi_mode, safe,
))));
}
UdfKind::SparkMultiply(r#gen::SparkMultiplyUdf { ansi_mode, safe }) => {
return Ok(Arc::new(ScalarUDF::from(SparkMultiply::new(
ansi_mode, safe,
))));
}
UdfKind::SparkDivide(r#gen::SparkDivideUdf { ansi_mode, safe }) => {
return Ok(Arc::new(ScalarUDF::from(SparkDivide::new(ansi_mode, safe))));
}
UdfKind::SparkMakeTimestampNtz(r#gen::SparkMakeTimestampNtzUdf { is_try }) => {
return Ok(Arc::new(ScalarUDF::from(SparkMakeTimestampNtz::new(
is_try,
Expand Down Expand Up @@ -2590,17 +2606,9 @@ impl PhysicalExtensionCodec for RemoteExecutionCodec {
"spark_to_utf8" => Ok(Arc::new(ScalarUDF::from(SparkToUtf8::new()))),
"spark_to_large_utf8" => Ok(Arc::new(ScalarUDF::from(SparkToLargeUtf8::new()))),
"spark_to_utf8_view" => Ok(Arc::new(ScalarUDF::from(SparkToUtf8View::new()))),
"spark_try_add" | "try_add" => Ok(Arc::new(ScalarUDF::from(SparkTryAdd::new()))),
"spark_try_divide" | "try_divide" => Ok(Arc::new(ScalarUDF::from(SparkTryDiv::new()))),
"spark_try_mod" | "try_mod" => Ok(Arc::new(ScalarUDF::from(SparkTryMod::new()))),
"spark_try_multiply" | "try_multiply" => {
Ok(Arc::new(ScalarUDF::from(SparkTryMult::new())))
}
"spark_version" | "version" => Ok(Arc::new(ScalarUDF::from(SparkVersion::new()))),
"spark_to_json" | "to_json" => Ok(Arc::new(ScalarUDF::from(SparkToJson::new()))),
"spark_try_subtract" | "try_subtract" => {
Ok(Arc::new(ScalarUDF::from(SparkTrySubtract::new())))
}
"spark_uniform" | "uniform" => Ok(Arc::new(ScalarUDF::from(SparkUniform::new()))),
"spark_width_bucket" | "width_bucket" => {
Ok(Arc::new(ScalarUDF::from(SparkWidthBucket::new())))
Expand Down Expand Up @@ -2708,14 +2716,10 @@ impl PhysicalExtensionCodec for RemoteExecutionCodec {
|| node_inner.is::<SparkToLargeUtf8>()
|| node_inner.is::<SparkToUtf8>()
|| node_inner.is::<SparkToUtf8View>()
|| node_inner.is::<SparkTryAdd>()
|| node_inner.is::<SparkTryAESDecrypt>()
|| node_inner.is::<SparkTryAESEncrypt>()
|| node_inner.is::<SparkTryDiv>()
|| node_inner.is::<SparkTryMod>()
|| node_inner.is::<SparkTryMult>()
|| node_inner.is::<SparkTryParseUrl>()
|| node_inner.is::<SparkTrySubtract>()
|| node_inner.is::<SparkTryToBinary>()
|| node_inner.is::<HllSketchEstimateFunction>()
|| node_inner.is::<HllUnionFunction>()
Expand Down Expand Up @@ -2878,6 +2882,26 @@ impl PhysicalExtensionCodec for RemoteExecutionCodec {
} else if let Some(func) = node.inner().downcast_ref::<SparkNegative>() {
let ansi_mode = func.ansi_mode();
UdfKind::SparkNegative(r#gen::SparkNegativeUdf { ansi_mode })
} else if let Some(func) = node.inner().downcast_ref::<SparkAdd>() {
UdfKind::SparkAdd(r#gen::SparkAddUdf {
ansi_mode: func.ansi_mode(),
safe: func.safe(),
})
} else if let Some(func) = node.inner().downcast_ref::<SparkSubtract>() {
UdfKind::SparkSubtract(r#gen::SparkSubtractUdf {
ansi_mode: func.ansi_mode(),
safe: func.safe(),
})
} else if let Some(func) = node.inner().downcast_ref::<SparkMultiply>() {
UdfKind::SparkMultiply(r#gen::SparkMultiplyUdf {
ansi_mode: func.ansi_mode(),
safe: func.safe(),
})
} else if let Some(func) = node.inner().downcast_ref::<SparkDivide>() {
UdfKind::SparkDivide(r#gen::SparkDivideUdf {
ansi_mode: func.ansi_mode(),
safe: func.safe(),
})
} else if let Some(func) = node.inner().downcast_ref::<SparkMakeTimestampNtz>() {
let is_try = func.is_try();
UdfKind::SparkMakeTimestampNtz(r#gen::SparkMakeTimestampNtzUdf { is_try })
Expand Down
8 changes: 4 additions & 4 deletions crates/sail-function/src/scalar/math/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,19 +3,19 @@ pub mod rand_poisson;
pub mod randn;
pub mod random;
pub mod spark_abs;
pub mod spark_add;
pub mod spark_bin;
pub mod spark_bround;
pub mod spark_ceil_floor;
pub mod spark_conv;
pub mod spark_div;
pub mod spark_divide;
pub mod spark_multiply;
pub mod spark_negative;
pub mod spark_pmod;
pub mod spark_signum;
pub mod spark_try_add;
pub mod spark_try_div;
pub mod spark_subtract;
pub mod spark_try_mod;
pub mod spark_try_mult;
pub mod spark_try_subtract;
pub mod spark_unhex;
pub mod spark_uniform;
mod utils;
Expand Down
203 changes: 203 additions & 0 deletions crates/sail-function/src/scalar/math/spark_add.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,203 @@
use std::sync::Arc;

use datafusion::arrow::array::{Array, ArrayRef, AsArray};
use datafusion::arrow::compute::kernels::numeric::{add, add_wrapping};
use datafusion::arrow::datatypes::IntervalUnit::{MonthDayNano, YearMonth};
use datafusion::arrow::datatypes::TimeUnit::Microsecond;
use datafusion::arrow::datatypes::{
DataType, Date32Type, DurationMicrosecondType, Int32Type, IntervalMonthDayNanoType,
IntervalYearMonthType, TimestampMicrosecondType,
};
use datafusion_common::Result;
use datafusion_expr::{
ColumnarValue, Operator, ScalarFunctionArgs, ScalarUDFImpl, Signature, Volatility,
};

use crate::error::invalid_arg_count_exec_err;
use crate::scalar::math::utils::decimal::{DecimalBinaryOp, decimal_binary_op};
use crate::scalar::math::utils::try_op::{
add_months, arith_input_types, arith_result_type, try_add_interval_monthdaynano,
try_arrow_arith, try_binary_op_date32_i32, try_op_date32_interval_yearmonth,
try_op_date32_monthdaynano, try_op_interval_yearmonth, try_op_timestamp_duration,
};

/// Spark `+` and `try_add`, unified. `safe = true` is `try_add` (any overflow →
/// NULL, ANSI-invariant); `safe = false` is `+` and honors `ansi_mode` (overflow
/// → `ARITHMETIC_OVERFLOW` under ANSI, two's-complement wrap for integrals /
/// NULL for decimals under non-ANSI). Numeric coercion delegates to
/// `BinaryTypeCoercer` for `Operator::Plus`; date/interval/timestamp arithmetic
/// (only routed here for the `safe` / `try_add` path) reuses the checked
/// `try_op` kernels.
#[derive(Debug, PartialEq, Eq, Hash)]
pub struct SparkAdd {
signature: Signature,
ansi_mode: bool,
safe: bool,
}

impl Default for SparkAdd {
fn default() -> Self {
Self::new(false, false)
}
}

impl SparkAdd {
pub fn new(ansi_mode: bool, safe: bool) -> Self {
Self {
signature: Signature::user_defined(Volatility::Immutable),
ansi_mode,
safe,
}
}

pub fn ansi_mode(&self) -> bool {
self.ansi_mode
}

pub fn safe(&self) -> bool {
self.safe
}

/// A decimal/precision overflow raises only for `+` under ANSI mode; `try_add`
/// and non-ANSI `+` turn it into NULL.
fn error_on_overflow(&self) -> bool {
self.ansi_mode && !self.safe
}
}

impl ScalarUDFImpl for SparkAdd {
fn name(&self) -> &str {
if self.safe { "try_add" } else { "spark_add" }
}

fn signature(&self) -> &Signature {
&self.signature
}

fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
match arg_types {
[DataType::Date32, _] | [_, DataType::Date32] => Ok(DataType::Date32),
[DataType::Interval(YearMonth), _] | [_, DataType::Interval(YearMonth)] => {
Ok(DataType::Interval(YearMonth))
}
[DataType::Interval(MonthDayNano), DataType::Int32]
| [DataType::Int32, DataType::Interval(MonthDayNano)]
| [DataType::Interval(MonthDayNano), DataType::Int64]
| [DataType::Int64, DataType::Interval(MonthDayNano)]
| [
DataType::Interval(MonthDayNano),
DataType::Interval(MonthDayNano),
] => Ok(DataType::Interval(MonthDayNano)),
[
DataType::Timestamp(Microsecond, _),
DataType::Duration(Microsecond),
] => Ok(DataType::Timestamp(Microsecond, None)),
[left, right] => arith_result_type(left, Operator::Plus, right),
_ => Err(invalid_arg_count_exec_err(
"spark_add",
(2, 2),
arg_types.len(),
)),
}
}

fn coerce_types(&self, arg_types: &[DataType]) -> Result<Vec<DataType>> {
let [left, right] = arg_types else {
return Err(invalid_arg_count_exec_err(
"spark_add",
(2, 2),
arg_types.len(),
));
};
if *left == DataType::Null {
return Ok(vec![right.clone(), right.clone()]);
} else if *right == DataType::Null {
return Ok(vec![left.clone(), left.clone()]);
}
let temporal = matches!(
(left, right),
(DataType::Date32, DataType::Int32)
| (DataType::Date32, DataType::Interval(YearMonth))
| (DataType::Date32, DataType::Interval(MonthDayNano))
| (DataType::Interval(YearMonth), DataType::Date32)
| (DataType::Interval(YearMonth), DataType::Interval(YearMonth))
| (
DataType::Timestamp(Microsecond, _),
DataType::Duration(Microsecond)
)
| (
DataType::Interval(MonthDayNano),
DataType::Interval(MonthDayNano)
)
);
if temporal {
return Ok(vec![left.clone(), right.clone()]);
}
let (left, right) = arith_input_types(left, Operator::Plus, right)?;
Ok(vec![left, right])
}

fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
let return_type = args.return_field.data_type().clone();
let args = ColumnarValue::values_to_arrays(&args.args)?;
let [left, right] = args.as_slice() else {
return Err(invalid_arg_count_exec_err("spark_add", (2, 2), args.len()));
};

let result: ArrayRef = match (left.data_type(), right.data_type()) {
// Date/interval/timestamp: only reached via the `safe` (try_add) path;
// regular `+` handles these natively in the plan builder.
(DataType::Date32, DataType::Int32) => {
let l = left.as_primitive::<Date32Type>();
let r = right.as_primitive::<Int32Type>();
Arc::new(try_binary_op_date32_i32(l, r, i32::checked_add))
}
(DataType::Date32, DataType::Interval(YearMonth)) => {
let l = left.as_primitive::<Date32Type>();
let r = right.as_primitive::<IntervalYearMonthType>();
Arc::new(try_op_date32_interval_yearmonth(l, r, add_months))
}
(DataType::Date32, DataType::Interval(MonthDayNano)) => {
let l = left.as_primitive::<Date32Type>();
let r = right.as_primitive::<IntervalMonthDayNanoType>();
Arc::new(try_op_date32_monthdaynano(l, r, |x| x))
}
(DataType::Interval(YearMonth), DataType::Interval(YearMonth)) => {
let l = left.as_primitive::<IntervalYearMonthType>();
let r = right.as_primitive::<IntervalYearMonthType>();
Arc::new(try_op_interval_yearmonth(l, r, i32::checked_add))
}
(DataType::Interval(MonthDayNano), DataType::Interval(MonthDayNano)) => {
let l = left.as_primitive::<IntervalMonthDayNanoType>();
let r = right.as_primitive::<IntervalMonthDayNanoType>();
Arc::new(try_add_interval_monthdaynano(l, r))
}
(DataType::Timestamp(Microsecond, _), DataType::Duration(Microsecond)) => {
let l = left.as_primitive::<TimestampMicrosecondType>();
let r = right.as_primitive::<DurationMicrosecondType>();
Arc::new(try_op_timestamp_duration(l, r, i64::checked_add))
}
// Decimal: Spark precision rules + overflow disposition per mode.
(DataType::Decimal128(..), DataType::Decimal128(..)) => decimal_binary_op(
left,
right,
DecimalBinaryOp::Add,
&return_type,
self.error_on_overflow(),
)?,
// Integral / float: safe → per-element NULL; ANSI → checked (error);
// non-ANSI → wrapping (float is unaffected by wrap/checked).
_ => {
if self.safe {
try_arrow_arith(left, right, &return_type, add)?
} else if self.ansi_mode {
add(left, right)?
} else {
add_wrapping(left, right)?
}
}
};

Ok(ColumnarValue::Array(result))
}
}
Loading
Loading