Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
243 changes: 237 additions & 6 deletions datafusion/expr-common/src/interval_arithmetic.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,8 +27,10 @@ use crate::type_coercion::binary::{BinaryTypeCoercer, comparison_coercion};
use arrow::compute::{CastOptions, cast_with_options};
use arrow::datatypes::{
DataType, IntervalDayTime, IntervalMonthDayNano, IntervalUnit,
MAX_DECIMAL32_FOR_EACH_PRECISION, MAX_DECIMAL64_FOR_EACH_PRECISION,
MAX_DECIMAL128_FOR_EACH_PRECISION, MAX_DECIMAL256_FOR_EACH_PRECISION,
MIN_DECIMAL128_FOR_EACH_PRECISION, MIN_DECIMAL256_FOR_EACH_PRECISION, TimeUnit,
MIN_DECIMAL32_FOR_EACH_PRECISION, MIN_DECIMAL64_FOR_EACH_PRECISION,
MIN_DECIMAL128_FOR_EACH_PRECISION, MIN_DECIMAL256_FOR_EACH_PRECISION, TimeUnit, i256,
};
use datafusion_common::rounding::{alter_fp_rounding_mode, next_down, next_up};
use datafusion_common::{
Expand All @@ -37,7 +39,14 @@ use datafusion_common::{
};

macro_rules! get_extreme_value {
($extreme:ident, $DECIMAL128_ARRAY:ident, $DECIMAL256_ARRAY:ident, $value:expr) => {
(
$extreme:ident,
$DECIMAL32_ARRAY:ident,
$DECIMAL64_ARRAY:ident,
$DECIMAL128_ARRAY:ident,
$DECIMAL256_ARRAY:ident,
$value:expr
) => {
match $value {
DataType::UInt8 => ScalarValue::UInt8(Some(u8::$extreme)),
DataType::UInt16 => ScalarValue::UInt16(Some(u16::$extreme)),
Expand Down Expand Up @@ -84,6 +93,16 @@ macro_rules! get_extreme_value {
DataType::Interval(IntervalUnit::MonthDayNano) => {
ScalarValue::IntervalMonthDayNano(Some(IntervalMonthDayNano::$extreme))
}
DataType::Decimal32(precision, scale) => ScalarValue::Decimal32(
Some($DECIMAL32_ARRAY[*precision as usize]),
*precision,
*scale,
),
DataType::Decimal64(precision, scale) => ScalarValue::Decimal64(
Some($DECIMAL64_ARRAY[*precision as usize]),
*precision,
*scale,
),
DataType::Decimal128(precision, scale) => ScalarValue::Decimal128(
Some($DECIMAL128_ARRAY[*precision as usize]),
*precision,
Expand All @@ -100,7 +119,15 @@ macro_rules! get_extreme_value {
}

macro_rules! value_transition {
($bound:ident, $direction:expr, $value:expr) => {
(
$bound:ident,
$direction:expr,
$DECIMAL32_ARRAY:ident,
$DECIMAL64_ARRAY:ident,
$DECIMAL128_ARRAY:ident,
$DECIMAL256_ARRAY:ident,
$value:expr
) => {
match $value {
UInt8(Some(value)) if value == u8::$bound => UInt8(None),
UInt16(Some(value)) if value == u16::$bound => UInt16(None),
Expand Down Expand Up @@ -147,6 +174,28 @@ macro_rules! value_transition {
{
IntervalMonthDayNano(None)
}
// Decimals saturate at the extreme value of their precision, not of
// the underlying integer type:
Decimal32(Some(value), precision, scale)
if value == $DECIMAL32_ARRAY[precision as usize] =>
{
Decimal32(None, precision, scale)
}
Decimal64(Some(value), precision, scale)
if value == $DECIMAL64_ARRAY[precision as usize] =>
{
Decimal64(None, precision, scale)
}
Decimal128(Some(value), precision, scale)
if value == $DECIMAL128_ARRAY[precision as usize] =>
{
Decimal128(None, precision, scale)
}
Decimal256(Some(value), precision, scale)
if value == $DECIMAL256_ARRAY[precision as usize] =>
{
Decimal256(None, precision, scale)
}
_ => next_value_helper::<$direction>($value),
}
};
Expand Down Expand Up @@ -1203,6 +1252,8 @@ fn handle_overflow<const UPPER: bool>(
(true, false) => {
get_extreme_value!(
MIN,
MIN_DECIMAL32_FOR_EACH_PRECISION,
MIN_DECIMAL64_FOR_EACH_PRECISION,
MIN_DECIMAL128_FOR_EACH_PRECISION,
MIN_DECIMAL256_FOR_EACH_PRECISION,
dt
Expand All @@ -1211,6 +1262,8 @@ fn handle_overflow<const UPPER: bool>(
(false, true) => {
get_extreme_value!(
MAX,
MAX_DECIMAL32_FOR_EACH_PRECISION,
MAX_DECIMAL64_FOR_EACH_PRECISION,
MAX_DECIMAL128_FOR_EACH_PRECISION,
MAX_DECIMAL256_FOR_EACH_PRECISION,
dt
Expand All @@ -1223,14 +1276,30 @@ fn handle_overflow<const UPPER: bool>(
// used without caution.
fn next_value(value: ScalarValue) -> ScalarValue {
use ScalarValue::*;
value_transition!(MAX, true, value)
value_transition!(
MAX,
true,
MAX_DECIMAL32_FOR_EACH_PRECISION,
MAX_DECIMAL64_FOR_EACH_PRECISION,
MAX_DECIMAL128_FOR_EACH_PRECISION,
MAX_DECIMAL256_FOR_EACH_PRECISION,
value
)
}

// This function should remain private since it may corrupt the an interval if
// used without caution.
fn prev_value(value: ScalarValue) -> ScalarValue {
use ScalarValue::*;
value_transition!(MIN, false, value)
value_transition!(
MIN,
false,
MIN_DECIMAL32_FOR_EACH_PRECISION,
MIN_DECIMAL64_FOR_EACH_PRECISION,
MIN_DECIMAL128_FOR_EACH_PRECISION,
MIN_DECIMAL256_FOR_EACH_PRECISION,
value
)
}

trait OneTrait: Sized + std::ops::Add + std::ops::Sub {
Expand All @@ -1241,6 +1310,12 @@ macro_rules! impl_OneTrait{
}
impl_OneTrait! {u8, u16, u32, u64, i8, i16, i32, i64, i128}

impl OneTrait for i256 {
fn one() -> Self {
i256::ONE
}
}

impl OneTrait for IntervalDayTime {
fn one() -> Self {
IntervalDayTime {
Expand Down Expand Up @@ -1295,6 +1370,23 @@ fn next_value_helper<const INC: bool>(value: ScalarValue) -> ScalarValue {
UInt16(Some(val)) => UInt16(Some(increment_decrement::<INC, u16>(val))),
UInt32(Some(val)) => UInt32(Some(increment_decrement::<INC, u32>(val))),
UInt64(Some(val)) => UInt64(Some(increment_decrement::<INC, u64>(val))),
// Decimals step by one unit in the last place, keeping precision and scale:
Decimal32(Some(val), precision, scale) => {
Decimal32(Some(increment_decrement::<INC, i32>(val)), precision, scale)
}
Decimal64(Some(val), precision, scale) => {
Decimal64(Some(increment_decrement::<INC, i64>(val)), precision, scale)
}
Decimal128(Some(val), precision, scale) => Decimal128(
Some(increment_decrement::<INC, i128>(val)),
precision,
scale,
),
Decimal256(Some(val), precision, scale) => Decimal256(
Some(increment_decrement::<INC, i256>(val)),
precision,
scale,
),
DurationSecond(Some(val)) => {
DurationSecond(Some(increment_decrement::<INC, i64>(val)))
}
Expand Down Expand Up @@ -2274,7 +2366,10 @@ mod tests {
};

use crate::interval_arithmetic::NullableInterval;
use arrow::datatypes::DataType;
use arrow::datatypes::{
DataType, MAX_DECIMAL128_FOR_EACH_PRECISION, MIN_DECIMAL128_FOR_EACH_PRECISION,
i256,
};
use datafusion_common::rounding::{next_down, next_up};
use datafusion_common::{Result, ScalarValue};

Expand Down Expand Up @@ -2365,6 +2460,124 @@ mod tests {
Ok(())
}

#[test]
fn test_next_prev_value_decimal() {
use ScalarValue::*;

// Decimals step by one unit in the last place, keeping precision and scale:
let cases = vec![
(Decimal32(Some(2400), 9, 2), Decimal32(Some(2401), 9, 2)),
(Decimal64(Some(2400), 18, 2), Decimal64(Some(2401), 18, 2)),
(Decimal128(Some(2400), 15, 2), Decimal128(Some(2401), 15, 2)),
(
Decimal256(Some(i256::from(2400)), 40, 2),
Decimal256(Some(i256::from(2401)), 40, 2),
),
(Decimal128(Some(-1), 3, 2), Decimal128(Some(0), 3, 2)),
];
for (value, next) in cases {
assert_eq!(next_value(value.clone()), next);
assert_eq!(prev_value(next), value);
}

// Values at the extremes of their precision (e.g. -9.99 and 9.99 for
// precision 3, scale 2) become unbounded:
let d128_max = MAX_DECIMAL128_FOR_EACH_PRECISION[38];
let d128_min = MIN_DECIMAL128_FOR_EACH_PRECISION[38];
let cases = vec![
(
Decimal32(Some(-999), 3, 2),
Decimal32(Some(999), 3, 2),
Decimal32(Some(-998), 3, 2),
Decimal32(Some(998), 3, 2),
Decimal32(None, 3, 2),
),
(
Decimal64(Some(-999), 3, 2),
Decimal64(Some(999), 3, 2),
Decimal64(Some(-998), 3, 2),
Decimal64(Some(998), 3, 2),
Decimal64(None, 3, 2),
),
(
Decimal128(Some(-999), 3, 2),
Decimal128(Some(999), 3, 2),
Decimal128(Some(-998), 3, 2),
Decimal128(Some(998), 3, 2),
Decimal128(None, 3, 2),
),
(
Decimal256(Some(i256::from(-999)), 3, 2),
Decimal256(Some(i256::from(999)), 3, 2),
Decimal256(Some(i256::from(-998)), 3, 2),
Decimal256(Some(i256::from(998)), 3, 2),
Decimal256(None, 3, 2),
),
(
Decimal128(Some(d128_min), 38, 0),
Decimal128(Some(d128_max), 38, 0),
Decimal128(Some(d128_min + 1), 38, 0),
Decimal128(Some(d128_max - 1), 38, 0),
Decimal128(None, 38, 0),
),
];
for (min, max, after_min, before_max, inf) in cases {
assert_eq!(next_value(max.clone()), inf);
assert_eq!(prev_value(max), before_max);

assert_eq!(prev_value(min.clone()), inf);
assert_eq!(next_value(min), after_min);

assert_eq!(next_value(inf.clone()), inf);
assert_eq!(prev_value(inf.clone()), inf);
}
}

#[test]
fn test_satisfy_strict_comparison_decimal() -> Result<()> {
use ScalarValue::*;

// x in [0.00, 50.00], constraint 24.00 > x tightens x to [0.00, 23.99]:
let x =
Interval::try_new(Decimal128(Some(0), 15, 2), Decimal128(Some(5000), 15, 2))?;
let literal = Interval::try_new(
Decimal128(Some(2400), 15, 2),
Decimal128(Some(2400), 15, 2),
)?;
let (_, x_lt) = satisfy_greater(&literal, &x, true)?.unwrap();
assert_eq!(
x_lt,
Interval::try_new(Decimal128(Some(0), 15, 2), Decimal128(Some(2399), 15, 2))?
);

// Constraint x > 24.00 tightens x to [24.01, 50.00]:
let (x_gt, _) = satisfy_greater(&x, &literal, true)?.unwrap();
assert_eq!(
x_gt,
Interval::try_new(
Decimal128(Some(2401), 15, 2),
Decimal128(Some(5000), 15, 2),
)?
);

// Non-strict comparisons keep the closed bound:
let (_, x_le) = satisfy_greater(&literal, &x, false)?.unwrap();
assert_eq!(
x_le,
Interval::try_new(Decimal128(Some(0), 15, 2), Decimal128(Some(2400), 15, 2))?
);

// A strict bound at the precision's extreme conservatively becomes
// unbounded instead of stepping past it (x > 9.99 for Decimal128(3, 2)):
let x = Interval::make_unbounded(&DataType::Decimal128(3, 2))?;
let max =
Interval::try_new(Decimal128(Some(999), 3, 2), Decimal128(Some(999), 3, 2))?;
let (x_gt, _) = satisfy_greater(&x, &max, true)?.unwrap();
assert_eq!(x_gt, Interval::make_unbounded(&DataType::Decimal128(3, 2))?);

Ok(())
}

#[test]
fn test_new_interval() -> Result<()> {
use ScalarValue::*;
Expand Down Expand Up @@ -4064,6 +4277,24 @@ mod tests {
ScalarValue::Decimal128(Some(99999999999999999999999999999999999999), 38, 35)
);

// Test Decimal32/Decimal64 underflow handling:
let op = Operator::Minus;
let dt = DataType::Decimal32(3, 2);
let lhs = ScalarValue::Decimal32(Some(-999), 3, 2);
let rhs = ScalarValue::Decimal32(Some(1), 3, 2);
let result = handle_overflow::<true>(&dt, op, &lhs, &rhs);
assert_eq!(result, ScalarValue::Decimal32(Some(-999), 3, 2));
let result = handle_overflow::<false>(&dt, op, &lhs, &rhs);
assert_eq!(result, ScalarValue::Decimal32(None, 3, 2));

let dt = DataType::Decimal64(3, 2);
let lhs = ScalarValue::Decimal64(Some(-999), 3, 2);
let rhs = ScalarValue::Decimal64(Some(1), 3, 2);
let result = handle_overflow::<true>(&dt, op, &lhs, &rhs);
assert_eq!(result, ScalarValue::Decimal64(Some(-999), 3, 2));
let result = handle_overflow::<false>(&dt, op, &lhs, &rhs);
assert_eq!(result, ScalarValue::Decimal64(None, 3, 2));

Ok(())
}

Expand Down
Loading