Skip to main content

arrow_arith/
numeric.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! Defines numeric arithmetic kernels on [`PrimitiveArray`], such as [`add`]
19
20use std::cmp::Ordering;
21use std::fmt::Formatter;
22use std::sync::Arc;
23
24use arrow_array::cast::AsArray;
25use arrow_array::temporal_conversions::{NANOSECONDS, SECONDS_IN_DAY};
26use arrow_array::timezone::Tz;
27use arrow_array::types::*;
28use arrow_array::*;
29use arrow_buffer::{ArrowNativeType, IntervalDayTime, IntervalMonthDayNano};
30use arrow_schema::{ArrowError, DataType, IntervalUnit, TimeUnit};
31use num_traits::ToPrimitive;
32
33use crate::arity::{binary, try_binary};
34
35/// Perform `lhs + rhs`, returning an error on overflow
36pub fn add(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
37    arithmetic_op(Op::Add, lhs, rhs)
38}
39
40/// Perform `lhs + rhs`, wrapping on overflow for [`DataType::is_integer`]
41pub fn add_wrapping(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
42    arithmetic_op(Op::AddWrapping, lhs, rhs)
43}
44
45/// Perform `lhs - rhs`, returning an error on overflow
46pub fn sub(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
47    arithmetic_op(Op::Sub, lhs, rhs)
48}
49
50/// Perform `lhs - rhs`, wrapping on overflow for [`DataType::is_integer`]
51pub fn sub_wrapping(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
52    arithmetic_op(Op::SubWrapping, lhs, rhs)
53}
54
55/// Perform `lhs * rhs`, returning an error on overflow
56pub fn mul(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
57    arithmetic_op(Op::Mul, lhs, rhs)
58}
59
60/// Perform `lhs * rhs`, wrapping on overflow for [`DataType::is_integer`]
61pub fn mul_wrapping(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
62    arithmetic_op(Op::MulWrapping, lhs, rhs)
63}
64
65/// Perform `lhs / rhs`
66///
67/// Overflow or division by zero will result in an error, with exception to
68/// floating point numbers, which instead follow the IEEE 754 rules
69pub fn div(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
70    arithmetic_op(Op::Div, lhs, rhs)
71}
72
73/// Perform `lhs % rhs`
74///
75/// Division by zero will result in an error, with exception to
76/// floating point numbers, which instead follow the IEEE 754 rules
77///
78/// `signed_integer::MIN % -1` will not result in an error but return 0
79pub fn rem(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
80    arithmetic_op(Op::Rem, lhs, rhs)
81}
82
83macro_rules! neg_checked {
84    ($t:ty, $a:ident) => {{
85        let array = $a
86            .as_primitive::<$t>()
87            .try_unary::<_, $t, _>(|x| x.neg_checked())?;
88        Ok(Arc::new(array))
89    }};
90}
91
92macro_rules! neg_wrapping {
93    ($t:ty, $a:ident) => {{
94        let array = $a.as_primitive::<$t>().unary::<_, $t>(|x| x.neg_wrapping());
95        Ok(Arc::new(array))
96    }};
97}
98
99/// Negates each element of  `array`, returning an error on overflow
100///
101/// Note: negation of unsigned arrays is not supported and will return in an error,
102/// for wrapping unsigned negation consider using [`neg_wrapping`][neg_wrapping()]
103pub fn neg(array: &dyn Array) -> Result<ArrayRef, ArrowError> {
104    use DataType::*;
105    use IntervalUnit::*;
106    use TimeUnit::*;
107
108    match array.data_type() {
109        Int8 => neg_checked!(Int8Type, array),
110        Int16 => neg_checked!(Int16Type, array),
111        Int32 => neg_checked!(Int32Type, array),
112        Int64 => neg_checked!(Int64Type, array),
113        Float16 => neg_wrapping!(Float16Type, array),
114        Float32 => neg_wrapping!(Float32Type, array),
115        Float64 => neg_wrapping!(Float64Type, array),
116        Decimal32(p, s) => {
117            let a = array
118                .as_primitive::<Decimal32Type>()
119                .try_unary::<_, Decimal32Type, _>(|x| x.neg_checked())?;
120
121            Ok(Arc::new(a.with_precision_and_scale(*p, *s)?))
122        }
123        Decimal64(p, s) => {
124            let a = array
125                .as_primitive::<Decimal64Type>()
126                .try_unary::<_, Decimal64Type, _>(|x| x.neg_checked())?;
127
128            Ok(Arc::new(a.with_precision_and_scale(*p, *s)?))
129        }
130        Decimal128(p, s) => {
131            let a = array
132                .as_primitive::<Decimal128Type>()
133                .try_unary::<_, Decimal128Type, _>(|x| x.neg_checked())?;
134
135            Ok(Arc::new(a.with_precision_and_scale(*p, *s)?))
136        }
137        Decimal256(p, s) => {
138            let a = array
139                .as_primitive::<Decimal256Type>()
140                .try_unary::<_, Decimal256Type, _>(|x| x.neg_checked())?;
141
142            Ok(Arc::new(a.with_precision_and_scale(*p, *s)?))
143        }
144        Duration(Second) => neg_checked!(DurationSecondType, array),
145        Duration(Millisecond) => neg_checked!(DurationMillisecondType, array),
146        Duration(Microsecond) => neg_checked!(DurationMicrosecondType, array),
147        Duration(Nanosecond) => neg_checked!(DurationNanosecondType, array),
148        Interval(YearMonth) => neg_checked!(IntervalYearMonthType, array),
149        Interval(DayTime) => {
150            let a = array
151                .as_primitive::<IntervalDayTimeType>()
152                .try_unary::<_, IntervalDayTimeType, ArrowError>(|x| {
153                    let (days, ms) = IntervalDayTimeType::to_parts(x);
154                    Ok(IntervalDayTimeType::make_value(
155                        days.neg_checked()?,
156                        ms.neg_checked()?,
157                    ))
158                })?;
159            Ok(Arc::new(a))
160        }
161        Interval(MonthDayNano) => {
162            let a = array
163                .as_primitive::<IntervalMonthDayNanoType>()
164                .try_unary::<_, IntervalMonthDayNanoType, ArrowError>(|x| {
165                    let (months, days, nanos) = IntervalMonthDayNanoType::to_parts(x);
166                    Ok(IntervalMonthDayNanoType::make_value(
167                        months.neg_checked()?,
168                        days.neg_checked()?,
169                        nanos.neg_checked()?,
170                    ))
171                })?;
172            Ok(Arc::new(a))
173        }
174        t => Err(ArrowError::InvalidArgumentError(format!(
175            "Invalid arithmetic operation: !{t}"
176        ))),
177    }
178}
179
180/// Negates each element of  `array`, wrapping on overflow for [`DataType::is_integer`]
181pub fn neg_wrapping(array: &dyn Array) -> Result<ArrayRef, ArrowError> {
182    downcast_integer! {
183        array.data_type() => (neg_wrapping, array),
184        _ => neg(array),
185    }
186}
187
188/// An enumeration of arithmetic operations
189///
190/// This allows sharing the type dispatch logic across the various kernels
191#[derive(Debug, Copy, Clone)]
192enum Op {
193    AddWrapping,
194    Add,
195    SubWrapping,
196    Sub,
197    MulWrapping,
198    Mul,
199    Div,
200    Rem,
201}
202
203impl std::fmt::Display for Op {
204    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
205        match self {
206            Op::AddWrapping | Op::Add => write!(f, "+"),
207            Op::SubWrapping | Op::Sub => write!(f, "-"),
208            Op::MulWrapping | Op::Mul => write!(f, "*"),
209            Op::Div => write!(f, "/"),
210            Op::Rem => write!(f, "%"),
211        }
212    }
213}
214
215impl Op {
216    fn commutative(&self) -> bool {
217        matches!(
218            self,
219            Self::Add | Self::AddWrapping | Self::Mul | Self::MulWrapping
220        )
221    }
222}
223
224/// Dispatch the given `op` to the appropriate specialized kernel
225fn arithmetic_op(op: Op, lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
226    use DataType::*;
227    use IntervalUnit::*;
228    use TimeUnit::*;
229
230    macro_rules! integer_helper {
231        ($t:ty, $op:ident, $l:ident, $l_scalar:ident, $r:ident, $r_scalar:ident) => {
232            integer_op::<$t>($op, $l, $l_scalar, $r, $r_scalar)
233        };
234    }
235
236    let (l, l_scalar) = lhs.get();
237    let (r, r_scalar) = rhs.get();
238    downcast_integer! {
239        l.data_type(), r.data_type() => (integer_helper, op, l, l_scalar, r, r_scalar),
240        (Float16, Float16) => float_op::<Float16Type>(op, l, l_scalar, r, r_scalar),
241        (Float32, Float32) => float_op::<Float32Type>(op, l, l_scalar, r, r_scalar),
242        (Float64, Float64) => float_op::<Float64Type>(op, l, l_scalar, r, r_scalar),
243        (Timestamp(Second, _), _) => timestamp_op::<TimestampSecondType>(op, l, l_scalar, r, r_scalar),
244        (Timestamp(Millisecond, _), _) => timestamp_op::<TimestampMillisecondType>(op, l, l_scalar, r, r_scalar),
245        (Timestamp(Microsecond, _), _) => timestamp_op::<TimestampMicrosecondType>(op, l, l_scalar, r, r_scalar),
246        (Timestamp(Nanosecond, _), _) => timestamp_op::<TimestampNanosecondType>(op, l, l_scalar, r, r_scalar),
247        (Duration(Second), Duration(Second)) => duration_op::<DurationSecondType>(op, l, l_scalar, r, r_scalar),
248        (Duration(Millisecond), Duration(Millisecond)) => duration_op::<DurationMillisecondType>(op, l, l_scalar, r, r_scalar),
249        (Duration(Microsecond), Duration(Microsecond)) => duration_op::<DurationMicrosecondType>(op, l, l_scalar, r, r_scalar),
250        (Duration(Nanosecond), Duration(Nanosecond)) => duration_op::<DurationNanosecondType>(op, l, l_scalar, r, r_scalar),
251        (Interval(YearMonth), Interval(YearMonth) | Int64) => interval_op::<IntervalYearMonthType>(op, l, l_scalar, r, r_scalar),
252        (Interval(DayTime), Interval(DayTime) | Int64) => interval_op::<IntervalDayTimeType>(op, l, l_scalar, r, r_scalar),
253        (Interval(MonthDayNano), Interval(MonthDayNano) | Int64) => interval_op::<IntervalMonthDayNanoType>(op, l, l_scalar, r, r_scalar),
254        (Interval(MonthDayNano), Float64) => interval_f64_op(op, l, l_scalar, r, r_scalar),
255        (Date32, _) => date_op::<Date32Type>(op, l, l_scalar, r, r_scalar),
256        (Date64, _) => date_op::<Date64Type>(op, l, l_scalar, r, r_scalar),
257        (Decimal32(_, _), Decimal32(_, _)) => decimal_op::<Decimal32Type>(op, l, l_scalar, r, r_scalar),
258        (Decimal64(_, _), Decimal64(_, _)) => decimal_op::<Decimal64Type>(op, l, l_scalar, r, r_scalar),
259        (Decimal128(_, _), Decimal128(_, _)) => decimal_op::<Decimal128Type>(op, l, l_scalar, r, r_scalar),
260        (Decimal256(_, _), Decimal256(_, _)) => decimal_op::<Decimal256Type>(op, l, l_scalar, r, r_scalar),
261        (l_t, r_t) => match (l_t, r_t) {
262            (Duration(_) | Interval(_), Date32 | Date64 | Timestamp(_, _)) if op.commutative() => {
263                arithmetic_op(op, rhs, lhs)
264            }
265            (Int64, Interval(_)) | (Float64, Interval(MonthDayNano))
266                if matches!(op, Op::Mul) =>
267            {
268                arithmetic_op(op, rhs, lhs)
269            }
270            _ => Err(ArrowError::InvalidArgumentError(
271              format!("Invalid arithmetic operation: {l_t} {op} {r_t}")
272            ))
273        }
274    }
275}
276
277/// Perform an infallible binary operation on potentially scalar inputs
278macro_rules! op {
279    ($l:ident, $l_s:expr, $r:ident, $r_s:expr, $op:expr) => {
280        match ($l_s, $r_s) {
281            (true, true) | (false, false) => binary($l, $r, |$l, $r| $op)?,
282            (true, false) => match ($l.null_count() == 0).then(|| $l.value(0)) {
283                None => PrimitiveArray::new_null($r.len()),
284                Some($l) => $r.unary(|$r| $op),
285            },
286            (false, true) => match ($r.null_count() == 0).then(|| $r.value(0)) {
287                None => PrimitiveArray::new_null($l.len()),
288                Some($r) => $l.unary(|$l| $op),
289            },
290        }
291    };
292}
293
294/// Same as `op` but with a type hint for the returned array
295macro_rules! op_ref {
296    ($t:ty, $l:ident, $l_s:expr, $r:ident, $r_s:expr, $op:expr) => {{
297        let array: PrimitiveArray<$t> = op!($l, $l_s, $r, $r_s, $op);
298        Arc::new(array)
299    }};
300}
301
302/// Perform a fallible binary operation on potentially scalar inputs
303macro_rules! try_op {
304    ($l:ident, $l_s:expr, $r:ident, $r_s:expr, $op:expr) => {
305        match ($l_s, $r_s) {
306            (true, true) | (false, false) => try_binary($l, $r, |$l, $r| $op)?,
307            (true, false) => match ($l.null_count() == 0).then(|| $l.value(0)) {
308                None => PrimitiveArray::new_null($r.len()),
309                Some($l) => $r.try_unary(|$r| $op)?,
310            },
311            (false, true) => match ($r.null_count() == 0).then(|| $r.value(0)) {
312                None => PrimitiveArray::new_null($l.len()),
313                Some($r) => $l.try_unary(|$l| $op)?,
314            },
315        }
316    };
317}
318
319/// Same as `try_op` but with a type hint for the returned array
320macro_rules! try_op_ref {
321    ($t:ty, $l:ident, $l_s:expr, $r:ident, $r_s:expr, $op:expr) => {{
322        let array: PrimitiveArray<$t> = try_op!($l, $l_s, $r, $r_s, $op);
323        Arc::new(array)
324    }};
325}
326
327/// Perform an arithmetic operation on integers
328fn integer_op<T: ArrowPrimitiveType>(
329    op: Op,
330    l: &dyn Array,
331    l_s: bool,
332    r: &dyn Array,
333    r_s: bool,
334) -> Result<ArrayRef, ArrowError> {
335    let l = l.as_primitive::<T>();
336    let r = r.as_primitive::<T>();
337    let array: PrimitiveArray<T> = match op {
338        Op::AddWrapping => op!(l, l_s, r, r_s, l.add_wrapping(r)),
339        Op::Add => try_op!(l, l_s, r, r_s, l.add_checked(r)),
340        Op::SubWrapping => op!(l, l_s, r, r_s, l.sub_wrapping(r)),
341        Op::Sub => try_op!(l, l_s, r, r_s, l.sub_checked(r)),
342        Op::MulWrapping => op!(l, l_s, r, r_s, l.mul_wrapping(r)),
343        Op::Mul => try_op!(l, l_s, r, r_s, l.mul_checked(r)),
344        Op::Div => try_op!(l, l_s, r, r_s, l.div_checked(r)),
345        Op::Rem => try_op!(l, l_s, r, r_s, {
346            if r.is_zero() {
347                Err(ArrowError::DivideByZero)
348            } else {
349                Ok(l.mod_wrapping(r))
350            }
351        }),
352    };
353    Ok(Arc::new(array))
354}
355
356/// Perform an arithmetic operation on floats
357fn float_op<T: ArrowPrimitiveType>(
358    op: Op,
359    l: &dyn Array,
360    l_s: bool,
361    r: &dyn Array,
362    r_s: bool,
363) -> Result<ArrayRef, ArrowError> {
364    let l = l.as_primitive::<T>();
365    let r = r.as_primitive::<T>();
366    let array: PrimitiveArray<T> = match op {
367        Op::AddWrapping | Op::Add => op!(l, l_s, r, r_s, l.add_wrapping(r)),
368        Op::SubWrapping | Op::Sub => op!(l, l_s, r, r_s, l.sub_wrapping(r)),
369        Op::MulWrapping | Op::Mul => op!(l, l_s, r, r_s, l.mul_wrapping(r)),
370        Op::Div => op!(l, l_s, r, r_s, l.div_wrapping(r)),
371        Op::Rem => op!(l, l_s, r, r_s, l.mod_wrapping(r)),
372    };
373    Ok(Arc::new(array))
374}
375
376/// Arithmetic trait for timestamp arrays
377trait TimestampOp: ArrowTimestampType {
378    type Duration: ArrowPrimitiveType<Native = i64>;
379
380    fn add_year_month(timestamp: i64, delta: i32, tz: Tz) -> Option<i64>;
381    fn add_day_time(timestamp: i64, delta: IntervalDayTime, tz: Tz) -> Option<i64>;
382    fn add_month_day_nano(timestamp: i64, delta: IntervalMonthDayNano, tz: Tz) -> Option<i64>;
383
384    fn sub_year_month(timestamp: i64, delta: i32, tz: Tz) -> Option<i64>;
385    fn sub_day_time(timestamp: i64, delta: IntervalDayTime, tz: Tz) -> Option<i64>;
386    fn sub_month_day_nano(timestamp: i64, delta: IntervalMonthDayNano, tz: Tz) -> Option<i64>;
387}
388
389macro_rules! timestamp {
390    ($t:ty, $d:ty) => {
391        impl TimestampOp for $t {
392            type Duration = $d;
393
394            fn add_year_month(left: i64, right: i32, tz: Tz) -> Option<i64> {
395                Self::add_year_months(left, right, tz)
396            }
397
398            fn add_day_time(left: i64, right: IntervalDayTime, tz: Tz) -> Option<i64> {
399                Self::add_day_time(left, right, tz)
400            }
401
402            fn add_month_day_nano(left: i64, right: IntervalMonthDayNano, tz: Tz) -> Option<i64> {
403                Self::add_month_day_nano(left, right, tz)
404            }
405
406            fn sub_year_month(left: i64, right: i32, tz: Tz) -> Option<i64> {
407                Self::subtract_year_months(left, right, tz)
408            }
409
410            fn sub_day_time(left: i64, right: IntervalDayTime, tz: Tz) -> Option<i64> {
411                Self::subtract_day_time(left, right, tz)
412            }
413
414            fn sub_month_day_nano(left: i64, right: IntervalMonthDayNano, tz: Tz) -> Option<i64> {
415                Self::subtract_month_day_nano(left, right, tz)
416            }
417        }
418    };
419}
420timestamp!(TimestampSecondType, DurationSecondType);
421timestamp!(TimestampMillisecondType, DurationMillisecondType);
422timestamp!(TimestampMicrosecondType, DurationMicrosecondType);
423timestamp!(TimestampNanosecondType, DurationNanosecondType);
424
425/// Perform arithmetic operation on a timestamp array
426fn timestamp_op<T: TimestampOp>(
427    op: Op,
428    l: &dyn Array,
429    l_s: bool,
430    r: &dyn Array,
431    r_s: bool,
432) -> Result<ArrayRef, ArrowError> {
433    use DataType::*;
434    use IntervalUnit::*;
435
436    let l = l.as_primitive::<T>();
437    let l_tz: Tz = l.timezone().unwrap_or("+00:00").parse()?;
438
439    let array: PrimitiveArray<T> = match (op, r.data_type()) {
440        (Op::Sub | Op::SubWrapping, Timestamp(unit, _)) if unit == &T::UNIT => {
441            let r = r.as_primitive::<T>();
442            return Ok(try_op_ref!(T::Duration, l, l_s, r, r_s, l.sub_checked(r)));
443        }
444
445        (Op::Add | Op::AddWrapping, Duration(unit)) if unit == &T::UNIT => {
446            let r = r.as_primitive::<T::Duration>();
447            try_op!(l, l_s, r, r_s, l.add_checked(r))
448        }
449        (Op::Sub | Op::SubWrapping, Duration(unit)) if unit == &T::UNIT => {
450            let r = r.as_primitive::<T::Duration>();
451            try_op!(l, l_s, r, r_s, l.sub_checked(r))
452        }
453
454        (Op::Add | Op::AddWrapping, Interval(YearMonth)) => {
455            let r = r.as_primitive::<IntervalYearMonthType>();
456            try_op!(
457                l,
458                l_s,
459                r,
460                r_s,
461                T::add_year_month(l, r, l_tz).ok_or(ArrowError::ComputeError(
462                    "Timestamp out of range".to_string()
463                ))
464            )
465        }
466        (Op::Sub | Op::SubWrapping, Interval(YearMonth)) => {
467            let r = r.as_primitive::<IntervalYearMonthType>();
468            try_op!(
469                l,
470                l_s,
471                r,
472                r_s,
473                T::sub_year_month(l, r, l_tz).ok_or(ArrowError::ComputeError(
474                    "Timestamp out of range".to_string()
475                ))
476            )
477        }
478
479        (Op::Add | Op::AddWrapping, Interval(DayTime)) => {
480            let r = r.as_primitive::<IntervalDayTimeType>();
481            try_op!(
482                l,
483                l_s,
484                r,
485                r_s,
486                T::add_day_time(l, r, l_tz).ok_or(ArrowError::ComputeError(
487                    "Timestamp out of range".to_string()
488                ))
489            )
490        }
491        (Op::Sub | Op::SubWrapping, Interval(DayTime)) => {
492            let r = r.as_primitive::<IntervalDayTimeType>();
493            try_op!(
494                l,
495                l_s,
496                r,
497                r_s,
498                T::sub_day_time(l, r, l_tz).ok_or(ArrowError::ComputeError(
499                    "Timestamp out of range".to_string()
500                ))
501            )
502        }
503
504        (Op::Add | Op::AddWrapping, Interval(MonthDayNano)) => {
505            let r = r.as_primitive::<IntervalMonthDayNanoType>();
506            try_op!(
507                l,
508                l_s,
509                r,
510                r_s,
511                T::add_month_day_nano(l, r, l_tz).ok_or(ArrowError::ComputeError(
512                    "Timestamp out of range".to_string()
513                ))
514            )
515        }
516        (Op::Sub | Op::SubWrapping, Interval(MonthDayNano)) => {
517            let r = r.as_primitive::<IntervalMonthDayNanoType>();
518            try_op!(
519                l,
520                l_s,
521                r,
522                r_s,
523                T::sub_month_day_nano(l, r, l_tz).ok_or(ArrowError::ComputeError(
524                    "Timestamp out of range".to_string()
525                ))
526            )
527        }
528        _ => {
529            return Err(ArrowError::InvalidArgumentError(format!(
530                "Invalid timestamp arithmetic operation: {} {op} {}",
531                l.data_type(),
532                r.data_type()
533            )));
534        }
535    };
536    Ok(Arc::new(array.with_timezone_opt(l.timezone())))
537}
538
539/// Arithmetic trait for date arrays
540trait DateOp: ArrowTemporalType {
541    fn add_year_month(timestamp: Self::Native, delta: i32) -> Result<Self::Native, ArrowError>;
542    fn add_day_time(
543        timestamp: Self::Native,
544        delta: IntervalDayTime,
545    ) -> Result<Self::Native, ArrowError>;
546    fn add_month_day_nano(
547        timestamp: Self::Native,
548        delta: IntervalMonthDayNano,
549    ) -> Result<Self::Native, ArrowError>;
550
551    fn sub_year_month(timestamp: Self::Native, delta: i32) -> Result<Self::Native, ArrowError>;
552    fn sub_day_time(
553        timestamp: Self::Native,
554        delta: IntervalDayTime,
555    ) -> Result<Self::Native, ArrowError>;
556    fn sub_month_day_nano(
557        timestamp: Self::Native,
558        delta: IntervalMonthDayNano,
559    ) -> Result<Self::Native, ArrowError>;
560}
561
562macro_rules! date {
563    ($t:ty) => {
564        impl DateOp for $t {
565            fn add_year_month(left: Self::Native, right: i32) -> Result<Self::Native, ArrowError> {
566                Self::add_year_months_opt(left, right).ok_or_else(|| {
567                    ArrowError::ComputeError(format!(
568                        "Date arithmetic overflow: {left} + {right} months"
569                    ))
570                })
571            }
572
573            fn add_day_time(
574                left: Self::Native,
575                right: IntervalDayTime,
576            ) -> Result<Self::Native, ArrowError> {
577                Self::add_day_time_opt(left, right).ok_or_else(|| {
578                    ArrowError::ComputeError(format!(
579                        "Date arithmetic overflow: {left} + {right:?}"
580                    ))
581                })
582            }
583
584            fn add_month_day_nano(
585                left: Self::Native,
586                right: IntervalMonthDayNano,
587            ) -> Result<Self::Native, ArrowError> {
588                Self::add_month_day_nano_opt(left, right).ok_or_else(|| {
589                    ArrowError::ComputeError(format!(
590                        "Date arithmetic overflow: {left} + {right:?}"
591                    ))
592                })
593            }
594
595            fn sub_year_month(left: Self::Native, right: i32) -> Result<Self::Native, ArrowError> {
596                Self::subtract_year_months_opt(left, right).ok_or_else(|| {
597                    ArrowError::ComputeError(format!(
598                        "Date arithmetic overflow: {left} - {right} months"
599                    ))
600                })
601            }
602
603            fn sub_day_time(
604                left: Self::Native,
605                right: IntervalDayTime,
606            ) -> Result<Self::Native, ArrowError> {
607                Self::subtract_day_time_opt(left, right).ok_or_else(|| {
608                    ArrowError::ComputeError(format!(
609                        "Date arithmetic overflow: {left} - {right:?}"
610                    ))
611                })
612            }
613
614            fn sub_month_day_nano(
615                left: Self::Native,
616                right: IntervalMonthDayNano,
617            ) -> Result<Self::Native, ArrowError> {
618                Self::subtract_month_day_nano_opt(left, right).ok_or_else(|| {
619                    ArrowError::ComputeError(format!(
620                        "Date arithmetic overflow: {left} - {right:?}"
621                    ))
622                })
623            }
624        }
625    };
626}
627
628date!(Date32Type);
629date!(Date64Type);
630
631/// Arithmetic trait for interval arrays
632trait IntervalOp: ArrowPrimitiveType {
633    fn add(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError>;
634    fn sub(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError>;
635    fn mul_i64(left: Self::Native, right: i64) -> Result<Self::Native, ArrowError>;
636}
637
638fn mul_i32_i64(left: i32, right: i64) -> Result<i32, ArrowError> {
639    let value = i64::from(left).mul_checked(right)?;
640    i32::try_from(value).map_err(|_| {
641        ArrowError::ArithmeticOverflow(format!("Overflow happened on: {left} * {right}"))
642    })
643}
644
645impl IntervalOp for IntervalYearMonthType {
646    fn add(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
647        left.add_checked(right)
648    }
649
650    fn sub(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
651        left.sub_checked(right)
652    }
653
654    fn mul_i64(left: Self::Native, right: i64) -> Result<Self::Native, ArrowError> {
655        mul_i32_i64(left, right)
656    }
657}
658
659impl IntervalOp for IntervalDayTimeType {
660    fn add(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
661        let (l_days, l_ms) = Self::to_parts(left);
662        let (r_days, r_ms) = Self::to_parts(right);
663        let days = l_days.add_checked(r_days)?;
664        let ms = l_ms.add_checked(r_ms)?;
665        Ok(Self::make_value(days, ms))
666    }
667
668    fn sub(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
669        let (l_days, l_ms) = Self::to_parts(left);
670        let (r_days, r_ms) = Self::to_parts(right);
671        let days = l_days.sub_checked(r_days)?;
672        let ms = l_ms.sub_checked(r_ms)?;
673        Ok(Self::make_value(days, ms))
674    }
675
676    fn mul_i64(left: Self::Native, right: i64) -> Result<Self::Native, ArrowError> {
677        let (days, ms) = Self::to_parts(left);
678        Ok(Self::make_value(
679            mul_i32_i64(days, right)?,
680            mul_i32_i64(ms, right)?,
681        ))
682    }
683}
684
685impl IntervalOp for IntervalMonthDayNanoType {
686    fn add(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
687        let (l_months, l_days, l_nanos) = Self::to_parts(left);
688        let (r_months, r_days, r_nanos) = Self::to_parts(right);
689        let months = l_months.add_checked(r_months)?;
690        let days = l_days.add_checked(r_days)?;
691        let nanos = l_nanos.add_checked(r_nanos)?;
692        Ok(Self::make_value(months, days, nanos))
693    }
694
695    fn sub(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
696        let (l_months, l_days, l_nanos) = Self::to_parts(left);
697        let (r_months, r_days, r_nanos) = Self::to_parts(right);
698        let months = l_months.sub_checked(r_months)?;
699        let days = l_days.sub_checked(r_days)?;
700        let nanos = l_nanos.sub_checked(r_nanos)?;
701        Ok(Self::make_value(months, days, nanos))
702    }
703
704    fn mul_i64(left: Self::Native, right: i64) -> Result<Self::Native, ArrowError> {
705        let (months, days, nanos) = Self::to_parts(left);
706        Ok(Self::make_value(
707            mul_i32_i64(months, right)?,
708            mul_i32_i64(days, right)?,
709            nanos.mul_checked(right)?,
710        ))
711    }
712}
713
714fn interval_mul_op<T: IntervalOp>(
715    interval: &dyn Array,
716    interval_scalar: bool,
717    factor: &dyn Array,
718    factor_scalar: bool,
719) -> Result<ArrayRef, ArrowError> {
720    let interval = interval.as_primitive::<T>();
721    let factor = factor.as_primitive::<Int64Type>();
722    Ok(try_op_ref!(
723        T,
724        interval,
725        interval_scalar,
726        factor,
727        factor_scalar,
728        T::mul_i64(interval, factor)
729    ))
730}
731
732/// Multiplies an `IntervalMonthDayNano` by an `f64`, mirroring DuckDB's
733/// `interval_t` layout of months, days, and a sub-day component (nanoseconds in
734/// Arrow, microseconds in DuckDB).
735/// <https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/src/include/duckdb/common/types/interval.hpp#L24-L27>
736///
737/// Algorithm:
738///
739/// 1. use checked integer multiplication when `factor` fits in `i64` (early return).
740/// 2. multiply months and days separately
741/// 3. cascade remainders: convert fractional months to days using 30
742///    days per month, then fractional days to a sub-day value using 24 hours
743///    per day.
744/// 4. combine the cascaded remainder with the scaled input nanoseconds and round ties-to-even at nanosecond precision.
745/// 5. return an overflow error if any output component is out of
746///    range.
747fn interval_mul_f64(
748    interval: IntervalMonthDayNano,
749    factor: f64,
750) -> Result<IntervalMonthDayNano, ArrowError> {
751    const DAYS_PER_MONTH: f64 = 30.;
752    const NANOS_PER_SECOND: f64 = NANOSECONDS as f64;
753    const SECONDS_PER_DAY: f64 = SECONDS_IN_DAY as f64;
754
755    // Keep integral factors exact instead of round-tripping i64 nanoseconds through f64.
756    if factor.fract() == 0.
757        && let Some(factor) = ToPrimitive::to_i64(&factor)
758    {
759        return IntervalMonthDayNanoType::mul_i64(interval, factor);
760    }
761
762    // Based on DuckDB's INTERVAL * DOUBLE implementation, which is referenced from PostgreSQL's interval_mul:
763    // https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/src/function/scalar/operator/multiply.cpp#L48-L123
764    // PostgreSQL's interval_mul:
765    // https://github.com/postgres/postgres/blob/78758d37306cd89ab060f00cb06f249018d5b8da/src/backend/utils/adt/timestamp.c#L3627-L3744
766    let overflow =
767        |component| ArrowError::ArithmeticOverflow(format!("Overflow in interval {component}"));
768    let timestamp_round =
769        |value: f64| (value * NANOS_PER_SECOND).round_ties_even() / NANOS_PER_SECOND;
770
771    let months_product = f64::from(interval.months) * factor;
772    if !months_product.is_finite()
773        || months_product < f64::from(i32::MIN)
774        || months_product > f64::from(i32::MAX)
775    {
776        return Err(overflow("months"));
777    }
778    let months = months_product.to_i32().ok_or_else(|| overflow("months"))?;
779
780    let days_product = f64::from(interval.days) * factor;
781    if !days_product.is_finite()
782        || days_product < f64::from(i32::MIN)
783        || days_product > f64::from(i32::MAX)
784    {
785        return Err(overflow("days"));
786    }
787    let mut days = days_product.to_i32().ok_or_else(|| overflow("days"))?;
788
789    let month_remainder = timestamp_round(months_product.fract() * DAYS_PER_MONTH);
790    let month_remainder_days = month_remainder
791        .to_i32()
792        .ok_or_else(|| overflow("month remainder"))?;
793    let mut seconds_remainder =
794        timestamp_round((days_product.fract() + month_remainder.fract()) * SECONDS_PER_DAY);
795
796    if seconds_remainder.abs() >= SECONDS_PER_DAY {
797        let remainder_days = (seconds_remainder / SECONDS_PER_DAY)
798            .to_i32()
799            .ok_or_else(|| overflow("day remainder"))?;
800        days = days
801            .checked_add(remainder_days)
802            .ok_or_else(|| overflow("days"))?;
803        seconds_remainder -= f64::from(remainder_days) * SECONDS_PER_DAY;
804    }
805    days = days
806        .checked_add(month_remainder_days)
807        .ok_or_else(|| overflow("days"))?;
808
809    let nanoseconds = ((interval.nanoseconds as f64) * factor
810        + seconds_remainder * NANOS_PER_SECOND)
811        .round_ties_even();
812    let nanoseconds = ToPrimitive::to_i64(&nanoseconds).ok_or_else(|| {
813        ArrowError::ArithmeticOverflow(format!("Overflow in interval nanoseconds: {nanoseconds}"))
814    })?;
815
816    Ok(IntervalMonthDayNano::new(months, days, nanoseconds))
817}
818
819fn interval_f64_op(
820    op: Op,
821    interval: &dyn Array,
822    interval_scalar: bool,
823    factor: &dyn Array,
824    factor_scalar: bool,
825) -> Result<ArrayRef, ArrowError> {
826    let interval = interval.as_primitive::<IntervalMonthDayNanoType>();
827    let factor = factor.as_primitive::<Float64Type>();
828    Ok(try_op_ref!(
829        IntervalMonthDayNanoType,
830        interval,
831        interval_scalar,
832        factor,
833        factor_scalar,
834        {
835            match op {
836                Op::Mul => interval_mul_f64(interval, factor),
837                Op::Div if factor == 0. => Err(ArrowError::DivideByZero),
838                // DuckDB defines interval division as multiplication by the reciprocal:
839                // https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/src/function/scalar/operator/arithmetic.cpp#L1102-L1110
840                Op::Div => interval_mul_f64(interval, 1. / factor),
841                _ => Err(ArrowError::InvalidArgumentError(format!(
842                    "Invalid interval arithmetic operation: Interval(MonthDayNano) {op} Float64"
843                ))),
844            }
845        }
846    ))
847}
848
849/// Perform arithmetic operation on an interval array
850fn interval_op<T: IntervalOp>(
851    op: Op,
852    l: &dyn Array,
853    l_s: bool,
854    r: &dyn Array,
855    r_s: bool,
856) -> Result<ArrayRef, ArrowError> {
857    match (op, r.data_type()) {
858        (Op::Add | Op::AddWrapping, data_type) if data_type == l.data_type() => {
859            let l = l.as_primitive::<T>();
860            let r = r.as_primitive::<T>();
861            Ok(try_op_ref!(T, l, l_s, r, r_s, T::add(l, r)))
862        }
863        (Op::Sub | Op::SubWrapping, data_type) if data_type == l.data_type() => {
864            let l = l.as_primitive::<T>();
865            let r = r.as_primitive::<T>();
866            Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub(l, r)))
867        }
868        (Op::Mul, DataType::Int64) => interval_mul_op::<T>(l, l_s, r, r_s),
869        _ => Err(ArrowError::InvalidArgumentError(format!(
870            "Invalid interval arithmetic operation: {} {op} {}",
871            l.data_type(),
872            r.data_type()
873        ))),
874    }
875}
876
877fn duration_op<T: ArrowPrimitiveType>(
878    op: Op,
879    l: &dyn Array,
880    l_s: bool,
881    r: &dyn Array,
882    r_s: bool,
883) -> Result<ArrayRef, ArrowError> {
884    let l = l.as_primitive::<T>();
885    let r = r.as_primitive::<T>();
886    match op {
887        Op::Add | Op::AddWrapping => Ok(try_op_ref!(T, l, l_s, r, r_s, l.add_checked(r))),
888        Op::Sub | Op::SubWrapping => Ok(try_op_ref!(T, l, l_s, r, r_s, l.sub_checked(r))),
889        _ => Err(ArrowError::InvalidArgumentError(format!(
890            "Invalid duration arithmetic operation: {} {op} {}",
891            l.data_type(),
892            r.data_type()
893        ))),
894    }
895}
896
897/// Perform arithmetic operation on a date array
898fn date_op<T: DateOp>(
899    op: Op,
900    l: &dyn Array,
901    l_s: bool,
902    r: &dyn Array,
903    r_s: bool,
904) -> Result<ArrayRef, ArrowError> {
905    use DataType::*;
906    use IntervalUnit::*;
907
908    const NUM_SECONDS_IN_DAY: i64 = 60 * 60 * 24;
909
910    let r_t = r.data_type();
911    match (T::DATA_TYPE, op, r_t) {
912        (Date32, Op::Sub | Op::SubWrapping, Date32) => {
913            let l = l.as_primitive::<Date32Type>();
914            let r = r.as_primitive::<Date32Type>();
915            return Ok(op_ref!(
916                DurationSecondType,
917                l,
918                l_s,
919                r,
920                r_s,
921                ((l as i64) - (r as i64)) * NUM_SECONDS_IN_DAY
922            ));
923        }
924        (Date64, Op::Sub | Op::SubWrapping, Date64) => {
925            let l = l.as_primitive::<Date64Type>();
926            let r = r.as_primitive::<Date64Type>();
927            let result = try_op_ref!(DurationMillisecondType, l, l_s, r, r_s, l.sub_checked(r));
928            return Ok(result);
929        }
930        _ => {}
931    }
932
933    let l = l.as_primitive::<T>();
934    match (op, r_t) {
935        (Op::Add | Op::AddWrapping, Interval(YearMonth)) => {
936            let r = r.as_primitive::<IntervalYearMonthType>();
937            Ok(try_op_ref!(T, l, l_s, r, r_s, T::add_year_month(l, r)))
938        }
939        (Op::Sub | Op::SubWrapping, Interval(YearMonth)) => {
940            let r = r.as_primitive::<IntervalYearMonthType>();
941            Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub_year_month(l, r)))
942        }
943
944        (Op::Add | Op::AddWrapping, Interval(DayTime)) => {
945            let r = r.as_primitive::<IntervalDayTimeType>();
946            Ok(try_op_ref!(T, l, l_s, r, r_s, T::add_day_time(l, r)))
947        }
948        (Op::Sub | Op::SubWrapping, Interval(DayTime)) => {
949            let r = r.as_primitive::<IntervalDayTimeType>();
950            Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub_day_time(l, r)))
951        }
952
953        (Op::Add | Op::AddWrapping, Interval(MonthDayNano)) => {
954            let r = r.as_primitive::<IntervalMonthDayNanoType>();
955            Ok(try_op_ref!(T, l, l_s, r, r_s, T::add_month_day_nano(l, r)))
956        }
957        (Op::Sub | Op::SubWrapping, Interval(MonthDayNano)) => {
958            let r = r.as_primitive::<IntervalMonthDayNanoType>();
959            Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub_month_day_nano(l, r)))
960        }
961
962        _ => Err(ArrowError::InvalidArgumentError(format!(
963            "Invalid date arithmetic operation: {} {op} {}",
964            l.data_type(),
965            r.data_type()
966        ))),
967    }
968}
969
970/// Divides `l * 10^mul_pow` by `r` a digit at a time, without forming the scaled numerator.
971/// Used when scaling `l` would overflow `T::Native`, which it does well before the quotient
972/// does.
973///
974/// Runs on magnitudes and restores the sign last, so it truncates toward zero.
975fn scaled_div<T: DecimalType>(
976    l: T::Native,
977    r: T::Native,
978    mul_pow: i8,
979) -> Result<T::Native, ArrowError> {
980    let zero = T::Native::ZERO;
981    let negative = l.is_lt(zero) != r.is_lt(zero);
982    let dividend = abs_checked::<T>(l)?;
983    let divisor = abs_checked::<T>(r)?;
984
985    let mut quotient = dividend.div_checked(divisor)?;
986    let mut remainder = dividend.mod_checked(divisor)?;
987    for _ in 0..mul_pow {
988        // `remainder * 10` overflows for divisors near `T::Native::MAX`, so add the remainder
989        // ten times and take the divisor off whenever the running sum reaches it.
990        let mut carried = zero;
991        let mut digit = zero;
992        for _ in 0..10 {
993            let headroom = divisor.sub_wrapping(remainder);
994            if carried.is_lt(headroom) {
995                carried = carried.add_wrapping(remainder);
996            } else {
997                carried = carried.sub_wrapping(headroom);
998                digit = digit.add_wrapping(T::Native::ONE);
999            }
1000        }
1001        quotient = quotient
1002            .mul_checked(T::Native::usize_as(10))?
1003            .add_checked(digit)?;
1004        remainder = carried;
1005    }
1006
1007    if negative {
1008        quotient.neg_checked()
1009    } else {
1010        Ok(quotient)
1011    }
1012}
1013
1014fn abs_checked<T: DecimalType>(value: T::Native) -> Result<T::Native, ArrowError> {
1015    if value.is_lt(T::Native::ZERO) {
1016        value.neg_checked()
1017    } else {
1018        Ok(value)
1019    }
1020}
1021
1022/// Perform arithmetic operation on decimal arrays
1023fn decimal_op<T: DecimalType>(
1024    op: Op,
1025    l: &dyn Array,
1026    l_s: bool,
1027    r: &dyn Array,
1028    r_s: bool,
1029) -> Result<ArrayRef, ArrowError> {
1030    let l = l.as_primitive::<T>();
1031    let r = r.as_primitive::<T>();
1032
1033    let (p1, s1, p2, s2) = match (l.data_type(), r.data_type()) {
1034        (DataType::Decimal32(p1, s1), DataType::Decimal32(p2, s2)) => (p1, s1, p2, s2),
1035        (DataType::Decimal64(p1, s1), DataType::Decimal64(p2, s2)) => (p1, s1, p2, s2),
1036        (DataType::Decimal128(p1, s1), DataType::Decimal128(p2, s2)) => (p1, s1, p2, s2),
1037        (DataType::Decimal256(p1, s1), DataType::Decimal256(p2, s2)) => (p1, s1, p2, s2),
1038        _ => unreachable!(),
1039    };
1040
1041    // Follow the Hive decimal arithmetic rules
1042    // https://cwiki.apache.org/confluence/download/attachments/27362075/Hive_Decimal_Precision_Scale_Support.pdf
1043    let array: PrimitiveArray<T> = match op {
1044        Op::Add | Op::AddWrapping | Op::Sub | Op::SubWrapping => {
1045            let (p1, s1, p2, s2) = (*p1 as i16, *s1 as i16, *p2 as i16, *s2 as i16);
1046            // max(s1, s2)
1047            let result_scale = s1.max(s2);
1048
1049            // max(s1, s2) + max(p1-s1, p2-s2) + 1
1050            let result_precision =
1051                (result_scale + (p1 - s1).max(p2 - s2) + 1).min(T::MAX_PRECISION as i16) as u8;
1052
1053            let l_mul = T::Native::usize_as(10).pow_checked((result_scale - s1) as _)?;
1054            let r_mul = T::Native::usize_as(10).pow_checked((result_scale - s2) as _)?;
1055
1056            match op {
1057                // Equal scales make both decimal multipliers one.
1058                Op::Add | Op::AddWrapping if s1 == s2 => {
1059                    try_op!(l, l_s, r, r_s, l.add_checked(r))
1060                }
1061                Op::Sub | Op::SubWrapping if s1 == s2 => {
1062                    try_op!(l, l_s, r, r_s, l.sub_checked(r))
1063                }
1064                Op::Add | Op::AddWrapping => {
1065                    try_op!(
1066                        l,
1067                        l_s,
1068                        r,
1069                        r_s,
1070                        l.mul_checked(l_mul)?.add_checked(r.mul_checked(r_mul)?)
1071                    )
1072                }
1073                Op::Sub | Op::SubWrapping => {
1074                    try_op!(
1075                        l,
1076                        l_s,
1077                        r,
1078                        r_s,
1079                        l.mul_checked(l_mul)?.sub_checked(r.mul_checked(r_mul)?)
1080                    )
1081                }
1082                _ => unreachable!(),
1083            }
1084            .with_precision_and_scale(result_precision, result_scale as i8)?
1085        }
1086        Op::Mul | Op::MulWrapping => {
1087            let result_precision = p1.saturating_add(p2 + 1).min(T::MAX_PRECISION);
1088            let result_scale = *s1 as i16 + *s2 as i16;
1089            if result_scale > T::MAX_SCALE as i16 {
1090                // SQL standard says that if the resulting scale of a multiply operation goes
1091                // beyond the maximum, rounding is not acceptable and thus an error occurs
1092                return Err(ArrowError::InvalidArgumentError(format!(
1093                    "Output scale of {} {op} {} would exceed max scale of {}",
1094                    l.data_type(),
1095                    r.data_type(),
1096                    T::MAX_SCALE
1097                )));
1098            }
1099            let result_scale = i8::try_from(result_scale).map_err(|_| {
1100                ArrowError::InvalidArgumentError(format!(
1101                    "Output scale of {} {op} {} would be less than min scale of {}",
1102                    l.data_type(),
1103                    r.data_type(),
1104                    i8::MIN
1105                ))
1106            })?;
1107
1108            try_op!(l, l_s, r, r_s, l.mul_checked(r))
1109                .with_precision_and_scale(result_precision, result_scale)?
1110        }
1111
1112        Op::Div => {
1113            // Follow postgres and MySQL adding a fixed scale increment of 4
1114            // s1 + 4
1115            let result_scale = s1.saturating_add(4).min(T::MAX_SCALE);
1116            let mul_pow = result_scale - s1 + s2;
1117
1118            // p1 - s1 + s2 + result_scale
1119            let result_precision = (mul_pow.saturating_add(*p1 as i8) as u8).min(T::MAX_PRECISION);
1120
1121            let (l_mul, r_mul) = match mul_pow.cmp(&0) {
1122                Ordering::Greater => (
1123                    T::Native::usize_as(10).pow_checked(mul_pow as _)?,
1124                    T::Native::ONE,
1125                ),
1126                Ordering::Equal => (T::Native::ONE, T::Native::ONE),
1127                Ordering::Less => (
1128                    T::Native::ONE,
1129                    T::Native::usize_as(10).pow_checked(mul_pow.neg_wrapping() as _)?,
1130                ),
1131            };
1132
1133            try_op!(
1134                l,
1135                l_s,
1136                r,
1137                r_s,
1138                match l.mul_checked(l_mul) {
1139                    Ok(scaled) => scaled.div_checked(r.mul_checked(r_mul)?),
1140                    Err(_) => scaled_div::<T>(l, r, mul_pow),
1141                }
1142            )
1143            .with_precision_and_scale(result_precision, result_scale)?
1144        }
1145
1146        Op::Rem => {
1147            let (p1, s1, p2, s2) = (*p1 as i16, *s1 as i16, *p2 as i16, *s2 as i16);
1148            // max(s1, s2)
1149            let result_scale = s1.max(s2);
1150            // min(p1-s1, p2 -s2) + max( s1,s2 )
1151            let result_precision =
1152                (result_scale + (p1 - s1).min(p2 - s2)).min(T::MAX_PRECISION as i16) as u8;
1153
1154            let l_mul = T::Native::usize_as(10).pow_checked((result_scale - s1) as _)?;
1155            let r_mul = T::Native::usize_as(10).pow_checked((result_scale - s2) as _)?;
1156
1157            try_op!(
1158                l,
1159                l_s,
1160                r,
1161                r_s,
1162                l.mul_checked(l_mul)?.mod_checked(r.mul_checked(r_mul)?)
1163            )
1164            .with_precision_and_scale(result_precision, result_scale as i8)?
1165        }
1166    };
1167
1168    Ok(Arc::new(array))
1169}
1170
1171#[cfg(test)]
1172mod tests {
1173    use super::*;
1174    use arrow_array::temporal_conversions::{as_date, as_datetime};
1175    use arrow_buffer::{ScalarBuffer, i256};
1176    use chrono::{DateTime, NaiveDate};
1177
1178    // The valid date range of NaiveDate is from January 1, -262143 to December 31, 262142 (Gregorian calendar).
1179    const MAX_VALID_DATE: NaiveDate = NaiveDate::from_ymd_opt(262142, 12, 31).unwrap();
1180    const MIN_VALID_DATE: NaiveDate = NaiveDate::from_ymd_opt(-262143, 1, 1).unwrap();
1181    const MAX_VALID_MILLIS: i64 = date_to_millis(MAX_VALID_DATE);
1182    const MIN_VALID_MILLIS: i64 = date_to_millis(MIN_VALID_DATE);
1183    const MAX_VALID_DAYS: i32 = date_to_days(MAX_VALID_DATE);
1184    const MIN_VALID_DAYS: i32 = date_to_days(MIN_VALID_DATE);
1185    const EPOCH: NaiveDate = NaiveDate::from_ymd_opt(1970, 1, 1).unwrap();
1186    const YEAR_2000: NaiveDate = NaiveDate::from_ymd_opt(2000, 1, 1).unwrap();
1187
1188    const fn date_to_millis(date: NaiveDate) -> i64 {
1189        date.signed_duration_since(EPOCH).num_milliseconds()
1190    }
1191
1192    const fn date_to_days(date: NaiveDate) -> i32 {
1193        date.signed_duration_since(EPOCH).num_days() as i32
1194    }
1195
1196    fn test_neg_primitive<T: ArrowPrimitiveType>(
1197        input: &[T::Native],
1198        out: Result<&[T::Native], &str>,
1199    ) {
1200        let a = PrimitiveArray::<T>::new(ScalarBuffer::from(input.to_vec()), None);
1201        match out {
1202            Ok(expected) => {
1203                let result = neg(&a).unwrap();
1204                assert_eq!(result.as_primitive::<T>().values(), expected);
1205            }
1206            Err(e) => {
1207                let err = neg(&a).unwrap_err().to_string();
1208                assert_eq!(e, err);
1209            }
1210        }
1211    }
1212
1213    #[test]
1214    fn test_neg() {
1215        let input = &[1, -5, 2, 693, 3929];
1216        let output = &[-1, 5, -2, -693, -3929];
1217        test_neg_primitive::<Int32Type>(input, Ok(output));
1218
1219        let input = &[1, -5, 2, 693, 3929];
1220        let output = &[-1, 5, -2, -693, -3929];
1221        test_neg_primitive::<Int64Type>(input, Ok(output));
1222        test_neg_primitive::<DurationSecondType>(input, Ok(output));
1223        test_neg_primitive::<DurationMillisecondType>(input, Ok(output));
1224        test_neg_primitive::<DurationMicrosecondType>(input, Ok(output));
1225        test_neg_primitive::<DurationNanosecondType>(input, Ok(output));
1226
1227        let input = &[f32::MAX, f32::MIN, f32::INFINITY, 1.3, 0.5];
1228        let output = &[f32::MIN, f32::MAX, f32::NEG_INFINITY, -1.3, -0.5];
1229        test_neg_primitive::<Float32Type>(input, Ok(output));
1230
1231        test_neg_primitive::<Int32Type>(
1232            &[i32::MIN],
1233            Err("Arithmetic overflow: Overflow happened on: - -2147483648"),
1234        );
1235        test_neg_primitive::<Int64Type>(
1236            &[i64::MIN],
1237            Err("Arithmetic overflow: Overflow happened on: - -9223372036854775808"),
1238        );
1239        test_neg_primitive::<DurationSecondType>(
1240            &[i64::MIN],
1241            Err("Arithmetic overflow: Overflow happened on: - -9223372036854775808"),
1242        );
1243
1244        let r = neg_wrapping(&Int32Array::from(vec![i32::MIN])).unwrap();
1245        assert_eq!(r.as_primitive::<Int32Type>().value(0), i32::MIN);
1246
1247        let r = neg_wrapping(&Int64Array::from(vec![i64::MIN])).unwrap();
1248        assert_eq!(r.as_primitive::<Int64Type>().value(0), i64::MIN);
1249
1250        let err = neg_wrapping(&DurationSecondArray::from(vec![i64::MIN]))
1251            .unwrap_err()
1252            .to_string();
1253
1254        assert_eq!(
1255            err,
1256            "Arithmetic overflow: Overflow happened on: - -9223372036854775808"
1257        );
1258
1259        let a = Decimal32Array::from(vec![1, 3, -44, 2, 4])
1260            .with_precision_and_scale(9, 6)
1261            .unwrap();
1262
1263        let r = neg(&a).unwrap();
1264        assert_eq!(r.data_type(), a.data_type());
1265        assert_eq!(
1266            r.as_primitive::<Decimal32Type>().values(),
1267            &[-1, -3, 44, -2, -4]
1268        );
1269
1270        let a = Decimal64Array::from(vec![1, 3, -44, 2, 4])
1271            .with_precision_and_scale(9, 6)
1272            .unwrap();
1273
1274        let r = neg(&a).unwrap();
1275        assert_eq!(r.data_type(), a.data_type());
1276        assert_eq!(
1277            r.as_primitive::<Decimal64Type>().values(),
1278            &[-1, -3, 44, -2, -4]
1279        );
1280
1281        let a = Decimal128Array::from(vec![1, 3, -44, 2, 4])
1282            .with_precision_and_scale(9, 6)
1283            .unwrap();
1284
1285        let r = neg(&a).unwrap();
1286        assert_eq!(r.data_type(), a.data_type());
1287        assert_eq!(
1288            r.as_primitive::<Decimal128Type>().values(),
1289            &[-1, -3, 44, -2, -4]
1290        );
1291
1292        let a = Decimal256Array::from(vec![
1293            i256::from_i128(342),
1294            i256::from_i128(-4949),
1295            i256::from_i128(3),
1296        ])
1297        .with_precision_and_scale(9, 6)
1298        .unwrap();
1299
1300        let r = neg(&a).unwrap();
1301        assert_eq!(r.data_type(), a.data_type());
1302        assert_eq!(
1303            r.as_primitive::<Decimal256Type>().values(),
1304            &[
1305                i256::from_i128(-342),
1306                i256::from_i128(4949),
1307                i256::from_i128(-3),
1308            ]
1309        );
1310
1311        let a = IntervalYearMonthArray::from(vec![
1312            IntervalYearMonthType::make_value(2, 4),
1313            IntervalYearMonthType::make_value(2, -4),
1314            IntervalYearMonthType::make_value(-3, -5),
1315        ]);
1316        let r = neg(&a).unwrap();
1317        assert_eq!(
1318            r.as_primitive::<IntervalYearMonthType>().values(),
1319            &[
1320                IntervalYearMonthType::make_value(-2, -4),
1321                IntervalYearMonthType::make_value(-2, 4),
1322                IntervalYearMonthType::make_value(3, 5),
1323            ]
1324        );
1325
1326        let a = IntervalDayTimeArray::from(vec![
1327            IntervalDayTimeType::make_value(2, 4),
1328            IntervalDayTimeType::make_value(2, -4),
1329            IntervalDayTimeType::make_value(-3, -5),
1330        ]);
1331        let r = neg(&a).unwrap();
1332        assert_eq!(
1333            r.as_primitive::<IntervalDayTimeType>().values(),
1334            &[
1335                IntervalDayTimeType::make_value(-2, -4),
1336                IntervalDayTimeType::make_value(-2, 4),
1337                IntervalDayTimeType::make_value(3, 5),
1338            ]
1339        );
1340
1341        let a = IntervalMonthDayNanoArray::from(vec![
1342            IntervalMonthDayNanoType::make_value(2, 4, 5953394),
1343            IntervalMonthDayNanoType::make_value(2, -4, -45839),
1344            IntervalMonthDayNanoType::make_value(-3, -5, 6944),
1345        ]);
1346        let r = neg(&a).unwrap();
1347        assert_eq!(
1348            r.as_primitive::<IntervalMonthDayNanoType>().values(),
1349            &[
1350                IntervalMonthDayNanoType::make_value(-2, -4, -5953394),
1351                IntervalMonthDayNanoType::make_value(-2, 4, 45839),
1352                IntervalMonthDayNanoType::make_value(3, 5, -6944),
1353            ]
1354        );
1355    }
1356
1357    #[test]
1358    fn test_integer() {
1359        let a = Int32Array::from(vec![4, 3, 5, -6, 100]);
1360        let b = Int32Array::from(vec![6, 2, 5, -7, 3]);
1361        let result = add(&a, &b).unwrap();
1362        assert_eq!(
1363            result.as_ref(),
1364            &Int32Array::from(vec![10, 5, 10, -13, 103])
1365        );
1366        let result = sub(&a, &b).unwrap();
1367        assert_eq!(result.as_ref(), &Int32Array::from(vec![-2, 1, 0, 1, 97]));
1368        let result = div(&a, &b).unwrap();
1369        assert_eq!(result.as_ref(), &Int32Array::from(vec![0, 1, 1, 0, 33]));
1370        let result = mul(&a, &b).unwrap();
1371        assert_eq!(result.as_ref(), &Int32Array::from(vec![24, 6, 25, 42, 300]));
1372        let result = rem(&a, &b).unwrap();
1373        assert_eq!(result.as_ref(), &Int32Array::from(vec![4, 1, 0, -6, 1]));
1374
1375        let a = Int8Array::from(vec![Some(2), None, Some(45)]);
1376        let b = Int8Array::from(vec![Some(5), Some(3), None]);
1377        let result = add(&a, &b).unwrap();
1378        assert_eq!(result.as_ref(), &Int8Array::from(vec![Some(7), None, None]));
1379
1380        let a = UInt8Array::from(vec![56, 5, 3]);
1381        let b = UInt8Array::from(vec![200, 2, 5]);
1382        let err = add(&a, &b).unwrap_err().to_string();
1383        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 56 + 200");
1384        let result = add_wrapping(&a, &b).unwrap();
1385        assert_eq!(result.as_ref(), &UInt8Array::from(vec![0, 7, 8]));
1386
1387        let a = UInt8Array::from(vec![34, 5, 3]);
1388        let b = UInt8Array::from(vec![200, 2, 5]);
1389        let err = sub(&a, &b).unwrap_err().to_string();
1390        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 34 - 200");
1391        let result = sub_wrapping(&a, &b).unwrap();
1392        assert_eq!(result.as_ref(), &UInt8Array::from(vec![90, 3, 254]));
1393
1394        let a = UInt8Array::from(vec![34, 5, 3]);
1395        let b = UInt8Array::from(vec![200, 2, 5]);
1396        let err = mul(&a, &b).unwrap_err().to_string();
1397        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 34 * 200");
1398        let result = mul_wrapping(&a, &b).unwrap();
1399        assert_eq!(result.as_ref(), &UInt8Array::from(vec![144, 10, 15]));
1400
1401        let a = Int16Array::from(vec![i16::MIN]);
1402        let b = Int16Array::from(vec![-1]);
1403        let err = div(&a, &b).unwrap_err().to_string();
1404        assert_eq!(
1405            err,
1406            "Arithmetic overflow: Overflow happened on: -32768 / -1"
1407        );
1408
1409        let a = Int16Array::from(vec![i16::MIN]);
1410        let b = Int16Array::from(vec![-1]);
1411        let result = rem(&a, &b).unwrap();
1412        assert_eq!(result.as_ref(), &Int16Array::from(vec![0]));
1413
1414        let a = Int16Array::from(vec![21]);
1415        let b = Int16Array::from(vec![0]);
1416        let err = div(&a, &b).unwrap_err().to_string();
1417        assert_eq!(err, "Divide by zero error");
1418
1419        let a = Int16Array::from(vec![21]);
1420        let b = Int16Array::from(vec![0]);
1421        let err = rem(&a, &b).unwrap_err().to_string();
1422        assert_eq!(err, "Divide by zero error");
1423    }
1424
1425    #[test]
1426    fn test_float() {
1427        let a = Float32Array::from(vec![1., f32::MAX, 6., -4., -1., 0.]);
1428        let b = Float32Array::from(vec![1., f32::MAX, f32::MAX, -3., 45., 0.]);
1429        let result = add(&a, &b).unwrap();
1430        assert_eq!(
1431            result.as_ref(),
1432            &Float32Array::from(vec![2., f32::INFINITY, f32::MAX, -7., 44.0, 0.])
1433        );
1434
1435        let result = sub(&a, &b).unwrap();
1436        assert_eq!(
1437            result.as_ref(),
1438            &Float32Array::from(vec![0., 0., f32::MIN, -1., -46., 0.])
1439        );
1440
1441        let result = mul(&a, &b).unwrap();
1442        assert_eq!(
1443            result.as_ref(),
1444            &Float32Array::from(vec![1., f32::INFINITY, f32::INFINITY, 12., -45., 0.])
1445        );
1446
1447        let result = div(&a, &b).unwrap();
1448        let r = result.as_primitive::<Float32Type>();
1449        assert_eq!(r.value(0), 1.);
1450        assert_eq!(r.value(1), 1.);
1451        assert!(r.value(2) < f32::EPSILON);
1452        assert_eq!(r.value(3), -4. / -3.);
1453        assert!(r.value(5).is_nan());
1454
1455        let result = rem(&a, &b).unwrap();
1456        let r = result.as_primitive::<Float32Type>();
1457        assert_eq!(&r.values()[..5], &[0., 0., 6., -1., -1.]);
1458        assert!(r.value(5).is_nan());
1459    }
1460
1461    #[test]
1462    fn test_decimal() {
1463        // 0.015 7.842 -0.577 0.334 -0.078 0.003
1464        let a = Decimal128Array::from(vec![15, 0, -577, 334, -78, 3])
1465            .with_precision_and_scale(12, 3)
1466            .unwrap();
1467
1468        // 5.4 0 -35.6 0.3 0.6 7.45
1469        let b = Decimal128Array::from(vec![54, 34, -356, 3, 6, 745])
1470            .with_precision_and_scale(12, 1)
1471            .unwrap();
1472
1473        let result = add(&a, &b).unwrap();
1474        assert_eq!(result.data_type(), &DataType::Decimal128(15, 3));
1475        assert_eq!(
1476            result.as_primitive::<Decimal128Type>().values(),
1477            &[5415, 3400, -36177, 634, 522, 74503]
1478        );
1479
1480        let result = sub(&a, &b).unwrap();
1481        assert_eq!(result.data_type(), &DataType::Decimal128(15, 3));
1482        assert_eq!(
1483            result.as_primitive::<Decimal128Type>().values(),
1484            &[-5385, -3400, 35023, 34, -678, -74497]
1485        );
1486
1487        let result = mul(&a, &b).unwrap();
1488        assert_eq!(result.data_type(), &DataType::Decimal128(25, 4));
1489        assert_eq!(
1490            result.as_primitive::<Decimal128Type>().values(),
1491            &[810, 0, 205412, 1002, -468, 2235]
1492        );
1493
1494        let result = div(&a, &b).unwrap();
1495        assert_eq!(result.data_type(), &DataType::Decimal128(17, 7));
1496        assert_eq!(
1497            result.as_primitive::<Decimal128Type>().values(),
1498            &[27777, 0, 162078, 11133333, -1300000, 402]
1499        );
1500
1501        let result = rem(&a, &b).unwrap();
1502        assert_eq!(result.data_type(), &DataType::Decimal128(12, 3));
1503        assert_eq!(
1504            result.as_primitive::<Decimal128Type>().values(),
1505            &[15, 0, -577, 34, -78, 3]
1506        );
1507
1508        let a = Decimal128Array::from(vec![1])
1509            .with_precision_and_scale(3, 3)
1510            .unwrap();
1511        let b = Decimal128Array::from(vec![1])
1512            .with_precision_and_scale(37, 37)
1513            .unwrap();
1514        let err = mul(&a, &b).unwrap_err().to_string();
1515        assert_eq!(
1516            err,
1517            "Invalid argument error: Output scale of Decimal128(3, 3) * Decimal128(37, 37) would exceed max scale of 38"
1518        );
1519
1520        let a = Decimal128Array::from(vec![1])
1521            .with_precision_and_scale(3, -2)
1522            .unwrap();
1523        let err = add(&a, &b).unwrap_err().to_string();
1524        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 10 ^ 39");
1525
1526        let a = Decimal128Array::from(vec![10])
1527            .with_precision_and_scale(3, -1)
1528            .unwrap();
1529        let err = add(&a, &b).unwrap_err().to_string();
1530        assert_eq!(
1531            err,
1532            "Arithmetic overflow: Overflow happened on: 10 * 100000000000000000000000000000000000000"
1533        );
1534
1535        let b = Decimal128Array::from(vec![0])
1536            .with_precision_and_scale(1, 1)
1537            .unwrap();
1538        let err = div(&a, &b).unwrap_err().to_string();
1539        assert_eq!(err, "Divide by zero error");
1540        let err = rem(&a, &b).unwrap_err().to_string();
1541        assert_eq!(err, "Divide by zero error");
1542    }
1543
1544    #[test]
1545    fn test_decimal256_add_sub_negative_scale_metadata() {
1546        let a = Decimal256Array::from(vec![Some(i256::ONE), Some(i256::MINUS_ONE), None])
1547            .with_precision_and_scale(76, -52)
1548            .unwrap();
1549        let b = Decimal256Array::from(vec![i256::from_i128(2), i256::ONE, i256::ONE])
1550            .with_precision_and_scale(76, -52)
1551            .unwrap();
1552        let expected =
1553            Decimal256Array::from(vec![Some(i256::from_i128(3)), Some(i256::ZERO), None])
1554                .with_precision_and_scale(76, -52)
1555                .unwrap();
1556        assert_eq!(add(&a, &b).unwrap().as_ref(), &expected);
1557
1558        let expected =
1559            Decimal256Array::from(vec![Some(i256::MINUS_ONE), Some(i256::from_i128(-2)), None])
1560                .with_precision_and_scale(76, -52)
1561                .unwrap();
1562        assert_eq!(sub(&a, &b).unwrap().as_ref(), &expected);
1563    }
1564
1565    #[test]
1566    fn test_decimal256_add_sub_minimum_scale() {
1567        let a = Decimal256Array::from(vec![i256::ONE])
1568            .with_precision_and_scale(20, i8::MIN)
1569            .unwrap();
1570        let b = Decimal256Array::from(vec![i256::from_i128(2)])
1571            .with_precision_and_scale(20, i8::MIN)
1572            .unwrap();
1573        let expected = Decimal256Array::from(vec![i256::from_i128(3)])
1574            .with_precision_and_scale(21, i8::MIN)
1575            .unwrap();
1576        assert_eq!(add(&a, &b).unwrap().as_ref(), &expected);
1577
1578        let expected = Decimal256Array::from(vec![i256::MINUS_ONE])
1579            .with_precision_and_scale(21, i8::MIN)
1580            .unwrap();
1581        assert_eq!(sub(&a, &b).unwrap().as_ref(), &expected);
1582    }
1583
1584    #[test]
1585    fn test_decimal256_remainder_negative_scale_metadata() {
1586        let a = Decimal256Array::from(vec![Some(i256::ONE), Some(i256::MINUS_ONE), None])
1587            .with_precision_and_scale(76, -52)
1588            .unwrap();
1589        let b = Decimal256Array::from(vec![i256::from_i128(2), i256::ONE, i256::ONE])
1590            .with_precision_and_scale(76, -52)
1591            .unwrap();
1592        let expected = Decimal256Array::from(vec![Some(i256::ONE), Some(i256::ZERO), None])
1593            .with_precision_and_scale(76, -52)
1594            .unwrap();
1595        assert_eq!(rem(&a, &b).unwrap().as_ref(), &expected);
1596    }
1597
1598    #[test]
1599    fn test_decimal256_adjacent_minimum_scales() {
1600        let a = Decimal256Array::from(vec![i256::ONE])
1601            .with_precision_and_scale(20, i8::MIN)
1602            .unwrap();
1603        let b = Decimal256Array::from(vec![i256::from_i128(2)])
1604            .with_precision_and_scale(20, i8::MIN + 1)
1605            .unwrap();
1606        let expected = Decimal256Array::from(vec![i256::from_i128(12)])
1607            .with_precision_and_scale(22, -127)
1608            .unwrap();
1609        assert_eq!(add(&a, &b).unwrap().as_ref(), &expected);
1610
1611        let expected = Decimal256Array::from(vec![i256::from_i128(8)])
1612            .with_precision_and_scale(22, -127)
1613            .unwrap();
1614        assert_eq!(sub(&a, &b).unwrap().as_ref(), &expected);
1615
1616        let expected = Decimal256Array::from(vec![i256::ZERO])
1617            .with_precision_and_scale(20, -127)
1618            .unwrap();
1619        assert_eq!(rem(&a, &b).unwrap().as_ref(), &expected);
1620    }
1621
1622    #[test]
1623    fn test_decimal256_extreme_scale_difference_overflow() {
1624        let a = Decimal256Array::from(vec![i256::ONE])
1625            .with_precision_and_scale(20, i8::MIN)
1626            .unwrap();
1627        let b = Decimal256Array::from(vec![i256::from_i128(2)])
1628            .with_precision_and_scale(76, 76)
1629            .unwrap();
1630        let expected = "Arithmetic overflow: Overflow happened on: 10 ^ 204";
1631        assert_eq!(add(&a, &b).unwrap_err().to_string(), expected);
1632        assert_eq!(sub(&a, &b).unwrap_err().to_string(), expected);
1633        assert_eq!(rem(&a, &b).unwrap_err().to_string(), expected);
1634    }
1635
1636    #[test]
1637    fn test_decimal256_multiply_minimum_scale() {
1638        let a = Decimal256Array::from(vec![i256::ONE])
1639            .with_precision_and_scale(76, -64)
1640            .unwrap();
1641        let expected = Decimal256Array::from(vec![i256::ONE])
1642            .with_precision_and_scale(76, i8::MIN)
1643            .unwrap();
1644        assert_eq!(mul(&a, &a).unwrap().as_ref(), &expected);
1645
1646        let b = a.clone().with_precision_and_scale(76, -65).unwrap();
1647        assert_eq!(
1648            mul(&a, &b).unwrap_err().to_string(),
1649            "Invalid argument error: Output scale of Decimal256(76, -64) * Decimal256(76, -65) would be less than min scale of -128"
1650        );
1651    }
1652
1653    #[test]
1654    fn test_decimal256_div_wide_intermediate() {
1655        // Dividing two scale-37 values needs l * 10^41, which is 79 digits and does not
1656        // fit in an i256, even though the 41-digit quotient does.
1657        let a = Decimal256Array::from(vec![i256::from_i128(
1658            60096743305738933273387748827369321010i128,
1659        )])
1660        .with_precision_and_scale(38, 37)
1661        .unwrap();
1662        let b = Decimal256Array::from(vec![i256::from_i128(
1663            60096763826458053191384497987259478584i128,
1664        )])
1665        .with_precision_and_scale(38, 37)
1666        .unwrap();
1667
1668        let result = div(&a, &b).unwrap();
1669        assert_eq!(result.data_type(), &DataType::Decimal256(76, 41));
1670        assert_eq!(
1671            result.as_primitive::<Decimal256Type>().value(0),
1672            i256::from_string("99999965853869970143724273117679321341339").unwrap()
1673        );
1674
1675        // Truncation stays toward zero on either side of the fallback.
1676        let neg_a = neg(&a).unwrap();
1677        let result = div(neg_a.as_primitive::<Decimal256Type>(), &b).unwrap();
1678        assert_eq!(
1679            result.as_primitive::<Decimal256Type>().value(0),
1680            i256::from_string("-99999965853869970143724273117679321341339").unwrap()
1681        );
1682
1683        let neg_b = neg(&b).unwrap();
1684        let result = div(&a, neg_b.as_primitive::<Decimal256Type>()).unwrap();
1685        assert_eq!(
1686            result.as_primitive::<Decimal256Type>().value(0),
1687            i256::from_string("-99999965853869970143724273117679321341339").unwrap()
1688        );
1689
1690        let zero = Decimal256Array::from(vec![i256::ZERO])
1691            .with_precision_and_scale(38, 37)
1692            .unwrap();
1693        let err = div(&a, &zero).unwrap_err().to_string();
1694        assert_eq!(err, "Divide by zero error");
1695    }
1696
1697    #[test]
1698    fn test_decimal256_div_divisor_near_max() {
1699        // A divisor past i256::MAX / 10 leaves no room to scale the running remainder either.
1700        let a = Decimal256Array::from(vec![
1701            i256::from_string(
1702                "5900000000000000000000000000000000000000000000000000000000000000000000000000",
1703            )
1704            .unwrap(),
1705        ])
1706        .with_precision_and_scale(76, 37)
1707        .unwrap();
1708        let b = Decimal256Array::from(vec![
1709            i256::from_string(
1710                "6000000000000000000000000000000000000000000000000000000000000000000000000000",
1711            )
1712            .unwrap(),
1713        ])
1714        .with_precision_and_scale(76, 37)
1715        .unwrap();
1716
1717        let result = div(&a, &b).unwrap();
1718        assert_eq!(
1719            result.as_primitive::<Decimal256Type>().value(0),
1720            i256::from_string("98333333333333333333333333333333333333333").unwrap()
1721        );
1722    }
1723
1724    #[test]
1725    fn test_decimal128_div_wide_intermediate() {
1726        // Same overflow one type down: 3.0 / 6.0 at scale 37 needs l * 10^38, 76 digits in i128.
1727        let a = Decimal128Array::from(vec![30000000000000000000000000000000000000i128])
1728            .with_precision_and_scale(38, 37)
1729            .unwrap();
1730        let b = Decimal128Array::from(vec![60000000000000000000000000000000000000i128])
1731            .with_precision_and_scale(38, 37)
1732            .unwrap();
1733
1734        let result = div(&a, &b).unwrap();
1735        assert_eq!(result.data_type(), &DataType::Decimal128(38, 38));
1736        assert_eq!(
1737            result.as_primitive::<Decimal128Type>().value(0),
1738            50000000000000000000000000000000000000i128
1739        );
1740    }
1741
1742    #[test]
1743    fn test_scaled_div_agrees_with_direct_division() {
1744        for (l, r, mul_pow) in [
1745            (7i128, 3i128, 4i8),
1746            (-7, 3, 4),
1747            (7, -3, 4),
1748            (-7, -3, 4),
1749            (1, 999_999_999, 9),
1750            (i128::MAX / 10, 7, 1),
1751            (i128::MAX / 10, -i128::MAX / 11, 1),
1752            (0, 5, 6),
1753        ] {
1754            let scaled = l * 10i128.pow(mul_pow as u32);
1755            assert_eq!(
1756                scaled_div::<Decimal128Type>(l, r, mul_pow).unwrap(),
1757                scaled / r,
1758                "{l} * 10^{mul_pow} / {r}"
1759            );
1760        }
1761    }
1762
1763    #[test]
1764    fn test_decimal256_same_scale_add_sub() {
1765        let lhs = Decimal256Array::from(vec![
1766            Some(i256::from_parts(u128::MAX, 0)),
1767            Some(i256::MINUS_ONE),
1768            None,
1769        ])
1770        .with_precision_and_scale(70, 2)
1771        .unwrap();
1772        let rhs = Decimal256Array::from(vec![Some(i256::ONE), Some(i256::ONE), Some(i256::MAX)])
1773            .with_precision_and_scale(70, 2)
1774            .unwrap();
1775
1776        let expected =
1777            Decimal256Array::from(vec![Some(i256::from_parts(0, 1)), Some(i256::ZERO), None])
1778                .with_precision_and_scale(71, 2)
1779                .unwrap();
1780        for operation in [add, add_wrapping] {
1781            let result = operation(&lhs, &rhs).unwrap();
1782            assert_eq!(result.as_primitive::<Decimal256Type>(), &expected);
1783        }
1784
1785        let expected = Decimal256Array::from(vec![
1786            Some(i256::from_parts(u128::MAX - 1, 0)),
1787            Some(i256::from_i128(-2)),
1788            None,
1789        ])
1790        .with_precision_and_scale(71, 2)
1791        .unwrap();
1792        for operation in [sub, sub_wrapping] {
1793            let result = operation(&lhs, &rhs).unwrap();
1794            assert_eq!(result.as_primitive::<Decimal256Type>(), &expected);
1795        }
1796
1797        let lhs = Decimal256Array::from(vec![i256::MAX])
1798            .with_precision_and_scale(76, 0)
1799            .unwrap();
1800        let rhs = Decimal256Array::from(vec![i256::ONE])
1801            .with_precision_and_scale(76, 0)
1802            .unwrap();
1803        for operation in [add, add_wrapping] {
1804            assert_eq!(
1805                operation(&lhs, &rhs).unwrap_err().to_string(),
1806                format!(
1807                    "Arithmetic overflow: Overflow happened on: {:?} + {:?}",
1808                    i256::MAX,
1809                    i256::ONE
1810                )
1811            );
1812        }
1813
1814        let lhs = Decimal256Array::from(vec![i256::MIN])
1815            .with_precision_and_scale(76, 0)
1816            .unwrap();
1817        for operation in [sub, sub_wrapping] {
1818            assert_eq!(
1819                operation(&lhs, &rhs).unwrap_err().to_string(),
1820                format!(
1821                    "Arithmetic overflow: Overflow happened on: {:?} - {:?}",
1822                    i256::MIN,
1823                    i256::ONE
1824                )
1825            );
1826        }
1827    }
1828
1829    fn test_timestamp_impl<T: TimestampOp>() {
1830        let a = PrimitiveArray::<T>::new(vec![2000000, 434030324, 53943340].into(), None);
1831        let b = PrimitiveArray::<T>::new(vec![329593, 59349, 694994].into(), None);
1832
1833        let result = sub(&a, &b).unwrap();
1834        assert_eq!(
1835            result.as_primitive::<T::Duration>().values(),
1836            &[1670407, 433970975, 53248346]
1837        );
1838
1839        let r2 = add(&b, &result.as_ref()).unwrap();
1840        assert_eq!(r2.as_ref(), &a);
1841
1842        let r3 = add(&result.as_ref(), &b).unwrap();
1843        assert_eq!(r3.as_ref(), &a);
1844
1845        let format_array = |x: &dyn Array| -> Vec<String> {
1846            x.as_primitive::<T>()
1847                .values()
1848                .into_iter()
1849                .map(|x| as_datetime::<T>(*x).unwrap().to_string())
1850                .collect()
1851        };
1852
1853        let values = vec![
1854            "1970-01-01T00:00:00Z",
1855            "2010-04-01T04:00:20Z",
1856            "1960-01-30T04:23:20Z",
1857        ]
1858        .into_iter()
1859        .map(|x| {
1860            T::from_naive_datetime(DateTime::parse_from_rfc3339(x).unwrap().naive_utc(), None)
1861                .unwrap()
1862        })
1863        .collect();
1864
1865        let a = PrimitiveArray::<T>::new(values, None);
1866        let b = IntervalYearMonthArray::from(vec![
1867            IntervalYearMonthType::make_value(5, 34),
1868            IntervalYearMonthType::make_value(-2, 4),
1869            IntervalYearMonthType::make_value(7, -4),
1870        ]);
1871        let r4 = add(&a, &b).unwrap();
1872        assert_eq!(
1873            &format_array(r4.as_ref()),
1874            &[
1875                "1977-11-01 00:00:00".to_string(),
1876                "2008-08-01 04:00:20".to_string(),
1877                "1966-09-30 04:23:20".to_string()
1878            ]
1879        );
1880
1881        let r5 = sub(&r4, &b).unwrap();
1882        assert_eq!(r5.as_ref(), &a);
1883
1884        let b = IntervalDayTimeArray::from(vec![
1885            IntervalDayTimeType::make_value(5, 454000),
1886            IntervalDayTimeType::make_value(-34, 0),
1887            IntervalDayTimeType::make_value(7, -4000),
1888        ]);
1889        let r6 = add(&a, &b).unwrap();
1890        assert_eq!(
1891            &format_array(r6.as_ref()),
1892            &[
1893                "1970-01-06 00:07:34".to_string(),
1894                "2010-02-26 04:00:20".to_string(),
1895                "1960-02-06 04:23:16".to_string()
1896            ]
1897        );
1898
1899        let r7 = sub(&r6, &b).unwrap();
1900        assert_eq!(r7.as_ref(), &a);
1901
1902        let b = IntervalMonthDayNanoArray::from(vec![
1903            IntervalMonthDayNanoType::make_value(344, 34, -43_000_000_000),
1904            IntervalMonthDayNanoType::make_value(-593, -33, 13_000_000_000),
1905            IntervalMonthDayNanoType::make_value(5, 2, 493_000_000_000),
1906        ]);
1907        let r8 = add(&a, &b).unwrap();
1908        assert_eq!(
1909            &format_array(r8.as_ref()),
1910            &[
1911                "1998-10-04 23:59:17".to_string(),
1912                "1960-09-29 04:00:33".to_string(),
1913                "1960-07-02 04:31:33".to_string()
1914            ]
1915        );
1916
1917        let r9 = sub(&r8, &b).unwrap();
1918        // Note: subtraction is not the inverse of addition for intervals
1919        assert_eq!(
1920            &format_array(r9.as_ref()),
1921            &[
1922                "1970-01-02 00:00:00".to_string(),
1923                "2010-04-02 04:00:20".to_string(),
1924                "1960-01-31 04:23:20".to_string()
1925            ]
1926        );
1927    }
1928
1929    #[test]
1930    fn test_timestamp() {
1931        test_timestamp_impl::<TimestampSecondType>();
1932        test_timestamp_impl::<TimestampMillisecondType>();
1933        test_timestamp_impl::<TimestampMicrosecondType>();
1934        test_timestamp_impl::<TimestampNanosecondType>();
1935    }
1936
1937    #[test]
1938    fn test_interval() {
1939        let a = IntervalYearMonthArray::from(vec![
1940            IntervalYearMonthType::make_value(32, 4),
1941            IntervalYearMonthType::make_value(32, 4),
1942        ]);
1943        let b = IntervalYearMonthArray::from(vec![
1944            IntervalYearMonthType::make_value(-4, 6),
1945            IntervalYearMonthType::make_value(-3, 23),
1946        ]);
1947        let result = add(&a, &b).unwrap();
1948        assert_eq!(
1949            result.as_ref(),
1950            &IntervalYearMonthArray::from(vec![
1951                IntervalYearMonthType::make_value(28, 10),
1952                IntervalYearMonthType::make_value(29, 27)
1953            ])
1954        );
1955        let result = sub(&a, &b).unwrap();
1956        assert_eq!(
1957            result.as_ref(),
1958            &IntervalYearMonthArray::from(vec![
1959                IntervalYearMonthType::make_value(36, -2),
1960                IntervalYearMonthType::make_value(35, -19)
1961            ])
1962        );
1963
1964        let a = IntervalDayTimeArray::from(vec![
1965            IntervalDayTimeType::make_value(32, 4),
1966            IntervalDayTimeType::make_value(32, 4),
1967        ]);
1968        let b = IntervalDayTimeArray::from(vec![
1969            IntervalDayTimeType::make_value(-4, 6),
1970            IntervalDayTimeType::make_value(-3, 23),
1971        ]);
1972        let result = add(&a, &b).unwrap();
1973        assert_eq!(
1974            result.as_ref(),
1975            &IntervalDayTimeArray::from(vec![
1976                IntervalDayTimeType::make_value(28, 10),
1977                IntervalDayTimeType::make_value(29, 27)
1978            ])
1979        );
1980        let result = sub(&a, &b).unwrap();
1981        assert_eq!(
1982            result.as_ref(),
1983            &IntervalDayTimeArray::from(vec![
1984                IntervalDayTimeType::make_value(36, -2),
1985                IntervalDayTimeType::make_value(35, -19)
1986            ])
1987        );
1988        let a = IntervalMonthDayNanoArray::from(vec![
1989            IntervalMonthDayNanoType::make_value(32, 4, 4000000000000),
1990            IntervalMonthDayNanoType::make_value(32, 4, 45463000000000000),
1991        ]);
1992        let b = IntervalMonthDayNanoArray::from(vec![
1993            IntervalMonthDayNanoType::make_value(-4, 6, 46000000000000),
1994            IntervalMonthDayNanoType::make_value(-3, 23, 3564000000000000),
1995        ]);
1996        let result = add(&a, &b).unwrap();
1997        assert_eq!(
1998            result.as_ref(),
1999            &IntervalMonthDayNanoArray::from(vec![
2000                IntervalMonthDayNanoType::make_value(28, 10, 50000000000000),
2001                IntervalMonthDayNanoType::make_value(29, 27, 49027000000000000)
2002            ])
2003        );
2004        let result = sub(&a, &b).unwrap();
2005        assert_eq!(
2006            result.as_ref(),
2007            &IntervalMonthDayNanoArray::from(vec![
2008                IntervalMonthDayNanoType::make_value(36, -2, -42000000000000),
2009                IntervalMonthDayNanoType::make_value(35, -19, 41899000000000000)
2010            ])
2011        );
2012        let a = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNano::MAX]);
2013        let b = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNano::ONE]);
2014        let err = add(&a, &b).unwrap_err().to_string();
2015        assert_eq!(
2016            err,
2017            "Arithmetic overflow: Overflow happened on: 2147483647 + 1"
2018        );
2019    }
2020
2021    #[test]
2022    fn test_interval_mul_i64() {
2023        let interval = IntervalYearMonthArray::from(vec![16, 5, 0]);
2024        let factor = Int64Array::from(vec![3, -2, i64::MAX]);
2025        let expected = IntervalYearMonthArray::from(vec![48, -10, 0]);
2026        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
2027        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
2028
2029        let interval = IntervalDayTimeArray::from(vec![
2030            Some(IntervalDayTimeType::make_value(10, 2 * 60 * 60 * 1000)),
2031            None,
2032        ]);
2033        let factor = Int64Array::new_scalar(3);
2034        let expected = IntervalDayTimeArray::from(vec![
2035            Some(IntervalDayTimeType::make_value(30, 6 * 60 * 60 * 1000)),
2036            None,
2037        ]);
2038        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
2039        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
2040
2041        let null_factor = Scalar::new(Int64Array::new_null(1));
2042        let expected = IntervalDayTimeArray::new_null(interval.len());
2043        assert_eq!(mul(&interval, &null_factor).unwrap().as_ref(), &expected);
2044        assert_eq!(mul(&null_factor, &interval).unwrap().as_ref(), &expected);
2045
2046        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2047            12,
2048            15,
2049            5_000_000_000,
2050        ));
2051        let factor = Int64Array::from(vec![2, 0, -1]);
2052        let expected = IntervalMonthDayNanoArray::from(vec![
2053            IntervalMonthDayNanoType::make_value(24, 30, 10_000_000_000),
2054            IntervalMonthDayNanoType::make_value(0, 0, 0),
2055            IntervalMonthDayNanoType::make_value(-12, -15, -5_000_000_000),
2056        ]);
2057        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
2058        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
2059
2060        let float_factor = Float64Array::new_scalar(2.);
2061        assert!(mul_wrapping(&float_factor, &interval).is_err());
2062        assert!(mul_wrapping(&factor, &interval).is_err());
2063    }
2064
2065    #[test]
2066    fn test_interval_mul_div_f64() {
2067        const HOUR_NANOS: i64 = 3_600_000_000_000;
2068        const MINUTE_NANOS: i64 = 60_000_000_000;
2069
2070        // Adapted from DuckDB's interval multiplication tests:
2071        // https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/test/sql/function/interval/test_interval_muldiv.test#L1-L99
2072        // DuckDB's cases come from PostgreSQL's interval regression tests:
2073        // https://github.com/postgres/postgres/blob/78758d37306cd89ab060f00cb06f249018d5b8da/src/test/regress/sql/interval.sql#L118-L164
2074        let interval = IntervalMonthDayNanoArray::from(vec![
2075            IntervalMonthDayNanoType::make_value(41, 12, 360 * HOUR_NANOS),
2076            IntervalMonthDayNanoType::make_value(-41, -12, 360 * HOUR_NANOS),
2077            IntervalMonthDayNanoType::make_value(1, 1, 0),
2078            IntervalMonthDayNanoType::make_value(0, 0, 1),
2079            IntervalMonthDayNanoType::make_value(0, 0, 3),
2080            IntervalMonthDayNanoType::make_value(0, 0, -1),
2081            IntervalMonthDayNanoType::make_value(0, 0, -3),
2082        ]);
2083        let factor = Float64Array::from(vec![0.3, 0.3, 1.5, 0.5, 0.5, 0.5, 0.5]);
2084        let expected = IntervalMonthDayNanoArray::from(vec![
2085            IntervalMonthDayNanoType::make_value(12, 12, 122 * HOUR_NANOS + 24 * MINUTE_NANOS),
2086            IntervalMonthDayNanoType::make_value(-12, -12, 93 * HOUR_NANOS + 36 * MINUTE_NANOS),
2087            IntervalMonthDayNanoType::make_value(1, 16, 12 * HOUR_NANOS),
2088            IntervalMonthDayNanoType::make_value(0, 0, 0),
2089            IntervalMonthDayNanoType::make_value(0, 0, 2),
2090            IntervalMonthDayNanoType::make_value(0, 0, 0),
2091            IntervalMonthDayNanoType::make_value(0, 0, -2),
2092        ]);
2093        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
2094        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
2095
2096        let interval = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
2097            9,
2098            -27,
2099            45_296 * NANOSECONDS,
2100        )]);
2101        let factor = Float64Array::new_scalar(0.3);
2102        let expected = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
2103            2,
2104            13,
2105            4_948_800_000_000,
2106        )]);
2107        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
2108
2109        let interval = IntervalMonthDayNanoArray::from(vec![
2110            IntervalMonthDayNanoType::make_value(0, 1, 0),
2111            IntervalMonthDayNanoType::make_value(4, 0, 0),
2112            IntervalMonthDayNanoType::make_value(1, 1, 0),
2113            IntervalMonthDayNanoType::make_value(0, 0, (1_i64 << 53) - 1),
2114            IntervalMonthDayNanoType::make_value(0, 0, i64::MAX),
2115            IntervalMonthDayNanoType::make_value(1, 1, 1),
2116            IntervalMonthDayNanoType::make_value(1, 1, 1),
2117            IntervalMonthDayNanoType::make_value(1, 0, 0),
2118        ]);
2119        let factor = Float64Array::from(vec![
2120            3.,
2121            5.,
2122            2.,
2123            0.7,
2124            1.,
2125            f64::INFINITY,
2126            f64::NEG_INFINITY,
2127            -2.,
2128        ]);
2129        let expected = IntervalMonthDayNanoArray::from(vec![
2130            IntervalMonthDayNanoType::make_value(0, 0, 8 * HOUR_NANOS),
2131            IntervalMonthDayNanoType::make_value(0, 24, 0),
2132            IntervalMonthDayNanoType::make_value(0, 15, 12 * HOUR_NANOS),
2133            IntervalMonthDayNanoType::make_value(0, 0, 12_867_427_506_772_844),
2134            IntervalMonthDayNanoType::make_value(0, 0, i64::MAX),
2135            IntervalMonthDayNanoType::make_value(0, 0, 0),
2136            IntervalMonthDayNanoType::make_value(0, 0, 0),
2137            IntervalMonthDayNanoType::make_value(0, -15, 0),
2138        ]);
2139        assert_eq!(div(&interval, &factor).unwrap().as_ref(), &expected);
2140
2141        let null_factor = Scalar::new(Float64Array::new_null(1));
2142        assert_eq!(
2143            mul(&interval, &null_factor).unwrap().as_ref(),
2144            &IntervalMonthDayNanoArray::new_null(interval.len())
2145        );
2146    }
2147
2148    #[test]
2149    fn test_interval_mul_div_f64_errors() {
2150        let factor = Float64Array::new_scalar(2.);
2151        let year_month = IntervalYearMonthArray::new_scalar(1);
2152        let day_time = IntervalDayTimeArray::new_scalar(IntervalDayTime::new(1, 1));
2153        for interval in [&year_month as &dyn Datum, &day_time] {
2154            assert!(mul(interval, &factor).is_err());
2155            assert!(mul(&factor, interval).is_err());
2156            assert!(div(interval, &factor).is_err());
2157        }
2158
2159        let interval =
2160            IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(1, 1, 1));
2161
2162        assert!(matches!(
2163            add(&interval, &factor),
2164            Err(ArrowError::InvalidArgumentError(_))
2165        ));
2166
2167        let zero = Float64Array::new_scalar(-0.);
2168        assert!(matches!(
2169            div(&interval, &zero),
2170            Err(ArrowError::DivideByZero)
2171        ));
2172
2173        assert!(div(&factor, &interval).is_err());
2174
2175        for factor in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
2176            let factor = Float64Array::new_scalar(factor);
2177            assert!(matches!(
2178                mul(&interval, &factor),
2179                Err(ArrowError::ArithmeticOverflow(_))
2180            ));
2181        }
2182
2183        let nan = Float64Array::new_scalar(f64::NAN);
2184        assert!(matches!(
2185            div(&interval, &nan),
2186            Err(ArrowError::ArithmeticOverflow(_))
2187        ));
2188
2189        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2190            i32::MAX,
2191            0,
2192            0,
2193        ));
2194        assert!(matches!(
2195            mul(&interval, &factor),
2196            Err(ArrowError::ArithmeticOverflow(_))
2197        ));
2198
2199        let factor = Float64Array::new_scalar(1.5);
2200        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2201            0,
2202            0,
2203            i64::MAX,
2204        ));
2205        assert!(matches!(
2206            mul(&interval, &factor),
2207            Err(ArrowError::ArithmeticOverflow(_))
2208        ));
2209
2210        let factor = Float64Array::new_scalar(1.000_000_000_4);
2211        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2212            i32::MIN,
2213            0,
2214            0,
2215        ));
2216        assert!(matches!(
2217            mul(&interval, &factor),
2218            Err(ArrowError::ArithmeticOverflow(_))
2219        ));
2220
2221        let factor = Float64Array::new_scalar(0.999_999_999);
2222        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2223            1,
2224            i32::MAX,
2225            0,
2226        ));
2227        assert!(matches!(
2228            mul(&interval, &factor),
2229            Err(ArrowError::ArithmeticOverflow(_))
2230        ));
2231    }
2232
2233    #[test]
2234    fn test_interval_mul_i64_overflow() {
2235        let interval = IntervalYearMonthArray::from(vec![i32::MAX]);
2236        let factor = Int64Array::from(vec![2]);
2237        assert_eq!(
2238            mul(&interval, &factor).unwrap_err().to_string(),
2239            "Arithmetic overflow: Overflow happened on: 2147483647 * 2"
2240        );
2241
2242        let interval = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
2243            0,
2244            0,
2245            i64::MAX,
2246        )]);
2247        assert_eq!(
2248            mul(&interval, &factor).unwrap_err().to_string(),
2249            "Arithmetic overflow: Overflow happened on: 9223372036854775807 * 2"
2250        );
2251    }
2252
2253    fn test_duration_impl<T: ArrowPrimitiveType<Native = i64>>() {
2254        let a = PrimitiveArray::<T>::new(vec![1000, 4394, -3944].into(), None);
2255        let b = PrimitiveArray::<T>::new(vec![4, -5, -243].into(), None);
2256
2257        let result = add(&a, &b).unwrap();
2258        assert_eq!(result.as_primitive::<T>().values(), &[1004, 4389, -4187]);
2259        let result = sub(&a, &b).unwrap();
2260        assert_eq!(result.as_primitive::<T>().values(), &[996, 4399, -3701]);
2261
2262        let err = mul(&a, &b).unwrap_err().to_string();
2263        assert!(
2264            err.contains("Invalid duration arithmetic operation"),
2265            "{err}"
2266        );
2267
2268        let err = div(&a, &b).unwrap_err().to_string();
2269        assert!(
2270            err.contains("Invalid duration arithmetic operation"),
2271            "{err}"
2272        );
2273
2274        let err = rem(&a, &b).unwrap_err().to_string();
2275        assert!(
2276            err.contains("Invalid duration arithmetic operation"),
2277            "{err}"
2278        );
2279
2280        let a = PrimitiveArray::<T>::new(vec![i64::MAX].into(), None);
2281        let b = PrimitiveArray::<T>::new(vec![1].into(), None);
2282        let err = add(&a, &b).unwrap_err().to_string();
2283        assert_eq!(
2284            err,
2285            "Arithmetic overflow: Overflow happened on: 9223372036854775807 + 1"
2286        );
2287    }
2288
2289    #[test]
2290    fn test_duration() {
2291        test_duration_impl::<DurationSecondType>();
2292        test_duration_impl::<DurationMillisecondType>();
2293        test_duration_impl::<DurationMicrosecondType>();
2294        test_duration_impl::<DurationNanosecondType>();
2295    }
2296
2297    fn test_date_impl<T: ArrowPrimitiveType, F>(f: F)
2298    where
2299        F: Fn(NaiveDate) -> T::Native,
2300        T::Native: TryInto<i64>,
2301    {
2302        let a = PrimitiveArray::<T>::new(
2303            vec![
2304                f(NaiveDate::from_ymd_opt(1979, 1, 30).unwrap()),
2305                f(NaiveDate::from_ymd_opt(2010, 4, 3).unwrap()),
2306                f(NaiveDate::from_ymd_opt(2008, 2, 29).unwrap()),
2307            ]
2308            .into(),
2309            None,
2310        );
2311
2312        let b = IntervalYearMonthArray::from(vec![
2313            IntervalYearMonthType::make_value(34, 2),
2314            IntervalYearMonthType::make_value(3, -3),
2315            IntervalYearMonthType::make_value(-12, 4),
2316        ]);
2317
2318        let format_array = |x: &dyn Array| -> Vec<String> {
2319            x.as_primitive::<T>()
2320                .values()
2321                .into_iter()
2322                .map(|x| {
2323                    as_date::<T>((*x).try_into().ok().unwrap())
2324                        .unwrap()
2325                        .to_string()
2326                })
2327                .collect()
2328        };
2329
2330        let result = add(&a, &b).unwrap();
2331        assert_eq!(
2332            &format_array(result.as_ref()),
2333            &[
2334                "2013-03-30".to_string(),
2335                "2013-01-03".to_string(),
2336                "1996-06-29".to_string(),
2337            ]
2338        );
2339        let result = sub(&result, &b).unwrap();
2340        assert_eq!(result.as_ref(), &a);
2341
2342        let b = IntervalDayTimeArray::from(vec![
2343            IntervalDayTimeType::make_value(34, 2),
2344            IntervalDayTimeType::make_value(3, -3),
2345            IntervalDayTimeType::make_value(-12, 4),
2346        ]);
2347
2348        let result = add(&a, &b).unwrap();
2349        assert_eq!(
2350            &format_array(result.as_ref()),
2351            &[
2352                "1979-03-05".to_string(),
2353                "2010-04-06".to_string(),
2354                "2008-02-17".to_string(),
2355            ]
2356        );
2357        let result = sub(&result, &b).unwrap();
2358        assert_eq!(result.as_ref(), &a);
2359
2360        let b = IntervalMonthDayNanoArray::from(vec![
2361            IntervalMonthDayNanoType::make_value(34, 2, -34353534),
2362            IntervalMonthDayNanoType::make_value(3, -3, 2443),
2363            IntervalMonthDayNanoType::make_value(-12, 4, 2323242423232),
2364        ]);
2365
2366        let result = add(&a, &b).unwrap();
2367        assert_eq!(
2368            &format_array(result.as_ref()),
2369            &[
2370                "1981-12-02".to_string(),
2371                "2010-06-30".to_string(),
2372                "2007-03-04".to_string(),
2373            ]
2374        );
2375        let result = sub(&result, &b).unwrap();
2376        assert_eq!(
2377            &format_array(result.as_ref()),
2378            &[
2379                "1979-01-31".to_string(),
2380                "2010-04-02".to_string(),
2381                "2008-02-29".to_string(),
2382            ]
2383        );
2384    }
2385
2386    #[test]
2387    fn test_date() {
2388        test_date_impl::<Date32Type, _>(Date32Type::from_naive_date);
2389        test_date_impl::<Date64Type, _>(Date64Type::from_naive_date);
2390
2391        let a = Date32Array::from(vec![i32::MIN, i32::MAX, 23, 7684]);
2392        let b = Date32Array::from(vec![i32::MIN, i32::MIN, -2, 45]);
2393        let result = sub(&a, &b).unwrap();
2394        assert_eq!(
2395            result.as_primitive::<DurationSecondType>().values(),
2396            &[0, 371085174288000, 2160000, 660009600]
2397        );
2398
2399        let a = Date64Array::from(vec![4343, 76676, 3434]);
2400        let b = Date64Array::from(vec![3, -5, 5]);
2401        let result = sub(&a, &b).unwrap();
2402        assert_eq!(
2403            result.as_primitive::<DurationMillisecondType>().values(),
2404            &[4340, 76681, 3429]
2405        );
2406
2407        let a = Date64Array::from(vec![i64::MAX]);
2408        let b = Date64Array::from(vec![-1]);
2409        let err = sub(&a, &b).unwrap_err().to_string();
2410        assert_eq!(
2411            err,
2412            "Arithmetic overflow: Overflow happened on: 9223372036854775807 - -1"
2413        );
2414    }
2415
2416    #[test]
2417    fn test_date32_to_naive_date_opt_boundaries() {
2418        assert_eq!(MAX_VALID_DAYS, 95026236);
2419        assert_eq!(MIN_VALID_DAYS, -96465292);
2420
2421        // Valid boundary dates work
2422        assert!(Date32Type::to_naive_date_opt(MAX_VALID_DAYS).is_some());
2423        assert!(Date32Type::to_naive_date_opt(MIN_VALID_DAYS).is_some());
2424
2425        // Beyond boundaries fail
2426        assert!(Date32Type::to_naive_date_opt(MAX_VALID_DAYS + 1).is_none());
2427        assert!(Date32Type::to_naive_date_opt(MIN_VALID_DAYS - 1).is_none());
2428
2429        // Extreme values fail
2430        assert!(Date32Type::to_naive_date_opt(i32::MAX).is_none());
2431        assert!(Date32Type::to_naive_date_opt(i32::MIN).is_none());
2432
2433        // Common values work
2434        assert!(Date32Type::to_naive_date_opt(0).is_some());
2435        assert!(Date32Type::to_naive_date_opt(date_to_days(YEAR_2000)).is_some());
2436    }
2437
2438    #[test]
2439    fn test_date64_to_naive_date_opt_boundaries() {
2440        const MS_PER_DAY: i64 = 24 * 60 * 60 * 1000;
2441
2442        // Verify boundary millisecond values
2443        assert_eq!(MAX_VALID_MILLIS, 8210266790400000i64);
2444        assert_eq!(MIN_VALID_MILLIS, -8334601228800000i64);
2445
2446        // Valid boundary dates work
2447        assert!(Date64Type::to_naive_date_opt(MAX_VALID_MILLIS).is_some());
2448        assert!(Date64Type::to_naive_date_opt(MIN_VALID_MILLIS).is_some());
2449
2450        // Beyond boundaries fail
2451        assert!(Date64Type::to_naive_date_opt(MAX_VALID_MILLIS + MS_PER_DAY).is_none());
2452        assert!(Date64Type::to_naive_date_opt(MIN_VALID_MILLIS - MS_PER_DAY).is_none());
2453
2454        // Extreme values fail
2455        assert!(Date64Type::to_naive_date_opt(i64::MAX).is_none());
2456        assert!(Date64Type::to_naive_date_opt(i64::MIN).is_none());
2457
2458        // Common values work
2459        assert!(Date64Type::to_naive_date_opt(0).is_some());
2460        assert!(Date64Type::to_naive_date_opt(date_to_millis(YEAR_2000)).is_some());
2461    }
2462
2463    macro_rules! test_year_month_ops {
2464        ($type:ty, $date_fn:expr) => {{
2465            let date = $date_fn(YEAR_2000);
2466
2467            // Normal operations succeed
2468            assert!(
2469                <$type>::add_year_months_opt(date, 120).is_some(),
2470                "add_year_months: normal add"
2471            );
2472            assert!(
2473                <$type>::add_year_months_opt(date, 0).is_some(),
2474                "add_year_months: zero interval"
2475            );
2476            assert!(
2477                <$type>::subtract_year_months_opt(date, 120).is_some(),
2478                "subtract_year_months: normal subtract"
2479            );
2480            assert!(
2481                <$type>::subtract_year_months_opt(date, 0).is_some(),
2482                "subtract_year_months: zero interval"
2483            );
2484
2485            // Large but valid years work
2486            let large_year = $date_fn(NaiveDate::from_ymd_opt(5000, 1, 1).unwrap());
2487            let neg_year = $date_fn(NaiveDate::from_ymd_opt(-5000, 12, 31).unwrap());
2488            assert!(
2489                <$type>::add_year_months_opt(large_year, 12).is_some(),
2490                "add_year_months: large year"
2491            );
2492            assert!(
2493                <$type>::add_year_months_opt(neg_year, -12).is_some(),
2494                "add_year_months: negative year"
2495            );
2496            assert!(
2497                <$type>::subtract_year_months_opt(large_year, 12).is_some(),
2498                "subtract_year_months: large year"
2499            );
2500            assert!(
2501                <$type>::subtract_year_months_opt(neg_year, -12).is_some(),
2502                "subtract_year_months: negative year"
2503            );
2504
2505            // Overflow handling
2506            assert!(
2507                <$type>::subtract_year_months_opt($date_fn(MIN_VALID_DATE), 1).is_none(),
2508                "subtract_year_months: overflow days from min"
2509            );
2510            assert!(
2511                <$type>::subtract_year_months_opt($date_fn(MAX_VALID_DATE), -1).is_none(),
2512                "subtract_year_months: overflow neg days from max"
2513            );
2514            assert!(
2515                <$type>::add_year_months_opt($date_fn(MAX_VALID_DATE), 1).is_none(),
2516                "add_year_months: overflow days"
2517            );
2518            assert!(
2519                <$type>::add_year_months_opt($date_fn(MIN_VALID_DATE), -1).is_none(),
2520                "add_year_months: overflow neg days"
2521            );
2522        }};
2523    }
2524
2525    #[test]
2526    fn test_date_year_month_operations() {
2527        test_year_month_ops!(Date32Type, date_to_days);
2528        test_year_month_ops!(Date64Type, date_to_millis);
2529    }
2530
2531    macro_rules! test_day_time_ops {
2532        ($type:ty, $date_fn:expr) => {{
2533            let date = $date_fn(YEAR_2000);
2534
2535            // Moderate intervals succeed
2536            assert!(
2537                <$type>::add_day_time_opt(date, IntervalDayTime::new(30, 0)).is_some(),
2538                "add_day_time: +30 days"
2539            );
2540            assert!(
2541                <$type>::add_day_time_opt(date, IntervalDayTime::new(-30, 0)).is_some(),
2542                "add_day_time: -30 days"
2543            );
2544            assert!(
2545                <$type>::add_day_time_opt(date, IntervalDayTime::new(1000, 12345)).is_some(),
2546                "add_day_time: normal"
2547            );
2548            assert!(
2549                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(30, 0)).is_some(),
2550                "subtract_day_time: +30 days"
2551            );
2552            assert!(
2553                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(-30, 0)).is_some(),
2554                "subtract_day_time: -30 days"
2555            );
2556            assert!(
2557                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(1000, 12345)).is_some(),
2558                "subtract_day_time: normal"
2559            );
2560
2561            // Overflow handling - subtract
2562            assert!(
2563                <$type>::subtract_day_time_opt(
2564                    $date_fn(MIN_VALID_DATE),
2565                    IntervalDayTime::new(1, 0)
2566                )
2567                .is_none(),
2568                "subtract_day_time: overflow days from min"
2569            );
2570            assert!(
2571                <$type>::subtract_day_time_opt(
2572                    $date_fn(MAX_VALID_DATE),
2573                    IntervalDayTime::new(-1, 0)
2574                )
2575                .is_none(),
2576                "subtract_day_time: overflow neg days from max"
2577            );
2578
2579            // Overflow handling - add
2580            assert!(
2581                <$type>::add_day_time_opt($date_fn(MAX_VALID_DATE), IntervalDayTime::new(1, 0))
2582                    .is_none(),
2583                "add_day_time: overflow days"
2584            );
2585            assert!(
2586                <$type>::add_day_time_opt($date_fn(MIN_VALID_DATE), IntervalDayTime::new(-1, 0))
2587                    .is_none(),
2588                "add_day_time: overflow neg days"
2589            );
2590
2591            // Extreme intervals fail
2592            assert!(
2593                <$type>::add_day_time_opt(
2594                    $date_fn(EPOCH),
2595                    IntervalDayTime::new(i32::MAX, i32::MAX)
2596                )
2597                .is_none(),
2598                "add_day_time: max interval"
2599            );
2600            assert!(
2601                <$type>::add_day_time_opt(
2602                    $date_fn(EPOCH),
2603                    IntervalDayTime::new(i32::MIN, i32::MIN)
2604                )
2605                .is_none(),
2606                "add_day_time: min interval"
2607            );
2608            assert!(
2609                <$type>::subtract_day_time_opt(
2610                    $date_fn(EPOCH),
2611                    IntervalDayTime::new(i32::MAX, i32::MAX)
2612                )
2613                .is_none(),
2614                "subtract_day_time: max interval"
2615            );
2616            assert!(
2617                <$type>::subtract_day_time_opt(
2618                    $date_fn(EPOCH),
2619                    IntervalDayTime::new(i32::MIN, i32::MIN)
2620                )
2621                .is_none(),
2622                "subtract_day_time: min interval"
2623            );
2624        }};
2625    }
2626
2627    #[test]
2628    fn test_date_day_time_operations() {
2629        test_day_time_ops!(Date32Type, date_to_days);
2630        test_day_time_ops!(Date64Type, date_to_millis);
2631    }
2632
2633    macro_rules! test_month_day_nano_ops {
2634        ($type:ty, $date_fn:expr) => {{
2635            let date = $date_fn(YEAR_2000);
2636            let zero = IntervalMonthDayNano::new(0, 0, 0);
2637
2638            // Normal operations succeed
2639            assert!(
2640                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(1, 30, 0))
2641                    .is_some(),
2642                "add_month_day_nano: +1mo +30d"
2643            );
2644            assert!(
2645                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(-1, -30, 0))
2646                    .is_some(),
2647                "add_month_day_nano: -1mo -30d"
2648            );
2649            assert!(
2650                <$type>::add_month_day_nano_opt(date, zero).is_some(),
2651                "add_month_day_nano: zero interval"
2652            );
2653            assert!(
2654                <$type>::add_month_day_nano_opt(
2655                    date,
2656                    IntervalMonthDayNano::new(2, 10, 123_456_789_000)
2657                )
2658                .is_some(),
2659                "add_month_day_nano: normal"
2660            );
2661            assert!(
2662                <$type>::subtract_month_day_nano_opt(date, IntervalMonthDayNano::new(1, 30, 0))
2663                    .is_some(),
2664                "subtract_month_day_nano: +1mo +30d"
2665            );
2666            assert!(
2667                <$type>::subtract_month_day_nano_opt(date, IntervalMonthDayNano::new(-1, -30, 0))
2668                    .is_some(),
2669                "subtract_month_day_nano: -1mo -30d"
2670            );
2671            assert!(
2672                <$type>::subtract_month_day_nano_opt(date, zero).is_some(),
2673                "subtract_month_day_nano: zero interval"
2674            );
2675            assert!(
2676                <$type>::subtract_month_day_nano_opt(
2677                    date,
2678                    IntervalMonthDayNano::new(2, 10, 123_456_789_000)
2679                )
2680                .is_some(),
2681                "subtract_month_day_nano: normal"
2682            );
2683
2684            // Overflow handling - subtract
2685            assert!(
2686                <$type>::subtract_month_day_nano_opt(
2687                    $date_fn(MIN_VALID_DATE),
2688                    IntervalMonthDayNano::new(0, 1, 0)
2689                )
2690                .is_none(),
2691                "subtract_month_day_nano: overflow days from min"
2692            );
2693            assert!(
2694                <$type>::subtract_month_day_nano_opt(
2695                    $date_fn(MAX_VALID_DATE),
2696                    IntervalMonthDayNano::new(0, -1, 0)
2697                )
2698                .is_none(),
2699                "subtract_month_day_nano: overflow neg days from max"
2700            );
2701
2702            // Overflow handling - add
2703            assert!(
2704                <$type>::add_month_day_nano_opt(
2705                    $date_fn(MAX_VALID_DATE),
2706                    IntervalMonthDayNano::new(0, 1, 0)
2707                )
2708                .is_none(),
2709                "add_month_day_nano: overflow days"
2710            );
2711            assert!(
2712                <$type>::add_month_day_nano_opt(
2713                    $date_fn(MIN_VALID_DATE),
2714                    IntervalMonthDayNano::new(0, -1, 0)
2715                )
2716                .is_none(),
2717                "add_month_day_nano: overflow neg days"
2718            );
2719
2720            // Nanosecond precision works
2721            assert!(
2722                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(0, 0, 999_999_999))
2723                    .is_some(),
2724                "add_month_day_nano: nanos"
2725            );
2726            assert!(
2727                <$type>::subtract_month_day_nano_opt(
2728                    date,
2729                    IntervalMonthDayNano::new(0, 0, 999_999_999)
2730                )
2731                .is_some(),
2732                "subtract_month_day_nano: nanos"
2733            );
2734            // 1 day in nanos
2735            assert!(
2736                <$type>::add_month_day_nano_opt(
2737                    date,
2738                    IntervalMonthDayNano::new(0, 0, 86_400_000_000_000)
2739                )
2740                .is_some(),
2741                "add_month_day_nano: 1 day nanos"
2742            );
2743            assert!(
2744                <$type>::subtract_month_day_nano_opt(
2745                    date,
2746                    IntervalMonthDayNano::new(0, 0, 86_400_000_000_000)
2747                )
2748                .is_some(),
2749                "subtract_month_day_nano: 1 day nanos"
2750            );
2751        }};
2752    }
2753
2754    #[test]
2755    fn test_date_month_day_nano_operations() {
2756        test_month_day_nano_ops!(Date32Type, date_to_days);
2757        test_month_day_nano_ops!(Date64Type, date_to_millis);
2758    }
2759}