Jefffrey commented on code in PR #11302:
URL: https://github.com/apache/arrow-rs/pull/11302#discussion_r4177981497
##########
arrow-cast/src/cast/decimal.rs:
##########
@@ -675,44 +675,106 @@ where
D: DecimalType + ArrowPrimitiveType,
<D as ArrowPrimitiveType>::Native: DecimalCast,
{
+ validate_decimal_precision_and_scale::<D>(precision, scale)?;
let mul = 10_f64.powi(scale as i32);
if cast_options.safe {
array
- .unary_opt::<_, D>(|v| {
- single_float_to_decimal::<D>(v.as_(), mul)
- .filter(|v| D::is_valid_decimal_precision(*v, precision))
- })
+ .unary_opt::<_, D>(|v| float_to_decimal_checked::<D>(v.as_(), mul,
precision).ok())
.with_precision_and_scale(precision, scale)
.map(|a| Arc::new(a) as ArrayRef)
} else {
array
.try_unary::<_, D, _>(|v| {
- let v = single_float_to_decimal::<D>(v.as_(),
mul).ok_or_else(|| {
- ArrowError::CastError(format!(
- "Cannot cast to {}({}, {}). Overflowing on {:?}",
- D::PREFIX,
- precision,
- scale,
- v
- ))
- })?;
- D::validate_decimal_precision(v, precision, scale).map(|()| v)
+ float_to_decimal_checked::<D>(v.as_(), mul, precision)
+ .map_err(|e| e.into_arrow_error(v, precision, scale))
})?
.with_precision_and_scale(precision, scale)
.map(|a| Arc::new(a) as ArrayRef)
}
}
+/// Converts a floating point value to the unscaled integer representation of
a decimal.
+///
+/// Scales the input by `10^scale` and rounds half away from zero. Returns an
error
+/// if the precision or scale is invalid, or the result does not fit the native
+/// type or the requested precision. Non-finite inputs are rejected.
+///
+/// ```
+/// use arrow_array::types::Decimal32Type;
+/// use arrow_cast::cast::float_to_decimal;
+///
+/// assert_eq!(float_to_decimal::<Decimal32Type>(12.345, 5, 2).unwrap(), 1235);
+/// assert!(float_to_decimal::<Decimal32Type>(12345.678, 5, 2).is_err());
+/// ```
+#[inline]
+pub fn float_to_decimal<D>(input: f64, precision: u8, scale: i8) ->
Result<D::Native, ArrowError>
+where
+ D: DecimalType,
+ D::Native: DecimalCast,
+{
+ validate_decimal_precision_and_scale::<D>(precision, scale)?;
+ float_to_decimal_checked::<D>(input, 10_f64.powi(scale as i32), precision)
+ .map_err(|e| e.into_arrow_error(input, precision, scale))
+}
+
+// Keep failure reporting allocation-free for safe array casts, which discard
errors.
+enum FloatToDecimalError<D: DecimalType> {
+ NativeOverflow,
+ PrecisionOverflow(D::Native),
+}
+
+impl<D: DecimalType> FloatToDecimalError<D> {
+ fn into_arrow_error(self, input: impl std::fmt::Debug, precision: u8,
scale: i8) -> ArrowError {
+ match self {
+ Self::NativeOverflow => ArrowError::CastError(format!(
+ "Cannot cast to {}({}, {}). Overflowing on {:?}",
+ D::PREFIX,
+ precision,
+ scale,
+ input
+ )),
+ Self::PrecisionOverflow(value) => {
+ D::validate_decimal_precision(value, precision,
scale).unwrap_err()
+ }
+ }
+ }
+}
+
+// `mul` is precomputed once for array casts. Precision must already be
validated.
+#[inline]
+fn float_to_decimal_checked<D>(
+ input: f64,
+ mul: f64,
+ precision: u8,
+) -> Result<D::Native, FloatToDecimalError<D>>
+where
+ D: DecimalType,
+ D::Native: DecimalCast,
+{
+ let value =
+ D::Native::from_f64((mul *
input).round()).ok_or(FloatToDecimalError::NativeOverflow)?;
+ if D::is_valid_decimal_precision(value, precision) {
+ Ok(value)
+ } else {
+ Err(FloatToDecimalError::PrecisionOverflow(value))
+ }
+}
+
/// Cast a single floating point value to a decimal native with the given
multiple.
-/// Returns `None` if the value cannot be represented with the requested
precision.
+/// Returns `None` if the scaled and rounded value does not fit the maximum
precision
+/// of the decimal type. The caller must separately validate any smaller
target precision.
+///
+/// Unlike earlier versions, this also rejects values that fit the native
integer
+/// type but exceed the decimal type's maximum precision.
+#[deprecated(since = "60.0.0", note = "Use `float_to_decimal` instead")]
#[inline(always)]
pub fn single_float_to_decimal<D>(input: f64, mul: f64) -> Option<D::Native>
where
D: DecimalType + ArrowPrimitiveType,
<D as ArrowPrimitiveType>::Native: DecimalCast,
{
- D::Native::from_f64((mul * input).round())
+ float_to_decimal_checked::<D>(input, mul, D::MAX_PRECISION).ok()
Review Comment:
perhaps we should just leave this function unchanged if we deprecate it
anyway
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]