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            // max(s1, s2)
1046            let result_scale = *s1.max(s2);
1047
1048            // max(s1, s2) + max(p1-s1, p2-s2) + 1
1049            let result_precision =
1050                (result_scale.saturating_add((*p1 as i8 - s1).max(*p2 as i8 - s2)) as u8)
1051                    .saturating_add(1)
1052                    .min(T::MAX_PRECISION);
1053
1054            let l_mul = T::Native::usize_as(10).pow_checked((result_scale - s1) as _)?;
1055            let r_mul = T::Native::usize_as(10).pow_checked((result_scale - s2) as _)?;
1056
1057            match op {
1058                // Equal scales make both decimal multipliers one.
1059                Op::Add | Op::AddWrapping if s1 == s2 => {
1060                    try_op!(l, l_s, r, r_s, l.add_checked(r))
1061                }
1062                Op::Sub | Op::SubWrapping if s1 == s2 => {
1063                    try_op!(l, l_s, r, r_s, l.sub_checked(r))
1064                }
1065                Op::Add | Op::AddWrapping => {
1066                    try_op!(
1067                        l,
1068                        l_s,
1069                        r,
1070                        r_s,
1071                        l.mul_checked(l_mul)?.add_checked(r.mul_checked(r_mul)?)
1072                    )
1073                }
1074                Op::Sub | Op::SubWrapping => {
1075                    try_op!(
1076                        l,
1077                        l_s,
1078                        r,
1079                        r_s,
1080                        l.mul_checked(l_mul)?.sub_checked(r.mul_checked(r_mul)?)
1081                    )
1082                }
1083                _ => unreachable!(),
1084            }
1085            .with_precision_and_scale(result_precision, result_scale)?
1086        }
1087        Op::Mul | Op::MulWrapping => {
1088            let result_precision = p1.saturating_add(p2 + 1).min(T::MAX_PRECISION);
1089            let result_scale = s1.saturating_add(*s2);
1090            if result_scale > T::MAX_SCALE {
1091                // SQL standard says that if the resulting scale of a multiply operation goes
1092                // beyond the maximum, rounding is not acceptable and thus an error occurs
1093                return Err(ArrowError::InvalidArgumentError(format!(
1094                    "Output scale of {} {op} {} would exceed max scale of {}",
1095                    l.data_type(),
1096                    r.data_type(),
1097                    T::MAX_SCALE
1098                )));
1099            }
1100
1101            try_op!(l, l_s, r, r_s, l.mul_checked(r))
1102                .with_precision_and_scale(result_precision, result_scale)?
1103        }
1104
1105        Op::Div => {
1106            // Follow postgres and MySQL adding a fixed scale increment of 4
1107            // s1 + 4
1108            let result_scale = s1.saturating_add(4).min(T::MAX_SCALE);
1109            let mul_pow = result_scale - s1 + s2;
1110
1111            // p1 - s1 + s2 + result_scale
1112            let result_precision = (mul_pow.saturating_add(*p1 as i8) as u8).min(T::MAX_PRECISION);
1113
1114            let (l_mul, r_mul) = match mul_pow.cmp(&0) {
1115                Ordering::Greater => (
1116                    T::Native::usize_as(10).pow_checked(mul_pow as _)?,
1117                    T::Native::ONE,
1118                ),
1119                Ordering::Equal => (T::Native::ONE, T::Native::ONE),
1120                Ordering::Less => (
1121                    T::Native::ONE,
1122                    T::Native::usize_as(10).pow_checked(mul_pow.neg_wrapping() as _)?,
1123                ),
1124            };
1125
1126            try_op!(
1127                l,
1128                l_s,
1129                r,
1130                r_s,
1131                match l.mul_checked(l_mul) {
1132                    Ok(scaled) => scaled.div_checked(r.mul_checked(r_mul)?),
1133                    Err(_) => scaled_div::<T>(l, r, mul_pow),
1134                }
1135            )
1136            .with_precision_and_scale(result_precision, result_scale)?
1137        }
1138
1139        Op::Rem => {
1140            // max(s1, s2)
1141            let result_scale = *s1.max(s2);
1142            // min(p1-s1, p2 -s2) + max( s1,s2 )
1143            let result_precision =
1144                (result_scale.saturating_add((*p1 as i8 - s1).min(*p2 as i8 - s2)) as u8)
1145                    .min(T::MAX_PRECISION);
1146
1147            let l_mul = T::Native::usize_as(10).pow_wrapping((result_scale - s1) as _);
1148            let r_mul = T::Native::usize_as(10).pow_wrapping((result_scale - s2) as _);
1149
1150            try_op!(
1151                l,
1152                l_s,
1153                r,
1154                r_s,
1155                l.mul_checked(l_mul)?.mod_checked(r.mul_checked(r_mul)?)
1156            )
1157            .with_precision_and_scale(result_precision, result_scale)?
1158        }
1159    };
1160
1161    Ok(Arc::new(array))
1162}
1163
1164#[cfg(test)]
1165mod tests {
1166    use super::*;
1167    use arrow_array::temporal_conversions::{as_date, as_datetime};
1168    use arrow_buffer::{ScalarBuffer, i256};
1169    use chrono::{DateTime, NaiveDate};
1170
1171    // The valid date range of NaiveDate is from January 1, -262143 to December 31, 262142 (Gregorian calendar).
1172    const MAX_VALID_DATE: NaiveDate = NaiveDate::from_ymd_opt(262142, 12, 31).unwrap();
1173    const MIN_VALID_DATE: NaiveDate = NaiveDate::from_ymd_opt(-262143, 1, 1).unwrap();
1174    const MAX_VALID_MILLIS: i64 = date_to_millis(MAX_VALID_DATE);
1175    const MIN_VALID_MILLIS: i64 = date_to_millis(MIN_VALID_DATE);
1176    const MAX_VALID_DAYS: i32 = date_to_days(MAX_VALID_DATE);
1177    const MIN_VALID_DAYS: i32 = date_to_days(MIN_VALID_DATE);
1178    const EPOCH: NaiveDate = NaiveDate::from_ymd_opt(1970, 1, 1).unwrap();
1179    const YEAR_2000: NaiveDate = NaiveDate::from_ymd_opt(2000, 1, 1).unwrap();
1180
1181    const fn date_to_millis(date: NaiveDate) -> i64 {
1182        date.signed_duration_since(EPOCH).num_milliseconds()
1183    }
1184
1185    const fn date_to_days(date: NaiveDate) -> i32 {
1186        date.signed_duration_since(EPOCH).num_days() as i32
1187    }
1188
1189    fn test_neg_primitive<T: ArrowPrimitiveType>(
1190        input: &[T::Native],
1191        out: Result<&[T::Native], &str>,
1192    ) {
1193        let a = PrimitiveArray::<T>::new(ScalarBuffer::from(input.to_vec()), None);
1194        match out {
1195            Ok(expected) => {
1196                let result = neg(&a).unwrap();
1197                assert_eq!(result.as_primitive::<T>().values(), expected);
1198            }
1199            Err(e) => {
1200                let err = neg(&a).unwrap_err().to_string();
1201                assert_eq!(e, err);
1202            }
1203        }
1204    }
1205
1206    #[test]
1207    fn test_neg() {
1208        let input = &[1, -5, 2, 693, 3929];
1209        let output = &[-1, 5, -2, -693, -3929];
1210        test_neg_primitive::<Int32Type>(input, Ok(output));
1211
1212        let input = &[1, -5, 2, 693, 3929];
1213        let output = &[-1, 5, -2, -693, -3929];
1214        test_neg_primitive::<Int64Type>(input, Ok(output));
1215        test_neg_primitive::<DurationSecondType>(input, Ok(output));
1216        test_neg_primitive::<DurationMillisecondType>(input, Ok(output));
1217        test_neg_primitive::<DurationMicrosecondType>(input, Ok(output));
1218        test_neg_primitive::<DurationNanosecondType>(input, Ok(output));
1219
1220        let input = &[f32::MAX, f32::MIN, f32::INFINITY, 1.3, 0.5];
1221        let output = &[f32::MIN, f32::MAX, f32::NEG_INFINITY, -1.3, -0.5];
1222        test_neg_primitive::<Float32Type>(input, Ok(output));
1223
1224        test_neg_primitive::<Int32Type>(
1225            &[i32::MIN],
1226            Err("Arithmetic overflow: Overflow happened on: - -2147483648"),
1227        );
1228        test_neg_primitive::<Int64Type>(
1229            &[i64::MIN],
1230            Err("Arithmetic overflow: Overflow happened on: - -9223372036854775808"),
1231        );
1232        test_neg_primitive::<DurationSecondType>(
1233            &[i64::MIN],
1234            Err("Arithmetic overflow: Overflow happened on: - -9223372036854775808"),
1235        );
1236
1237        let r = neg_wrapping(&Int32Array::from(vec![i32::MIN])).unwrap();
1238        assert_eq!(r.as_primitive::<Int32Type>().value(0), i32::MIN);
1239
1240        let r = neg_wrapping(&Int64Array::from(vec![i64::MIN])).unwrap();
1241        assert_eq!(r.as_primitive::<Int64Type>().value(0), i64::MIN);
1242
1243        let err = neg_wrapping(&DurationSecondArray::from(vec![i64::MIN]))
1244            .unwrap_err()
1245            .to_string();
1246
1247        assert_eq!(
1248            err,
1249            "Arithmetic overflow: Overflow happened on: - -9223372036854775808"
1250        );
1251
1252        let a = Decimal32Array::from(vec![1, 3, -44, 2, 4])
1253            .with_precision_and_scale(9, 6)
1254            .unwrap();
1255
1256        let r = neg(&a).unwrap();
1257        assert_eq!(r.data_type(), a.data_type());
1258        assert_eq!(
1259            r.as_primitive::<Decimal32Type>().values(),
1260            &[-1, -3, 44, -2, -4]
1261        );
1262
1263        let a = Decimal64Array::from(vec![1, 3, -44, 2, 4])
1264            .with_precision_and_scale(9, 6)
1265            .unwrap();
1266
1267        let r = neg(&a).unwrap();
1268        assert_eq!(r.data_type(), a.data_type());
1269        assert_eq!(
1270            r.as_primitive::<Decimal64Type>().values(),
1271            &[-1, -3, 44, -2, -4]
1272        );
1273
1274        let a = Decimal128Array::from(vec![1, 3, -44, 2, 4])
1275            .with_precision_and_scale(9, 6)
1276            .unwrap();
1277
1278        let r = neg(&a).unwrap();
1279        assert_eq!(r.data_type(), a.data_type());
1280        assert_eq!(
1281            r.as_primitive::<Decimal128Type>().values(),
1282            &[-1, -3, 44, -2, -4]
1283        );
1284
1285        let a = Decimal256Array::from(vec![
1286            i256::from_i128(342),
1287            i256::from_i128(-4949),
1288            i256::from_i128(3),
1289        ])
1290        .with_precision_and_scale(9, 6)
1291        .unwrap();
1292
1293        let r = neg(&a).unwrap();
1294        assert_eq!(r.data_type(), a.data_type());
1295        assert_eq!(
1296            r.as_primitive::<Decimal256Type>().values(),
1297            &[
1298                i256::from_i128(-342),
1299                i256::from_i128(4949),
1300                i256::from_i128(-3),
1301            ]
1302        );
1303
1304        let a = IntervalYearMonthArray::from(vec![
1305            IntervalYearMonthType::make_value(2, 4),
1306            IntervalYearMonthType::make_value(2, -4),
1307            IntervalYearMonthType::make_value(-3, -5),
1308        ]);
1309        let r = neg(&a).unwrap();
1310        assert_eq!(
1311            r.as_primitive::<IntervalYearMonthType>().values(),
1312            &[
1313                IntervalYearMonthType::make_value(-2, -4),
1314                IntervalYearMonthType::make_value(-2, 4),
1315                IntervalYearMonthType::make_value(3, 5),
1316            ]
1317        );
1318
1319        let a = IntervalDayTimeArray::from(vec![
1320            IntervalDayTimeType::make_value(2, 4),
1321            IntervalDayTimeType::make_value(2, -4),
1322            IntervalDayTimeType::make_value(-3, -5),
1323        ]);
1324        let r = neg(&a).unwrap();
1325        assert_eq!(
1326            r.as_primitive::<IntervalDayTimeType>().values(),
1327            &[
1328                IntervalDayTimeType::make_value(-2, -4),
1329                IntervalDayTimeType::make_value(-2, 4),
1330                IntervalDayTimeType::make_value(3, 5),
1331            ]
1332        );
1333
1334        let a = IntervalMonthDayNanoArray::from(vec![
1335            IntervalMonthDayNanoType::make_value(2, 4, 5953394),
1336            IntervalMonthDayNanoType::make_value(2, -4, -45839),
1337            IntervalMonthDayNanoType::make_value(-3, -5, 6944),
1338        ]);
1339        let r = neg(&a).unwrap();
1340        assert_eq!(
1341            r.as_primitive::<IntervalMonthDayNanoType>().values(),
1342            &[
1343                IntervalMonthDayNanoType::make_value(-2, -4, -5953394),
1344                IntervalMonthDayNanoType::make_value(-2, 4, 45839),
1345                IntervalMonthDayNanoType::make_value(3, 5, -6944),
1346            ]
1347        );
1348    }
1349
1350    #[test]
1351    fn test_integer() {
1352        let a = Int32Array::from(vec![4, 3, 5, -6, 100]);
1353        let b = Int32Array::from(vec![6, 2, 5, -7, 3]);
1354        let result = add(&a, &b).unwrap();
1355        assert_eq!(
1356            result.as_ref(),
1357            &Int32Array::from(vec![10, 5, 10, -13, 103])
1358        );
1359        let result = sub(&a, &b).unwrap();
1360        assert_eq!(result.as_ref(), &Int32Array::from(vec![-2, 1, 0, 1, 97]));
1361        let result = div(&a, &b).unwrap();
1362        assert_eq!(result.as_ref(), &Int32Array::from(vec![0, 1, 1, 0, 33]));
1363        let result = mul(&a, &b).unwrap();
1364        assert_eq!(result.as_ref(), &Int32Array::from(vec![24, 6, 25, 42, 300]));
1365        let result = rem(&a, &b).unwrap();
1366        assert_eq!(result.as_ref(), &Int32Array::from(vec![4, 1, 0, -6, 1]));
1367
1368        let a = Int8Array::from(vec![Some(2), None, Some(45)]);
1369        let b = Int8Array::from(vec![Some(5), Some(3), None]);
1370        let result = add(&a, &b).unwrap();
1371        assert_eq!(result.as_ref(), &Int8Array::from(vec![Some(7), None, None]));
1372
1373        let a = UInt8Array::from(vec![56, 5, 3]);
1374        let b = UInt8Array::from(vec![200, 2, 5]);
1375        let err = add(&a, &b).unwrap_err().to_string();
1376        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 56 + 200");
1377        let result = add_wrapping(&a, &b).unwrap();
1378        assert_eq!(result.as_ref(), &UInt8Array::from(vec![0, 7, 8]));
1379
1380        let a = UInt8Array::from(vec![34, 5, 3]);
1381        let b = UInt8Array::from(vec![200, 2, 5]);
1382        let err = sub(&a, &b).unwrap_err().to_string();
1383        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 34 - 200");
1384        let result = sub_wrapping(&a, &b).unwrap();
1385        assert_eq!(result.as_ref(), &UInt8Array::from(vec![90, 3, 254]));
1386
1387        let a = UInt8Array::from(vec![34, 5, 3]);
1388        let b = UInt8Array::from(vec![200, 2, 5]);
1389        let err = mul(&a, &b).unwrap_err().to_string();
1390        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 34 * 200");
1391        let result = mul_wrapping(&a, &b).unwrap();
1392        assert_eq!(result.as_ref(), &UInt8Array::from(vec![144, 10, 15]));
1393
1394        let a = Int16Array::from(vec![i16::MIN]);
1395        let b = Int16Array::from(vec![-1]);
1396        let err = div(&a, &b).unwrap_err().to_string();
1397        assert_eq!(
1398            err,
1399            "Arithmetic overflow: Overflow happened on: -32768 / -1"
1400        );
1401
1402        let a = Int16Array::from(vec![i16::MIN]);
1403        let b = Int16Array::from(vec![-1]);
1404        let result = rem(&a, &b).unwrap();
1405        assert_eq!(result.as_ref(), &Int16Array::from(vec![0]));
1406
1407        let a = Int16Array::from(vec![21]);
1408        let b = Int16Array::from(vec![0]);
1409        let err = div(&a, &b).unwrap_err().to_string();
1410        assert_eq!(err, "Divide by zero error");
1411
1412        let a = Int16Array::from(vec![21]);
1413        let b = Int16Array::from(vec![0]);
1414        let err = rem(&a, &b).unwrap_err().to_string();
1415        assert_eq!(err, "Divide by zero error");
1416    }
1417
1418    #[test]
1419    fn test_float() {
1420        let a = Float32Array::from(vec![1., f32::MAX, 6., -4., -1., 0.]);
1421        let b = Float32Array::from(vec![1., f32::MAX, f32::MAX, -3., 45., 0.]);
1422        let result = add(&a, &b).unwrap();
1423        assert_eq!(
1424            result.as_ref(),
1425            &Float32Array::from(vec![2., f32::INFINITY, f32::MAX, -7., 44.0, 0.])
1426        );
1427
1428        let result = sub(&a, &b).unwrap();
1429        assert_eq!(
1430            result.as_ref(),
1431            &Float32Array::from(vec![0., 0., f32::MIN, -1., -46., 0.])
1432        );
1433
1434        let result = mul(&a, &b).unwrap();
1435        assert_eq!(
1436            result.as_ref(),
1437            &Float32Array::from(vec![1., f32::INFINITY, f32::INFINITY, 12., -45., 0.])
1438        );
1439
1440        let result = div(&a, &b).unwrap();
1441        let r = result.as_primitive::<Float32Type>();
1442        assert_eq!(r.value(0), 1.);
1443        assert_eq!(r.value(1), 1.);
1444        assert!(r.value(2) < f32::EPSILON);
1445        assert_eq!(r.value(3), -4. / -3.);
1446        assert!(r.value(5).is_nan());
1447
1448        let result = rem(&a, &b).unwrap();
1449        let r = result.as_primitive::<Float32Type>();
1450        assert_eq!(&r.values()[..5], &[0., 0., 6., -1., -1.]);
1451        assert!(r.value(5).is_nan());
1452    }
1453
1454    #[test]
1455    fn test_decimal() {
1456        // 0.015 7.842 -0.577 0.334 -0.078 0.003
1457        let a = Decimal128Array::from(vec![15, 0, -577, 334, -78, 3])
1458            .with_precision_and_scale(12, 3)
1459            .unwrap();
1460
1461        // 5.4 0 -35.6 0.3 0.6 7.45
1462        let b = Decimal128Array::from(vec![54, 34, -356, 3, 6, 745])
1463            .with_precision_and_scale(12, 1)
1464            .unwrap();
1465
1466        let result = add(&a, &b).unwrap();
1467        assert_eq!(result.data_type(), &DataType::Decimal128(15, 3));
1468        assert_eq!(
1469            result.as_primitive::<Decimal128Type>().values(),
1470            &[5415, 3400, -36177, 634, 522, 74503]
1471        );
1472
1473        let result = sub(&a, &b).unwrap();
1474        assert_eq!(result.data_type(), &DataType::Decimal128(15, 3));
1475        assert_eq!(
1476            result.as_primitive::<Decimal128Type>().values(),
1477            &[-5385, -3400, 35023, 34, -678, -74497]
1478        );
1479
1480        let result = mul(&a, &b).unwrap();
1481        assert_eq!(result.data_type(), &DataType::Decimal128(25, 4));
1482        assert_eq!(
1483            result.as_primitive::<Decimal128Type>().values(),
1484            &[810, 0, 205412, 1002, -468, 2235]
1485        );
1486
1487        let result = div(&a, &b).unwrap();
1488        assert_eq!(result.data_type(), &DataType::Decimal128(17, 7));
1489        assert_eq!(
1490            result.as_primitive::<Decimal128Type>().values(),
1491            &[27777, 0, 162078, 11133333, -1300000, 402]
1492        );
1493
1494        let result = rem(&a, &b).unwrap();
1495        assert_eq!(result.data_type(), &DataType::Decimal128(12, 3));
1496        assert_eq!(
1497            result.as_primitive::<Decimal128Type>().values(),
1498            &[15, 0, -577, 34, -78, 3]
1499        );
1500
1501        let a = Decimal128Array::from(vec![1])
1502            .with_precision_and_scale(3, 3)
1503            .unwrap();
1504        let b = Decimal128Array::from(vec![1])
1505            .with_precision_and_scale(37, 37)
1506            .unwrap();
1507        let err = mul(&a, &b).unwrap_err().to_string();
1508        assert_eq!(
1509            err,
1510            "Invalid argument error: Output scale of Decimal128(3, 3) * Decimal128(37, 37) would exceed max scale of 38"
1511        );
1512
1513        let a = Decimal128Array::from(vec![1])
1514            .with_precision_and_scale(3, -2)
1515            .unwrap();
1516        let err = add(&a, &b).unwrap_err().to_string();
1517        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 10 ^ 39");
1518
1519        let a = Decimal128Array::from(vec![10])
1520            .with_precision_and_scale(3, -1)
1521            .unwrap();
1522        let err = add(&a, &b).unwrap_err().to_string();
1523        assert_eq!(
1524            err,
1525            "Arithmetic overflow: Overflow happened on: 10 * 100000000000000000000000000000000000000"
1526        );
1527
1528        let b = Decimal128Array::from(vec![0])
1529            .with_precision_and_scale(1, 1)
1530            .unwrap();
1531        let err = div(&a, &b).unwrap_err().to_string();
1532        assert_eq!(err, "Divide by zero error");
1533        let err = rem(&a, &b).unwrap_err().to_string();
1534        assert_eq!(err, "Divide by zero error");
1535    }
1536
1537    #[test]
1538    fn test_decimal256_div_wide_intermediate() {
1539        // Dividing two scale-37 values needs l * 10^41, which is 79 digits and does not
1540        // fit in an i256, even though the 41-digit quotient does.
1541        let a = Decimal256Array::from(vec![i256::from_i128(
1542            60096743305738933273387748827369321010i128,
1543        )])
1544        .with_precision_and_scale(38, 37)
1545        .unwrap();
1546        let b = Decimal256Array::from(vec![i256::from_i128(
1547            60096763826458053191384497987259478584i128,
1548        )])
1549        .with_precision_and_scale(38, 37)
1550        .unwrap();
1551
1552        let result = div(&a, &b).unwrap();
1553        assert_eq!(result.data_type(), &DataType::Decimal256(76, 41));
1554        assert_eq!(
1555            result.as_primitive::<Decimal256Type>().value(0),
1556            i256::from_string("99999965853869970143724273117679321341339").unwrap()
1557        );
1558
1559        // Truncation stays toward zero on either side of the fallback.
1560        let neg_a = neg(&a).unwrap();
1561        let result = div(neg_a.as_primitive::<Decimal256Type>(), &b).unwrap();
1562        assert_eq!(
1563            result.as_primitive::<Decimal256Type>().value(0),
1564            i256::from_string("-99999965853869970143724273117679321341339").unwrap()
1565        );
1566
1567        let neg_b = neg(&b).unwrap();
1568        let result = div(&a, neg_b.as_primitive::<Decimal256Type>()).unwrap();
1569        assert_eq!(
1570            result.as_primitive::<Decimal256Type>().value(0),
1571            i256::from_string("-99999965853869970143724273117679321341339").unwrap()
1572        );
1573
1574        let zero = Decimal256Array::from(vec![i256::ZERO])
1575            .with_precision_and_scale(38, 37)
1576            .unwrap();
1577        let err = div(&a, &zero).unwrap_err().to_string();
1578        assert_eq!(err, "Divide by zero error");
1579    }
1580
1581    #[test]
1582    fn test_decimal256_div_divisor_near_max() {
1583        // A divisor past i256::MAX / 10 leaves no room to scale the running remainder either.
1584        let a = Decimal256Array::from(vec![
1585            i256::from_string(
1586                "5900000000000000000000000000000000000000000000000000000000000000000000000000",
1587            )
1588            .unwrap(),
1589        ])
1590        .with_precision_and_scale(76, 37)
1591        .unwrap();
1592        let b = Decimal256Array::from(vec![
1593            i256::from_string(
1594                "6000000000000000000000000000000000000000000000000000000000000000000000000000",
1595            )
1596            .unwrap(),
1597        ])
1598        .with_precision_and_scale(76, 37)
1599        .unwrap();
1600
1601        let result = div(&a, &b).unwrap();
1602        assert_eq!(
1603            result.as_primitive::<Decimal256Type>().value(0),
1604            i256::from_string("98333333333333333333333333333333333333333").unwrap()
1605        );
1606    }
1607
1608    #[test]
1609    fn test_decimal128_div_wide_intermediate() {
1610        // Same overflow one type down: 3.0 / 6.0 at scale 37 needs l * 10^38, 76 digits in i128.
1611        let a = Decimal128Array::from(vec![30000000000000000000000000000000000000i128])
1612            .with_precision_and_scale(38, 37)
1613            .unwrap();
1614        let b = Decimal128Array::from(vec![60000000000000000000000000000000000000i128])
1615            .with_precision_and_scale(38, 37)
1616            .unwrap();
1617
1618        let result = div(&a, &b).unwrap();
1619        assert_eq!(result.data_type(), &DataType::Decimal128(38, 38));
1620        assert_eq!(
1621            result.as_primitive::<Decimal128Type>().value(0),
1622            50000000000000000000000000000000000000i128
1623        );
1624    }
1625
1626    #[test]
1627    fn test_scaled_div_agrees_with_direct_division() {
1628        for (l, r, mul_pow) in [
1629            (7i128, 3i128, 4i8),
1630            (-7, 3, 4),
1631            (7, -3, 4),
1632            (-7, -3, 4),
1633            (1, 999_999_999, 9),
1634            (i128::MAX / 10, 7, 1),
1635            (i128::MAX / 10, -i128::MAX / 11, 1),
1636            (0, 5, 6),
1637        ] {
1638            let scaled = l * 10i128.pow(mul_pow as u32);
1639            assert_eq!(
1640                scaled_div::<Decimal128Type>(l, r, mul_pow).unwrap(),
1641                scaled / r,
1642                "{l} * 10^{mul_pow} / {r}"
1643            );
1644        }
1645    }
1646
1647    #[test]
1648    fn test_decimal256_same_scale_add_sub() {
1649        let lhs = Decimal256Array::from(vec![
1650            Some(i256::from_parts(u128::MAX, 0)),
1651            Some(i256::MINUS_ONE),
1652            None,
1653        ])
1654        .with_precision_and_scale(70, 2)
1655        .unwrap();
1656        let rhs = Decimal256Array::from(vec![Some(i256::ONE), Some(i256::ONE), Some(i256::MAX)])
1657            .with_precision_and_scale(70, 2)
1658            .unwrap();
1659
1660        let expected =
1661            Decimal256Array::from(vec![Some(i256::from_parts(0, 1)), Some(i256::ZERO), None])
1662                .with_precision_and_scale(71, 2)
1663                .unwrap();
1664        for operation in [add, add_wrapping] {
1665            let result = operation(&lhs, &rhs).unwrap();
1666            assert_eq!(result.as_primitive::<Decimal256Type>(), &expected);
1667        }
1668
1669        let expected = Decimal256Array::from(vec![
1670            Some(i256::from_parts(u128::MAX - 1, 0)),
1671            Some(i256::from_i128(-2)),
1672            None,
1673        ])
1674        .with_precision_and_scale(71, 2)
1675        .unwrap();
1676        for operation in [sub, sub_wrapping] {
1677            let result = operation(&lhs, &rhs).unwrap();
1678            assert_eq!(result.as_primitive::<Decimal256Type>(), &expected);
1679        }
1680
1681        let lhs = Decimal256Array::from(vec![i256::MAX])
1682            .with_precision_and_scale(76, 0)
1683            .unwrap();
1684        let rhs = Decimal256Array::from(vec![i256::ONE])
1685            .with_precision_and_scale(76, 0)
1686            .unwrap();
1687        for operation in [add, add_wrapping] {
1688            assert_eq!(
1689                operation(&lhs, &rhs).unwrap_err().to_string(),
1690                format!(
1691                    "Arithmetic overflow: Overflow happened on: {:?} + {:?}",
1692                    i256::MAX,
1693                    i256::ONE
1694                )
1695            );
1696        }
1697
1698        let lhs = Decimal256Array::from(vec![i256::MIN])
1699            .with_precision_and_scale(76, 0)
1700            .unwrap();
1701        for operation in [sub, sub_wrapping] {
1702            assert_eq!(
1703                operation(&lhs, &rhs).unwrap_err().to_string(),
1704                format!(
1705                    "Arithmetic overflow: Overflow happened on: {:?} - {:?}",
1706                    i256::MIN,
1707                    i256::ONE
1708                )
1709            );
1710        }
1711    }
1712
1713    fn test_timestamp_impl<T: TimestampOp>() {
1714        let a = PrimitiveArray::<T>::new(vec![2000000, 434030324, 53943340].into(), None);
1715        let b = PrimitiveArray::<T>::new(vec![329593, 59349, 694994].into(), None);
1716
1717        let result = sub(&a, &b).unwrap();
1718        assert_eq!(
1719            result.as_primitive::<T::Duration>().values(),
1720            &[1670407, 433970975, 53248346]
1721        );
1722
1723        let r2 = add(&b, &result.as_ref()).unwrap();
1724        assert_eq!(r2.as_ref(), &a);
1725
1726        let r3 = add(&result.as_ref(), &b).unwrap();
1727        assert_eq!(r3.as_ref(), &a);
1728
1729        let format_array = |x: &dyn Array| -> Vec<String> {
1730            x.as_primitive::<T>()
1731                .values()
1732                .into_iter()
1733                .map(|x| as_datetime::<T>(*x).unwrap().to_string())
1734                .collect()
1735        };
1736
1737        let values = vec![
1738            "1970-01-01T00:00:00Z",
1739            "2010-04-01T04:00:20Z",
1740            "1960-01-30T04:23:20Z",
1741        ]
1742        .into_iter()
1743        .map(|x| {
1744            T::from_naive_datetime(DateTime::parse_from_rfc3339(x).unwrap().naive_utc(), None)
1745                .unwrap()
1746        })
1747        .collect();
1748
1749        let a = PrimitiveArray::<T>::new(values, None);
1750        let b = IntervalYearMonthArray::from(vec![
1751            IntervalYearMonthType::make_value(5, 34),
1752            IntervalYearMonthType::make_value(-2, 4),
1753            IntervalYearMonthType::make_value(7, -4),
1754        ]);
1755        let r4 = add(&a, &b).unwrap();
1756        assert_eq!(
1757            &format_array(r4.as_ref()),
1758            &[
1759                "1977-11-01 00:00:00".to_string(),
1760                "2008-08-01 04:00:20".to_string(),
1761                "1966-09-30 04:23:20".to_string()
1762            ]
1763        );
1764
1765        let r5 = sub(&r4, &b).unwrap();
1766        assert_eq!(r5.as_ref(), &a);
1767
1768        let b = IntervalDayTimeArray::from(vec![
1769            IntervalDayTimeType::make_value(5, 454000),
1770            IntervalDayTimeType::make_value(-34, 0),
1771            IntervalDayTimeType::make_value(7, -4000),
1772        ]);
1773        let r6 = add(&a, &b).unwrap();
1774        assert_eq!(
1775            &format_array(r6.as_ref()),
1776            &[
1777                "1970-01-06 00:07:34".to_string(),
1778                "2010-02-26 04:00:20".to_string(),
1779                "1960-02-06 04:23:16".to_string()
1780            ]
1781        );
1782
1783        let r7 = sub(&r6, &b).unwrap();
1784        assert_eq!(r7.as_ref(), &a);
1785
1786        let b = IntervalMonthDayNanoArray::from(vec![
1787            IntervalMonthDayNanoType::make_value(344, 34, -43_000_000_000),
1788            IntervalMonthDayNanoType::make_value(-593, -33, 13_000_000_000),
1789            IntervalMonthDayNanoType::make_value(5, 2, 493_000_000_000),
1790        ]);
1791        let r8 = add(&a, &b).unwrap();
1792        assert_eq!(
1793            &format_array(r8.as_ref()),
1794            &[
1795                "1998-10-04 23:59:17".to_string(),
1796                "1960-09-29 04:00:33".to_string(),
1797                "1960-07-02 04:31:33".to_string()
1798            ]
1799        );
1800
1801        let r9 = sub(&r8, &b).unwrap();
1802        // Note: subtraction is not the inverse of addition for intervals
1803        assert_eq!(
1804            &format_array(r9.as_ref()),
1805            &[
1806                "1970-01-02 00:00:00".to_string(),
1807                "2010-04-02 04:00:20".to_string(),
1808                "1960-01-31 04:23:20".to_string()
1809            ]
1810        );
1811    }
1812
1813    #[test]
1814    fn test_timestamp() {
1815        test_timestamp_impl::<TimestampSecondType>();
1816        test_timestamp_impl::<TimestampMillisecondType>();
1817        test_timestamp_impl::<TimestampMicrosecondType>();
1818        test_timestamp_impl::<TimestampNanosecondType>();
1819    }
1820
1821    #[test]
1822    fn test_interval() {
1823        let a = IntervalYearMonthArray::from(vec![
1824            IntervalYearMonthType::make_value(32, 4),
1825            IntervalYearMonthType::make_value(32, 4),
1826        ]);
1827        let b = IntervalYearMonthArray::from(vec![
1828            IntervalYearMonthType::make_value(-4, 6),
1829            IntervalYearMonthType::make_value(-3, 23),
1830        ]);
1831        let result = add(&a, &b).unwrap();
1832        assert_eq!(
1833            result.as_ref(),
1834            &IntervalYearMonthArray::from(vec![
1835                IntervalYearMonthType::make_value(28, 10),
1836                IntervalYearMonthType::make_value(29, 27)
1837            ])
1838        );
1839        let result = sub(&a, &b).unwrap();
1840        assert_eq!(
1841            result.as_ref(),
1842            &IntervalYearMonthArray::from(vec![
1843                IntervalYearMonthType::make_value(36, -2),
1844                IntervalYearMonthType::make_value(35, -19)
1845            ])
1846        );
1847
1848        let a = IntervalDayTimeArray::from(vec![
1849            IntervalDayTimeType::make_value(32, 4),
1850            IntervalDayTimeType::make_value(32, 4),
1851        ]);
1852        let b = IntervalDayTimeArray::from(vec![
1853            IntervalDayTimeType::make_value(-4, 6),
1854            IntervalDayTimeType::make_value(-3, 23),
1855        ]);
1856        let result = add(&a, &b).unwrap();
1857        assert_eq!(
1858            result.as_ref(),
1859            &IntervalDayTimeArray::from(vec![
1860                IntervalDayTimeType::make_value(28, 10),
1861                IntervalDayTimeType::make_value(29, 27)
1862            ])
1863        );
1864        let result = sub(&a, &b).unwrap();
1865        assert_eq!(
1866            result.as_ref(),
1867            &IntervalDayTimeArray::from(vec![
1868                IntervalDayTimeType::make_value(36, -2),
1869                IntervalDayTimeType::make_value(35, -19)
1870            ])
1871        );
1872        let a = IntervalMonthDayNanoArray::from(vec![
1873            IntervalMonthDayNanoType::make_value(32, 4, 4000000000000),
1874            IntervalMonthDayNanoType::make_value(32, 4, 45463000000000000),
1875        ]);
1876        let b = IntervalMonthDayNanoArray::from(vec![
1877            IntervalMonthDayNanoType::make_value(-4, 6, 46000000000000),
1878            IntervalMonthDayNanoType::make_value(-3, 23, 3564000000000000),
1879        ]);
1880        let result = add(&a, &b).unwrap();
1881        assert_eq!(
1882            result.as_ref(),
1883            &IntervalMonthDayNanoArray::from(vec![
1884                IntervalMonthDayNanoType::make_value(28, 10, 50000000000000),
1885                IntervalMonthDayNanoType::make_value(29, 27, 49027000000000000)
1886            ])
1887        );
1888        let result = sub(&a, &b).unwrap();
1889        assert_eq!(
1890            result.as_ref(),
1891            &IntervalMonthDayNanoArray::from(vec![
1892                IntervalMonthDayNanoType::make_value(36, -2, -42000000000000),
1893                IntervalMonthDayNanoType::make_value(35, -19, 41899000000000000)
1894            ])
1895        );
1896        let a = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNano::MAX]);
1897        let b = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNano::ONE]);
1898        let err = add(&a, &b).unwrap_err().to_string();
1899        assert_eq!(
1900            err,
1901            "Arithmetic overflow: Overflow happened on: 2147483647 + 1"
1902        );
1903    }
1904
1905    #[test]
1906    fn test_interval_mul_i64() {
1907        let interval = IntervalYearMonthArray::from(vec![16, 5, 0]);
1908        let factor = Int64Array::from(vec![3, -2, i64::MAX]);
1909        let expected = IntervalYearMonthArray::from(vec![48, -10, 0]);
1910        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
1911        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
1912
1913        let interval = IntervalDayTimeArray::from(vec![
1914            Some(IntervalDayTimeType::make_value(10, 2 * 60 * 60 * 1000)),
1915            None,
1916        ]);
1917        let factor = Int64Array::new_scalar(3);
1918        let expected = IntervalDayTimeArray::from(vec![
1919            Some(IntervalDayTimeType::make_value(30, 6 * 60 * 60 * 1000)),
1920            None,
1921        ]);
1922        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
1923        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
1924
1925        let null_factor = Scalar::new(Int64Array::new_null(1));
1926        let expected = IntervalDayTimeArray::new_null(interval.len());
1927        assert_eq!(mul(&interval, &null_factor).unwrap().as_ref(), &expected);
1928        assert_eq!(mul(&null_factor, &interval).unwrap().as_ref(), &expected);
1929
1930        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
1931            12,
1932            15,
1933            5_000_000_000,
1934        ));
1935        let factor = Int64Array::from(vec![2, 0, -1]);
1936        let expected = IntervalMonthDayNanoArray::from(vec![
1937            IntervalMonthDayNanoType::make_value(24, 30, 10_000_000_000),
1938            IntervalMonthDayNanoType::make_value(0, 0, 0),
1939            IntervalMonthDayNanoType::make_value(-12, -15, -5_000_000_000),
1940        ]);
1941        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
1942        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
1943
1944        let float_factor = Float64Array::new_scalar(2.);
1945        assert!(mul_wrapping(&float_factor, &interval).is_err());
1946        assert!(mul_wrapping(&factor, &interval).is_err());
1947    }
1948
1949    #[test]
1950    fn test_interval_mul_div_f64() {
1951        const HOUR_NANOS: i64 = 3_600_000_000_000;
1952        const MINUTE_NANOS: i64 = 60_000_000_000;
1953
1954        // Adapted from DuckDB's interval multiplication tests:
1955        // https://github.com/duckdb/duckdb/blob/21aca0424f1faf78b593b1e6fbfdd4846624c987/test/sql/function/interval/test_interval_muldiv.test#L1-L99
1956        // DuckDB's cases come from PostgreSQL's interval regression tests:
1957        // https://github.com/postgres/postgres/blob/78758d37306cd89ab060f00cb06f249018d5b8da/src/test/regress/sql/interval.sql#L118-L164
1958        let interval = IntervalMonthDayNanoArray::from(vec![
1959            IntervalMonthDayNanoType::make_value(41, 12, 360 * HOUR_NANOS),
1960            IntervalMonthDayNanoType::make_value(-41, -12, 360 * HOUR_NANOS),
1961            IntervalMonthDayNanoType::make_value(1, 1, 0),
1962            IntervalMonthDayNanoType::make_value(0, 0, 1),
1963            IntervalMonthDayNanoType::make_value(0, 0, 3),
1964            IntervalMonthDayNanoType::make_value(0, 0, -1),
1965            IntervalMonthDayNanoType::make_value(0, 0, -3),
1966        ]);
1967        let factor = Float64Array::from(vec![0.3, 0.3, 1.5, 0.5, 0.5, 0.5, 0.5]);
1968        let expected = IntervalMonthDayNanoArray::from(vec![
1969            IntervalMonthDayNanoType::make_value(12, 12, 122 * HOUR_NANOS + 24 * MINUTE_NANOS),
1970            IntervalMonthDayNanoType::make_value(-12, -12, 93 * HOUR_NANOS + 36 * MINUTE_NANOS),
1971            IntervalMonthDayNanoType::make_value(1, 16, 12 * HOUR_NANOS),
1972            IntervalMonthDayNanoType::make_value(0, 0, 0),
1973            IntervalMonthDayNanoType::make_value(0, 0, 2),
1974            IntervalMonthDayNanoType::make_value(0, 0, 0),
1975            IntervalMonthDayNanoType::make_value(0, 0, -2),
1976        ]);
1977        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
1978        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
1979
1980        let interval = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
1981            9,
1982            -27,
1983            45_296 * NANOSECONDS,
1984        )]);
1985        let factor = Float64Array::new_scalar(0.3);
1986        let expected = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
1987            2,
1988            13,
1989            4_948_800_000_000,
1990        )]);
1991        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
1992
1993        let interval = IntervalMonthDayNanoArray::from(vec![
1994            IntervalMonthDayNanoType::make_value(0, 1, 0),
1995            IntervalMonthDayNanoType::make_value(4, 0, 0),
1996            IntervalMonthDayNanoType::make_value(1, 1, 0),
1997            IntervalMonthDayNanoType::make_value(0, 0, (1_i64 << 53) - 1),
1998            IntervalMonthDayNanoType::make_value(0, 0, i64::MAX),
1999            IntervalMonthDayNanoType::make_value(1, 1, 1),
2000            IntervalMonthDayNanoType::make_value(1, 1, 1),
2001            IntervalMonthDayNanoType::make_value(1, 0, 0),
2002        ]);
2003        let factor = Float64Array::from(vec![
2004            3.,
2005            5.,
2006            2.,
2007            0.7,
2008            1.,
2009            f64::INFINITY,
2010            f64::NEG_INFINITY,
2011            -2.,
2012        ]);
2013        let expected = IntervalMonthDayNanoArray::from(vec![
2014            IntervalMonthDayNanoType::make_value(0, 0, 8 * HOUR_NANOS),
2015            IntervalMonthDayNanoType::make_value(0, 24, 0),
2016            IntervalMonthDayNanoType::make_value(0, 15, 12 * HOUR_NANOS),
2017            IntervalMonthDayNanoType::make_value(0, 0, 12_867_427_506_772_844),
2018            IntervalMonthDayNanoType::make_value(0, 0, i64::MAX),
2019            IntervalMonthDayNanoType::make_value(0, 0, 0),
2020            IntervalMonthDayNanoType::make_value(0, 0, 0),
2021            IntervalMonthDayNanoType::make_value(0, -15, 0),
2022        ]);
2023        assert_eq!(div(&interval, &factor).unwrap().as_ref(), &expected);
2024
2025        let null_factor = Scalar::new(Float64Array::new_null(1));
2026        assert_eq!(
2027            mul(&interval, &null_factor).unwrap().as_ref(),
2028            &IntervalMonthDayNanoArray::new_null(interval.len())
2029        );
2030    }
2031
2032    #[test]
2033    fn test_interval_mul_div_f64_errors() {
2034        let factor = Float64Array::new_scalar(2.);
2035        let year_month = IntervalYearMonthArray::new_scalar(1);
2036        let day_time = IntervalDayTimeArray::new_scalar(IntervalDayTime::new(1, 1));
2037        for interval in [&year_month as &dyn Datum, &day_time] {
2038            assert!(mul(interval, &factor).is_err());
2039            assert!(mul(&factor, interval).is_err());
2040            assert!(div(interval, &factor).is_err());
2041        }
2042
2043        let interval =
2044            IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(1, 1, 1));
2045
2046        assert!(matches!(
2047            add(&interval, &factor),
2048            Err(ArrowError::InvalidArgumentError(_))
2049        ));
2050
2051        let zero = Float64Array::new_scalar(-0.);
2052        assert!(matches!(
2053            div(&interval, &zero),
2054            Err(ArrowError::DivideByZero)
2055        ));
2056
2057        assert!(div(&factor, &interval).is_err());
2058
2059        for factor in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
2060            let factor = Float64Array::new_scalar(factor);
2061            assert!(matches!(
2062                mul(&interval, &factor),
2063                Err(ArrowError::ArithmeticOverflow(_))
2064            ));
2065        }
2066
2067        let nan = Float64Array::new_scalar(f64::NAN);
2068        assert!(matches!(
2069            div(&interval, &nan),
2070            Err(ArrowError::ArithmeticOverflow(_))
2071        ));
2072
2073        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2074            i32::MAX,
2075            0,
2076            0,
2077        ));
2078        assert!(matches!(
2079            mul(&interval, &factor),
2080            Err(ArrowError::ArithmeticOverflow(_))
2081        ));
2082
2083        let factor = Float64Array::new_scalar(1.5);
2084        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2085            0,
2086            0,
2087            i64::MAX,
2088        ));
2089        assert!(matches!(
2090            mul(&interval, &factor),
2091            Err(ArrowError::ArithmeticOverflow(_))
2092        ));
2093
2094        let factor = Float64Array::new_scalar(1.000_000_000_4);
2095        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2096            i32::MIN,
2097            0,
2098            0,
2099        ));
2100        assert!(matches!(
2101            mul(&interval, &factor),
2102            Err(ArrowError::ArithmeticOverflow(_))
2103        ));
2104
2105        let factor = Float64Array::new_scalar(0.999_999_999);
2106        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
2107            1,
2108            i32::MAX,
2109            0,
2110        ));
2111        assert!(matches!(
2112            mul(&interval, &factor),
2113            Err(ArrowError::ArithmeticOverflow(_))
2114        ));
2115    }
2116
2117    #[test]
2118    fn test_interval_mul_i64_overflow() {
2119        let interval = IntervalYearMonthArray::from(vec![i32::MAX]);
2120        let factor = Int64Array::from(vec![2]);
2121        assert_eq!(
2122            mul(&interval, &factor).unwrap_err().to_string(),
2123            "Arithmetic overflow: Overflow happened on: 2147483647 * 2"
2124        );
2125
2126        let interval = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
2127            0,
2128            0,
2129            i64::MAX,
2130        )]);
2131        assert_eq!(
2132            mul(&interval, &factor).unwrap_err().to_string(),
2133            "Arithmetic overflow: Overflow happened on: 9223372036854775807 * 2"
2134        );
2135    }
2136
2137    fn test_duration_impl<T: ArrowPrimitiveType<Native = i64>>() {
2138        let a = PrimitiveArray::<T>::new(vec![1000, 4394, -3944].into(), None);
2139        let b = PrimitiveArray::<T>::new(vec![4, -5, -243].into(), None);
2140
2141        let result = add(&a, &b).unwrap();
2142        assert_eq!(result.as_primitive::<T>().values(), &[1004, 4389, -4187]);
2143        let result = sub(&a, &b).unwrap();
2144        assert_eq!(result.as_primitive::<T>().values(), &[996, 4399, -3701]);
2145
2146        let err = mul(&a, &b).unwrap_err().to_string();
2147        assert!(
2148            err.contains("Invalid duration arithmetic operation"),
2149            "{err}"
2150        );
2151
2152        let err = div(&a, &b).unwrap_err().to_string();
2153        assert!(
2154            err.contains("Invalid duration arithmetic operation"),
2155            "{err}"
2156        );
2157
2158        let err = rem(&a, &b).unwrap_err().to_string();
2159        assert!(
2160            err.contains("Invalid duration arithmetic operation"),
2161            "{err}"
2162        );
2163
2164        let a = PrimitiveArray::<T>::new(vec![i64::MAX].into(), None);
2165        let b = PrimitiveArray::<T>::new(vec![1].into(), None);
2166        let err = add(&a, &b).unwrap_err().to_string();
2167        assert_eq!(
2168            err,
2169            "Arithmetic overflow: Overflow happened on: 9223372036854775807 + 1"
2170        );
2171    }
2172
2173    #[test]
2174    fn test_duration() {
2175        test_duration_impl::<DurationSecondType>();
2176        test_duration_impl::<DurationMillisecondType>();
2177        test_duration_impl::<DurationMicrosecondType>();
2178        test_duration_impl::<DurationNanosecondType>();
2179    }
2180
2181    fn test_date_impl<T: ArrowPrimitiveType, F>(f: F)
2182    where
2183        F: Fn(NaiveDate) -> T::Native,
2184        T::Native: TryInto<i64>,
2185    {
2186        let a = PrimitiveArray::<T>::new(
2187            vec![
2188                f(NaiveDate::from_ymd_opt(1979, 1, 30).unwrap()),
2189                f(NaiveDate::from_ymd_opt(2010, 4, 3).unwrap()),
2190                f(NaiveDate::from_ymd_opt(2008, 2, 29).unwrap()),
2191            ]
2192            .into(),
2193            None,
2194        );
2195
2196        let b = IntervalYearMonthArray::from(vec![
2197            IntervalYearMonthType::make_value(34, 2),
2198            IntervalYearMonthType::make_value(3, -3),
2199            IntervalYearMonthType::make_value(-12, 4),
2200        ]);
2201
2202        let format_array = |x: &dyn Array| -> Vec<String> {
2203            x.as_primitive::<T>()
2204                .values()
2205                .into_iter()
2206                .map(|x| {
2207                    as_date::<T>((*x).try_into().ok().unwrap())
2208                        .unwrap()
2209                        .to_string()
2210                })
2211                .collect()
2212        };
2213
2214        let result = add(&a, &b).unwrap();
2215        assert_eq!(
2216            &format_array(result.as_ref()),
2217            &[
2218                "2013-03-30".to_string(),
2219                "2013-01-03".to_string(),
2220                "1996-06-29".to_string(),
2221            ]
2222        );
2223        let result = sub(&result, &b).unwrap();
2224        assert_eq!(result.as_ref(), &a);
2225
2226        let b = IntervalDayTimeArray::from(vec![
2227            IntervalDayTimeType::make_value(34, 2),
2228            IntervalDayTimeType::make_value(3, -3),
2229            IntervalDayTimeType::make_value(-12, 4),
2230        ]);
2231
2232        let result = add(&a, &b).unwrap();
2233        assert_eq!(
2234            &format_array(result.as_ref()),
2235            &[
2236                "1979-03-05".to_string(),
2237                "2010-04-06".to_string(),
2238                "2008-02-17".to_string(),
2239            ]
2240        );
2241        let result = sub(&result, &b).unwrap();
2242        assert_eq!(result.as_ref(), &a);
2243
2244        let b = IntervalMonthDayNanoArray::from(vec![
2245            IntervalMonthDayNanoType::make_value(34, 2, -34353534),
2246            IntervalMonthDayNanoType::make_value(3, -3, 2443),
2247            IntervalMonthDayNanoType::make_value(-12, 4, 2323242423232),
2248        ]);
2249
2250        let result = add(&a, &b).unwrap();
2251        assert_eq!(
2252            &format_array(result.as_ref()),
2253            &[
2254                "1981-12-02".to_string(),
2255                "2010-06-30".to_string(),
2256                "2007-03-04".to_string(),
2257            ]
2258        );
2259        let result = sub(&result, &b).unwrap();
2260        assert_eq!(
2261            &format_array(result.as_ref()),
2262            &[
2263                "1979-01-31".to_string(),
2264                "2010-04-02".to_string(),
2265                "2008-02-29".to_string(),
2266            ]
2267        );
2268    }
2269
2270    #[test]
2271    fn test_date() {
2272        test_date_impl::<Date32Type, _>(Date32Type::from_naive_date);
2273        test_date_impl::<Date64Type, _>(Date64Type::from_naive_date);
2274
2275        let a = Date32Array::from(vec![i32::MIN, i32::MAX, 23, 7684]);
2276        let b = Date32Array::from(vec![i32::MIN, i32::MIN, -2, 45]);
2277        let result = sub(&a, &b).unwrap();
2278        assert_eq!(
2279            result.as_primitive::<DurationSecondType>().values(),
2280            &[0, 371085174288000, 2160000, 660009600]
2281        );
2282
2283        let a = Date64Array::from(vec![4343, 76676, 3434]);
2284        let b = Date64Array::from(vec![3, -5, 5]);
2285        let result = sub(&a, &b).unwrap();
2286        assert_eq!(
2287            result.as_primitive::<DurationMillisecondType>().values(),
2288            &[4340, 76681, 3429]
2289        );
2290
2291        let a = Date64Array::from(vec![i64::MAX]);
2292        let b = Date64Array::from(vec![-1]);
2293        let err = sub(&a, &b).unwrap_err().to_string();
2294        assert_eq!(
2295            err,
2296            "Arithmetic overflow: Overflow happened on: 9223372036854775807 - -1"
2297        );
2298    }
2299
2300    #[test]
2301    fn test_date32_to_naive_date_opt_boundaries() {
2302        assert_eq!(MAX_VALID_DAYS, 95026236);
2303        assert_eq!(MIN_VALID_DAYS, -96465292);
2304
2305        // Valid boundary dates work
2306        assert!(Date32Type::to_naive_date_opt(MAX_VALID_DAYS).is_some());
2307        assert!(Date32Type::to_naive_date_opt(MIN_VALID_DAYS).is_some());
2308
2309        // Beyond boundaries fail
2310        assert!(Date32Type::to_naive_date_opt(MAX_VALID_DAYS + 1).is_none());
2311        assert!(Date32Type::to_naive_date_opt(MIN_VALID_DAYS - 1).is_none());
2312
2313        // Extreme values fail
2314        assert!(Date32Type::to_naive_date_opt(i32::MAX).is_none());
2315        assert!(Date32Type::to_naive_date_opt(i32::MIN).is_none());
2316
2317        // Common values work
2318        assert!(Date32Type::to_naive_date_opt(0).is_some());
2319        assert!(Date32Type::to_naive_date_opt(date_to_days(YEAR_2000)).is_some());
2320    }
2321
2322    #[test]
2323    fn test_date64_to_naive_date_opt_boundaries() {
2324        const MS_PER_DAY: i64 = 24 * 60 * 60 * 1000;
2325
2326        // Verify boundary millisecond values
2327        assert_eq!(MAX_VALID_MILLIS, 8210266790400000i64);
2328        assert_eq!(MIN_VALID_MILLIS, -8334601228800000i64);
2329
2330        // Valid boundary dates work
2331        assert!(Date64Type::to_naive_date_opt(MAX_VALID_MILLIS).is_some());
2332        assert!(Date64Type::to_naive_date_opt(MIN_VALID_MILLIS).is_some());
2333
2334        // Beyond boundaries fail
2335        assert!(Date64Type::to_naive_date_opt(MAX_VALID_MILLIS + MS_PER_DAY).is_none());
2336        assert!(Date64Type::to_naive_date_opt(MIN_VALID_MILLIS - MS_PER_DAY).is_none());
2337
2338        // Extreme values fail
2339        assert!(Date64Type::to_naive_date_opt(i64::MAX).is_none());
2340        assert!(Date64Type::to_naive_date_opt(i64::MIN).is_none());
2341
2342        // Common values work
2343        assert!(Date64Type::to_naive_date_opt(0).is_some());
2344        assert!(Date64Type::to_naive_date_opt(date_to_millis(YEAR_2000)).is_some());
2345    }
2346
2347    macro_rules! test_year_month_ops {
2348        ($type:ty, $date_fn:expr) => {{
2349            let date = $date_fn(YEAR_2000);
2350
2351            // Normal operations succeed
2352            assert!(
2353                <$type>::add_year_months_opt(date, 120).is_some(),
2354                "add_year_months: normal add"
2355            );
2356            assert!(
2357                <$type>::add_year_months_opt(date, 0).is_some(),
2358                "add_year_months: zero interval"
2359            );
2360            assert!(
2361                <$type>::subtract_year_months_opt(date, 120).is_some(),
2362                "subtract_year_months: normal subtract"
2363            );
2364            assert!(
2365                <$type>::subtract_year_months_opt(date, 0).is_some(),
2366                "subtract_year_months: zero interval"
2367            );
2368
2369            // Large but valid years work
2370            let large_year = $date_fn(NaiveDate::from_ymd_opt(5000, 1, 1).unwrap());
2371            let neg_year = $date_fn(NaiveDate::from_ymd_opt(-5000, 12, 31).unwrap());
2372            assert!(
2373                <$type>::add_year_months_opt(large_year, 12).is_some(),
2374                "add_year_months: large year"
2375            );
2376            assert!(
2377                <$type>::add_year_months_opt(neg_year, -12).is_some(),
2378                "add_year_months: negative year"
2379            );
2380            assert!(
2381                <$type>::subtract_year_months_opt(large_year, 12).is_some(),
2382                "subtract_year_months: large year"
2383            );
2384            assert!(
2385                <$type>::subtract_year_months_opt(neg_year, -12).is_some(),
2386                "subtract_year_months: negative year"
2387            );
2388
2389            // Overflow handling
2390            assert!(
2391                <$type>::subtract_year_months_opt($date_fn(MIN_VALID_DATE), 1).is_none(),
2392                "subtract_year_months: overflow days from min"
2393            );
2394            assert!(
2395                <$type>::subtract_year_months_opt($date_fn(MAX_VALID_DATE), -1).is_none(),
2396                "subtract_year_months: overflow neg days from max"
2397            );
2398            assert!(
2399                <$type>::add_year_months_opt($date_fn(MAX_VALID_DATE), 1).is_none(),
2400                "add_year_months: overflow days"
2401            );
2402            assert!(
2403                <$type>::add_year_months_opt($date_fn(MIN_VALID_DATE), -1).is_none(),
2404                "add_year_months: overflow neg days"
2405            );
2406        }};
2407    }
2408
2409    #[test]
2410    fn test_date_year_month_operations() {
2411        test_year_month_ops!(Date32Type, date_to_days);
2412        test_year_month_ops!(Date64Type, date_to_millis);
2413    }
2414
2415    macro_rules! test_day_time_ops {
2416        ($type:ty, $date_fn:expr) => {{
2417            let date = $date_fn(YEAR_2000);
2418
2419            // Moderate intervals succeed
2420            assert!(
2421                <$type>::add_day_time_opt(date, IntervalDayTime::new(30, 0)).is_some(),
2422                "add_day_time: +30 days"
2423            );
2424            assert!(
2425                <$type>::add_day_time_opt(date, IntervalDayTime::new(-30, 0)).is_some(),
2426                "add_day_time: -30 days"
2427            );
2428            assert!(
2429                <$type>::add_day_time_opt(date, IntervalDayTime::new(1000, 12345)).is_some(),
2430                "add_day_time: normal"
2431            );
2432            assert!(
2433                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(30, 0)).is_some(),
2434                "subtract_day_time: +30 days"
2435            );
2436            assert!(
2437                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(-30, 0)).is_some(),
2438                "subtract_day_time: -30 days"
2439            );
2440            assert!(
2441                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(1000, 12345)).is_some(),
2442                "subtract_day_time: normal"
2443            );
2444
2445            // Overflow handling - subtract
2446            assert!(
2447                <$type>::subtract_day_time_opt(
2448                    $date_fn(MIN_VALID_DATE),
2449                    IntervalDayTime::new(1, 0)
2450                )
2451                .is_none(),
2452                "subtract_day_time: overflow days from min"
2453            );
2454            assert!(
2455                <$type>::subtract_day_time_opt(
2456                    $date_fn(MAX_VALID_DATE),
2457                    IntervalDayTime::new(-1, 0)
2458                )
2459                .is_none(),
2460                "subtract_day_time: overflow neg days from max"
2461            );
2462
2463            // Overflow handling - add
2464            assert!(
2465                <$type>::add_day_time_opt($date_fn(MAX_VALID_DATE), IntervalDayTime::new(1, 0))
2466                    .is_none(),
2467                "add_day_time: overflow days"
2468            );
2469            assert!(
2470                <$type>::add_day_time_opt($date_fn(MIN_VALID_DATE), IntervalDayTime::new(-1, 0))
2471                    .is_none(),
2472                "add_day_time: overflow neg days"
2473            );
2474
2475            // Extreme intervals fail
2476            assert!(
2477                <$type>::add_day_time_opt(
2478                    $date_fn(EPOCH),
2479                    IntervalDayTime::new(i32::MAX, i32::MAX)
2480                )
2481                .is_none(),
2482                "add_day_time: max interval"
2483            );
2484            assert!(
2485                <$type>::add_day_time_opt(
2486                    $date_fn(EPOCH),
2487                    IntervalDayTime::new(i32::MIN, i32::MIN)
2488                )
2489                .is_none(),
2490                "add_day_time: min interval"
2491            );
2492            assert!(
2493                <$type>::subtract_day_time_opt(
2494                    $date_fn(EPOCH),
2495                    IntervalDayTime::new(i32::MAX, i32::MAX)
2496                )
2497                .is_none(),
2498                "subtract_day_time: max interval"
2499            );
2500            assert!(
2501                <$type>::subtract_day_time_opt(
2502                    $date_fn(EPOCH),
2503                    IntervalDayTime::new(i32::MIN, i32::MIN)
2504                )
2505                .is_none(),
2506                "subtract_day_time: min interval"
2507            );
2508        }};
2509    }
2510
2511    #[test]
2512    fn test_date_day_time_operations() {
2513        test_day_time_ops!(Date32Type, date_to_days);
2514        test_day_time_ops!(Date64Type, date_to_millis);
2515    }
2516
2517    macro_rules! test_month_day_nano_ops {
2518        ($type:ty, $date_fn:expr) => {{
2519            let date = $date_fn(YEAR_2000);
2520            let zero = IntervalMonthDayNano::new(0, 0, 0);
2521
2522            // Normal operations succeed
2523            assert!(
2524                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(1, 30, 0))
2525                    .is_some(),
2526                "add_month_day_nano: +1mo +30d"
2527            );
2528            assert!(
2529                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(-1, -30, 0))
2530                    .is_some(),
2531                "add_month_day_nano: -1mo -30d"
2532            );
2533            assert!(
2534                <$type>::add_month_day_nano_opt(date, zero).is_some(),
2535                "add_month_day_nano: zero interval"
2536            );
2537            assert!(
2538                <$type>::add_month_day_nano_opt(
2539                    date,
2540                    IntervalMonthDayNano::new(2, 10, 123_456_789_000)
2541                )
2542                .is_some(),
2543                "add_month_day_nano: normal"
2544            );
2545            assert!(
2546                <$type>::subtract_month_day_nano_opt(date, IntervalMonthDayNano::new(1, 30, 0))
2547                    .is_some(),
2548                "subtract_month_day_nano: +1mo +30d"
2549            );
2550            assert!(
2551                <$type>::subtract_month_day_nano_opt(date, IntervalMonthDayNano::new(-1, -30, 0))
2552                    .is_some(),
2553                "subtract_month_day_nano: -1mo -30d"
2554            );
2555            assert!(
2556                <$type>::subtract_month_day_nano_opt(date, zero).is_some(),
2557                "subtract_month_day_nano: zero interval"
2558            );
2559            assert!(
2560                <$type>::subtract_month_day_nano_opt(
2561                    date,
2562                    IntervalMonthDayNano::new(2, 10, 123_456_789_000)
2563                )
2564                .is_some(),
2565                "subtract_month_day_nano: normal"
2566            );
2567
2568            // Overflow handling - subtract
2569            assert!(
2570                <$type>::subtract_month_day_nano_opt(
2571                    $date_fn(MIN_VALID_DATE),
2572                    IntervalMonthDayNano::new(0, 1, 0)
2573                )
2574                .is_none(),
2575                "subtract_month_day_nano: overflow days from min"
2576            );
2577            assert!(
2578                <$type>::subtract_month_day_nano_opt(
2579                    $date_fn(MAX_VALID_DATE),
2580                    IntervalMonthDayNano::new(0, -1, 0)
2581                )
2582                .is_none(),
2583                "subtract_month_day_nano: overflow neg days from max"
2584            );
2585
2586            // Overflow handling - add
2587            assert!(
2588                <$type>::add_month_day_nano_opt(
2589                    $date_fn(MAX_VALID_DATE),
2590                    IntervalMonthDayNano::new(0, 1, 0)
2591                )
2592                .is_none(),
2593                "add_month_day_nano: overflow days"
2594            );
2595            assert!(
2596                <$type>::add_month_day_nano_opt(
2597                    $date_fn(MIN_VALID_DATE),
2598                    IntervalMonthDayNano::new(0, -1, 0)
2599                )
2600                .is_none(),
2601                "add_month_day_nano: overflow neg days"
2602            );
2603
2604            // Nanosecond precision works
2605            assert!(
2606                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(0, 0, 999_999_999))
2607                    .is_some(),
2608                "add_month_day_nano: nanos"
2609            );
2610            assert!(
2611                <$type>::subtract_month_day_nano_opt(
2612                    date,
2613                    IntervalMonthDayNano::new(0, 0, 999_999_999)
2614                )
2615                .is_some(),
2616                "subtract_month_day_nano: nanos"
2617            );
2618            // 1 day in nanos
2619            assert!(
2620                <$type>::add_month_day_nano_opt(
2621                    date,
2622                    IntervalMonthDayNano::new(0, 0, 86_400_000_000_000)
2623                )
2624                .is_some(),
2625                "add_month_day_nano: 1 day nanos"
2626            );
2627            assert!(
2628                <$type>::subtract_month_day_nano_opt(
2629                    date,
2630                    IntervalMonthDayNano::new(0, 0, 86_400_000_000_000)
2631                )
2632                .is_some(),
2633                "subtract_month_day_nano: 1 day nanos"
2634            );
2635        }};
2636    }
2637
2638    #[test]
2639    fn test_date_month_day_nano_operations() {
2640        test_month_day_nano_ops!(Date32Type, date_to_days);
2641        test_month_day_nano_ops!(Date64Type, date_to_millis);
2642    }
2643}