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
89 changes: 60 additions & 29 deletions datafusion/functions-aggregate-common/src/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ use datafusion_common::hash_utils::RandomState;
use datafusion_common::utils::SingleRowListArrayBuilder;
use datafusion_common::utils::memory::estimate_memory_size;
use datafusion_common::{
HashSet, Result, ScalarValue, exec_err, internal_datafusion_err,
HashSet, Result, ScalarValue, exec_datafusion_err, exec_err, internal_datafusion_err,
};
use datafusion_expr_common::accumulator::Accumulator;
use datafusion_physical_expr_common::sort_expr::{LexOrdering, PhysicalSortExpr};
Expand Down Expand Up @@ -97,10 +97,9 @@ impl<T: ArrowNativeTypeOp> Eq for Hashable<T> {}
/// For example, the precision is 3, the max of value is `999` and the min
/// value is `-999`
pub struct DecimalAverager<T: DecimalType> {
/// scale factor for sum values (10^sum_scale)
sum_mul: T::Native,
/// scale factor for target (10^target_scale)
target_mul: T::Native,
/// factor that rescales sum values to the target scale
/// (`10^(target_scale - sum_scale)`)
scale_mul: T::Native,
/// the output precision
target_precision: u8,
/// the output scale
Expand All @@ -120,31 +119,29 @@ impl<T: DecimalType> DecimalAverager<T> {
target_precision: u8,
target_scale: i8,
) -> Result<Self> {
let sum_mul = T::Native::from_usize(10_usize)
.map(|b| b.pow_wrapping(sum_scale as u32))
.ok_or_else(|| {
internal_datafusion_err!("Failed to compute sum_mul in DecimalAverager")
})?;

let target_mul = T::Native::from_usize(10_usize)
.map(|b| b.pow_wrapping(target_scale as u32))
.ok_or_else(|| {
internal_datafusion_err!(
"Failed to compute target_mul in DecimalAverager"
)
})?;

if target_mul >= sum_mul {
Ok(Self {
sum_mul,
target_mul,
target_precision,
target_scale,
})
} else {
// Only the ratio `10^target_scale / 10^sum_scale` is needed, and the
// scale difference is non-negative even when both scales are negative
// (e.g. `Decimal128(10, -2)`).
let scale_diff = i16::from(target_scale) - i16::from(sum_scale);
if scale_diff < 0 {
// can't convert the lit decimal to the returned data type
exec_err!("Arithmetic Overflow in AvgAccumulator")
return exec_err!("Arithmetic Overflow in AvgAccumulator");
}

let ten = T::Native::from_usize(10_usize).ok_or_else(|| {
internal_datafusion_err!("Failed to compute scale_mul in DecimalAverager")
})?;
let scale_mul = ten.pow_checked(scale_diff as u32).map_err(|e| {
exec_datafusion_err!(
"Arithmetic Overflow in AvgAccumulator: cannot rescale from scale {sum_scale} to {target_scale}: {e}"
)
})?;

Ok(Self {
scale_mul,
target_precision,
target_scale,
})
}

/// Returns the `sum`/`count` as a i128/i256 Decimal128/Decimal256 with
Expand All @@ -155,7 +152,7 @@ impl<T: DecimalType> DecimalAverager<T> {
/// * count: total count, stored as a i128/i256 (*NOT* a Decimal128/Decimal256 value)
#[inline(always)]
pub fn avg(&self, sum: T::Native, count: T::Native) -> Result<T::Native> {
if let Ok(value) = sum.mul_checked(self.target_mul.div_wrapping(self.sum_mul)) {
if let Ok(value) = sum.mul_checked(self.scale_mul) {
let new_value = value.div_wrapping(count);

let validate = T::validate_decimal_precision(
Expand Down Expand Up @@ -264,3 +261,37 @@ impl<T: ArrowPrimitiveType> GenericDistinctBuffer<T> {
estimate_memory_size::<T::Native>(num_elements, fixed_size).unwrap()
}
}

#[cfg(test)]
mod tests {
use super::*;
use arrow::datatypes::Decimal128Type;

#[test]
fn decimal_averager_negative_scale() {
// Decimal128(10, -2) averaged into Decimal128(14, 2): sums are
// rescaled by 10^(2 - -2)
let averager = DecimalAverager::<Decimal128Type>::try_new(-2, 14, 2).unwrap();
// 100 + 200 at scale -2 is 30000; their average 15000 at scale 2 is 1500000
assert_eq!(averager.avg(300, 2).unwrap(), 1_500_000);
}

#[test]
fn decimal_averager_unrepresentable_multiplier_keeps_cause() {
// 10^255 does not fit in an i128; the error keeps the arrow overflow
// detail instead of dropping it
let err = DecimalAverager::<Decimal128Type>::try_new(-128, 38, 127)
.err()
.expect("10^255 overflows i128");
let msg = err.to_string();
assert!(
msg.contains("Arithmetic Overflow in AvgAccumulator"),
"unexpected error: {msg}"
);
assert!(
msg.contains("cannot rescale from scale -128 to 127"),
"unexpected error: {msg}"
);
assert!(msg.contains("10 ^ 255"), "unexpected error: {msg}");
}
}
26 changes: 26 additions & 0 deletions datafusion/sqllogictest/test_files/aggregate.slt
Original file line number Diff line number Diff line change
Expand Up @@ -3479,6 +3479,32 @@ select avg(c1), arrow_typeof(avg(c1)) from t;
statement ok
drop table t;

# avg_decimal_negative_scale
statement ok
create table t (k int, c1 decimal(10, -2)) as values
(1, arrow_cast(10000, 'Decimal128(10, -2)')),
(1, arrow_cast(20000, 'Decimal128(10, -2)')),
(2, arrow_cast(20000, 'Decimal128(10, -2)'));

query RT
select avg(c1), arrow_typeof(avg(c1)) from t where k = 1;
----
15000.00 Decimal128(14, 2)

query IR rowsort
select k, avg(c1) from t group by k;
----
1 15000.00
2 20000.00

query R
select avg(distinct c1) from t;
----
15000.00

statement ok
drop table t;

# avg_decimal_with_nulls
statement ok
create table t (c1 decimal(10, 0)) as values (1), (NULL), (3), (4), (5);
Expand Down