diff --git a/arrow-arith/src/numeric.rs b/arrow-arith/src/numeric.rs index cc94006fae70..d0248dec1e37 100644 --- a/arrow-arith/src/numeric.rs +++ b/arrow-arith/src/numeric.rs @@ -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}; @@ -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 + ) } } @@ -243,15 +248,10 @@ fn arithmetic_op(op: Op, lhs: &dyn Datum, rhs: &dyn Datum) -> Result duration_op::(op, l, l_scalar, r, r_scalar), (Duration(Microsecond), Duration(Microsecond)) => duration_op::(op, l, l_scalar, r, r_scalar), (Duration(Nanosecond), Duration(Nanosecond)) => duration_op::(op, l, l_scalar, r, r_scalar), - (Interval(YearMonth), Int64) if matches!(op, Op::Mul) => interval_mul_op::(l, l_scalar, r, r_scalar), - (Interval(DayTime), Int64) if matches!(op, Op::Mul) => interval_mul_op::(l, l_scalar, r, r_scalar), - (Interval(MonthDayNano), Int64) if matches!(op, Op::Mul) => interval_mul_op::(l, l_scalar, r, r_scalar), - (Int64, Interval(YearMonth)) if matches!(op, Op::Mul) => interval_mul_op::(r, r_scalar, l, l_scalar), - (Int64, Interval(DayTime)) if matches!(op, Op::Mul) => interval_mul_op::(r, r_scalar, l, l_scalar), - (Int64, Interval(MonthDayNano)) if matches!(op, Op::Mul) => interval_mul_op::(r, r_scalar, l, l_scalar), - (Interval(YearMonth), Interval(YearMonth)) => interval_op::(op, l, l_scalar, r, r_scalar), - (Interval(DayTime), Interval(DayTime)) => interval_op::(op, l, l_scalar, r, r_scalar), - (Interval(MonthDayNano), Interval(MonthDayNano)) => interval_op::(op, l, l_scalar, r, r_scalar), + (Interval(YearMonth), Interval(YearMonth) | Int64) => interval_op::(op, l, l_scalar, r, r_scalar), + (Interval(DayTime), Interval(DayTime) | Int64) => interval_op::(op, l, l_scalar, r, r_scalar), + (Interval(MonthDayNano), Interval(MonthDayNano) | Int64) => interval_op::(op, l, l_scalar, r, r_scalar), + (Interval(MonthDayNano), Float64) => interval_f64_op(op, l, l_scalar, r, r_scalar), (Date32, _) => date_op::(op, l, l_scalar, r, r_scalar), (Date64, _) => date_op::(op, l, l_scalar, r, r_scalar), (Decimal32(_, _), Decimal32(_, _)) => decimal_op::(op, l, l_scalar, r, r_scalar), @@ -262,6 +262,11 @@ fn arithmetic_op(op: Op, lhs: &dyn Datum, rhs: &dyn Datum) -> Result { 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}") )) @@ -724,6 +729,125 @@ fn interval_mul_op( )) } +/// 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). +/// +/// +/// 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 { + 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"))?; + + 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 { + let interval = interval.as_primitive::(); + let factor = factor.as_primitive::(); + 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( op: Op, @@ -732,11 +856,18 @@ fn interval_op( r: &dyn Array, r_s: bool, ) -> Result { - let l = l.as_primitive::(); - let r = r.as_primitive::(); - 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::(); + let r = r.as_primitive::(); + 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::(); + let r = r.as_primitive::(); + Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub(l, r))) + } + (Op::Mul, DataType::Int64) => interval_mul_op::(l, l_s, r, r_s), _ => Err(ArrowError::InvalidArgumentError(format!( "Invalid interval arithmetic operation: {} {op} {}", l.data_type(), @@ -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]);