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::timezone::Tz;
26use arrow_array::types::*;
27use arrow_array::*;
28use arrow_buffer::{ArrowNativeType, IntervalDayTime, IntervalMonthDayNano};
29use arrow_schema::{ArrowError, DataType, IntervalUnit, TimeUnit};
30
31use crate::arity::{binary, try_binary};
32
33/// Perform `lhs + rhs`, returning an error on overflow
34pub fn add(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
35    arithmetic_op(Op::Add, lhs, rhs)
36}
37
38/// Perform `lhs + rhs`, wrapping on overflow for [`DataType::is_integer`]
39pub fn add_wrapping(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
40    arithmetic_op(Op::AddWrapping, lhs, rhs)
41}
42
43/// Perform `lhs - rhs`, returning an error on overflow
44pub fn sub(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
45    arithmetic_op(Op::Sub, lhs, rhs)
46}
47
48/// Perform `lhs - rhs`, wrapping on overflow for [`DataType::is_integer`]
49pub fn sub_wrapping(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
50    arithmetic_op(Op::SubWrapping, lhs, rhs)
51}
52
53/// Perform `lhs * rhs`, returning an error on overflow
54pub fn mul(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
55    arithmetic_op(Op::Mul, lhs, rhs)
56}
57
58/// Perform `lhs * rhs`, wrapping on overflow for [`DataType::is_integer`]
59pub fn mul_wrapping(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
60    arithmetic_op(Op::MulWrapping, lhs, rhs)
61}
62
63/// Perform `lhs / rhs`
64///
65/// Overflow or division by zero will result in an error, with exception to
66/// floating point numbers, which instead follow the IEEE 754 rules
67pub fn div(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
68    arithmetic_op(Op::Div, lhs, rhs)
69}
70
71/// Perform `lhs % rhs`
72///
73/// Division by zero will result in an error, with exception to
74/// floating point numbers, which instead follow the IEEE 754 rules
75///
76/// `signed_integer::MIN % -1` will not result in an error but return 0
77pub fn rem(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
78    arithmetic_op(Op::Rem, lhs, rhs)
79}
80
81macro_rules! neg_checked {
82    ($t:ty, $a:ident) => {{
83        let array = $a
84            .as_primitive::<$t>()
85            .try_unary::<_, $t, _>(|x| x.neg_checked())?;
86        Ok(Arc::new(array))
87    }};
88}
89
90macro_rules! neg_wrapping {
91    ($t:ty, $a:ident) => {{
92        let array = $a.as_primitive::<$t>().unary::<_, $t>(|x| x.neg_wrapping());
93        Ok(Arc::new(array))
94    }};
95}
96
97/// Negates each element of  `array`, returning an error on overflow
98///
99/// Note: negation of unsigned arrays is not supported and will return in an error,
100/// for wrapping unsigned negation consider using [`neg_wrapping`][neg_wrapping()]
101pub fn neg(array: &dyn Array) -> Result<ArrayRef, ArrowError> {
102    use DataType::*;
103    use IntervalUnit::*;
104    use TimeUnit::*;
105
106    match array.data_type() {
107        Int8 => neg_checked!(Int8Type, array),
108        Int16 => neg_checked!(Int16Type, array),
109        Int32 => neg_checked!(Int32Type, array),
110        Int64 => neg_checked!(Int64Type, array),
111        Float16 => neg_wrapping!(Float16Type, array),
112        Float32 => neg_wrapping!(Float32Type, array),
113        Float64 => neg_wrapping!(Float64Type, array),
114        Decimal32(p, s) => {
115            let a = array
116                .as_primitive::<Decimal32Type>()
117                .try_unary::<_, Decimal32Type, _>(|x| x.neg_checked())?;
118
119            Ok(Arc::new(a.with_precision_and_scale(*p, *s)?))
120        }
121        Decimal64(p, s) => {
122            let a = array
123                .as_primitive::<Decimal64Type>()
124                .try_unary::<_, Decimal64Type, _>(|x| x.neg_checked())?;
125
126            Ok(Arc::new(a.with_precision_and_scale(*p, *s)?))
127        }
128        Decimal128(p, s) => {
129            let a = array
130                .as_primitive::<Decimal128Type>()
131                .try_unary::<_, Decimal128Type, _>(|x| x.neg_checked())?;
132
133            Ok(Arc::new(a.with_precision_and_scale(*p, *s)?))
134        }
135        Decimal256(p, s) => {
136            let a = array
137                .as_primitive::<Decimal256Type>()
138                .try_unary::<_, Decimal256Type, _>(|x| x.neg_checked())?;
139
140            Ok(Arc::new(a.with_precision_and_scale(*p, *s)?))
141        }
142        Duration(Second) => neg_checked!(DurationSecondType, array),
143        Duration(Millisecond) => neg_checked!(DurationMillisecondType, array),
144        Duration(Microsecond) => neg_checked!(DurationMicrosecondType, array),
145        Duration(Nanosecond) => neg_checked!(DurationNanosecondType, array),
146        Interval(YearMonth) => neg_checked!(IntervalYearMonthType, array),
147        Interval(DayTime) => {
148            let a = array
149                .as_primitive::<IntervalDayTimeType>()
150                .try_unary::<_, IntervalDayTimeType, ArrowError>(|x| {
151                    let (days, ms) = IntervalDayTimeType::to_parts(x);
152                    Ok(IntervalDayTimeType::make_value(
153                        days.neg_checked()?,
154                        ms.neg_checked()?,
155                    ))
156                })?;
157            Ok(Arc::new(a))
158        }
159        Interval(MonthDayNano) => {
160            let a = array
161                .as_primitive::<IntervalMonthDayNanoType>()
162                .try_unary::<_, IntervalMonthDayNanoType, ArrowError>(|x| {
163                    let (months, days, nanos) = IntervalMonthDayNanoType::to_parts(x);
164                    Ok(IntervalMonthDayNanoType::make_value(
165                        months.neg_checked()?,
166                        days.neg_checked()?,
167                        nanos.neg_checked()?,
168                    ))
169                })?;
170            Ok(Arc::new(a))
171        }
172        t => Err(ArrowError::InvalidArgumentError(format!(
173            "Invalid arithmetic operation: !{t}"
174        ))),
175    }
176}
177
178/// Negates each element of  `array`, wrapping on overflow for [`DataType::is_integer`]
179pub fn neg_wrapping(array: &dyn Array) -> Result<ArrayRef, ArrowError> {
180    downcast_integer! {
181        array.data_type() => (neg_wrapping, array),
182        _ => neg(array),
183    }
184}
185
186/// An enumeration of arithmetic operations
187///
188/// This allows sharing the type dispatch logic across the various kernels
189#[derive(Debug, Copy, Clone)]
190enum Op {
191    AddWrapping,
192    Add,
193    SubWrapping,
194    Sub,
195    MulWrapping,
196    Mul,
197    Div,
198    Rem,
199}
200
201impl std::fmt::Display for Op {
202    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
203        match self {
204            Op::AddWrapping | Op::Add => write!(f, "+"),
205            Op::SubWrapping | Op::Sub => write!(f, "-"),
206            Op::MulWrapping | Op::Mul => write!(f, "*"),
207            Op::Div => write!(f, "/"),
208            Op::Rem => write!(f, "%"),
209        }
210    }
211}
212
213impl Op {
214    fn commutative(&self) -> bool {
215        matches!(self, Self::Add | Self::AddWrapping)
216    }
217}
218
219/// Dispatch the given `op` to the appropriate specialized kernel
220fn arithmetic_op(op: Op, lhs: &dyn Datum, rhs: &dyn Datum) -> Result<ArrayRef, ArrowError> {
221    use DataType::*;
222    use IntervalUnit::*;
223    use TimeUnit::*;
224
225    macro_rules! integer_helper {
226        ($t:ty, $op:ident, $l:ident, $l_scalar:ident, $r:ident, $r_scalar:ident) => {
227            integer_op::<$t>($op, $l, $l_scalar, $r, $r_scalar)
228        };
229    }
230
231    let (l, l_scalar) = lhs.get();
232    let (r, r_scalar) = rhs.get();
233    downcast_integer! {
234        l.data_type(), r.data_type() => (integer_helper, op, l, l_scalar, r, r_scalar),
235        (Float16, Float16) => float_op::<Float16Type>(op, l, l_scalar, r, r_scalar),
236        (Float32, Float32) => float_op::<Float32Type>(op, l, l_scalar, r, r_scalar),
237        (Float64, Float64) => float_op::<Float64Type>(op, l, l_scalar, r, r_scalar),
238        (Timestamp(Second, _), _) => timestamp_op::<TimestampSecondType>(op, l, l_scalar, r, r_scalar),
239        (Timestamp(Millisecond, _), _) => timestamp_op::<TimestampMillisecondType>(op, l, l_scalar, r, r_scalar),
240        (Timestamp(Microsecond, _), _) => timestamp_op::<TimestampMicrosecondType>(op, l, l_scalar, r, r_scalar),
241        (Timestamp(Nanosecond, _), _) => timestamp_op::<TimestampNanosecondType>(op, l, l_scalar, r, r_scalar),
242        (Duration(Second), Duration(Second)) => duration_op::<DurationSecondType>(op, l, l_scalar, r, r_scalar),
243        (Duration(Millisecond), Duration(Millisecond)) => duration_op::<DurationMillisecondType>(op, l, l_scalar, r, r_scalar),
244        (Duration(Microsecond), Duration(Microsecond)) => duration_op::<DurationMicrosecondType>(op, l, l_scalar, r, r_scalar),
245        (Duration(Nanosecond), Duration(Nanosecond)) => duration_op::<DurationNanosecondType>(op, l, l_scalar, r, r_scalar),
246        (Interval(YearMonth), Int64) if matches!(op, Op::Mul) => interval_mul_op::<IntervalYearMonthType>(l, l_scalar, r, r_scalar),
247        (Interval(DayTime), Int64) if matches!(op, Op::Mul) => interval_mul_op::<IntervalDayTimeType>(l, l_scalar, r, r_scalar),
248        (Interval(MonthDayNano), Int64) if matches!(op, Op::Mul) => interval_mul_op::<IntervalMonthDayNanoType>(l, l_scalar, r, r_scalar),
249        (Int64, Interval(YearMonth)) if matches!(op, Op::Mul) => interval_mul_op::<IntervalYearMonthType>(r, r_scalar, l, l_scalar),
250        (Int64, Interval(DayTime)) if matches!(op, Op::Mul) => interval_mul_op::<IntervalDayTimeType>(r, r_scalar, l, l_scalar),
251        (Int64, Interval(MonthDayNano)) if matches!(op, Op::Mul) => interval_mul_op::<IntervalMonthDayNanoType>(r, r_scalar, l, l_scalar),
252        (Interval(YearMonth), Interval(YearMonth)) => interval_op::<IntervalYearMonthType>(op, l, l_scalar, r, r_scalar),
253        (Interval(DayTime), Interval(DayTime)) => interval_op::<IntervalDayTimeType>(op, l, l_scalar, r, r_scalar),
254        (Interval(MonthDayNano), Interval(MonthDayNano)) => interval_op::<IntervalMonthDayNanoType>(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            _ => Err(ArrowError::InvalidArgumentError(
266              format!("Invalid arithmetic operation: {l_t} {op} {r_t}")
267            ))
268        }
269    }
270}
271
272/// Perform an infallible binary operation on potentially scalar inputs
273macro_rules! op {
274    ($l:ident, $l_s:expr, $r:ident, $r_s:expr, $op:expr) => {
275        match ($l_s, $r_s) {
276            (true, true) | (false, false) => binary($l, $r, |$l, $r| $op)?,
277            (true, false) => match ($l.null_count() == 0).then(|| $l.value(0)) {
278                None => PrimitiveArray::new_null($r.len()),
279                Some($l) => $r.unary(|$r| $op),
280            },
281            (false, true) => match ($r.null_count() == 0).then(|| $r.value(0)) {
282                None => PrimitiveArray::new_null($l.len()),
283                Some($r) => $l.unary(|$l| $op),
284            },
285        }
286    };
287}
288
289/// Same as `op` but with a type hint for the returned array
290macro_rules! op_ref {
291    ($t:ty, $l:ident, $l_s:expr, $r:ident, $r_s:expr, $op:expr) => {{
292        let array: PrimitiveArray<$t> = op!($l, $l_s, $r, $r_s, $op);
293        Arc::new(array)
294    }};
295}
296
297/// Perform a fallible binary operation on potentially scalar inputs
298macro_rules! try_op {
299    ($l:ident, $l_s:expr, $r:ident, $r_s:expr, $op:expr) => {
300        match ($l_s, $r_s) {
301            (true, true) | (false, false) => try_binary($l, $r, |$l, $r| $op)?,
302            (true, false) => match ($l.null_count() == 0).then(|| $l.value(0)) {
303                None => PrimitiveArray::new_null($r.len()),
304                Some($l) => $r.try_unary(|$r| $op)?,
305            },
306            (false, true) => match ($r.null_count() == 0).then(|| $r.value(0)) {
307                None => PrimitiveArray::new_null($l.len()),
308                Some($r) => $l.try_unary(|$l| $op)?,
309            },
310        }
311    };
312}
313
314/// Same as `try_op` but with a type hint for the returned array
315macro_rules! try_op_ref {
316    ($t:ty, $l:ident, $l_s:expr, $r:ident, $r_s:expr, $op:expr) => {{
317        let array: PrimitiveArray<$t> = try_op!($l, $l_s, $r, $r_s, $op);
318        Arc::new(array)
319    }};
320}
321
322/// Perform an arithmetic operation on integers
323fn integer_op<T: ArrowPrimitiveType>(
324    op: Op,
325    l: &dyn Array,
326    l_s: bool,
327    r: &dyn Array,
328    r_s: bool,
329) -> Result<ArrayRef, ArrowError> {
330    let l = l.as_primitive::<T>();
331    let r = r.as_primitive::<T>();
332    let array: PrimitiveArray<T> = match op {
333        Op::AddWrapping => op!(l, l_s, r, r_s, l.add_wrapping(r)),
334        Op::Add => try_op!(l, l_s, r, r_s, l.add_checked(r)),
335        Op::SubWrapping => op!(l, l_s, r, r_s, l.sub_wrapping(r)),
336        Op::Sub => try_op!(l, l_s, r, r_s, l.sub_checked(r)),
337        Op::MulWrapping => op!(l, l_s, r, r_s, l.mul_wrapping(r)),
338        Op::Mul => try_op!(l, l_s, r, r_s, l.mul_checked(r)),
339        Op::Div => try_op!(l, l_s, r, r_s, l.div_checked(r)),
340        Op::Rem => try_op!(l, l_s, r, r_s, {
341            if r.is_zero() {
342                Err(ArrowError::DivideByZero)
343            } else {
344                Ok(l.mod_wrapping(r))
345            }
346        }),
347    };
348    Ok(Arc::new(array))
349}
350
351/// Perform an arithmetic operation on floats
352fn float_op<T: ArrowPrimitiveType>(
353    op: Op,
354    l: &dyn Array,
355    l_s: bool,
356    r: &dyn Array,
357    r_s: bool,
358) -> Result<ArrayRef, ArrowError> {
359    let l = l.as_primitive::<T>();
360    let r = r.as_primitive::<T>();
361    let array: PrimitiveArray<T> = match op {
362        Op::AddWrapping | Op::Add => op!(l, l_s, r, r_s, l.add_wrapping(r)),
363        Op::SubWrapping | Op::Sub => op!(l, l_s, r, r_s, l.sub_wrapping(r)),
364        Op::MulWrapping | Op::Mul => op!(l, l_s, r, r_s, l.mul_wrapping(r)),
365        Op::Div => op!(l, l_s, r, r_s, l.div_wrapping(r)),
366        Op::Rem => op!(l, l_s, r, r_s, l.mod_wrapping(r)),
367    };
368    Ok(Arc::new(array))
369}
370
371/// Arithmetic trait for timestamp arrays
372trait TimestampOp: ArrowTimestampType {
373    type Duration: ArrowPrimitiveType<Native = i64>;
374
375    fn add_year_month(timestamp: i64, delta: i32, tz: Tz) -> Option<i64>;
376    fn add_day_time(timestamp: i64, delta: IntervalDayTime, tz: Tz) -> Option<i64>;
377    fn add_month_day_nano(timestamp: i64, delta: IntervalMonthDayNano, tz: Tz) -> Option<i64>;
378
379    fn sub_year_month(timestamp: i64, delta: i32, tz: Tz) -> Option<i64>;
380    fn sub_day_time(timestamp: i64, delta: IntervalDayTime, tz: Tz) -> Option<i64>;
381    fn sub_month_day_nano(timestamp: i64, delta: IntervalMonthDayNano, tz: Tz) -> Option<i64>;
382}
383
384macro_rules! timestamp {
385    ($t:ty, $d:ty) => {
386        impl TimestampOp for $t {
387            type Duration = $d;
388
389            fn add_year_month(left: i64, right: i32, tz: Tz) -> Option<i64> {
390                Self::add_year_months(left, right, tz)
391            }
392
393            fn add_day_time(left: i64, right: IntervalDayTime, tz: Tz) -> Option<i64> {
394                Self::add_day_time(left, right, tz)
395            }
396
397            fn add_month_day_nano(left: i64, right: IntervalMonthDayNano, tz: Tz) -> Option<i64> {
398                Self::add_month_day_nano(left, right, tz)
399            }
400
401            fn sub_year_month(left: i64, right: i32, tz: Tz) -> Option<i64> {
402                Self::subtract_year_months(left, right, tz)
403            }
404
405            fn sub_day_time(left: i64, right: IntervalDayTime, tz: Tz) -> Option<i64> {
406                Self::subtract_day_time(left, right, tz)
407            }
408
409            fn sub_month_day_nano(left: i64, right: IntervalMonthDayNano, tz: Tz) -> Option<i64> {
410                Self::subtract_month_day_nano(left, right, tz)
411            }
412        }
413    };
414}
415timestamp!(TimestampSecondType, DurationSecondType);
416timestamp!(TimestampMillisecondType, DurationMillisecondType);
417timestamp!(TimestampMicrosecondType, DurationMicrosecondType);
418timestamp!(TimestampNanosecondType, DurationNanosecondType);
419
420/// Perform arithmetic operation on a timestamp array
421fn timestamp_op<T: TimestampOp>(
422    op: Op,
423    l: &dyn Array,
424    l_s: bool,
425    r: &dyn Array,
426    r_s: bool,
427) -> Result<ArrayRef, ArrowError> {
428    use DataType::*;
429    use IntervalUnit::*;
430
431    let l = l.as_primitive::<T>();
432    let l_tz: Tz = l.timezone().unwrap_or("+00:00").parse()?;
433
434    let array: PrimitiveArray<T> = match (op, r.data_type()) {
435        (Op::Sub | Op::SubWrapping, Timestamp(unit, _)) if unit == &T::UNIT => {
436            let r = r.as_primitive::<T>();
437            return Ok(try_op_ref!(T::Duration, l, l_s, r, r_s, l.sub_checked(r)));
438        }
439
440        (Op::Add | Op::AddWrapping, Duration(unit)) if unit == &T::UNIT => {
441            let r = r.as_primitive::<T::Duration>();
442            try_op!(l, l_s, r, r_s, l.add_checked(r))
443        }
444        (Op::Sub | Op::SubWrapping, Duration(unit)) if unit == &T::UNIT => {
445            let r = r.as_primitive::<T::Duration>();
446            try_op!(l, l_s, r, r_s, l.sub_checked(r))
447        }
448
449        (Op::Add | Op::AddWrapping, Interval(YearMonth)) => {
450            let r = r.as_primitive::<IntervalYearMonthType>();
451            try_op!(
452                l,
453                l_s,
454                r,
455                r_s,
456                T::add_year_month(l, r, l_tz).ok_or(ArrowError::ComputeError(
457                    "Timestamp out of range".to_string()
458                ))
459            )
460        }
461        (Op::Sub | Op::SubWrapping, Interval(YearMonth)) => {
462            let r = r.as_primitive::<IntervalYearMonthType>();
463            try_op!(
464                l,
465                l_s,
466                r,
467                r_s,
468                T::sub_year_month(l, r, l_tz).ok_or(ArrowError::ComputeError(
469                    "Timestamp out of range".to_string()
470                ))
471            )
472        }
473
474        (Op::Add | Op::AddWrapping, Interval(DayTime)) => {
475            let r = r.as_primitive::<IntervalDayTimeType>();
476            try_op!(
477                l,
478                l_s,
479                r,
480                r_s,
481                T::add_day_time(l, r, l_tz).ok_or(ArrowError::ComputeError(
482                    "Timestamp out of range".to_string()
483                ))
484            )
485        }
486        (Op::Sub | Op::SubWrapping, Interval(DayTime)) => {
487            let r = r.as_primitive::<IntervalDayTimeType>();
488            try_op!(
489                l,
490                l_s,
491                r,
492                r_s,
493                T::sub_day_time(l, r, l_tz).ok_or(ArrowError::ComputeError(
494                    "Timestamp out of range".to_string()
495                ))
496            )
497        }
498
499        (Op::Add | Op::AddWrapping, Interval(MonthDayNano)) => {
500            let r = r.as_primitive::<IntervalMonthDayNanoType>();
501            try_op!(
502                l,
503                l_s,
504                r,
505                r_s,
506                T::add_month_day_nano(l, r, l_tz).ok_or(ArrowError::ComputeError(
507                    "Timestamp out of range".to_string()
508                ))
509            )
510        }
511        (Op::Sub | Op::SubWrapping, Interval(MonthDayNano)) => {
512            let r = r.as_primitive::<IntervalMonthDayNanoType>();
513            try_op!(
514                l,
515                l_s,
516                r,
517                r_s,
518                T::sub_month_day_nano(l, r, l_tz).ok_or(ArrowError::ComputeError(
519                    "Timestamp out of range".to_string()
520                ))
521            )
522        }
523        _ => {
524            return Err(ArrowError::InvalidArgumentError(format!(
525                "Invalid timestamp arithmetic operation: {} {op} {}",
526                l.data_type(),
527                r.data_type()
528            )));
529        }
530    };
531    Ok(Arc::new(array.with_timezone_opt(l.timezone())))
532}
533
534/// Arithmetic trait for date arrays
535trait DateOp: ArrowTemporalType {
536    fn add_year_month(timestamp: Self::Native, delta: i32) -> Result<Self::Native, ArrowError>;
537    fn add_day_time(
538        timestamp: Self::Native,
539        delta: IntervalDayTime,
540    ) -> Result<Self::Native, ArrowError>;
541    fn add_month_day_nano(
542        timestamp: Self::Native,
543        delta: IntervalMonthDayNano,
544    ) -> Result<Self::Native, ArrowError>;
545
546    fn sub_year_month(timestamp: Self::Native, delta: i32) -> Result<Self::Native, ArrowError>;
547    fn sub_day_time(
548        timestamp: Self::Native,
549        delta: IntervalDayTime,
550    ) -> Result<Self::Native, ArrowError>;
551    fn sub_month_day_nano(
552        timestamp: Self::Native,
553        delta: IntervalMonthDayNano,
554    ) -> Result<Self::Native, ArrowError>;
555}
556
557macro_rules! date {
558    ($t:ty) => {
559        impl DateOp for $t {
560            fn add_year_month(left: Self::Native, right: i32) -> Result<Self::Native, ArrowError> {
561                Self::add_year_months_opt(left, right).ok_or_else(|| {
562                    ArrowError::ComputeError(format!(
563                        "Date arithmetic overflow: {left} + {right} months"
564                    ))
565                })
566            }
567
568            fn add_day_time(
569                left: Self::Native,
570                right: IntervalDayTime,
571            ) -> Result<Self::Native, ArrowError> {
572                Self::add_day_time_opt(left, right).ok_or_else(|| {
573                    ArrowError::ComputeError(format!(
574                        "Date arithmetic overflow: {left} + {right:?}"
575                    ))
576                })
577            }
578
579            fn add_month_day_nano(
580                left: Self::Native,
581                right: IntervalMonthDayNano,
582            ) -> Result<Self::Native, ArrowError> {
583                Self::add_month_day_nano_opt(left, right).ok_or_else(|| {
584                    ArrowError::ComputeError(format!(
585                        "Date arithmetic overflow: {left} + {right:?}"
586                    ))
587                })
588            }
589
590            fn sub_year_month(left: Self::Native, right: i32) -> Result<Self::Native, ArrowError> {
591                Self::subtract_year_months_opt(left, right).ok_or_else(|| {
592                    ArrowError::ComputeError(format!(
593                        "Date arithmetic overflow: {left} - {right} months"
594                    ))
595                })
596            }
597
598            fn sub_day_time(
599                left: Self::Native,
600                right: IntervalDayTime,
601            ) -> Result<Self::Native, ArrowError> {
602                Self::subtract_day_time_opt(left, right).ok_or_else(|| {
603                    ArrowError::ComputeError(format!(
604                        "Date arithmetic overflow: {left} - {right:?}"
605                    ))
606                })
607            }
608
609            fn sub_month_day_nano(
610                left: Self::Native,
611                right: IntervalMonthDayNano,
612            ) -> Result<Self::Native, ArrowError> {
613                Self::subtract_month_day_nano_opt(left, right).ok_or_else(|| {
614                    ArrowError::ComputeError(format!(
615                        "Date arithmetic overflow: {left} - {right:?}"
616                    ))
617                })
618            }
619        }
620    };
621}
622
623date!(Date32Type);
624date!(Date64Type);
625
626/// Arithmetic trait for interval arrays
627trait IntervalOp: ArrowPrimitiveType {
628    fn add(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError>;
629    fn sub(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError>;
630    fn mul_i64(left: Self::Native, right: i64) -> Result<Self::Native, ArrowError>;
631}
632
633fn mul_i32_i64(left: i32, right: i64) -> Result<i32, ArrowError> {
634    let value = i64::from(left).mul_checked(right)?;
635    i32::try_from(value).map_err(|_| {
636        ArrowError::ArithmeticOverflow(format!("Overflow happened on: {left} * {right}"))
637    })
638}
639
640impl IntervalOp for IntervalYearMonthType {
641    fn add(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
642        left.add_checked(right)
643    }
644
645    fn sub(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
646        left.sub_checked(right)
647    }
648
649    fn mul_i64(left: Self::Native, right: i64) -> Result<Self::Native, ArrowError> {
650        mul_i32_i64(left, right)
651    }
652}
653
654impl IntervalOp for IntervalDayTimeType {
655    fn add(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
656        let (l_days, l_ms) = Self::to_parts(left);
657        let (r_days, r_ms) = Self::to_parts(right);
658        let days = l_days.add_checked(r_days)?;
659        let ms = l_ms.add_checked(r_ms)?;
660        Ok(Self::make_value(days, ms))
661    }
662
663    fn sub(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
664        let (l_days, l_ms) = Self::to_parts(left);
665        let (r_days, r_ms) = Self::to_parts(right);
666        let days = l_days.sub_checked(r_days)?;
667        let ms = l_ms.sub_checked(r_ms)?;
668        Ok(Self::make_value(days, ms))
669    }
670
671    fn mul_i64(left: Self::Native, right: i64) -> Result<Self::Native, ArrowError> {
672        let (days, ms) = Self::to_parts(left);
673        Ok(Self::make_value(
674            mul_i32_i64(days, right)?,
675            mul_i32_i64(ms, right)?,
676        ))
677    }
678}
679
680impl IntervalOp for IntervalMonthDayNanoType {
681    fn add(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
682        let (l_months, l_days, l_nanos) = Self::to_parts(left);
683        let (r_months, r_days, r_nanos) = Self::to_parts(right);
684        let months = l_months.add_checked(r_months)?;
685        let days = l_days.add_checked(r_days)?;
686        let nanos = l_nanos.add_checked(r_nanos)?;
687        Ok(Self::make_value(months, days, nanos))
688    }
689
690    fn sub(left: Self::Native, right: Self::Native) -> Result<Self::Native, ArrowError> {
691        let (l_months, l_days, l_nanos) = Self::to_parts(left);
692        let (r_months, r_days, r_nanos) = Self::to_parts(right);
693        let months = l_months.sub_checked(r_months)?;
694        let days = l_days.sub_checked(r_days)?;
695        let nanos = l_nanos.sub_checked(r_nanos)?;
696        Ok(Self::make_value(months, days, nanos))
697    }
698
699    fn mul_i64(left: Self::Native, right: i64) -> Result<Self::Native, ArrowError> {
700        let (months, days, nanos) = Self::to_parts(left);
701        Ok(Self::make_value(
702            mul_i32_i64(months, right)?,
703            mul_i32_i64(days, right)?,
704            nanos.mul_checked(right)?,
705        ))
706    }
707}
708
709fn interval_mul_op<T: IntervalOp>(
710    interval: &dyn Array,
711    interval_scalar: bool,
712    factor: &dyn Array,
713    factor_scalar: bool,
714) -> Result<ArrayRef, ArrowError> {
715    let interval = interval.as_primitive::<T>();
716    let factor = factor.as_primitive::<Int64Type>();
717    Ok(try_op_ref!(
718        T,
719        interval,
720        interval_scalar,
721        factor,
722        factor_scalar,
723        T::mul_i64(interval, factor)
724    ))
725}
726
727/// Perform arithmetic operation on an interval array
728fn interval_op<T: IntervalOp>(
729    op: Op,
730    l: &dyn Array,
731    l_s: bool,
732    r: &dyn Array,
733    r_s: bool,
734) -> Result<ArrayRef, ArrowError> {
735    let l = l.as_primitive::<T>();
736    let r = r.as_primitive::<T>();
737    match op {
738        Op::Add | Op::AddWrapping => Ok(try_op_ref!(T, l, l_s, r, r_s, T::add(l, r))),
739        Op::Sub | Op::SubWrapping => Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub(l, r))),
740        _ => Err(ArrowError::InvalidArgumentError(format!(
741            "Invalid interval arithmetic operation: {} {op} {}",
742            l.data_type(),
743            r.data_type()
744        ))),
745    }
746}
747
748fn duration_op<T: ArrowPrimitiveType>(
749    op: Op,
750    l: &dyn Array,
751    l_s: bool,
752    r: &dyn Array,
753    r_s: bool,
754) -> Result<ArrayRef, ArrowError> {
755    let l = l.as_primitive::<T>();
756    let r = r.as_primitive::<T>();
757    match op {
758        Op::Add | Op::AddWrapping => Ok(try_op_ref!(T, l, l_s, r, r_s, l.add_checked(r))),
759        Op::Sub | Op::SubWrapping => Ok(try_op_ref!(T, l, l_s, r, r_s, l.sub_checked(r))),
760        _ => Err(ArrowError::InvalidArgumentError(format!(
761            "Invalid duration arithmetic operation: {} {op} {}",
762            l.data_type(),
763            r.data_type()
764        ))),
765    }
766}
767
768/// Perform arithmetic operation on a date array
769fn date_op<T: DateOp>(
770    op: Op,
771    l: &dyn Array,
772    l_s: bool,
773    r: &dyn Array,
774    r_s: bool,
775) -> Result<ArrayRef, ArrowError> {
776    use DataType::*;
777    use IntervalUnit::*;
778
779    const NUM_SECONDS_IN_DAY: i64 = 60 * 60 * 24;
780
781    let r_t = r.data_type();
782    match (T::DATA_TYPE, op, r_t) {
783        (Date32, Op::Sub | Op::SubWrapping, Date32) => {
784            let l = l.as_primitive::<Date32Type>();
785            let r = r.as_primitive::<Date32Type>();
786            return Ok(op_ref!(
787                DurationSecondType,
788                l,
789                l_s,
790                r,
791                r_s,
792                ((l as i64) - (r as i64)) * NUM_SECONDS_IN_DAY
793            ));
794        }
795        (Date64, Op::Sub | Op::SubWrapping, Date64) => {
796            let l = l.as_primitive::<Date64Type>();
797            let r = r.as_primitive::<Date64Type>();
798            let result = try_op_ref!(DurationMillisecondType, l, l_s, r, r_s, l.sub_checked(r));
799            return Ok(result);
800        }
801        _ => {}
802    }
803
804    let l = l.as_primitive::<T>();
805    match (op, r_t) {
806        (Op::Add | Op::AddWrapping, Interval(YearMonth)) => {
807            let r = r.as_primitive::<IntervalYearMonthType>();
808            Ok(try_op_ref!(T, l, l_s, r, r_s, T::add_year_month(l, r)))
809        }
810        (Op::Sub | Op::SubWrapping, Interval(YearMonth)) => {
811            let r = r.as_primitive::<IntervalYearMonthType>();
812            Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub_year_month(l, r)))
813        }
814
815        (Op::Add | Op::AddWrapping, Interval(DayTime)) => {
816            let r = r.as_primitive::<IntervalDayTimeType>();
817            Ok(try_op_ref!(T, l, l_s, r, r_s, T::add_day_time(l, r)))
818        }
819        (Op::Sub | Op::SubWrapping, Interval(DayTime)) => {
820            let r = r.as_primitive::<IntervalDayTimeType>();
821            Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub_day_time(l, r)))
822        }
823
824        (Op::Add | Op::AddWrapping, Interval(MonthDayNano)) => {
825            let r = r.as_primitive::<IntervalMonthDayNanoType>();
826            Ok(try_op_ref!(T, l, l_s, r, r_s, T::add_month_day_nano(l, r)))
827        }
828        (Op::Sub | Op::SubWrapping, Interval(MonthDayNano)) => {
829            let r = r.as_primitive::<IntervalMonthDayNanoType>();
830            Ok(try_op_ref!(T, l, l_s, r, r_s, T::sub_month_day_nano(l, r)))
831        }
832
833        _ => Err(ArrowError::InvalidArgumentError(format!(
834            "Invalid date arithmetic operation: {} {op} {}",
835            l.data_type(),
836            r.data_type()
837        ))),
838    }
839}
840
841/// Perform arithmetic operation on decimal arrays
842fn decimal_op<T: DecimalType>(
843    op: Op,
844    l: &dyn Array,
845    l_s: bool,
846    r: &dyn Array,
847    r_s: bool,
848) -> Result<ArrayRef, ArrowError> {
849    let l = l.as_primitive::<T>();
850    let r = r.as_primitive::<T>();
851
852    let (p1, s1, p2, s2) = match (l.data_type(), r.data_type()) {
853        (DataType::Decimal32(p1, s1), DataType::Decimal32(p2, s2)) => (p1, s1, p2, s2),
854        (DataType::Decimal64(p1, s1), DataType::Decimal64(p2, s2)) => (p1, s1, p2, s2),
855        (DataType::Decimal128(p1, s1), DataType::Decimal128(p2, s2)) => (p1, s1, p2, s2),
856        (DataType::Decimal256(p1, s1), DataType::Decimal256(p2, s2)) => (p1, s1, p2, s2),
857        _ => unreachable!(),
858    };
859
860    // Follow the Hive decimal arithmetic rules
861    // https://cwiki.apache.org/confluence/download/attachments/27362075/Hive_Decimal_Precision_Scale_Support.pdf
862    let array: PrimitiveArray<T> = match op {
863        Op::Add | Op::AddWrapping | Op::Sub | Op::SubWrapping => {
864            // max(s1, s2)
865            let result_scale = *s1.max(s2);
866
867            // max(s1, s2) + max(p1-s1, p2-s2) + 1
868            let result_precision =
869                (result_scale.saturating_add((*p1 as i8 - s1).max(*p2 as i8 - s2)) as u8)
870                    .saturating_add(1)
871                    .min(T::MAX_PRECISION);
872
873            let l_mul = T::Native::usize_as(10).pow_checked((result_scale - s1) as _)?;
874            let r_mul = T::Native::usize_as(10).pow_checked((result_scale - s2) as _)?;
875
876            match op {
877                // Equal scales make both decimal multipliers one.
878                Op::Add | Op::AddWrapping if s1 == s2 => {
879                    try_op!(l, l_s, r, r_s, l.add_checked(r))
880                }
881                Op::Sub | Op::SubWrapping if s1 == s2 => {
882                    try_op!(l, l_s, r, r_s, l.sub_checked(r))
883                }
884                Op::Add | Op::AddWrapping => {
885                    try_op!(
886                        l,
887                        l_s,
888                        r,
889                        r_s,
890                        l.mul_checked(l_mul)?.add_checked(r.mul_checked(r_mul)?)
891                    )
892                }
893                Op::Sub | Op::SubWrapping => {
894                    try_op!(
895                        l,
896                        l_s,
897                        r,
898                        r_s,
899                        l.mul_checked(l_mul)?.sub_checked(r.mul_checked(r_mul)?)
900                    )
901                }
902                _ => unreachable!(),
903            }
904            .with_precision_and_scale(result_precision, result_scale)?
905        }
906        Op::Mul | Op::MulWrapping => {
907            let result_precision = p1.saturating_add(p2 + 1).min(T::MAX_PRECISION);
908            let result_scale = s1.saturating_add(*s2);
909            if result_scale > T::MAX_SCALE {
910                // SQL standard says that if the resulting scale of a multiply operation goes
911                // beyond the maximum, rounding is not acceptable and thus an error occurs
912                return Err(ArrowError::InvalidArgumentError(format!(
913                    "Output scale of {} {op} {} would exceed max scale of {}",
914                    l.data_type(),
915                    r.data_type(),
916                    T::MAX_SCALE
917                )));
918            }
919
920            try_op!(l, l_s, r, r_s, l.mul_checked(r))
921                .with_precision_and_scale(result_precision, result_scale)?
922        }
923
924        Op::Div => {
925            // Follow postgres and MySQL adding a fixed scale increment of 4
926            // s1 + 4
927            let result_scale = s1.saturating_add(4).min(T::MAX_SCALE);
928            let mul_pow = result_scale - s1 + s2;
929
930            // p1 - s1 + s2 + result_scale
931            let result_precision = (mul_pow.saturating_add(*p1 as i8) as u8).min(T::MAX_PRECISION);
932
933            let (l_mul, r_mul) = match mul_pow.cmp(&0) {
934                Ordering::Greater => (
935                    T::Native::usize_as(10).pow_checked(mul_pow as _)?,
936                    T::Native::ONE,
937                ),
938                Ordering::Equal => (T::Native::ONE, T::Native::ONE),
939                Ordering::Less => (
940                    T::Native::ONE,
941                    T::Native::usize_as(10).pow_checked(mul_pow.neg_wrapping() as _)?,
942                ),
943            };
944
945            try_op!(
946                l,
947                l_s,
948                r,
949                r_s,
950                l.mul_checked(l_mul)?.div_checked(r.mul_checked(r_mul)?)
951            )
952            .with_precision_and_scale(result_precision, result_scale)?
953        }
954
955        Op::Rem => {
956            // max(s1, s2)
957            let result_scale = *s1.max(s2);
958            // min(p1-s1, p2 -s2) + max( s1,s2 )
959            let result_precision =
960                (result_scale.saturating_add((*p1 as i8 - s1).min(*p2 as i8 - s2)) as u8)
961                    .min(T::MAX_PRECISION);
962
963            let l_mul = T::Native::usize_as(10).pow_wrapping((result_scale - s1) as _);
964            let r_mul = T::Native::usize_as(10).pow_wrapping((result_scale - s2) as _);
965
966            try_op!(
967                l,
968                l_s,
969                r,
970                r_s,
971                l.mul_checked(l_mul)?.mod_checked(r.mul_checked(r_mul)?)
972            )
973            .with_precision_and_scale(result_precision, result_scale)?
974        }
975    };
976
977    Ok(Arc::new(array))
978}
979
980#[cfg(test)]
981mod tests {
982    use super::*;
983    use arrow_array::temporal_conversions::{as_date, as_datetime};
984    use arrow_buffer::{ScalarBuffer, i256};
985    use chrono::{DateTime, NaiveDate};
986
987    // The valid date range of NaiveDate is from January 1, -262143 to December 31, 262142 (Gregorian calendar).
988    const MAX_VALID_DATE: NaiveDate = NaiveDate::from_ymd_opt(262142, 12, 31).unwrap();
989    const MIN_VALID_DATE: NaiveDate = NaiveDate::from_ymd_opt(-262143, 1, 1).unwrap();
990    const MAX_VALID_MILLIS: i64 = date_to_millis(MAX_VALID_DATE);
991    const MIN_VALID_MILLIS: i64 = date_to_millis(MIN_VALID_DATE);
992    const MAX_VALID_DAYS: i32 = date_to_days(MAX_VALID_DATE);
993    const MIN_VALID_DAYS: i32 = date_to_days(MIN_VALID_DATE);
994    const EPOCH: NaiveDate = NaiveDate::from_ymd_opt(1970, 1, 1).unwrap();
995    const YEAR_2000: NaiveDate = NaiveDate::from_ymd_opt(2000, 1, 1).unwrap();
996
997    const fn date_to_millis(date: NaiveDate) -> i64 {
998        date.signed_duration_since(EPOCH).num_milliseconds()
999    }
1000
1001    const fn date_to_days(date: NaiveDate) -> i32 {
1002        date.signed_duration_since(EPOCH).num_days() as i32
1003    }
1004
1005    fn test_neg_primitive<T: ArrowPrimitiveType>(
1006        input: &[T::Native],
1007        out: Result<&[T::Native], &str>,
1008    ) {
1009        let a = PrimitiveArray::<T>::new(ScalarBuffer::from(input.to_vec()), None);
1010        match out {
1011            Ok(expected) => {
1012                let result = neg(&a).unwrap();
1013                assert_eq!(result.as_primitive::<T>().values(), expected);
1014            }
1015            Err(e) => {
1016                let err = neg(&a).unwrap_err().to_string();
1017                assert_eq!(e, err);
1018            }
1019        }
1020    }
1021
1022    #[test]
1023    fn test_neg() {
1024        let input = &[1, -5, 2, 693, 3929];
1025        let output = &[-1, 5, -2, -693, -3929];
1026        test_neg_primitive::<Int32Type>(input, Ok(output));
1027
1028        let input = &[1, -5, 2, 693, 3929];
1029        let output = &[-1, 5, -2, -693, -3929];
1030        test_neg_primitive::<Int64Type>(input, Ok(output));
1031        test_neg_primitive::<DurationSecondType>(input, Ok(output));
1032        test_neg_primitive::<DurationMillisecondType>(input, Ok(output));
1033        test_neg_primitive::<DurationMicrosecondType>(input, Ok(output));
1034        test_neg_primitive::<DurationNanosecondType>(input, Ok(output));
1035
1036        let input = &[f32::MAX, f32::MIN, f32::INFINITY, 1.3, 0.5];
1037        let output = &[f32::MIN, f32::MAX, f32::NEG_INFINITY, -1.3, -0.5];
1038        test_neg_primitive::<Float32Type>(input, Ok(output));
1039
1040        test_neg_primitive::<Int32Type>(
1041            &[i32::MIN],
1042            Err("Arithmetic overflow: Overflow happened on: - -2147483648"),
1043        );
1044        test_neg_primitive::<Int64Type>(
1045            &[i64::MIN],
1046            Err("Arithmetic overflow: Overflow happened on: - -9223372036854775808"),
1047        );
1048        test_neg_primitive::<DurationSecondType>(
1049            &[i64::MIN],
1050            Err("Arithmetic overflow: Overflow happened on: - -9223372036854775808"),
1051        );
1052
1053        let r = neg_wrapping(&Int32Array::from(vec![i32::MIN])).unwrap();
1054        assert_eq!(r.as_primitive::<Int32Type>().value(0), i32::MIN);
1055
1056        let r = neg_wrapping(&Int64Array::from(vec![i64::MIN])).unwrap();
1057        assert_eq!(r.as_primitive::<Int64Type>().value(0), i64::MIN);
1058
1059        let err = neg_wrapping(&DurationSecondArray::from(vec![i64::MIN]))
1060            .unwrap_err()
1061            .to_string();
1062
1063        assert_eq!(
1064            err,
1065            "Arithmetic overflow: Overflow happened on: - -9223372036854775808"
1066        );
1067
1068        let a = Decimal32Array::from(vec![1, 3, -44, 2, 4])
1069            .with_precision_and_scale(9, 6)
1070            .unwrap();
1071
1072        let r = neg(&a).unwrap();
1073        assert_eq!(r.data_type(), a.data_type());
1074        assert_eq!(
1075            r.as_primitive::<Decimal32Type>().values(),
1076            &[-1, -3, 44, -2, -4]
1077        );
1078
1079        let a = Decimal64Array::from(vec![1, 3, -44, 2, 4])
1080            .with_precision_and_scale(9, 6)
1081            .unwrap();
1082
1083        let r = neg(&a).unwrap();
1084        assert_eq!(r.data_type(), a.data_type());
1085        assert_eq!(
1086            r.as_primitive::<Decimal64Type>().values(),
1087            &[-1, -3, 44, -2, -4]
1088        );
1089
1090        let a = Decimal128Array::from(vec![1, 3, -44, 2, 4])
1091            .with_precision_and_scale(9, 6)
1092            .unwrap();
1093
1094        let r = neg(&a).unwrap();
1095        assert_eq!(r.data_type(), a.data_type());
1096        assert_eq!(
1097            r.as_primitive::<Decimal128Type>().values(),
1098            &[-1, -3, 44, -2, -4]
1099        );
1100
1101        let a = Decimal256Array::from(vec![
1102            i256::from_i128(342),
1103            i256::from_i128(-4949),
1104            i256::from_i128(3),
1105        ])
1106        .with_precision_and_scale(9, 6)
1107        .unwrap();
1108
1109        let r = neg(&a).unwrap();
1110        assert_eq!(r.data_type(), a.data_type());
1111        assert_eq!(
1112            r.as_primitive::<Decimal256Type>().values(),
1113            &[
1114                i256::from_i128(-342),
1115                i256::from_i128(4949),
1116                i256::from_i128(-3),
1117            ]
1118        );
1119
1120        let a = IntervalYearMonthArray::from(vec![
1121            IntervalYearMonthType::make_value(2, 4),
1122            IntervalYearMonthType::make_value(2, -4),
1123            IntervalYearMonthType::make_value(-3, -5),
1124        ]);
1125        let r = neg(&a).unwrap();
1126        assert_eq!(
1127            r.as_primitive::<IntervalYearMonthType>().values(),
1128            &[
1129                IntervalYearMonthType::make_value(-2, -4),
1130                IntervalYearMonthType::make_value(-2, 4),
1131                IntervalYearMonthType::make_value(3, 5),
1132            ]
1133        );
1134
1135        let a = IntervalDayTimeArray::from(vec![
1136            IntervalDayTimeType::make_value(2, 4),
1137            IntervalDayTimeType::make_value(2, -4),
1138            IntervalDayTimeType::make_value(-3, -5),
1139        ]);
1140        let r = neg(&a).unwrap();
1141        assert_eq!(
1142            r.as_primitive::<IntervalDayTimeType>().values(),
1143            &[
1144                IntervalDayTimeType::make_value(-2, -4),
1145                IntervalDayTimeType::make_value(-2, 4),
1146                IntervalDayTimeType::make_value(3, 5),
1147            ]
1148        );
1149
1150        let a = IntervalMonthDayNanoArray::from(vec![
1151            IntervalMonthDayNanoType::make_value(2, 4, 5953394),
1152            IntervalMonthDayNanoType::make_value(2, -4, -45839),
1153            IntervalMonthDayNanoType::make_value(-3, -5, 6944),
1154        ]);
1155        let r = neg(&a).unwrap();
1156        assert_eq!(
1157            r.as_primitive::<IntervalMonthDayNanoType>().values(),
1158            &[
1159                IntervalMonthDayNanoType::make_value(-2, -4, -5953394),
1160                IntervalMonthDayNanoType::make_value(-2, 4, 45839),
1161                IntervalMonthDayNanoType::make_value(3, 5, -6944),
1162            ]
1163        );
1164    }
1165
1166    #[test]
1167    fn test_integer() {
1168        let a = Int32Array::from(vec![4, 3, 5, -6, 100]);
1169        let b = Int32Array::from(vec![6, 2, 5, -7, 3]);
1170        let result = add(&a, &b).unwrap();
1171        assert_eq!(
1172            result.as_ref(),
1173            &Int32Array::from(vec![10, 5, 10, -13, 103])
1174        );
1175        let result = sub(&a, &b).unwrap();
1176        assert_eq!(result.as_ref(), &Int32Array::from(vec![-2, 1, 0, 1, 97]));
1177        let result = div(&a, &b).unwrap();
1178        assert_eq!(result.as_ref(), &Int32Array::from(vec![0, 1, 1, 0, 33]));
1179        let result = mul(&a, &b).unwrap();
1180        assert_eq!(result.as_ref(), &Int32Array::from(vec![24, 6, 25, 42, 300]));
1181        let result = rem(&a, &b).unwrap();
1182        assert_eq!(result.as_ref(), &Int32Array::from(vec![4, 1, 0, -6, 1]));
1183
1184        let a = Int8Array::from(vec![Some(2), None, Some(45)]);
1185        let b = Int8Array::from(vec![Some(5), Some(3), None]);
1186        let result = add(&a, &b).unwrap();
1187        assert_eq!(result.as_ref(), &Int8Array::from(vec![Some(7), None, None]));
1188
1189        let a = UInt8Array::from(vec![56, 5, 3]);
1190        let b = UInt8Array::from(vec![200, 2, 5]);
1191        let err = add(&a, &b).unwrap_err().to_string();
1192        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 56 + 200");
1193        let result = add_wrapping(&a, &b).unwrap();
1194        assert_eq!(result.as_ref(), &UInt8Array::from(vec![0, 7, 8]));
1195
1196        let a = UInt8Array::from(vec![34, 5, 3]);
1197        let b = UInt8Array::from(vec![200, 2, 5]);
1198        let err = sub(&a, &b).unwrap_err().to_string();
1199        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 34 - 200");
1200        let result = sub_wrapping(&a, &b).unwrap();
1201        assert_eq!(result.as_ref(), &UInt8Array::from(vec![90, 3, 254]));
1202
1203        let a = UInt8Array::from(vec![34, 5, 3]);
1204        let b = UInt8Array::from(vec![200, 2, 5]);
1205        let err = mul(&a, &b).unwrap_err().to_string();
1206        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 34 * 200");
1207        let result = mul_wrapping(&a, &b).unwrap();
1208        assert_eq!(result.as_ref(), &UInt8Array::from(vec![144, 10, 15]));
1209
1210        let a = Int16Array::from(vec![i16::MIN]);
1211        let b = Int16Array::from(vec![-1]);
1212        let err = div(&a, &b).unwrap_err().to_string();
1213        assert_eq!(
1214            err,
1215            "Arithmetic overflow: Overflow happened on: -32768 / -1"
1216        );
1217
1218        let a = Int16Array::from(vec![i16::MIN]);
1219        let b = Int16Array::from(vec![-1]);
1220        let result = rem(&a, &b).unwrap();
1221        assert_eq!(result.as_ref(), &Int16Array::from(vec![0]));
1222
1223        let a = Int16Array::from(vec![21]);
1224        let b = Int16Array::from(vec![0]);
1225        let err = div(&a, &b).unwrap_err().to_string();
1226        assert_eq!(err, "Divide by zero error");
1227
1228        let a = Int16Array::from(vec![21]);
1229        let b = Int16Array::from(vec![0]);
1230        let err = rem(&a, &b).unwrap_err().to_string();
1231        assert_eq!(err, "Divide by zero error");
1232    }
1233
1234    #[test]
1235    fn test_float() {
1236        let a = Float32Array::from(vec![1., f32::MAX, 6., -4., -1., 0.]);
1237        let b = Float32Array::from(vec![1., f32::MAX, f32::MAX, -3., 45., 0.]);
1238        let result = add(&a, &b).unwrap();
1239        assert_eq!(
1240            result.as_ref(),
1241            &Float32Array::from(vec![2., f32::INFINITY, f32::MAX, -7., 44.0, 0.])
1242        );
1243
1244        let result = sub(&a, &b).unwrap();
1245        assert_eq!(
1246            result.as_ref(),
1247            &Float32Array::from(vec![0., 0., f32::MIN, -1., -46., 0.])
1248        );
1249
1250        let result = mul(&a, &b).unwrap();
1251        assert_eq!(
1252            result.as_ref(),
1253            &Float32Array::from(vec![1., f32::INFINITY, f32::INFINITY, 12., -45., 0.])
1254        );
1255
1256        let result = div(&a, &b).unwrap();
1257        let r = result.as_primitive::<Float32Type>();
1258        assert_eq!(r.value(0), 1.);
1259        assert_eq!(r.value(1), 1.);
1260        assert!(r.value(2) < f32::EPSILON);
1261        assert_eq!(r.value(3), -4. / -3.);
1262        assert!(r.value(5).is_nan());
1263
1264        let result = rem(&a, &b).unwrap();
1265        let r = result.as_primitive::<Float32Type>();
1266        assert_eq!(&r.values()[..5], &[0., 0., 6., -1., -1.]);
1267        assert!(r.value(5).is_nan());
1268    }
1269
1270    #[test]
1271    fn test_decimal() {
1272        // 0.015 7.842 -0.577 0.334 -0.078 0.003
1273        let a = Decimal128Array::from(vec![15, 0, -577, 334, -78, 3])
1274            .with_precision_and_scale(12, 3)
1275            .unwrap();
1276
1277        // 5.4 0 -35.6 0.3 0.6 7.45
1278        let b = Decimal128Array::from(vec![54, 34, -356, 3, 6, 745])
1279            .with_precision_and_scale(12, 1)
1280            .unwrap();
1281
1282        let result = add(&a, &b).unwrap();
1283        assert_eq!(result.data_type(), &DataType::Decimal128(15, 3));
1284        assert_eq!(
1285            result.as_primitive::<Decimal128Type>().values(),
1286            &[5415, 3400, -36177, 634, 522, 74503]
1287        );
1288
1289        let result = sub(&a, &b).unwrap();
1290        assert_eq!(result.data_type(), &DataType::Decimal128(15, 3));
1291        assert_eq!(
1292            result.as_primitive::<Decimal128Type>().values(),
1293            &[-5385, -3400, 35023, 34, -678, -74497]
1294        );
1295
1296        let result = mul(&a, &b).unwrap();
1297        assert_eq!(result.data_type(), &DataType::Decimal128(25, 4));
1298        assert_eq!(
1299            result.as_primitive::<Decimal128Type>().values(),
1300            &[810, 0, 205412, 1002, -468, 2235]
1301        );
1302
1303        let result = div(&a, &b).unwrap();
1304        assert_eq!(result.data_type(), &DataType::Decimal128(17, 7));
1305        assert_eq!(
1306            result.as_primitive::<Decimal128Type>().values(),
1307            &[27777, 0, 162078, 11133333, -1300000, 402]
1308        );
1309
1310        let result = rem(&a, &b).unwrap();
1311        assert_eq!(result.data_type(), &DataType::Decimal128(12, 3));
1312        assert_eq!(
1313            result.as_primitive::<Decimal128Type>().values(),
1314            &[15, 0, -577, 34, -78, 3]
1315        );
1316
1317        let a = Decimal128Array::from(vec![1])
1318            .with_precision_and_scale(3, 3)
1319            .unwrap();
1320        let b = Decimal128Array::from(vec![1])
1321            .with_precision_and_scale(37, 37)
1322            .unwrap();
1323        let err = mul(&a, &b).unwrap_err().to_string();
1324        assert_eq!(
1325            err,
1326            "Invalid argument error: Output scale of Decimal128(3, 3) * Decimal128(37, 37) would exceed max scale of 38"
1327        );
1328
1329        let a = Decimal128Array::from(vec![1])
1330            .with_precision_and_scale(3, -2)
1331            .unwrap();
1332        let err = add(&a, &b).unwrap_err().to_string();
1333        assert_eq!(err, "Arithmetic overflow: Overflow happened on: 10 ^ 39");
1334
1335        let a = Decimal128Array::from(vec![10])
1336            .with_precision_and_scale(3, -1)
1337            .unwrap();
1338        let err = add(&a, &b).unwrap_err().to_string();
1339        assert_eq!(
1340            err,
1341            "Arithmetic overflow: Overflow happened on: 10 * 100000000000000000000000000000000000000"
1342        );
1343
1344        let b = Decimal128Array::from(vec![0])
1345            .with_precision_and_scale(1, 1)
1346            .unwrap();
1347        let err = div(&a, &b).unwrap_err().to_string();
1348        assert_eq!(err, "Divide by zero error");
1349        let err = rem(&a, &b).unwrap_err().to_string();
1350        assert_eq!(err, "Divide by zero error");
1351    }
1352
1353    #[test]
1354    fn test_decimal256_same_scale_add_sub() {
1355        let lhs = Decimal256Array::from(vec![
1356            Some(i256::from_parts(u128::MAX, 0)),
1357            Some(i256::MINUS_ONE),
1358            None,
1359        ])
1360        .with_precision_and_scale(70, 2)
1361        .unwrap();
1362        let rhs = Decimal256Array::from(vec![Some(i256::ONE), Some(i256::ONE), Some(i256::MAX)])
1363            .with_precision_and_scale(70, 2)
1364            .unwrap();
1365
1366        let expected =
1367            Decimal256Array::from(vec![Some(i256::from_parts(0, 1)), Some(i256::ZERO), None])
1368                .with_precision_and_scale(71, 2)
1369                .unwrap();
1370        for operation in [add, add_wrapping] {
1371            let result = operation(&lhs, &rhs).unwrap();
1372            assert_eq!(result.as_primitive::<Decimal256Type>(), &expected);
1373        }
1374
1375        let expected = Decimal256Array::from(vec![
1376            Some(i256::from_parts(u128::MAX - 1, 0)),
1377            Some(i256::from_i128(-2)),
1378            None,
1379        ])
1380        .with_precision_and_scale(71, 2)
1381        .unwrap();
1382        for operation in [sub, sub_wrapping] {
1383            let result = operation(&lhs, &rhs).unwrap();
1384            assert_eq!(result.as_primitive::<Decimal256Type>(), &expected);
1385        }
1386
1387        let lhs = Decimal256Array::from(vec![i256::MAX])
1388            .with_precision_and_scale(76, 0)
1389            .unwrap();
1390        let rhs = Decimal256Array::from(vec![i256::ONE])
1391            .with_precision_and_scale(76, 0)
1392            .unwrap();
1393        for operation in [add, add_wrapping] {
1394            assert_eq!(
1395                operation(&lhs, &rhs).unwrap_err().to_string(),
1396                format!(
1397                    "Arithmetic overflow: Overflow happened on: {:?} + {:?}",
1398                    i256::MAX,
1399                    i256::ONE
1400                )
1401            );
1402        }
1403
1404        let lhs = Decimal256Array::from(vec![i256::MIN])
1405            .with_precision_and_scale(76, 0)
1406            .unwrap();
1407        for operation in [sub, sub_wrapping] {
1408            assert_eq!(
1409                operation(&lhs, &rhs).unwrap_err().to_string(),
1410                format!(
1411                    "Arithmetic overflow: Overflow happened on: {:?} - {:?}",
1412                    i256::MIN,
1413                    i256::ONE
1414                )
1415            );
1416        }
1417    }
1418
1419    fn test_timestamp_impl<T: TimestampOp>() {
1420        let a = PrimitiveArray::<T>::new(vec![2000000, 434030324, 53943340].into(), None);
1421        let b = PrimitiveArray::<T>::new(vec![329593, 59349, 694994].into(), None);
1422
1423        let result = sub(&a, &b).unwrap();
1424        assert_eq!(
1425            result.as_primitive::<T::Duration>().values(),
1426            &[1670407, 433970975, 53248346]
1427        );
1428
1429        let r2 = add(&b, &result.as_ref()).unwrap();
1430        assert_eq!(r2.as_ref(), &a);
1431
1432        let r3 = add(&result.as_ref(), &b).unwrap();
1433        assert_eq!(r3.as_ref(), &a);
1434
1435        let format_array = |x: &dyn Array| -> Vec<String> {
1436            x.as_primitive::<T>()
1437                .values()
1438                .into_iter()
1439                .map(|x| as_datetime::<T>(*x).unwrap().to_string())
1440                .collect()
1441        };
1442
1443        let values = vec![
1444            "1970-01-01T00:00:00Z",
1445            "2010-04-01T04:00:20Z",
1446            "1960-01-30T04:23:20Z",
1447        ]
1448        .into_iter()
1449        .map(|x| {
1450            T::from_naive_datetime(DateTime::parse_from_rfc3339(x).unwrap().naive_utc(), None)
1451                .unwrap()
1452        })
1453        .collect();
1454
1455        let a = PrimitiveArray::<T>::new(values, None);
1456        let b = IntervalYearMonthArray::from(vec![
1457            IntervalYearMonthType::make_value(5, 34),
1458            IntervalYearMonthType::make_value(-2, 4),
1459            IntervalYearMonthType::make_value(7, -4),
1460        ]);
1461        let r4 = add(&a, &b).unwrap();
1462        assert_eq!(
1463            &format_array(r4.as_ref()),
1464            &[
1465                "1977-11-01 00:00:00".to_string(),
1466                "2008-08-01 04:00:20".to_string(),
1467                "1966-09-30 04:23:20".to_string()
1468            ]
1469        );
1470
1471        let r5 = sub(&r4, &b).unwrap();
1472        assert_eq!(r5.as_ref(), &a);
1473
1474        let b = IntervalDayTimeArray::from(vec![
1475            IntervalDayTimeType::make_value(5, 454000),
1476            IntervalDayTimeType::make_value(-34, 0),
1477            IntervalDayTimeType::make_value(7, -4000),
1478        ]);
1479        let r6 = add(&a, &b).unwrap();
1480        assert_eq!(
1481            &format_array(r6.as_ref()),
1482            &[
1483                "1970-01-06 00:07:34".to_string(),
1484                "2010-02-26 04:00:20".to_string(),
1485                "1960-02-06 04:23:16".to_string()
1486            ]
1487        );
1488
1489        let r7 = sub(&r6, &b).unwrap();
1490        assert_eq!(r7.as_ref(), &a);
1491
1492        let b = IntervalMonthDayNanoArray::from(vec![
1493            IntervalMonthDayNanoType::make_value(344, 34, -43_000_000_000),
1494            IntervalMonthDayNanoType::make_value(-593, -33, 13_000_000_000),
1495            IntervalMonthDayNanoType::make_value(5, 2, 493_000_000_000),
1496        ]);
1497        let r8 = add(&a, &b).unwrap();
1498        assert_eq!(
1499            &format_array(r8.as_ref()),
1500            &[
1501                "1998-10-04 23:59:17".to_string(),
1502                "1960-09-29 04:00:33".to_string(),
1503                "1960-07-02 04:31:33".to_string()
1504            ]
1505        );
1506
1507        let r9 = sub(&r8, &b).unwrap();
1508        // Note: subtraction is not the inverse of addition for intervals
1509        assert_eq!(
1510            &format_array(r9.as_ref()),
1511            &[
1512                "1970-01-02 00:00:00".to_string(),
1513                "2010-04-02 04:00:20".to_string(),
1514                "1960-01-31 04:23:20".to_string()
1515            ]
1516        );
1517    }
1518
1519    #[test]
1520    fn test_timestamp() {
1521        test_timestamp_impl::<TimestampSecondType>();
1522        test_timestamp_impl::<TimestampMillisecondType>();
1523        test_timestamp_impl::<TimestampMicrosecondType>();
1524        test_timestamp_impl::<TimestampNanosecondType>();
1525    }
1526
1527    #[test]
1528    fn test_interval() {
1529        let a = IntervalYearMonthArray::from(vec![
1530            IntervalYearMonthType::make_value(32, 4),
1531            IntervalYearMonthType::make_value(32, 4),
1532        ]);
1533        let b = IntervalYearMonthArray::from(vec![
1534            IntervalYearMonthType::make_value(-4, 6),
1535            IntervalYearMonthType::make_value(-3, 23),
1536        ]);
1537        let result = add(&a, &b).unwrap();
1538        assert_eq!(
1539            result.as_ref(),
1540            &IntervalYearMonthArray::from(vec![
1541                IntervalYearMonthType::make_value(28, 10),
1542                IntervalYearMonthType::make_value(29, 27)
1543            ])
1544        );
1545        let result = sub(&a, &b).unwrap();
1546        assert_eq!(
1547            result.as_ref(),
1548            &IntervalYearMonthArray::from(vec![
1549                IntervalYearMonthType::make_value(36, -2),
1550                IntervalYearMonthType::make_value(35, -19)
1551            ])
1552        );
1553
1554        let a = IntervalDayTimeArray::from(vec![
1555            IntervalDayTimeType::make_value(32, 4),
1556            IntervalDayTimeType::make_value(32, 4),
1557        ]);
1558        let b = IntervalDayTimeArray::from(vec![
1559            IntervalDayTimeType::make_value(-4, 6),
1560            IntervalDayTimeType::make_value(-3, 23),
1561        ]);
1562        let result = add(&a, &b).unwrap();
1563        assert_eq!(
1564            result.as_ref(),
1565            &IntervalDayTimeArray::from(vec![
1566                IntervalDayTimeType::make_value(28, 10),
1567                IntervalDayTimeType::make_value(29, 27)
1568            ])
1569        );
1570        let result = sub(&a, &b).unwrap();
1571        assert_eq!(
1572            result.as_ref(),
1573            &IntervalDayTimeArray::from(vec![
1574                IntervalDayTimeType::make_value(36, -2),
1575                IntervalDayTimeType::make_value(35, -19)
1576            ])
1577        );
1578        let a = IntervalMonthDayNanoArray::from(vec![
1579            IntervalMonthDayNanoType::make_value(32, 4, 4000000000000),
1580            IntervalMonthDayNanoType::make_value(32, 4, 45463000000000000),
1581        ]);
1582        let b = IntervalMonthDayNanoArray::from(vec![
1583            IntervalMonthDayNanoType::make_value(-4, 6, 46000000000000),
1584            IntervalMonthDayNanoType::make_value(-3, 23, 3564000000000000),
1585        ]);
1586        let result = add(&a, &b).unwrap();
1587        assert_eq!(
1588            result.as_ref(),
1589            &IntervalMonthDayNanoArray::from(vec![
1590                IntervalMonthDayNanoType::make_value(28, 10, 50000000000000),
1591                IntervalMonthDayNanoType::make_value(29, 27, 49027000000000000)
1592            ])
1593        );
1594        let result = sub(&a, &b).unwrap();
1595        assert_eq!(
1596            result.as_ref(),
1597            &IntervalMonthDayNanoArray::from(vec![
1598                IntervalMonthDayNanoType::make_value(36, -2, -42000000000000),
1599                IntervalMonthDayNanoType::make_value(35, -19, 41899000000000000)
1600            ])
1601        );
1602        let a = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNano::MAX]);
1603        let b = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNano::ONE]);
1604        let err = add(&a, &b).unwrap_err().to_string();
1605        assert_eq!(
1606            err,
1607            "Arithmetic overflow: Overflow happened on: 2147483647 + 1"
1608        );
1609    }
1610
1611    #[test]
1612    fn test_interval_mul_i64() {
1613        let interval = IntervalYearMonthArray::from(vec![16, 5, 0]);
1614        let factor = Int64Array::from(vec![3, -2, i64::MAX]);
1615        let expected = IntervalYearMonthArray::from(vec![48, -10, 0]);
1616        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
1617        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
1618
1619        let interval = IntervalDayTimeArray::from(vec![
1620            Some(IntervalDayTimeType::make_value(10, 2 * 60 * 60 * 1000)),
1621            None,
1622        ]);
1623        let factor = Int64Array::new_scalar(3);
1624        let expected = IntervalDayTimeArray::from(vec![
1625            Some(IntervalDayTimeType::make_value(30, 6 * 60 * 60 * 1000)),
1626            None,
1627        ]);
1628        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
1629        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
1630
1631        let null_factor = Scalar::new(Int64Array::new_null(1));
1632        let expected = IntervalDayTimeArray::new_null(interval.len());
1633        assert_eq!(mul(&interval, &null_factor).unwrap().as_ref(), &expected);
1634        assert_eq!(mul(&null_factor, &interval).unwrap().as_ref(), &expected);
1635
1636        let interval = IntervalMonthDayNanoArray::new_scalar(IntervalMonthDayNanoType::make_value(
1637            12,
1638            15,
1639            5_000_000_000,
1640        ));
1641        let factor = Int64Array::from(vec![2, 0, -1]);
1642        let expected = IntervalMonthDayNanoArray::from(vec![
1643            IntervalMonthDayNanoType::make_value(24, 30, 10_000_000_000),
1644            IntervalMonthDayNanoType::make_value(0, 0, 0),
1645            IntervalMonthDayNanoType::make_value(-12, -15, -5_000_000_000),
1646        ]);
1647        assert_eq!(mul(&interval, &factor).unwrap().as_ref(), &expected);
1648        assert_eq!(mul(&factor, &interval).unwrap().as_ref(), &expected);
1649
1650        let float_factor = Float64Array::new_scalar(2.);
1651        assert!(mul(&interval, &float_factor).is_err());
1652        assert!(mul_wrapping(&factor, &interval).is_err());
1653    }
1654
1655    #[test]
1656    fn test_interval_mul_i64_overflow() {
1657        let interval = IntervalYearMonthArray::from(vec![i32::MAX]);
1658        let factor = Int64Array::from(vec![2]);
1659        assert_eq!(
1660            mul(&interval, &factor).unwrap_err().to_string(),
1661            "Arithmetic overflow: Overflow happened on: 2147483647 * 2"
1662        );
1663
1664        let interval = IntervalMonthDayNanoArray::from(vec![IntervalMonthDayNanoType::make_value(
1665            0,
1666            0,
1667            i64::MAX,
1668        )]);
1669        assert_eq!(
1670            mul(&interval, &factor).unwrap_err().to_string(),
1671            "Arithmetic overflow: Overflow happened on: 9223372036854775807 * 2"
1672        );
1673    }
1674
1675    fn test_duration_impl<T: ArrowPrimitiveType<Native = i64>>() {
1676        let a = PrimitiveArray::<T>::new(vec![1000, 4394, -3944].into(), None);
1677        let b = PrimitiveArray::<T>::new(vec![4, -5, -243].into(), None);
1678
1679        let result = add(&a, &b).unwrap();
1680        assert_eq!(result.as_primitive::<T>().values(), &[1004, 4389, -4187]);
1681        let result = sub(&a, &b).unwrap();
1682        assert_eq!(result.as_primitive::<T>().values(), &[996, 4399, -3701]);
1683
1684        let err = mul(&a, &b).unwrap_err().to_string();
1685        assert!(
1686            err.contains("Invalid duration arithmetic operation"),
1687            "{err}"
1688        );
1689
1690        let err = div(&a, &b).unwrap_err().to_string();
1691        assert!(
1692            err.contains("Invalid duration arithmetic operation"),
1693            "{err}"
1694        );
1695
1696        let err = rem(&a, &b).unwrap_err().to_string();
1697        assert!(
1698            err.contains("Invalid duration arithmetic operation"),
1699            "{err}"
1700        );
1701
1702        let a = PrimitiveArray::<T>::new(vec![i64::MAX].into(), None);
1703        let b = PrimitiveArray::<T>::new(vec![1].into(), None);
1704        let err = add(&a, &b).unwrap_err().to_string();
1705        assert_eq!(
1706            err,
1707            "Arithmetic overflow: Overflow happened on: 9223372036854775807 + 1"
1708        );
1709    }
1710
1711    #[test]
1712    fn test_duration() {
1713        test_duration_impl::<DurationSecondType>();
1714        test_duration_impl::<DurationMillisecondType>();
1715        test_duration_impl::<DurationMicrosecondType>();
1716        test_duration_impl::<DurationNanosecondType>();
1717    }
1718
1719    fn test_date_impl<T: ArrowPrimitiveType, F>(f: F)
1720    where
1721        F: Fn(NaiveDate) -> T::Native,
1722        T::Native: TryInto<i64>,
1723    {
1724        let a = PrimitiveArray::<T>::new(
1725            vec![
1726                f(NaiveDate::from_ymd_opt(1979, 1, 30).unwrap()),
1727                f(NaiveDate::from_ymd_opt(2010, 4, 3).unwrap()),
1728                f(NaiveDate::from_ymd_opt(2008, 2, 29).unwrap()),
1729            ]
1730            .into(),
1731            None,
1732        );
1733
1734        let b = IntervalYearMonthArray::from(vec![
1735            IntervalYearMonthType::make_value(34, 2),
1736            IntervalYearMonthType::make_value(3, -3),
1737            IntervalYearMonthType::make_value(-12, 4),
1738        ]);
1739
1740        let format_array = |x: &dyn Array| -> Vec<String> {
1741            x.as_primitive::<T>()
1742                .values()
1743                .into_iter()
1744                .map(|x| {
1745                    as_date::<T>((*x).try_into().ok().unwrap())
1746                        .unwrap()
1747                        .to_string()
1748                })
1749                .collect()
1750        };
1751
1752        let result = add(&a, &b).unwrap();
1753        assert_eq!(
1754            &format_array(result.as_ref()),
1755            &[
1756                "2013-03-30".to_string(),
1757                "2013-01-03".to_string(),
1758                "1996-06-29".to_string(),
1759            ]
1760        );
1761        let result = sub(&result, &b).unwrap();
1762        assert_eq!(result.as_ref(), &a);
1763
1764        let b = IntervalDayTimeArray::from(vec![
1765            IntervalDayTimeType::make_value(34, 2),
1766            IntervalDayTimeType::make_value(3, -3),
1767            IntervalDayTimeType::make_value(-12, 4),
1768        ]);
1769
1770        let result = add(&a, &b).unwrap();
1771        assert_eq!(
1772            &format_array(result.as_ref()),
1773            &[
1774                "1979-03-05".to_string(),
1775                "2010-04-06".to_string(),
1776                "2008-02-17".to_string(),
1777            ]
1778        );
1779        let result = sub(&result, &b).unwrap();
1780        assert_eq!(result.as_ref(), &a);
1781
1782        let b = IntervalMonthDayNanoArray::from(vec![
1783            IntervalMonthDayNanoType::make_value(34, 2, -34353534),
1784            IntervalMonthDayNanoType::make_value(3, -3, 2443),
1785            IntervalMonthDayNanoType::make_value(-12, 4, 2323242423232),
1786        ]);
1787
1788        let result = add(&a, &b).unwrap();
1789        assert_eq!(
1790            &format_array(result.as_ref()),
1791            &[
1792                "1981-12-02".to_string(),
1793                "2010-06-30".to_string(),
1794                "2007-03-04".to_string(),
1795            ]
1796        );
1797        let result = sub(&result, &b).unwrap();
1798        assert_eq!(
1799            &format_array(result.as_ref()),
1800            &[
1801                "1979-01-31".to_string(),
1802                "2010-04-02".to_string(),
1803                "2008-02-29".to_string(),
1804            ]
1805        );
1806    }
1807
1808    #[test]
1809    fn test_date() {
1810        test_date_impl::<Date32Type, _>(Date32Type::from_naive_date);
1811        test_date_impl::<Date64Type, _>(Date64Type::from_naive_date);
1812
1813        let a = Date32Array::from(vec![i32::MIN, i32::MAX, 23, 7684]);
1814        let b = Date32Array::from(vec![i32::MIN, i32::MIN, -2, 45]);
1815        let result = sub(&a, &b).unwrap();
1816        assert_eq!(
1817            result.as_primitive::<DurationSecondType>().values(),
1818            &[0, 371085174288000, 2160000, 660009600]
1819        );
1820
1821        let a = Date64Array::from(vec![4343, 76676, 3434]);
1822        let b = Date64Array::from(vec![3, -5, 5]);
1823        let result = sub(&a, &b).unwrap();
1824        assert_eq!(
1825            result.as_primitive::<DurationMillisecondType>().values(),
1826            &[4340, 76681, 3429]
1827        );
1828
1829        let a = Date64Array::from(vec![i64::MAX]);
1830        let b = Date64Array::from(vec![-1]);
1831        let err = sub(&a, &b).unwrap_err().to_string();
1832        assert_eq!(
1833            err,
1834            "Arithmetic overflow: Overflow happened on: 9223372036854775807 - -1"
1835        );
1836    }
1837
1838    #[test]
1839    fn test_date32_to_naive_date_opt_boundaries() {
1840        assert_eq!(MAX_VALID_DAYS, 95026236);
1841        assert_eq!(MIN_VALID_DAYS, -96465292);
1842
1843        // Valid boundary dates work
1844        assert!(Date32Type::to_naive_date_opt(MAX_VALID_DAYS).is_some());
1845        assert!(Date32Type::to_naive_date_opt(MIN_VALID_DAYS).is_some());
1846
1847        // Beyond boundaries fail
1848        assert!(Date32Type::to_naive_date_opt(MAX_VALID_DAYS + 1).is_none());
1849        assert!(Date32Type::to_naive_date_opt(MIN_VALID_DAYS - 1).is_none());
1850
1851        // Extreme values fail
1852        assert!(Date32Type::to_naive_date_opt(i32::MAX).is_none());
1853        assert!(Date32Type::to_naive_date_opt(i32::MIN).is_none());
1854
1855        // Common values work
1856        assert!(Date32Type::to_naive_date_opt(0).is_some());
1857        assert!(Date32Type::to_naive_date_opt(date_to_days(YEAR_2000)).is_some());
1858    }
1859
1860    #[test]
1861    fn test_date64_to_naive_date_opt_boundaries() {
1862        const MS_PER_DAY: i64 = 24 * 60 * 60 * 1000;
1863
1864        // Verify boundary millisecond values
1865        assert_eq!(MAX_VALID_MILLIS, 8210266790400000i64);
1866        assert_eq!(MIN_VALID_MILLIS, -8334601228800000i64);
1867
1868        // Valid boundary dates work
1869        assert!(Date64Type::to_naive_date_opt(MAX_VALID_MILLIS).is_some());
1870        assert!(Date64Type::to_naive_date_opt(MIN_VALID_MILLIS).is_some());
1871
1872        // Beyond boundaries fail
1873        assert!(Date64Type::to_naive_date_opt(MAX_VALID_MILLIS + MS_PER_DAY).is_none());
1874        assert!(Date64Type::to_naive_date_opt(MIN_VALID_MILLIS - MS_PER_DAY).is_none());
1875
1876        // Extreme values fail
1877        assert!(Date64Type::to_naive_date_opt(i64::MAX).is_none());
1878        assert!(Date64Type::to_naive_date_opt(i64::MIN).is_none());
1879
1880        // Common values work
1881        assert!(Date64Type::to_naive_date_opt(0).is_some());
1882        assert!(Date64Type::to_naive_date_opt(date_to_millis(YEAR_2000)).is_some());
1883    }
1884
1885    macro_rules! test_year_month_ops {
1886        ($type:ty, $date_fn:expr) => {{
1887            let date = $date_fn(YEAR_2000);
1888
1889            // Normal operations succeed
1890            assert!(
1891                <$type>::add_year_months_opt(date, 120).is_some(),
1892                "add_year_months: normal add"
1893            );
1894            assert!(
1895                <$type>::add_year_months_opt(date, 0).is_some(),
1896                "add_year_months: zero interval"
1897            );
1898            assert!(
1899                <$type>::subtract_year_months_opt(date, 120).is_some(),
1900                "subtract_year_months: normal subtract"
1901            );
1902            assert!(
1903                <$type>::subtract_year_months_opt(date, 0).is_some(),
1904                "subtract_year_months: zero interval"
1905            );
1906
1907            // Large but valid years work
1908            let large_year = $date_fn(NaiveDate::from_ymd_opt(5000, 1, 1).unwrap());
1909            let neg_year = $date_fn(NaiveDate::from_ymd_opt(-5000, 12, 31).unwrap());
1910            assert!(
1911                <$type>::add_year_months_opt(large_year, 12).is_some(),
1912                "add_year_months: large year"
1913            );
1914            assert!(
1915                <$type>::add_year_months_opt(neg_year, -12).is_some(),
1916                "add_year_months: negative year"
1917            );
1918            assert!(
1919                <$type>::subtract_year_months_opt(large_year, 12).is_some(),
1920                "subtract_year_months: large year"
1921            );
1922            assert!(
1923                <$type>::subtract_year_months_opt(neg_year, -12).is_some(),
1924                "subtract_year_months: negative year"
1925            );
1926
1927            // Overflow handling
1928            assert!(
1929                <$type>::subtract_year_months_opt($date_fn(MIN_VALID_DATE), 1).is_none(),
1930                "subtract_year_months: overflow days from min"
1931            );
1932            assert!(
1933                <$type>::subtract_year_months_opt($date_fn(MAX_VALID_DATE), -1).is_none(),
1934                "subtract_year_months: overflow neg days from max"
1935            );
1936            assert!(
1937                <$type>::add_year_months_opt($date_fn(MAX_VALID_DATE), 1).is_none(),
1938                "add_year_months: overflow days"
1939            );
1940            assert!(
1941                <$type>::add_year_months_opt($date_fn(MIN_VALID_DATE), -1).is_none(),
1942                "add_year_months: overflow neg days"
1943            );
1944        }};
1945    }
1946
1947    #[test]
1948    fn test_date_year_month_operations() {
1949        test_year_month_ops!(Date32Type, date_to_days);
1950        test_year_month_ops!(Date64Type, date_to_millis);
1951    }
1952
1953    macro_rules! test_day_time_ops {
1954        ($type:ty, $date_fn:expr) => {{
1955            let date = $date_fn(YEAR_2000);
1956
1957            // Moderate intervals succeed
1958            assert!(
1959                <$type>::add_day_time_opt(date, IntervalDayTime::new(30, 0)).is_some(),
1960                "add_day_time: +30 days"
1961            );
1962            assert!(
1963                <$type>::add_day_time_opt(date, IntervalDayTime::new(-30, 0)).is_some(),
1964                "add_day_time: -30 days"
1965            );
1966            assert!(
1967                <$type>::add_day_time_opt(date, IntervalDayTime::new(1000, 12345)).is_some(),
1968                "add_day_time: normal"
1969            );
1970            assert!(
1971                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(30, 0)).is_some(),
1972                "subtract_day_time: +30 days"
1973            );
1974            assert!(
1975                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(-30, 0)).is_some(),
1976                "subtract_day_time: -30 days"
1977            );
1978            assert!(
1979                <$type>::subtract_day_time_opt(date, IntervalDayTime::new(1000, 12345)).is_some(),
1980                "subtract_day_time: normal"
1981            );
1982
1983            // Overflow handling - subtract
1984            assert!(
1985                <$type>::subtract_day_time_opt(
1986                    $date_fn(MIN_VALID_DATE),
1987                    IntervalDayTime::new(1, 0)
1988                )
1989                .is_none(),
1990                "subtract_day_time: overflow days from min"
1991            );
1992            assert!(
1993                <$type>::subtract_day_time_opt(
1994                    $date_fn(MAX_VALID_DATE),
1995                    IntervalDayTime::new(-1, 0)
1996                )
1997                .is_none(),
1998                "subtract_day_time: overflow neg days from max"
1999            );
2000
2001            // Overflow handling - add
2002            assert!(
2003                <$type>::add_day_time_opt($date_fn(MAX_VALID_DATE), IntervalDayTime::new(1, 0))
2004                    .is_none(),
2005                "add_day_time: overflow days"
2006            );
2007            assert!(
2008                <$type>::add_day_time_opt($date_fn(MIN_VALID_DATE), IntervalDayTime::new(-1, 0))
2009                    .is_none(),
2010                "add_day_time: overflow neg days"
2011            );
2012
2013            // Extreme intervals fail
2014            assert!(
2015                <$type>::add_day_time_opt(
2016                    $date_fn(EPOCH),
2017                    IntervalDayTime::new(i32::MAX, i32::MAX)
2018                )
2019                .is_none(),
2020                "add_day_time: max interval"
2021            );
2022            assert!(
2023                <$type>::add_day_time_opt(
2024                    $date_fn(EPOCH),
2025                    IntervalDayTime::new(i32::MIN, i32::MIN)
2026                )
2027                .is_none(),
2028                "add_day_time: min interval"
2029            );
2030            assert!(
2031                <$type>::subtract_day_time_opt(
2032                    $date_fn(EPOCH),
2033                    IntervalDayTime::new(i32::MAX, i32::MAX)
2034                )
2035                .is_none(),
2036                "subtract_day_time: max interval"
2037            );
2038            assert!(
2039                <$type>::subtract_day_time_opt(
2040                    $date_fn(EPOCH),
2041                    IntervalDayTime::new(i32::MIN, i32::MIN)
2042                )
2043                .is_none(),
2044                "subtract_day_time: min interval"
2045            );
2046        }};
2047    }
2048
2049    #[test]
2050    fn test_date_day_time_operations() {
2051        test_day_time_ops!(Date32Type, date_to_days);
2052        test_day_time_ops!(Date64Type, date_to_millis);
2053    }
2054
2055    macro_rules! test_month_day_nano_ops {
2056        ($type:ty, $date_fn:expr) => {{
2057            let date = $date_fn(YEAR_2000);
2058            let zero = IntervalMonthDayNano::new(0, 0, 0);
2059
2060            // Normal operations succeed
2061            assert!(
2062                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(1, 30, 0))
2063                    .is_some(),
2064                "add_month_day_nano: +1mo +30d"
2065            );
2066            assert!(
2067                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(-1, -30, 0))
2068                    .is_some(),
2069                "add_month_day_nano: -1mo -30d"
2070            );
2071            assert!(
2072                <$type>::add_month_day_nano_opt(date, zero).is_some(),
2073                "add_month_day_nano: zero interval"
2074            );
2075            assert!(
2076                <$type>::add_month_day_nano_opt(
2077                    date,
2078                    IntervalMonthDayNano::new(2, 10, 123_456_789_000)
2079                )
2080                .is_some(),
2081                "add_month_day_nano: normal"
2082            );
2083            assert!(
2084                <$type>::subtract_month_day_nano_opt(date, IntervalMonthDayNano::new(1, 30, 0))
2085                    .is_some(),
2086                "subtract_month_day_nano: +1mo +30d"
2087            );
2088            assert!(
2089                <$type>::subtract_month_day_nano_opt(date, IntervalMonthDayNano::new(-1, -30, 0))
2090                    .is_some(),
2091                "subtract_month_day_nano: -1mo -30d"
2092            );
2093            assert!(
2094                <$type>::subtract_month_day_nano_opt(date, zero).is_some(),
2095                "subtract_month_day_nano: zero interval"
2096            );
2097            assert!(
2098                <$type>::subtract_month_day_nano_opt(
2099                    date,
2100                    IntervalMonthDayNano::new(2, 10, 123_456_789_000)
2101                )
2102                .is_some(),
2103                "subtract_month_day_nano: normal"
2104            );
2105
2106            // Overflow handling - subtract
2107            assert!(
2108                <$type>::subtract_month_day_nano_opt(
2109                    $date_fn(MIN_VALID_DATE),
2110                    IntervalMonthDayNano::new(0, 1, 0)
2111                )
2112                .is_none(),
2113                "subtract_month_day_nano: overflow days from min"
2114            );
2115            assert!(
2116                <$type>::subtract_month_day_nano_opt(
2117                    $date_fn(MAX_VALID_DATE),
2118                    IntervalMonthDayNano::new(0, -1, 0)
2119                )
2120                .is_none(),
2121                "subtract_month_day_nano: overflow neg days from max"
2122            );
2123
2124            // Overflow handling - add
2125            assert!(
2126                <$type>::add_month_day_nano_opt(
2127                    $date_fn(MAX_VALID_DATE),
2128                    IntervalMonthDayNano::new(0, 1, 0)
2129                )
2130                .is_none(),
2131                "add_month_day_nano: overflow days"
2132            );
2133            assert!(
2134                <$type>::add_month_day_nano_opt(
2135                    $date_fn(MIN_VALID_DATE),
2136                    IntervalMonthDayNano::new(0, -1, 0)
2137                )
2138                .is_none(),
2139                "add_month_day_nano: overflow neg days"
2140            );
2141
2142            // Nanosecond precision works
2143            assert!(
2144                <$type>::add_month_day_nano_opt(date, IntervalMonthDayNano::new(0, 0, 999_999_999))
2145                    .is_some(),
2146                "add_month_day_nano: nanos"
2147            );
2148            assert!(
2149                <$type>::subtract_month_day_nano_opt(
2150                    date,
2151                    IntervalMonthDayNano::new(0, 0, 999_999_999)
2152                )
2153                .is_some(),
2154                "subtract_month_day_nano: nanos"
2155            );
2156            // 1 day in nanos
2157            assert!(
2158                <$type>::add_month_day_nano_opt(
2159                    date,
2160                    IntervalMonthDayNano::new(0, 0, 86_400_000_000_000)
2161                )
2162                .is_some(),
2163                "add_month_day_nano: 1 day nanos"
2164            );
2165            assert!(
2166                <$type>::subtract_month_day_nano_opt(
2167                    date,
2168                    IntervalMonthDayNano::new(0, 0, 86_400_000_000_000)
2169                )
2170                .is_some(),
2171                "subtract_month_day_nano: 1 day nanos"
2172            );
2173        }};
2174    }
2175
2176    #[test]
2177    fn test_date_month_day_nano_operations() {
2178        test_month_day_nano_ops!(Date32Type, date_to_days);
2179        test_month_day_nano_ops!(Date64Type, date_to_millis);
2180    }
2181}