Skip to content
Open
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
331 changes: 315 additions & 16 deletions arrow-arith/src/numeric.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,11 +22,13 @@ use std::fmt::Formatter;
use std::sync::Arc;

use arrow_array::cast::AsArray;
use arrow_array::temporal_conversions::{NANOSECONDS, SECONDS_IN_DAY};
use arrow_array::timezone::Tz;
use arrow_array::types::*;
use arrow_array::*;
use arrow_buffer::{ArrowNativeType, IntervalDayTime, IntervalMonthDayNano};
use arrow_schema::{ArrowError, DataType, IntervalUnit, TimeUnit};
use num_traits::ToPrimitive;

use crate::arity::{binary, try_binary};

Expand Down Expand Up @@ -212,7 +214,10 @@ impl std::fmt::Display for Op {

impl Op {
fn commutative(&self) -> bool {
matches!(self, Self::Add | Self::AddWrapping)
matches!(
self,
Self::Add | Self::AddWrapping | Self::Mul | Self::MulWrapping
)
}
}

Expand Down Expand Up @@ -243,15 +248,10 @@ fn arithmetic_op(op: Op, lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, A
(Duration(Millisecond), Duration(Millisecond)) => duration_op::<DurationMillisecondType>(op, l, l_scalar, r, r_scalar),
(Duration(Microsecond), Duration(Microsecond)) => duration_op::<DurationMicrosecondType>(op, l, l_scalar, r, r_scalar),
(Duration(Nanosecond), Duration(Nanosecond)) => duration_op::<DurationNanosecondType>(op, l, l_scalar, r, r_scalar),
(Interval(YearMonth), Int64) if matches!(op, Op::Mul) => interval_mul_op::<IntervalYearMonthType>(l, l_scalar, r, r_scalar),
(Interval(DayTime), Int64) if matches!(op, Op::Mul) => interval_mul_op::<IntervalDayTimeType>(l, l_scalar, r, r_scalar),
(Interval(MonthDayNano), Int64) if matches!(op, Op::Mul) => interval_mul_op::<IntervalMonthDayNanoType>(l, l_scalar, r, r_scalar),
(Int64, Interval(YearMonth)) if matches!(op, Op::Mul) => interval_mul_op::<IntervalYearMonthType>(r, r_scalar, l, l_scalar),
(Int64, Interval(DayTime)) if matches!(op, Op::Mul) => interval_mul_op::<IntervalDayTimeType>(r, r_scalar, l, l_scalar),
(Int64, Interval(MonthDayNano)) if matches!(op, Op::Mul) => interval_mul_op::<IntervalMonthDayNanoType>(r, r_scalar, l, l_scalar),
(Interval(YearMonth), Interval(YearMonth)) => interval_op::<IntervalYearMonthType>(op, l, l_scalar, r, r_scalar),
(Interval(DayTime), Interval(DayTime)) => interval_op::<IntervalDayTimeType>(op, l, l_scalar, r, r_scalar),
(Interval(MonthDayNano), Interval(MonthDayNano)) => interval_op::<IntervalMonthDayNanoType>(op, l, l_scalar, r, r_scalar),
(Interval(YearMonth), Interval(YearMonth) | Int64) => interval_op::<IntervalYearMonthType>(op, l, l_scalar, r, r_scalar),
(Interval(DayTime), Interval(DayTime) | Int64) => interval_op::<IntervalDayTimeType>(op, l, l_scalar, r, r_scalar),
(Interval(MonthDayNano), Interval(MonthDayNano) | Int64) => interval_op::<IntervalMonthDayNanoType>(op, l, l_scalar, r, r_scalar),
(Interval(MonthDayNano), Float64) => interval_f64_op(op, l, l_scalar, r, r_scalar),
(Date32, _) => date_op::<Date32Type>(op, l, l_scalar, r, r_scalar),
(Date64, _) => date_op::<Date64Type>(op, l, l_scalar, r, r_scalar),
(Decimal32(_, _), Decimal32(_, _)) => decimal_op::<Decimal32Type>(op, l, l_scalar, r, r_scalar),
Expand All @@ -262,6 +262,11 @@ fn arithmetic_op(op: Op, lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, A
(Duration(_) | Interval(_), Date32 | Date64 | Timestamp(_, _)) if op.commutative() => {
arithmetic_op(op, rhs, lhs)
}
(Int64, Interval(_)) | (Float64, Interval(MonthDayNano))
if matches!(op, Op::Mul) =>
{
arithmetic_op(op, rhs, lhs)
}
_ => Err(ArrowError::InvalidArgumentError(
format!("Invalid arithmetic operation: {l_t} {op} {r_t}")
))
Expand Down Expand Up @@ -724,6 +729,125 @@ fn interval_mul_op<T: IntervalOp>(
))
}

/// Multiplies an `IntervalMonthDayNano` by an `f64`, mirroring DuckDB's
/// `interval_t` layout of months, days, and a sub-day component (nanoseconds in
/// Arrow, microseconds in DuckDB).
/// <https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/src/include/duckdb/common/types/interval.hpp#L24-L27>
///
/// Algorithm:
///
/// 1. use checked integer multiplication when `factor` fits in `i64` (early return).
/// 2. multiply months and days separately and truncate floating point.
/// 3. cascade remainders: convert fractional months to days using 30
/// days per month, then fractional days to a sub-day value using 24 hours
/// per day.
/// 4. combine the cascaded remainder with the scaled input nanoseconds and round ties-to-even at nanosecond precision.
/// 5. return an overflow error if any output component is out of
/// range.
fn interval_mul_f64(
interval: IntervalMonthDayNano,
factor: f64,
) -> Result<IntervalMonthDayNano, ArrowError> {
const DAYS_PER_MONTH: f64 = 30.;
const NANOS_PER_SECOND: f64 = NANOSECONDS as f64;
const SECONDS_PER_DAY: f64 = SECONDS_IN_DAY as f64;

// Keep integral factors exact instead of round-tripping i64 nanoseconds through f64.
if factor.fract() == 0. {
if let Some(factor) = ToPrimitive::to_i64(&factor) {
return IntervalMonthDayNanoType::mul_i64(interval, factor);
}
}

// Based on DuckDB's INTERVAL * DOUBLE implementation, which is referenced from PostgreSQL's interval_mul:
// https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/src/function/scalar/operator/multiply.cpp#L48-L123
// PostgreSQL's interval_mul:
// https://github.com/postgres/postgres/blob/78758d37306cd89ab060f00cb06f249018d5b8da/src/backend/utils/adt/timestamp.c#L3627-L3744
let overflow =
|component| ArrowError::ArithmeticOverflow(format!("Overflow in interval {component}"));
let timestamp_round =
|value: f64| (value * NANOS_PER_SECOND).round_ties_even() / NANOS_PER_SECOND;

let months_product = f64::from(interval.months) * factor;
if !months_product.is_finite()
|| months_product < f64::from(i32::MIN)
|| months_product > f64::from(i32::MAX)
{
return Err(overflow("months"));
}
let months = months_product.to_i32().ok_or_else(|| overflow("months"))?;
Comment on lines +772 to +778

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

does to_i32() not check the range and non-finite numbers when converting? trying to understand why we have two checks here

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If only relying on to_i32(), i32::MIN * 1.000_000_000_4 truncates to valid i32::MIN, so it no longer overflow.

let factor = Float64Array::new_scalar(1.000_000_000_4);
let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
i32::MIN,
0,
0,
));
assert!(matches!(
mul(&interval, &factor),
Err(ArrowError::ArithmeticOverflow(_))
));


let days_product = f64::from(interval.days) * factor;
if !days_product.is_finite()
|| days_product < f64::from(i32::MIN)
|| days_product > f64::from(i32::MAX)
{
return Err(overflow("days"));
}
let mut days = days_product.to_i32().ok_or_else(|| overflow("days"))?;

let month_remainder = timestamp_round(months_product.fract() * DAYS_PER_MONTH);
let month_remainder_days = month_remainder
.to_i32()
.ok_or_else(|| overflow("month remainder"))?;
let mut seconds_remainder = timestamp_round(
(days_product - f64::from(days) + month_remainder - f64::from(month_remainder_days))
* SECONDS_PER_DAY,
);

if seconds_remainder.abs() >= SECONDS_PER_DAY {
let remainder_days = (seconds_remainder / SECONDS_PER_DAY)
.to_i32()
.ok_or_else(|| overflow("day remainder"))?;
days = days
.checked_add(remainder_days)
.ok_or_else(|| overflow("days"))?;
seconds_remainder -= f64::from(remainder_days) * SECONDS_PER_DAY;
}
days = days
.checked_add(month_remainder_days)
.ok_or_else(|| overflow("days"))?;

let nanoseconds = ((interval.nanoseconds as f64) * factor
+ seconds_remainder * NANOS_PER_SECOND)
.round_ties_even();
let nanoseconds = ToPrimitive::to_i64(&nanoseconds).ok_or_else(|| {
ArrowError::ArithmeticOverflow(format!("Overflow in interval nanoseconds: {nanoseconds}"))
})?;

Ok(IntervalMonthDayNano::new(months, days, nanoseconds))
}

fn interval_f64_op(
op: Op,
interval: &dyn Array,
interval_scalar: bool,
factor: &dyn Array,
factor_scalar: bool,
) -> Result<ArrayRef, ArrowError> {
let interval = interval.as_primitive::<IntervalMonthDayNanoType>();
let factor = factor.as_primitive::<Float64Type>();
Ok(try_op_ref!(
IntervalMonthDayNanoType,
interval,
interval_scalar,
factor,
factor_scalar,
{
match op {
Op::Mul => interval_mul_f64(interval, factor),
Op::Div if factor == 0. => Err(ArrowError::DivideByZero),
// DuckDB defines interval division as multiplication by the reciprocal:
// https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/src/function/scalar/operator/arithmetic.cpp#L1102-L1110
Op::Div => interval_mul_f64(interval, 1. / factor),
_ => Err(ArrowError::InvalidArgumentError(format!(
"Invalid interval arithmetic operation: Interval(MonthDayNano) {op} Float64"
))),
}
}
))
}

/// Perform arithmetic operation on an interval array
fn interval_op<T: IntervalOp>(
op: Op,
Expand All @@ -732,11 +856,18 @@ fn interval_op<T: IntervalOp>(
r: &dyn Array,
r_s: bool,
) -> Result<ArrayRef, ArrowError> {
let l = l.as_primitive::<T>();
let r = r.as_primitive::<T>();
match op {
Op::Add | Op::AddWrapping => Ok(try_op_ref!(T, l, l_s, r, r_s, T::add(l, r))),
Op::Sub | Op::SubWrapping => Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub(l, r))),
match (op, r.data_type()) {
(Op::Add | Op::AddWrapping, data_type) if data_type == l.data_type() => {
let l = l.as_primitive::<T>();
let r = r.as_primitive::<T>();
Ok(try_op_ref!(T, l, l_s, r, r_s, T::add(l, r)))
}
(Op::Sub | Op::SubWrapping, data_type) if data_type == l.data_type() => {
let l = l.as_primitive::<T>();
let r = r.as_primitive::<T>();
Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub(l, r)))
}
(Op::Mul, DataType::Int64) => interval_mul_op::<T>(l, l_s, r, r_s),
_ => Err(ArrowError::InvalidArgumentError(format!(
"Invalid interval arithmetic operation: {} {op} {}",
l.data_type(),
Expand Down Expand Up @@ -1648,10 +1779,178 @@ mod tests {
assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);

let float_factor = Float64Array::new_scalar(2.);
assert!(mul(&interval, &float_factor).is_err());
assert!(mul_wrapping(&float_factor, &interval).is_err());
assert!(mul_wrapping(&factor, &interval).is_err());
}

#[test]
fn test_interval_mul_div_f64() {
const HOUR_NANOS: i64 = 3_600_000_000_000;
const MINUTE_NANOS: i64 = 60_000_000_000;

// Adapted from DuckDB's interval multiplication tests:
// https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/test/sql/function/interval/test_interval_muldiv.test#L1-L99
// DuckDB's cases come from PostgreSQL's interval regression tests:
// https://github.com/postgres/postgres/blob/78758d37306cd89ab060f00cb06f249018d5b8da/src/test/regress/sql/interval.sql#L118-L164
let interval = IntervalMonthDayNanoArray::from(vec![
IntervalMonthDayNanoType::make_value(41, 12, 360 * HOUR_NANOS),
IntervalMonthDayNanoType::make_value(-41, -12, 360 * HOUR_NANOS),
IntervalMonthDayNanoType::make_value(1, 1, 0),
IntervalMonthDayNanoType::make_value(0, 0, 1),
IntervalMonthDayNanoType::make_value(0, 0, 3),
IntervalMonthDayNanoType::make_value(0, 0, -1),
IntervalMonthDayNanoType::make_value(0, 0, -3),
]);
let factor = Float64Array::from(vec![0.3, 0.3, 1.5, 0.5, 0.5, 0.5, 0.5]);
let expected = IntervalMonthDayNanoArray::from(vec![
IntervalMonthDayNanoType::make_value(12, 12, 122 * HOUR_NANOS + 24 * MINUTE_NANOS),
IntervalMonthDayNanoType::make_value(-12, -12, 93 * HOUR_NANOS + 36 * MINUTE_NANOS),
IntervalMonthDayNanoType::make_value(1, 16, 12 * HOUR_NANOS),
IntervalMonthDayNanoType::make_value(0, 0, 0),
IntervalMonthDayNanoType::make_value(0, 0, 2),
IntervalMonthDayNanoType::make_value(0, 0, 0),
IntervalMonthDayNanoType::make_value(0, 0, -2),
]);
assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);

let interval = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
9,
-27,
45_296 * NANOSECONDS,
)]);
let factor = Float64Array::new_scalar(0.3);
let expected = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
2,
13,
4_948_800_000_000,
)]);
assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);

let interval = IntervalMonthDayNanoArray::from(vec![
IntervalMonthDayNanoType::make_value(0, 1, 0),
IntervalMonthDayNanoType::make_value(4, 0, 0),
IntervalMonthDayNanoType::make_value(1, 1, 0),
IntervalMonthDayNanoType::make_value(0, 0, (1_i64 << 53) - 1),
IntervalMonthDayNanoType::make_value(0, 0, i64::MAX),
IntervalMonthDayNanoType::make_value(1, 1, 1),
IntervalMonthDayNanoType::make_value(1, 1, 1),
IntervalMonthDayNanoType::make_value(1, 0, 0),
]);
let factor = Float64Array::from(vec![
3.,
5.,
2.,
0.7,
1.,
f64::INFINITY,
f64::NEG_INFINITY,
-2.,
]);
let expected = IntervalMonthDayNanoArray::from(vec![
IntervalMonthDayNanoType::make_value(0, 0, 8 * HOUR_NANOS),
IntervalMonthDayNanoType::make_value(0, 24, 0),
IntervalMonthDayNanoType::make_value(0, 15, 12 * HOUR_NANOS),
IntervalMonthDayNanoType::make_value(0, 0, 12_867_427_506_772_844),
IntervalMonthDayNanoType::make_value(0, 0, i64::MAX),
IntervalMonthDayNanoType::make_value(0, 0, 0),
IntervalMonthDayNanoType::make_value(0, 0, 0),
IntervalMonthDayNanoType::make_value(0, -15, 0),
]);
assert_eq!(div(&interval, &factor).unwrap().as_ref(), &expected);

let null_factor = Scalar::new(Float64Array::new_null(1));
assert_eq!(
mul(&interval, &null_factor).unwrap().as_ref(),
&IntervalMonthDayNanoArray::new_null(interval.len())
);
}

#[test]
fn test_interval_mul_div_f64_errors() {
let factor = Float64Array::new_scalar(2.);
let year_month = IntervalYearMonthArray::new_scalar(1);
let day_time = IntervalDayTimeArray::new_scalar(IntervalDayTime::new(1, 1));
for interval in [&year_month as &dyn Datum, &day_time] {
assert!(mul(interval, &factor).is_err());
assert!(mul(&factor, interval).is_err());
assert!(div(interval, &factor).is_err());
}

let interval =
IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(1, 1, 1));

assert!(matches!(
add(&interval, &factor),
Err(ArrowError::InvalidArgumentError(_))
));

let zero = Float64Array::new_scalar(-0.);
assert!(matches!(
div(&interval, &zero),
Err(ArrowError::DivideByZero)
));

assert!(div(&factor, &interval).is_err());

for factor in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
let factor = Float64Array::new_scalar(factor);
assert!(matches!(
mul(&interval, &factor),
Err(ArrowError::ArithmeticOverflow(_))
));
}

let nan = Float64Array::new_scalar(f64::NAN);
assert!(matches!(
div(&interval, &nan),
Err(ArrowError::ArithmeticOverflow(_))
));

let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
i32::MAX,
0,
0,
));
assert!(matches!(
mul(&interval, &factor),
Err(ArrowError::ArithmeticOverflow(_))
));

let factor = Float64Array::new_scalar(1.5);
let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
0,
0,
i64::MAX,
));
assert!(matches!(
mul(&interval, &factor),
Err(ArrowError::ArithmeticOverflow(_))
));

let factor = Float64Array::new_scalar(1.000_000_000_4);
let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
i32::MIN,
0,
0,
));
assert!(matches!(
mul(&interval, &factor),
Err(ArrowError::ArithmeticOverflow(_))
));

let factor = Float64Array::new_scalar(0.999_999_999);
let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
1,
i32::MAX,
0,
));
assert!(matches!(
mul(&interval, &factor),
Err(ArrowError::ArithmeticOverflow(_))
));
}

#[test]
fn test_interval_mul_i64_overflow() {
let interval = IntervalYearMonthArray::from(vec![i32::MAX]);
Expand Down
Loading