diff --git a/vortex-array/src/arrays/decimal/compute/cast.rs b/vortex-array/src/arrays/decimal/compute/cast.rs index 8770a9bc8f4..dd643bdb5cf 100644 --- a/vortex-array/src/arrays/decimal/compute/cast.rs +++ b/vortex-array/src/arrays/decimal/compute/cast.rs @@ -22,15 +22,20 @@ use crate::array::ArrayView; use crate::arrays::Decimal; use crate::arrays::DecimalArray; use crate::arrays::PrimitiveArray; +use crate::arrays::decimal::DecimalArrayExt; use crate::dtype::BigCast; use crate::dtype::DType; use crate::dtype::DecimalDType; use crate::dtype::DecimalType; +use crate::dtype::IntegerPType; use crate::dtype::NativeDecimalType; use crate::dtype::Nullability; use crate::dtype::PType; +use crate::dtype::ToI256; use crate::dtype::i256; use crate::match_each_decimal_value_type; +use crate::match_each_integer_ptype; +use crate::scalar::DecimalToIntegerCast; use crate::scalar::DecimalValue; use crate::scalar_fn::fns::cast::CastKernel; use crate::scalar_fn::fns::cast::CastReduce; @@ -88,6 +93,13 @@ impl CastKernel for Decimal { array.dtype() ); }; + if let DType::Primitive(ptype, nullability) = dtype + && ptype.is_int() + { + return match_each_integer_ptype!(*ptype, |T| { + cast_to_integer::(array, *nullability, ctx).map(Some) + }); + } if let DType::Primitive(PType::F64, nullability) = dtype { let scale = from_decimal_dtype.scale(); return cast_to_f64(array, scale, *nullability, ctx).map(Some); @@ -151,6 +163,25 @@ impl CastKernel for Decimal { } } +fn cast_to_integer( + array: ArrayView<'_, Decimal>, + nullability: Nullability, + ctx: &mut ExecutionCtx, +) -> VortexResult { + let source_validity = array.validity()?; + let mask = source_validity.execute_mask(array.len(), ctx)?; + let validity = source_validity.cast_nullability(nullability, array.len(), ctx)?; + let cast = DecimalToIntegerCast::::new(array.decimal_dtype().scale()); + let buffer = match_each_decimal_value_type!(array.values_type(), |F| { + let values = array.buffer::(); + cast_decimal_buffer(values.as_slice(), &mask, |value: F| { + cast.cast(value.to_i256()?) + }) + .map_err(|index| cast.error(DecimalValue::from(values[index]).as_i256()))? + }); + Ok(PrimitiveArray::new(buffer, validity).into_array()) +} + fn cast_to_f64( array: ArrayView<'_, Decimal>, scale: i8, @@ -239,7 +270,7 @@ fn cast_decimal_buffer( ) -> Result, usize> where F: NativeDecimalType, - T: NativeDecimalType, + T: Copy + Default, { let mut buffer = BufferMut::::with_capacity(values.len()); match valid_values { @@ -450,6 +481,7 @@ fn upcast_decimal_buffer(from: Buffe mod tests { use rstest::rstest; use vortex_buffer::buffer; + use vortex_error::VortexResult; use super::upcast_decimal_values; use crate::Canonical; @@ -457,16 +489,236 @@ mod tests { use crate::VortexSessionExecute; use crate::array_session; use crate::arrays::DecimalArray; + use crate::arrays::PrimitiveArray; + use crate::assert_arrays_eq; use crate::builtins::ArrayBuiltins; use crate::compute::conformance::cast::test_cast_conformance; + use crate::dtype::BigCast; use crate::dtype::DType; use crate::dtype::DecimalDType; use crate::dtype::DecimalType; use crate::dtype::Nullability; use crate::dtype::PType; + use crate::dtype::i256; + use crate::match_each_decimal_value_type; + use crate::match_each_integer_ptype; use crate::scalar::Scalar; use crate::validity::Validity; + #[rstest] + #[case::workflow(0, vec![0, 1], vec![0, 1])] + #[case::fractional(1, vec![-199, -1, 0, 1, 199], vec![-19, 0, 0, 0, 19])] + #[case::negative_scale(-2, vec![-12, 0, 12], vec![-1200, 0, 1200])] + #[case::large_exact( + 0, + vec![1_786_639_777_684_000_001], + vec![1_786_639_777_684_000_001] + )] + #[case::large_scale(76, vec![-1, 1], vec![0, 0])] + #[case::zero_extreme_scale(-128, vec![0, 0], vec![0, 0])] + fn cast_decimal_to_integer_policy( + #[case] scale: i8, + #[case] values: Vec, + #[case] expected: Vec, + ) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let decimal_dtype = DecimalDType::new(76, scale); + let array = DecimalArray::new( + values.iter().copied().map(i256::from_i128).collect(), + decimal_dtype, + Validity::NonNullable, + ); + let target = DType::Primitive(PType::I64, Nullability::NonNullable); + let casted = array + .into_array() + .cast(target.clone())? + .execute::(&mut ctx)?; + assert_arrays_eq!( + casted, + PrimitiveArray::from_iter(expected.iter().copied()), + &mut ctx + ); + for (value, expected) in values.into_iter().zip(expected) { + let scalar = Scalar::decimal(value.into(), decimal_dtype, Nullability::NonNullable); + assert_eq!(scalar.cast(&target)?, Scalar::from(expected)); + } + Ok(()) + } + + #[rstest] + fn cast_decimal_to_integer_storage_and_nulls( + #[values( + DecimalType::I8, + DecimalType::I16, + DecimalType::I32, + DecimalType::I64, + DecimalType::I128, + DecimalType::I256 + )] + storage: DecimalType, + #[values( + PType::I8, PType::I16, PType::I32, PType::I64, PType::U8, PType::U16, PType::U32, + PType::U64 + )] + target: PType, + ) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let decimal_dtype = DecimalDType::new(3, 1); + let array = match_each_decimal_value_type!(storage, |F| { + DecimalArray::from_option_iter( + [Some(19i8), None, Some(123)].map(|value| value.and_then(::from)), + decimal_dtype, + ) + }); + let dtype = DType::Primitive(target, Nullability::Nullable); + let casted = array + .into_array() + .cast(dtype.clone())? + .execute::(&mut ctx)?; + match_each_integer_ptype!(target, |T| { + assert_arrays_eq!( + casted, + PrimitiveArray::from_option_iter([Some(1 as T), None, Some(12 as T)]), + &mut ctx + ); + for (value, expected) in [(19i8, 1 as T), (123, 12 as T)] { + let scalar = Scalar::decimal(value.into(), decimal_dtype, Nullability::Nullable); + assert_eq!( + scalar.cast(&dtype)?, + Scalar::primitive(expected, Nullability::Nullable) + ); + } + }); + Ok(()) + } + + #[rstest] + fn cast_decimal_to_integer_bounds( + #[values( + PType::I8, PType::I16, PType::I32, PType::I64, PType::U8, PType::U16, PType::U32, + PType::U64 + )] + target: PType, + ) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let decimal_dtype = DecimalDType::new(76, 1); + let dtype = DType::Primitive(target, Nullability::NonNullable); + match_each_integer_ptype!(target, |T| { + let ten = i256::from_i128(10); + let minimum = i256::from_i128(T::MIN as i128) * ten; + let maximum = i256::from_i128(T::MAX as i128) * ten; + let array = DecimalArray::new( + buffer![minimum, maximum], + decimal_dtype, + Validity::NonNullable, + ); + let casted = array + .into_array() + .cast(dtype.clone())? + .execute::(&mut ctx)?; + assert_arrays_eq!( + casted, + PrimitiveArray::from_iter([T::MIN, T::MAX]), + &mut ctx + ); + for (value, expected) in [(minimum, T::MIN), (maximum, T::MAX)] { + let scalar = Scalar::decimal(value.into(), decimal_dtype, Nullability::NonNullable); + assert_eq!( + scalar.cast(&dtype)?, + Scalar::primitive(expected, Nullability::NonNullable) + ); + } + // The scalar policy rejects values outside the range even when truncation would fit. + for value in [ + minimum - i256::ONE, + maximum + i256::ONE, + minimum - ten, + maximum + ten, + ] { + let array = DecimalArray::new( + buffer![i256::ZERO, value], + decimal_dtype, + Validity::NonNullable, + ); + let error = array + .into_array() + .cast(dtype.clone())? + .execute::(&mut ctx) + .unwrap_err(); + assert!(error.to_string().contains("out of range")); + let scalar = Scalar::decimal(value.into(), decimal_dtype, Nullability::NonNullable); + assert!(scalar.cast(&dtype).is_err()); + } + }); + Ok(()) + } + + #[test] + fn cast_decimal_to_integer_masks_and_empty() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let decimal_dtype = DecimalDType::new(21, 0); + let target = DType::Primitive(PType::I8, Nullability::Nullable); + for validity in [Validity::from_iter([false, true]), Validity::AllInvalid] { + let expected = PrimitiveArray::new(buffer![0i8, 12], validity.clone()); + let array = DecimalArray::new(buffer![999i128, 12], decimal_dtype, validity); + let casted = array + .into_array() + .cast(target.clone())? + .execute::(&mut ctx)?; + assert_arrays_eq!(casted, expected, &mut ctx); + } + let array = DecimalArray::new( + buffer![999i128, 12], + decimal_dtype, + Validity::from_iter([false, true]), + ); + assert!( + array + .into_array() + .cast(target.as_nonnullable())? + .execute::(&mut ctx) + .is_err() + ); + let empty = DecimalArray::from_option_iter([] as [Option; 0], decimal_dtype); + let casted = empty + .into_array() + .cast(target.as_nonnullable())? + .execute::(&mut ctx)?; + assert_arrays_eq!(casted, PrimitiveArray::from_iter([] as [i8; 0]), &mut ctx); + Ok(()) + } + + #[test] + fn cast_decimal_to_integer_wide_storage() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let factor = i256::from_i128(10).checked_pow(40).unwrap(); + let value = i256::from_i128(123) * factor + i256::ONE; + let decimal_dtype = DecimalDType::new(76, 40); + let target = DType::Primitive(PType::I64, Nullability::NonNullable); + let array = DecimalArray::new(buffer![value, -value], decimal_dtype, Validity::NonNullable); + let casted = array + .into_array() + .cast(target.clone())? + .execute::(&mut ctx)?; + assert_arrays_eq!(casted, PrimitiveArray::from_iter([123i64, -123]), &mut ctx); + for (value, expected) in [(value, 123i64), (-value, -123)] { + let scalar = Scalar::decimal(value.into(), decimal_dtype, Nullability::NonNullable); + assert_eq!(scalar.cast(&target)?, Scalar::from(expected)); + } + let huge = DecimalArray::new( + buffer![i256::ONE], + DecimalDType::new(1, -128), + Validity::NonNullable, + ); + assert!( + huge.into_array() + .cast(target)? + .execute::(&mut ctx) + .is_err() + ); + Ok(()) + } + #[test] fn cast_decimal_to_nullable() { let mut ctx = array_session().create_execution_ctx(); diff --git a/vortex-array/src/scalar/typed_view/decimal/integer_cast.rs b/vortex-array/src/scalar/typed_view/decimal/integer_cast.rs new file mode 100644 index 00000000000..c9ea14a58ce --- /dev/null +++ b/vortex-array/src/scalar/typed_view/decimal/integer_cast.rs @@ -0,0 +1,63 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +//! Shared decimal-to-integer policy for scalar and array casts. + +use num_traits::CheckedMul; +use vortex_error::VortexError; +use vortex_error::vortex_err; + +use crate::dtype::BigCast; +use crate::dtype::IntegerPType; +use crate::dtype::i256; + +/// Truncate toward zero, rejecting values outside the target range before truncation. +/// Integer arithmetic preserves exact values even beyond floating-point precision. +pub(crate) struct DecimalToIntegerCast { + scale: i8, + factor: Option, + minimum: T, + maximum: T, +} + +impl DecimalToIntegerCast { + pub(crate) fn new(scale: i8) -> Self { + Self { + scale, + factor: i256::from_i128(10).checked_pow(scale.unsigned_abs().into()), + minimum: T::min_value(), + maximum: T::max_value(), + } + } + + pub(crate) fn cast(&self, value: i256) -> Option { + if value == i256::ZERO { + return Some(T::default()); + } + // A negative scale can require a factor larger than i256. Any nonzero + // value then exceeds every primitive integer's range. + let factor = self.factor?; + let integer = if self.scale > 0 { + let integer = value / factor; + if value % factor != i256::ZERO + && ((value < i256::ZERO && integer == self.minimum.to_i256()?) + || (value > i256::ZERO && integer == self.maximum.to_i256()?)) + { + return None; + } + integer + } else { + value.checked_mul(&factor)? + }; + ::from(integer) + } + + pub(crate) fn error(&self, value: i256) -> VortexError { + vortex_err!( + "Decimal value {} at scale {} out of range for {}", + value, + self.scale, + T::PTYPE + ) + } +} diff --git a/vortex-array/src/scalar/typed_view/decimal/mod.rs b/vortex-array/src/scalar/typed_view/decimal/mod.rs index 149fc8aeadf..2274bb10274 100644 --- a/vortex-array/src/scalar/typed_view/decimal/mod.rs +++ b/vortex-array/src/scalar/typed_view/decimal/mod.rs @@ -5,11 +5,13 @@ mod arithmetic; mod dvalue; +mod integer_cast; mod scalar; pub(crate) use arithmetic::decimal_numeric_result_dtype; pub(crate) use arithmetic::decimal_numeric_work_dtype; pub use dvalue::DecimalValue; +pub(crate) use integer_cast::DecimalToIntegerCast; pub use scalar::DecimalScalar; #[cfg(test)] diff --git a/vortex-array/src/scalar/typed_view/decimal/scalar.rs b/vortex-array/src/scalar/typed_view/decimal/scalar.rs index 5ec63975e3d..4113b625ecd 100644 --- a/vortex-array/src/scalar/typed_view/decimal/scalar.rs +++ b/vortex-array/src/scalar/typed_view/decimal/scalar.rs @@ -11,12 +11,14 @@ use vortex_error::VortexResult; use vortex_error::vortex_bail; use vortex_error::vortex_err; +use super::DecimalToIntegerCast; use super::arithmetic::checked_decimal_numeric; use super::arithmetic::decimal_numeric_result_dtype; use crate::dtype::DType; use crate::dtype::DecimalDType; use crate::dtype::PType; use crate::match_each_decimal_value; +use crate::match_each_integer_ptype; use crate::scalar::DecimalValue; use crate::scalar::NumericOperator; use crate::scalar::Scalar; @@ -77,6 +79,18 @@ impl<'a> DecimalScalar<'a> { Ok(Scalar::null(dtype.clone())) } } + DType::Primitive(ptype, nullability) if ptype.is_int() => { + let Some(value) = self.decimal_value else { + return Ok(Scalar::null(dtype.clone())); + }; + match_each_integer_ptype!(*ptype, |T| { + let cast = DecimalToIntegerCast::::new(self.decimal_type.scale()); + let value = value.as_i256(); + cast.cast(value) + .map(|integer| Scalar::primitive(integer, *nullability)) + .ok_or_else(|| cast.error(value)) + }) + } DType::Primitive(ptype, nullability) => { // Cast decimal to primitive type if let Some(decimal_value) = &self.decimal_value { @@ -103,68 +117,13 @@ impl<'a> DecimalScalar<'a> { reason = "truncation is intentional - range checks happen after" )] let primitive_scalar = match ptype { - PType::U8 => { - let v = actual_value as u8; - if actual_value < 0.0 || actual_value > u8::MAX as f64 { - vortex_bail!("Decimal value {} out of range for u8", actual_value); - } - Scalar::primitive(v, *nullability) - } - PType::U16 => { - let v = actual_value as u16; - if actual_value < 0.0 || actual_value > u16::MAX as f64 { - vortex_bail!("Decimal value {} out of range for u16", actual_value); - } - Scalar::primitive(v, *nullability) - } - PType::U32 => { - let v = actual_value as u32; - if actual_value < 0.0 || actual_value > u32::MAX as f64 { - vortex_bail!("Decimal value {} out of range for u32", actual_value); - } - Scalar::primitive(v, *nullability) - } - PType::U64 => { - let v = actual_value as u64; - if actual_value < 0.0 || actual_value > u64::MAX as f64 { - vortex_bail!("Decimal value {} out of range for u64", actual_value); - } - Scalar::primitive(v, *nullability) - } - PType::I8 => { - let v = actual_value as i8; - if actual_value < i8::MIN as f64 || actual_value > i8::MAX as f64 { - vortex_bail!("Decimal value {} out of range for i8", actual_value); - } - Scalar::primitive(v, *nullability) - } - PType::I16 => { - let v = actual_value as i16; - if actual_value < i16::MIN as f64 || actual_value > i16::MAX as f64 { - vortex_bail!("Decimal value {} out of range for i16", actual_value); - } - Scalar::primitive(v, *nullability) - } - PType::I32 => { - let v = actual_value as i32; - if actual_value < i32::MIN as f64 || actual_value > i32::MAX as f64 { - vortex_bail!("Decimal value {} out of range for i32", actual_value); - } - Scalar::primitive(v, *nullability) - } - PType::I64 => { - let v = actual_value as i64; - if actual_value < i64::MIN as f64 || actual_value > i64::MAX as f64 { - vortex_bail!("Decimal value {} out of range for i64", actual_value); - } - Scalar::primitive(v, *nullability) - } PType::F16 => { use crate::dtype::half::f16; Scalar::primitive(f16::from_f64(actual_value), *nullability) } PType::F32 => Scalar::primitive(actual_value as f32, *nullability), PType::F64 => Scalar::primitive(actual_value, *nullability), + _ => unreachable!("integer casts are handled above"), }; Ok(primitive_scalar) } else { diff --git a/vortex-sqllogictest/slt/datafusion/decimal_to_integer.slt b/vortex-sqllogictest/slt/datafusion/decimal_to_integer.slt new file mode 100644 index 00000000000..a9d0dbcf575 --- /dev/null +++ b/vortex-sqllogictest/slt/datafusion/decimal_to_integer.slt @@ -0,0 +1,61 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright the Vortex contributors + +include ../setup.slt.no + +statement ok +SET vortex.projection_pushdown = true; + +statement ok +COPY (SELECT id, CAST(value AS DECIMAL(21,1)) AS value FROM (VALUES + (1, '0.0'), + (2, '1.0'), + (3, '-19.9'), + (4, '19.9'), + (5, '1786639777684000001.0'), + (6, '-9223372036854775808.0'), + (7, '9223372036854775807.0'), + (8, NULL) +) AS t(id, value)) TO '${WORK_DIR}/decimal_to_integer.vortex'; + +statement ok +CREATE VIEW decimals AS SELECT * FROM '${WORK_DIR}/decimal_to_integer.vortex'; + +# EXPLAIN verifies that the cast executes inside the Vortex scan. +query TT +EXPLAIN SELECT CAST(value AS BIGINT) AS integer_value FROM decimals; +---- +logical_plan +01)Projection: CAST(decimals.value AS Int64) AS integer_value +02)--SubqueryAlias: decimals +03)----TableScan: ${WORK_DIR}/decimal_to_integer.vortex projection=[value] +physical_plan DataSourceExec: file_groups={1 group: [[${WORK_DIR}/decimal_to_integer.vortex]]}, projection=[CAST(value@1 AS Int64) as integer_value], file_type=vortex + +# Fractions truncate toward zero; large integers remain exact and nulls survive. +query I +SELECT CAST(value AS BIGINT) FROM decimals ORDER BY id; +---- +0 +1 +-19 +19 +1786639777684000001 +-9223372036854775808 +9223372036854775807 +NULL + +# Repeat with projection pushdown disabled to check agreement with DataFusion. +statement ok +SET vortex.projection_pushdown = false; + +query I +SELECT CAST(value AS BIGINT) FROM decimals ORDER BY id; +---- +0 +1 +-19 +19 +1786639777684000001 +-9223372036854775808 +9223372036854775807 +NULL diff --git a/vortex-sqllogictest/slt/duckdb/decimal_to_integer.slt b/vortex-sqllogictest/slt/duckdb/decimal_to_integer.slt new file mode 100644 index 00000000000..5f04075b0fc --- /dev/null +++ b/vortex-sqllogictest/slt/duckdb/decimal_to_integer.slt @@ -0,0 +1,43 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright the Vortex contributors + +include ../setup.slt.no + +statement ok +COPY (SELECT id, CAST(value AS DECIMAL(21,1)) AS value FROM (VALUES + (1, '0.0'), + (2, '1.0'), + (3, '-19.9'), + (4, '19.9'), + (5, '1786639777684000001.0'), + (6, '-9223372036854775808.0'), + (7, '9223372036854775807.0'), + (8, NULL) +) AS t(id, value)) TO '${WORK_DIR}/decimal_to_integer.vortex'; + +statement ok +CREATE VIEW decimals AS SELECT * FROM '${WORK_DIR}/decimal_to_integer.vortex'; + +# DuckDB rounds decimal fractions instead of truncating them. Decimal casts must +# stay outside the Vortex scan even though Vortex now supports integer targets. +query TT +EXPLAIN SELECT CAST(value AS BIGINT) AS integer_value FROM decimals; +---- +:.*PROJECTION.* + +query TT +EXPLAIN SELECT CAST(value AS BIGINT) AS integer_value FROM decimals; +---- +:.*SELECT projections.* + +query I +SELECT CAST(value AS BIGINT) FROM decimals ORDER BY id; +---- +0 +1 +-20 +20 +1786639777684000001 +-9223372036854775808 +9223372036854775807 +NULL