use std::fmt::{Display, Formatter};
use num_bigint::{BigInt, Sign};
use num_integer::Integer;
use num_traits::{Signed, Zero};
use thiserror::Error;
const NULL_TARGET_MESSAGE: &str = "Cannot aggregate on null";
const ITERABLE_NULL_SUM_MESSAGE: &str = "Cannot aggregate on iterable containing nulls";
const ARRAY_NULL_MESSAGE: &str = "Cannot aggregate on array containing nulls";
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub struct BigDecimalValue {
unscaled_value: BigInt,
scale: i32,
}
impl BigDecimalValue {
#[must_use]
pub fn from_unscaled(unscaled_value: BigInt, scale: i32) -> Self {
Self {
unscaled_value,
scale,
}
}
pub fn parse(value: &str) -> Result<Self, AggregateError> {
parse_decimal(value)
}
#[must_use]
pub fn unscaled_value(&self) -> &BigInt {
&self.unscaled_value
}
#[must_use]
pub const fn scale(&self) -> i32 {
self.scale
}
#[must_use]
pub fn precision(&self) -> usize {
if self.unscaled_value.is_zero() {
1
} else {
self.unscaled_value.abs().to_str_radix(10).len()
}
}
#[must_use]
pub fn to_plain_string(&self) -> String {
let negative = self.unscaled_value.sign() == Sign::Minus;
let digits = self.unscaled_value.abs().to_str_radix(10);
let mut result = if self.scale <= 0 {
let zero_count = i64::from(self.scale)
.checked_neg()
.and_then(|count| usize::try_from(count).ok())
.expect("Java scale magnitude must fit this platform");
format!("{digits}{}", "0".repeat(zero_count))
} else {
let scale = usize::try_from(self.scale).expect("positive scale");
if scale >= digits.len() {
format!("0.{}{}", "0".repeat(scale - digits.len()), digits)
} else {
let split = digits.len() - scale;
format!("{}.{}", &digits[..split], &digits[split..])
}
};
if negative {
result.insert(0, '-');
}
result
}
fn zero() -> Self {
Self::from_unscaled(BigInt::ZERO, 0)
}
fn from_i64(value: i64) -> Self {
Self::from_unscaled(BigInt::from(value), 0)
}
fn from_big_integer(value: &BigInt) -> Self {
Self::from_unscaled(value.clone(), 0)
}
fn from_f64(value: f64) -> Result<Self, AggregateError> {
if !value.is_finite() {
return Err(AggregateError::NumberFormat {
value: double_string(value),
});
}
parse_decimal(&double_string(value))
}
pub(crate) fn from_f64_exact(value: f64) -> Option<Self> {
if !value.is_finite() {
return None;
}
if value == 0.0 {
return Some(Self::zero());
}
let bits = value.to_bits();
let negative = bits >> 63 != 0;
let exponent_bits = ((bits >> 52) & 0x7ff) as i32;
let fraction = bits & ((1_u64 << 52) - 1);
let (mut significand, binary_exponent) = if exponent_bits == 0 {
(fraction, -1074)
} else {
((1_u64 << 52) | fraction, exponent_bits - 1023 - 52)
};
let mut unscaled = if binary_exponent >= 0 {
BigInt::from(significand)
<< usize::try_from(binary_exponent).expect("non-negative binary exponent")
} else {
let mut scale = -binary_exponent;
while scale > 0 && significand & 1 == 0 {
significand >>= 1;
scale -= 1;
}
let scale_u32 = u32::try_from(scale).expect("IEEE-754 scale fits u32");
let mut value = BigInt::from(significand) * BigInt::from(5_u8).pow(scale_u32);
if negative {
value = -value;
}
return Some(Self::from_unscaled(value, scale));
};
if negative {
unscaled = -unscaled;
}
Some(Self::from_unscaled(unscaled, 0))
}
pub(crate) fn add_java(&self, other: &Self) -> Self {
let result_scale = self.scale.max(other.scale);
let left = rescale_unscaled(&self.unscaled_value, self.scale, result_scale);
let right = rescale_unscaled(&other.unscaled_value, other.scale, result_scale);
Self::from_unscaled(left + right, result_scale)
}
pub(crate) fn with_scale_half_even(&self, target_scale: i32) -> Self {
if target_scale >= self.scale {
return Self::from_unscaled(
rescale_unscaled(&self.unscaled_value, self.scale, target_scale),
target_scale,
);
}
let difference = u32::try_from(i64::from(self.scale) - i64::from(target_scale))
.expect("positive scale difference fits u32");
let divisor = BigInt::from(10_u8).pow(difference);
let quotient = &self.unscaled_value / &divisor;
let remainder = &self.unscaled_value % &divisor;
let doubled = remainder.abs() * 2_u8;
let increment = doubled > divisor || (doubled == divisor && quotient.is_odd());
let rounded = if increment {
quotient
+ if self.unscaled_value.sign() == Sign::Minus {
-BigInt::from(1_u8)
} else {
BigInt::from(1_u8)
}
} else {
quotient
};
Self::from_unscaled(rounded, target_scale)
}
pub(crate) fn subtract_java(&self, other: &Self) -> Self {
let result_scale = self.scale.max(other.scale);
let left = rescale_unscaled(&self.unscaled_value, self.scale, result_scale);
let right = rescale_unscaled(&other.unscaled_value, other.scale, result_scale);
Self::from_unscaled(left - right, result_scale)
}
pub(crate) fn compare_java(&self, other: &Self) -> std::cmp::Ordering {
let comparison_scale = self.scale.max(other.scale);
let left = rescale_unscaled(&self.unscaled_value, self.scale, comparison_scale);
let right = rescale_unscaled(&other.unscaled_value, other.scale, comparison_scale);
left.cmp(&right)
}
pub(crate) fn multiply_java(&self, other: &Self) -> Result<Self, BigDecimalArithmeticError> {
let scale = checked_scale(i64::from(self.scale) + i64::from(other.scale))
.map_err(|_| BigDecimalArithmeticError::ScaleOverflow)?;
Ok(Self::from_unscaled(
&self.unscaled_value * &other.unscaled_value,
scale,
))
}
pub(crate) fn divide_java(&self, divisor: &Self) -> Result<Self, BigDecimalArithmeticError> {
self.divide_exact(divisor).map_err(|error| match error {
DivisionError::ByZero => BigDecimalArithmeticError::DivisionByZero,
DivisionError::NonTerminating => BigDecimalArithmeticError::NonTerminating,
DivisionError::ScaleOverflow => BigDecimalArithmeticError::ScaleOverflow,
})
}
pub(crate) fn divide_java_half_up(
&self,
divisor: &Self,
scale: i32,
) -> Result<Self, BigDecimalArithmeticError> {
if divisor.unscaled_value.is_zero() {
return Err(BigDecimalArithmeticError::DivisionByZero);
}
let exponent = i64::from(divisor.scale) + i64::from(scale) - i64::from(self.scale);
let mut numerator = self.unscaled_value.clone();
let mut denominator = divisor.unscaled_value.clone();
if exponent >= 0 {
numerator *= BigInt::from(10_u8).pow(
u32::try_from(exponent).map_err(|_| BigDecimalArithmeticError::ScaleOverflow)?,
);
} else {
denominator *= BigInt::from(10_u8).pow(
u32::try_from(-exponent).map_err(|_| BigDecimalArithmeticError::ScaleOverflow)?,
);
}
let quotient = &numerator / &denominator;
let remainder = &numerator % &denominator;
let rounded = if remainder.abs() * 2_u8 >= denominator.abs() {
if numerator.sign() == denominator.sign() {
quotient + 1_u8
} else {
quotient - 1_u8
}
} else {
quotient
};
Ok(Self::from_unscaled(rounded, scale))
}
pub(crate) fn remainder_java(&self, divisor: &Self) -> Result<Self, BigDecimalArithmeticError> {
if divisor.unscaled_value.is_zero() {
return Err(BigDecimalArithmeticError::DivisionByZero);
}
let result_scale = self.scale.max(divisor.scale);
let left = rescale_unscaled(&self.unscaled_value, self.scale, result_scale);
let right = rescale_unscaled(&divisor.unscaled_value, divisor.scale, result_scale);
Ok(Self::from_unscaled(left % right, result_scale))
}
fn divide_exact(&self, divisor: &Self) -> Result<Self, DivisionError> {
if divisor.unscaled_value.is_zero() {
return Err(DivisionError::ByZero);
}
let gcd = self.unscaled_value.abs().gcd(&divisor.unscaled_value.abs());
let mut numerator = &self.unscaled_value / &gcd;
let mut denominator = &divisor.unscaled_value / &gcd;
if denominator.sign() == Sign::Minus {
numerator = -numerator;
denominator = -denominator;
}
let mut twos = 0_i64;
while (&denominator % 2_u8).is_zero() {
denominator /= 2_u8;
twos += 1;
}
let mut fives = 0_i64;
while (&denominator % 5_u8).is_zero() {
denominator /= 5_u8;
fives += 1;
}
if denominator != BigInt::from(1_u8) {
return Err(DivisionError::NonTerminating);
}
let decimal_places = twos.max(fives);
if twos < decimal_places {
let exponent =
u32::try_from(decimal_places - twos).expect("factor count fits Java scale");
numerator *= BigInt::from(2_u8).pow(exponent);
}
if fives < decimal_places {
let exponent =
u32::try_from(decimal_places - fives).expect("factor count fits Java scale");
numerator *= BigInt::from(5_u8).pow(exponent);
}
let preferred_scale = i64::from(self.scale) - i64::from(divisor.scale);
let scale = checked_scale(preferred_scale + decimal_places)
.map_err(|_| DivisionError::ScaleOverflow)?;
Ok(Self::from_unscaled(numerator, scale))
}
fn divide_half_up_by_positive_integer(&self, divisor: i64, scale: i32) -> Self {
let exponent = i64::from(scale) - i64::from(self.scale);
let mut numerator = self.unscaled_value.clone();
let exponent = u32::try_from(exponent)
.expect("AggregateUtils fallback scale is never below total scale");
numerator *= BigInt::from(10_u8).pow(exponent);
let denominator = BigInt::from(divisor);
let quotient = &numerator / &denominator;
let remainder = &numerator % &denominator;
let rounded = if remainder.abs() * 2_u8 >= denominator.abs() {
if numerator.sign() == Sign::Minus {
quotient - 1_u8
} else {
quotient + 1_u8
}
} else {
quotient
};
Self::from_unscaled(rounded, scale)
}
fn to_display_string(&self) -> String {
let negative = self.unscaled_value.sign() == Sign::Minus;
let digits = self.unscaled_value.abs().to_str_radix(10);
let adjusted_exponent = i64::try_from(self.precision()).expect("precision fits i64")
- i64::from(self.scale)
- 1;
if self.scale >= 0 && adjusted_exponent >= -6 {
return self.to_plain_string();
}
let mut result = String::new();
if negative {
result.push('-');
}
result.push_str(&digits[..1]);
if digits.len() > 1 {
result.push('.');
result.push_str(&digits[1..]);
}
result.push('E');
if adjusted_exponent >= 0 {
result.push('+');
}
result.push_str(&adjusted_exponent.to_string());
result
}
}
impl Display for BigDecimalValue {
fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
formatter.write_str(&self.to_display_string())
}
}
#[derive(Clone, Debug, PartialEq)]
pub enum NumberValue {
BigDecimal(BigDecimalValue),
BigInteger(BigInt),
Byte(i8),
Short(i16),
Integer(i32),
Long(i64),
Float(f32),
Double(f64),
Other {
class_name: String,
double_value: f64,
},
}
#[derive(Clone, Debug, PartialEq)]
pub enum AggregateObjectValue {
Null,
Number(NumberValue),
Other(String),
}
pub trait NumberIterableValue {
fn iter_java_numbers(&self) -> Box<dyn Iterator<Item = Option<&NumberValue>> + '_>;
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct NumberListValue {
values: Vec<Option<NumberValue>>,
}
impl NumberListValue {
#[must_use]
pub fn new(values: Vec<Option<NumberValue>>) -> Self {
Self { values }
}
#[must_use]
pub fn as_slice(&self) -> &[Option<NumberValue>] {
&self.values
}
}
impl NumberIterableValue for NumberListValue {
fn iter_java_numbers(&self) -> Box<dyn Iterator<Item = Option<&NumberValue>> + '_> {
Box::new(self.values.iter().map(Option::as_ref))
}
}
#[derive(Clone, Debug, Eq, Error, PartialEq)]
pub enum AggregateError {
#[error("{message}")]
IllegalArgument {
message: &'static str,
},
#[error("class {actual_class} cannot be cast to class java.lang.Number")]
ClassCast {
actual_class: String,
},
#[error("Character array is missing \"e\" notation exponential mark for value {value}")]
NumberFormat {
value: String,
},
#[error("{message}")]
Arithmetic {
message: String,
},
}
pub struct AggregateUtils;
impl AggregateUtils {
pub fn sum_iterable(
target: Option<&dyn NumberIterableValue>,
) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_iterable(target, ITERABLE_NULL_SUM_MESSAGE, false)
}
pub fn sum_objects(
target: Option<&[AggregateObjectValue]>,
) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_objects(target, false)
}
pub fn sum_bytes(target: Option<&[i8]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Bytes), false)
}
pub fn sum_shorts(target: Option<&[i16]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Shorts), false)
}
pub fn sum_ints(target: Option<&[i32]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Ints), false)
}
pub fn sum_longs(target: Option<&[i64]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Longs), false)
}
pub fn sum_floats(target: Option<&[f32]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Floats), false)
}
pub fn sum_doubles(target: Option<&[f64]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Doubles), false)
}
pub fn avg_iterable(
target: Option<&dyn NumberIterableValue>,
) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_iterable(target, ARRAY_NULL_MESSAGE, true)
}
pub fn avg_objects(
target: Option<&[AggregateObjectValue]>,
) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_objects(target, true)
}
pub fn avg_bytes(target: Option<&[i8]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Bytes), true)
}
pub fn avg_shorts(target: Option<&[i16]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Shorts), true)
}
pub fn avg_ints(target: Option<&[i32]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Ints), true)
}
pub fn avg_longs(target: Option<&[i64]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Longs), true)
}
pub fn avg_floats(target: Option<&[f32]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Floats), true)
}
pub fn avg_doubles(target: Option<&[f64]>) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_primitive(target.map(PrimitiveArray::Doubles), true)
}
pub(crate) fn sum_numbers(
target: Option<&[Option<NumberValue>]>,
) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_number_array(target, false)
}
pub(crate) fn avg_numbers(
target: Option<&[Option<NumberValue>]>,
) -> Result<Option<BigDecimalValue>, AggregateError> {
aggregate_number_array(target, true)
}
}
fn aggregate_iterable(
target: Option<&dyn NumberIterableValue>,
null_message: &'static str,
average: bool,
) -> Result<Option<BigDecimalValue>, AggregateError> {
let target = target.ok_or(AggregateError::IllegalArgument {
message: NULL_TARGET_MESSAGE,
})?;
for element in target.iter_java_numbers() {
if element.is_none() {
return Err(AggregateError::IllegalArgument {
message: null_message,
});
}
}
let mut total = BigDecimalValue::zero();
let mut size = 0_usize;
for element in target.iter_java_numbers() {
let number = element.ok_or(AggregateError::IllegalArgument {
message: null_message,
})?;
total = total.add_java(&to_big_decimal(number)?);
size += 1;
}
finish_aggregate(total, size, average)
}
fn aggregate_objects(
target: Option<&[AggregateObjectValue]>,
average: bool,
) -> Result<Option<BigDecimalValue>, AggregateError> {
let target = target.ok_or(AggregateError::IllegalArgument {
message: NULL_TARGET_MESSAGE,
})?;
if target
.iter()
.any(|element| matches!(element, AggregateObjectValue::Null))
{
return Err(AggregateError::IllegalArgument {
message: ARRAY_NULL_MESSAGE,
});
}
let mut total = BigDecimalValue::zero();
for element in target {
total = total.add_java(&to_big_decimal(object_number(element)?)?);
}
finish_aggregate(total, target.len(), average)
}
fn aggregate_number_array(
target: Option<&[Option<NumberValue>]>,
average: bool,
) -> Result<Option<BigDecimalValue>, AggregateError> {
let target = target.ok_or(AggregateError::IllegalArgument {
message: NULL_TARGET_MESSAGE,
})?;
let mut numbers = Vec::with_capacity(target.len());
for number in target {
match number {
Some(number) => numbers.push(number),
None => {
return Err(AggregateError::IllegalArgument {
message: ARRAY_NULL_MESSAGE,
});
}
}
}
let mut total = BigDecimalValue::zero();
for number in numbers {
total = total.add_java(&to_big_decimal(number)?);
}
finish_aggregate(total, target.len(), average)
}
enum PrimitiveArray<'a> {
Bytes(&'a [i8]),
Shorts(&'a [i16]),
Ints(&'a [i32]),
Longs(&'a [i64]),
Floats(&'a [f32]),
Doubles(&'a [f64]),
}
fn aggregate_primitive(
target: Option<PrimitiveArray<'_>>,
average: bool,
) -> Result<Option<BigDecimalValue>, AggregateError> {
let target = target.ok_or(AggregateError::IllegalArgument {
message: NULL_TARGET_MESSAGE,
})?;
let mut total = BigDecimalValue::zero();
let size = match target {
PrimitiveArray::Bytes(values) => {
for value in values {
total = total.add_java(&BigDecimalValue::from_i64(i64::from(*value)));
}
values.len()
}
PrimitiveArray::Shorts(values) => {
for value in values {
total = total.add_java(&BigDecimalValue::from_i64(i64::from(*value)));
}
values.len()
}
PrimitiveArray::Ints(values) => {
for value in values {
total = total.add_java(&BigDecimalValue::from_i64(i64::from(*value)));
}
values.len()
}
PrimitiveArray::Longs(values) => {
for value in values {
total = total.add_java(&BigDecimalValue::from_i64(*value));
}
values.len()
}
PrimitiveArray::Floats(values) => {
for value in values {
total = total.add_java(&BigDecimalValue::from_f64(f64::from(*value))?);
}
values.len()
}
PrimitiveArray::Doubles(values) => {
for value in values {
total = total.add_java(&BigDecimalValue::from_f64(*value)?);
}
values.len()
}
};
finish_aggregate(total, size, average)
}
fn finish_aggregate(
total: BigDecimalValue,
size: usize,
average: bool,
) -> Result<Option<BigDecimalValue>, AggregateError> {
if size == 0 {
return Ok(None);
}
if !average {
return Ok(Some(total));
}
let divisor_value = i64::try_from(size).map_err(|_| AggregateError::Arithmetic {
message: "BigDecimal divisor exceeds Java long range".to_owned(),
})?;
let divisor = BigDecimalValue::from_i64(divisor_value);
match total.divide_exact(&divisor) {
Ok(value) => Ok(Some(value)),
Err(DivisionError::NonTerminating) => {
let scale = total.scale.max(10);
Ok(Some(
total.divide_half_up_by_positive_integer(divisor_value, scale),
))
}
Err(error) => Err(AggregateError::from(error)),
}
}
fn to_big_decimal(number: &NumberValue) -> Result<BigDecimalValue, AggregateError> {
match number {
NumberValue::BigDecimal(value) => Ok(value.clone()),
NumberValue::BigInteger(value) => Ok(BigDecimalValue::from_big_integer(value)),
NumberValue::Byte(value) => Ok(BigDecimalValue::from_i64(i64::from(*value))),
NumberValue::Short(value) => Ok(BigDecimalValue::from_i64(i64::from(*value))),
NumberValue::Integer(value) => Ok(BigDecimalValue::from_i64(i64::from(*value))),
NumberValue::Long(value) => Ok(BigDecimalValue::from_i64(*value)),
NumberValue::Float(value) => BigDecimalValue::from_f64(f64::from(*value)),
NumberValue::Double(value) => BigDecimalValue::from_f64(*value),
NumberValue::Other { double_value, .. } => BigDecimalValue::from_f64(*double_value),
}
}
fn object_number(element: &AggregateObjectValue) -> Result<&NumberValue, AggregateError> {
match element {
AggregateObjectValue::Number(number) => Ok(number),
AggregateObjectValue::Other(actual_class) => Err(AggregateError::ClassCast {
actual_class: actual_class.clone(),
}),
AggregateObjectValue::Null => Err(AggregateError::IllegalArgument {
message: ARRAY_NULL_MESSAGE,
}),
}
}
fn rescale_unscaled(unscaled: &BigInt, source_scale: i32, target_scale: i32) -> BigInt {
if unscaled.is_zero() {
return BigInt::ZERO;
}
let exponent = i64::from(target_scale) - i64::from(source_scale);
let exponent = u32::try_from(exponent).expect("target scale is the maximum source scale");
unscaled * BigInt::from(10_u8).pow(exponent)
}
fn checked_scale(scale: i64) -> Result<i32, AggregateError> {
i32::try_from(scale).map_err(|_| AggregateError::Arithmetic {
message: "Underflow".to_owned(),
})
}
fn parse_decimal(value: &str) -> Result<BigDecimalValue, AggregateError> {
let (mantissa, exponent) = match value.find(['e', 'E']) {
Some(index) => {
let exponent =
value[index + 1..]
.parse::<i64>()
.map_err(|_| AggregateError::NumberFormat {
value: value.to_owned(),
})?;
(&value[..index], exponent)
}
None => (value, 0),
};
let negative = mantissa.starts_with('-');
let unsigned = mantissa.strip_prefix(['-', '+']).unwrap_or(mantissa);
let mut split = unsigned.split('.');
let integer = split.next().unwrap_or_default();
let fraction = split.next();
if split.next().is_some()
|| (integer.is_empty() && fraction.is_none_or(str::is_empty))
|| !integer.bytes().all(|byte| byte.is_ascii_digit())
|| fraction.is_some_and(|digits| !digits.bytes().all(|byte| byte.is_ascii_digit()))
{
return Err(AggregateError::NumberFormat {
value: value.to_owned(),
});
}
let fraction = fraction.unwrap_or_default();
let digits = format!("{integer}{fraction}");
let mut unscaled =
BigInt::parse_bytes(digits.as_bytes(), 10).expect("validated decimal digits");
if negative {
unscaled = -unscaled;
}
let scale =
checked_scale(i64::try_from(fraction.len()).expect("fraction length fits i64") - exponent)?;
Ok(BigDecimalValue::from_unscaled(unscaled, scale))
}
pub(crate) fn double_string(value: f64) -> String {
if value.is_nan() {
return "NaN".to_owned();
}
if value == f64::INFINITY {
return "Infinity".to_owned();
}
if value == f64::NEG_INFINITY {
return "-Infinity".to_owned();
}
if value == 0.0 {
return if value.is_sign_negative() {
"-0.0".to_owned()
} else {
"0.0".to_owned()
};
}
if value.to_bits() == 1 {
return "4.9E-324".to_owned();
}
if value.to_bits() == (1_u64 << 63) | 1 {
return "-4.9E-324".to_owned();
}
let mut buffer = ryu::Buffer::new();
let raw = buffer.format(value);
let negative = raw.starts_with('-');
let unsigned = raw.strip_prefix('-').unwrap_or(raw);
let (mantissa, explicit_exponent) = match unsigned.find(['e', 'E']) {
Some(index) => (
&unsigned[..index],
unsigned[index + 1..].parse::<i32>().expect("ryu exponent"),
),
None => (unsigned, 0),
};
let fraction_length = mantissa
.find('.')
.map_or(0, |index| mantissa.len() - index - 1);
let mut digits = mantissa.replace('.', "");
let mut exponent = explicit_exponent - i32::try_from(fraction_length).expect("fraction");
while digits.starts_with('0') && digits.len() > 1 {
digits.remove(0);
}
while digits.ends_with('0') && digits.len() > 1 {
digits.pop();
exponent += 1;
}
let decimal_position = i32::try_from(digits.len()).expect("digits") + exponent;
let absolute = value.abs();
let mut result = if (1.0e-3..1.0e7).contains(&absolute) {
if decimal_position <= 0 {
format!(
"0.{}{}",
"0".repeat(usize::try_from(-decimal_position).expect("zero count")),
digits
)
} else if usize::try_from(decimal_position).expect("position") >= digits.len() {
format!(
"{}{}.0",
digits,
"0".repeat(usize::try_from(decimal_position).expect("position") - digits.len())
)
} else {
let split = usize::try_from(decimal_position).expect("position");
format!("{}.{}", &digits[..split], &digits[split..])
}
} else {
let scientific_exponent = decimal_position - 1;
let fraction = if digits.len() == 1 { "0" } else { &digits[1..] };
format!("{}.{fraction}E{scientific_exponent}", &digits[..1])
};
if negative {
result.insert(0, '-');
}
result
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum DivisionError {
ByZero,
NonTerminating,
ScaleOverflow,
}
#[derive(Clone, Copy, Debug, Eq, Error, PartialEq)]
pub(crate) enum BigDecimalArithmeticError {
#[error("Division by zero")]
DivisionByZero,
#[error("Non-terminating decimal expansion")]
NonTerminating,
#[error("Underflow")]
ScaleOverflow,
}
impl From<DivisionError> for AggregateError {
fn from(error: DivisionError) -> Self {
let message = match error {
DivisionError::ByZero => "Division by zero",
DivisionError::NonTerminating => "Non-terminating decimal expansion",
DivisionError::ScaleOverflow => "Underflow",
};
Self::Arithmetic {
message: message.to_owned(),
}
}
}
#[cfg(test)]
mod tests {
use std::cell::Cell;
use num_bigint::BigInt;
use super::{
AggregateError, AggregateObjectValue, AggregateUtils, BigDecimalValue, NumberIterableValue,
NumberListValue, NumberValue, double_string,
};
#[test]
fn preserves_big_decimal_representation_and_java_display_boundaries() {
let values = [
("1", "1", "1", 0),
("1.00", "1.00", "1.00", 2),
("1E+7", "1E+7", "10000000", -7),
("0.000001", "0.000001", "0.000001", 6),
("0.0000001", "1E-7", "0.0000001", 7),
("-0.0", "0.0", "0.0", 1),
];
for (source, display, plain, scale) in values {
let value = BigDecimalValue::parse(source).expect("decimal");
assert_eq!(value.to_string(), display);
assert_eq!(value.to_plain_string(), plain);
assert_eq!(value.scale(), scale);
assert!(value.precision() >= 1);
}
let explicit = BigDecimalValue::from_unscaled(BigInt::from(123), 2);
assert_eq!(explicit.unscaled_value(), &BigInt::from(123));
assert_eq!(explicit.to_string(), "1.23");
}
#[test]
fn formats_java_double_thresholds_special_values_and_signed_zero() {
assert_eq!(double_string(1.0), "1.0");
assert_eq!(double_string(9_999_999.0), "9999999.0");
assert_eq!(double_string(10_000_000.0), "1.0E7");
assert_eq!(double_string(0.001), "0.001");
assert_eq!(double_string(0.0001), "1.0E-4");
assert_eq!(double_string(-0.0), "-0.0");
assert_eq!(double_string(f64::NAN), "NaN");
assert_eq!(double_string(f64::INFINITY), "Infinity");
assert_eq!(double_string(f64::NEG_INFINITY), "-Infinity");
}
#[test]
fn sums_all_number_runtime_types_and_preserves_scale() {
let numbers = NumberListValue::new(vec![
Some(NumberValue::BigDecimal(
BigDecimalValue::parse("1.20").expect("decimal"),
)),
Some(NumberValue::BigInteger(BigInt::from(2))),
Some(NumberValue::Byte(3)),
Some(NumberValue::Short(4)),
Some(NumberValue::Integer(5)),
Some(NumberValue::Long(6)),
Some(NumberValue::Float(0.5)),
Some(NumberValue::Double(0.25)),
Some(NumberValue::Other {
class_name: "example.CustomNumber".to_owned(),
double_value: 0.05,
}),
]);
assert_eq!(numbers.as_slice().len(), 9);
let result = AggregateUtils::sum_iterable(Some(&numbers))
.expect("sum")
.expect("value");
assert_eq!(result.to_string(), "22.00");
assert_eq!(result.scale(), 2);
}
#[test]
fn averages_exactly_or_with_java_half_up_scale() {
let exact = AggregateUtils::avg_ints(Some(&[1, 2]))
.expect("average")
.expect("value");
assert_eq!(exact.to_string(), "1.5");
let repeating = AggregateUtils::avg_ints(Some(&[1, 1, 2]))
.expect("average")
.expect("value");
assert_eq!(repeating.to_string(), "1.3333333333");
let negative = AggregateUtils::avg_ints(Some(&[-1, -1, -2]))
.expect("average")
.expect("value");
assert_eq!(negative.to_string(), "-1.3333333333");
let scaled = AggregateUtils::avg_objects(Some(&[
AggregateObjectValue::Number(NumberValue::BigDecimal(
BigDecimalValue::parse("1.000000000000").expect("decimal"),
)),
AggregateObjectValue::Number(NumberValue::Integer(2)),
AggregateObjectValue::Number(NumberValue::Integer(2)),
]))
.expect("average")
.expect("value");
assert_eq!(scaled.to_string(), "1.666666666667");
assert_eq!(scaled.scale(), 12);
}
#[test]
fn preserves_null_validation_order_and_error_categories() {
assert_eq!(
AggregateUtils::sum_ints(None),
Err(AggregateError::IllegalArgument {
message: "Cannot aggregate on null"
})
);
let null_numbers = NumberListValue::new(vec![Some(NumberValue::Integer(1)), None]);
assert_eq!(
AggregateUtils::sum_iterable(Some(&null_numbers)),
Err(AggregateError::IllegalArgument {
message: "Cannot aggregate on iterable containing nulls"
})
);
assert_eq!(
AggregateUtils::avg_iterable(Some(&null_numbers)),
Err(AggregateError::IllegalArgument {
message: "Cannot aggregate on array containing nulls"
})
);
let objects = [
AggregateObjectValue::Other("java.lang.String".to_owned()),
AggregateObjectValue::Null,
];
assert_eq!(
AggregateUtils::sum_objects(Some(&objects)),
Err(AggregateError::IllegalArgument {
message: "Cannot aggregate on array containing nulls"
})
);
assert_eq!(
AggregateUtils::sum_objects(Some(&[AggregateObjectValue::Other(
"java.lang.String".to_owned()
)])),
Err(AggregateError::ClassCast {
actual_class: "java.lang.String".to_owned()
})
);
assert_eq!(
AggregateUtils::sum_doubles(Some(&[f64::NAN])),
Err(AggregateError::NumberFormat {
value: "NaN".to_owned()
})
);
let nan_numbers = NumberListValue::new(vec![Some(NumberValue::Double(f64::NAN))]);
assert_eq!(
AggregateUtils::sum_iterable(Some(&nan_numbers)),
Err(AggregateError::NumberFormat {
value: "NaN".to_owned()
})
);
assert_eq!(
AggregateUtils::sum_objects(Some(&[AggregateObjectValue::Number(
NumberValue::Double(f64::NAN)
)])),
Err(AggregateError::NumberFormat {
value: "NaN".to_owned()
})
);
}
#[test]
fn invokes_java_iterable_twice_and_handles_empty_targets() {
struct CountingIterable {
values: Vec<Option<NumberValue>>,
iterations: Cell<usize>,
}
impl NumberIterableValue for CountingIterable {
fn iter_java_numbers(&self) -> Box<dyn Iterator<Item = Option<&NumberValue>> + '_> {
self.iterations.set(self.iterations.get() + 1);
Box::new(self.values.iter().map(Option::as_ref))
}
}
let iterable = CountingIterable {
values: vec![Some(NumberValue::Integer(1))],
iterations: Cell::new(0),
};
assert_eq!(
AggregateUtils::sum_iterable(Some(&iterable))
.expect("sum")
.expect("value")
.to_string(),
"1"
);
assert_eq!(iterable.iterations.get(), 2);
assert_eq!(AggregateUtils::sum_ints(Some(&[])).expect("sum"), None);
assert_eq!(AggregateUtils::avg_ints(Some(&[])).expect("avg"), None);
assert_eq!(AggregateUtils::sum_objects(Some(&[])).expect("sum"), None);
}
#[test]
fn exercises_every_primitive_overload_and_float_failure() {
assert_eq!(
AggregateUtils::sum_bytes(Some(&[1, 2]))
.expect("sum")
.expect("value")
.to_string(),
"3"
);
assert_eq!(
AggregateUtils::sum_shorts(Some(&[1, 2]))
.expect("sum")
.expect("value")
.to_string(),
"3"
);
assert_eq!(
AggregateUtils::sum_longs(Some(&[1, 2]))
.expect("sum")
.expect("value")
.to_string(),
"3"
);
assert_eq!(
AggregateUtils::sum_floats(Some(&[0.5, 0.25]))
.expect("sum")
.expect("value")
.to_string(),
"0.75"
);
assert_eq!(
AggregateUtils::avg_bytes(Some(&[1, 2]))
.expect("avg")
.expect("value")
.to_string(),
"1.5"
);
assert_eq!(
AggregateUtils::avg_shorts(Some(&[1, 2]))
.expect("avg")
.expect("value")
.to_string(),
"1.5"
);
assert_eq!(
AggregateUtils::avg_longs(Some(&[1, 2]))
.expect("avg")
.expect("value")
.to_string(),
"1.5"
);
assert_eq!(
AggregateUtils::avg_floats(Some(&[0.5, 0.25]))
.expect("avg")
.expect("value")
.to_string(),
"0.375"
);
assert_eq!(
AggregateUtils::sum_doubles(Some(&[0.5, 0.25]))
.expect("sum")
.expect("value")
.to_string(),
"0.75"
);
assert_eq!(
AggregateUtils::avg_doubles(Some(&[0.5, 0.25]))
.expect("avg")
.expect("value")
.to_string(),
"0.375"
);
assert!(AggregateUtils::sum_floats(Some(&[f32::INFINITY])).is_err());
}
#[test]
fn rejects_malformed_decimal_and_covers_arithmetic_errors() {
assert_eq!(
BigDecimalValue::parse(".1")
.expect("leading point")
.to_string(),
"0.1"
);
assert_eq!(
BigDecimalValue::parse("1.")
.expect("trailing point")
.to_string(),
"1"
);
for malformed in ["", ".", "1.2.3", "1e", "NaN"] {
assert!(BigDecimalValue::parse(malformed).is_err(), "{malformed}");
}
assert!(BigDecimalValue::parse("1e2147483649").is_err());
assert_eq!(
AggregateError::from(super::DivisionError::ByZero).to_string(),
"Division by zero"
);
assert_eq!(
AggregateError::from(super::DivisionError::NonTerminating).to_string(),
"Non-terminating decimal expansion"
);
assert_eq!(
AggregateError::from(super::DivisionError::ScaleOverflow).to_string(),
"Underflow"
);
}
#[test]
fn covers_exact_division_rounding_and_scale_boundaries() {
let one = BigDecimalValue::parse("1").expect("one");
let zero = BigDecimalValue::parse("0").expect("zero");
assert_eq!(one.divide_exact(&zero), Err(super::DivisionError::ByZero));
assert_eq!(
one.divide_exact(&BigDecimalValue::parse("-8").expect("negative divisor"))
.expect("exact")
.to_string(),
"-0.125"
);
assert_eq!(
one.divide_exact(&BigDecimalValue::parse("8").expect("eight"))
.expect("exact")
.to_string(),
"0.125"
);
assert_eq!(
one.divide_exact(&BigDecimalValue::parse("125").expect("one hundred twenty five"))
.expect("exact")
.to_string(),
"0.008"
);
assert_eq!(
BigDecimalValue::from_unscaled(BigInt::from(1), i32::MAX)
.divide_exact(&BigDecimalValue::from_unscaled(BigInt::from(1), i32::MIN)),
Err(super::DivisionError::ScaleOverflow)
);
assert_eq!(
BigDecimalValue::parse("1")
.expect("one")
.divide_half_up_by_positive_integer(3, 0)
.to_string(),
"0"
);
assert_eq!(
BigDecimalValue::parse("2")
.expect("two")
.divide_half_up_by_positive_integer(3, 0)
.to_string(),
"1"
);
assert_eq!(
BigDecimalValue::parse("-1")
.expect("negative one")
.divide_half_up_by_positive_integer(2, 0)
.to_string(),
"-1"
);
assert!(super::finish_aggregate(BigDecimalValue::zero(), usize::MAX, true).is_err());
}
#[test]
fn covers_internal_validation_and_stateful_second_iteration() {
assert!(AggregateUtils::sum_iterable(None).is_err());
assert_eq!(
super::object_number(&AggregateObjectValue::Null),
Err(AggregateError::IllegalArgument {
message: "Cannot aggregate on array containing nulls"
})
);
struct ChangesAfterValidation {
calls: Cell<usize>,
number: NumberValue,
}
impl NumberIterableValue for ChangesAfterValidation {
fn iter_java_numbers(&self) -> Box<dyn Iterator<Item = Option<&NumberValue>> + '_> {
let call = self.calls.get();
self.calls.set(call + 1);
if call == 0 {
Box::new(std::iter::once(Some(&self.number)))
} else {
Box::new(std::iter::once(None))
}
}
}
let iterable = ChangesAfterValidation {
calls: Cell::new(0),
number: NumberValue::Integer(1),
};
assert_eq!(
AggregateUtils::sum_iterable(Some(&iterable)),
Err(AggregateError::IllegalArgument {
message: "Cannot aggregate on iterable containing nulls"
})
);
let negative_scientific = BigDecimalValue::parse("-12E+7").expect("scientific");
assert_eq!(negative_scientific.to_string(), "-1.2E+8");
assert_eq!(
BigDecimalValue::parse("+1.5E+2")
.expect("positive exponent")
.to_string(),
"1.5E+2"
);
}
#[test]
fn covers_every_array_null_and_number_array_validation_path() {
assert!(AggregateUtils::sum_bytes(None).is_err());
assert!(AggregateUtils::avg_bytes(None).is_err());
assert!(AggregateUtils::sum_shorts(None).is_err());
assert!(AggregateUtils::avg_shorts(None).is_err());
assert!(AggregateUtils::avg_ints(None).is_err());
assert!(AggregateUtils::sum_longs(None).is_err());
assert!(AggregateUtils::avg_longs(None).is_err());
assert!(AggregateUtils::sum_floats(None).is_err());
assert!(AggregateUtils::avg_floats(None).is_err());
assert!(AggregateUtils::sum_doubles(None).is_err());
assert!(AggregateUtils::avg_doubles(None).is_err());
assert!(AggregateUtils::avg_floats(Some(&[f32::NAN])).is_err());
assert!(AggregateUtils::avg_doubles(Some(&[f64::NAN])).is_err());
assert!(AggregateUtils::sum_numbers(None).is_err());
assert!(AggregateUtils::avg_numbers(None).is_err());
let null_numbers = [None];
assert!(AggregateUtils::sum_numbers(Some(&null_numbers)).is_err());
assert!(AggregateUtils::avg_numbers(Some(&null_numbers)).is_err());
let nan_numbers = [Some(NumberValue::Double(f64::NAN))];
assert!(AggregateUtils::sum_numbers(Some(&nan_numbers)).is_err());
assert!(AggregateUtils::avg_numbers(Some(&nan_numbers)).is_err());
let overflow_average =
[
AggregateObjectValue::Number(NumberValue::BigDecimal(
BigDecimalValue::from_unscaled(BigInt::from(1), i32::MAX),
)),
AggregateObjectValue::Number(NumberValue::BigDecimal(
BigDecimalValue::from_unscaled(BigInt::from(0), i32::MAX),
)),
];
assert_eq!(
AggregateUtils::avg_objects(Some(&overflow_average)),
Err(AggregateError::Arithmetic {
message: "Underflow".to_owned()
})
);
}
}