Skip to content
Open
23 changes: 12 additions & 11 deletions encodings/decimal-byte-parts/src/decimal_byte_parts/assemble.rs
Original file line number Diff line number Diff line change
Expand Up @@ -87,8 +87,9 @@ fn assemble_narrow_decimal(
) -> VortexResult<ArrayRef> {
// TODO(mk): Broadcast a constant MSP directly instead of materializing its buffer.
let msp = msp.clone().execute::<PrimitiveArray>(exec_ctx)?;
// Parts written before storage was bounded by precision may be wider than it allows.
Ok(match_each_signed_integer_ptype!(msp.ptype(), |P| {
DecimalArray::new(msp.to_buffer::<P>(), decimal_dtype, validity).into_array()
DecimalArray::try_new_narrowed(msp.to_buffer::<P>(), decimal_dtype, validity)?.into_array()
}))
}

Expand Down Expand Up @@ -118,16 +119,16 @@ fn assemble_wide_decimal_from_arrays(
Ok(match_each_signed_integer_ptype!(msp.ptype(), |Msp| {
let msp = msp.as_slice::<Msp>();
match lower.as_slice() {
[first] => DecimalArray::new(
[first] => DecimalArray::try_new_narrowed(
assemble_wide_decimal::<i128, Msp, 1>(
msp,
first.as_slice::<u64>().iter().map(|&word| [word]),
),
decimal_dtype,
validity,
)
)?
.into_array(),
[first, second] => DecimalArray::new(
[first, second] => DecimalArray::try_new_narrowed(
assemble_wide_decimal::<i256, Msp, 2>(
msp,
first
Expand All @@ -138,9 +139,9 @@ fn assemble_wide_decimal_from_arrays(
),
decimal_dtype,
validity,
)
)?
.into_array(),
[first, second, third] => DecimalArray::new(
[first, second, third] => DecimalArray::try_new_narrowed(
assemble_wide_decimal::<i256, Msp, 3>(
msp,
first
Expand All @@ -152,7 +153,7 @@ fn assemble_wide_decimal_from_arrays(
),
decimal_dtype,
validity,
)
)?
.into_array(),
_ => vortex_bail!("expected between one and {MAX_LOWER_PARTS} lower parts"),
}
Expand Down Expand Up @@ -368,14 +369,14 @@ mod tests {
validity: Validity,
) -> VortexResult<()> {
let mut ctx = array_session().create_execution_ctx();
let decimal = DecimalArray::new(buffer![1i32, 2, 3], DecimalDType::new(2, 0), validity);
let decimal = DecimalArray::new(buffer![1i8, 2, 3], DecimalDType::new(2, 0), validity);
let parts = split_decimal(&decimal, &mut ctx)?;
assert!(parts.lower_parts.is_empty());
assert_eq!(parts.msp.dtype().as_ptype(), PType::I32);
assert_eq!(parts.msp.dtype().as_ptype(), PType::I8);
let msp = parts.msp.execute::<PrimitiveArray>(&mut ctx)?;
assert_eq!(
msp.as_slice::<i32>().as_ptr(),
decimal.buffer::<i32>().as_ptr()
msp.as_slice::<i8>().as_ptr(),
decimal.buffer::<i8>().as_ptr()
);
assert_arrays_eq!(decimal.clone(), round_trip(decimal)?, &mut ctx);
Ok(())
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -172,7 +172,7 @@ mod tests {
)?
.into_array();
let canonical = DecimalArray::new(
Buffer::from_iter(values.iter().map(|v| *v as i128)),
Buffer::from_iter(values.iter().map(|v| *v as i64)),
decimal_dtype,
validity,
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -39,18 +39,18 @@ use crate::decimal_byte_parts::MAX_LOWER_PARTS;
#[case::one_lower_part(DecimalByteParts::try_new_with_lower_parts(
msp(), vec![lower_part()], DecimalDType::new(38, 2),
))]
#[case::wider_i64_storage(DecimalByteParts::encode(
&DecimalArray::new(buffer![-99i64, 0, 99], DecimalDType::new(2, 0), Validity::NonNullable),
#[case::i64_storage(DecimalByteParts::encode(
&DecimalArray::new(buffer![-99i64, 0, 99], DecimalDType::new(18, 0), Validity::NonNullable),
&mut array_session().create_execution_ctx(),
))]
#[case::wider_i128_storage(DecimalByteParts::encode(
&DecimalArray::new(buffer![-99i128, 0, 99], DecimalDType::new(2, 0), Validity::NonNullable),
#[case::i128_storage(DecimalByteParts::encode(
&DecimalArray::new(buffer![-99i128, 0, 99], DecimalDType::new(38, 0), Validity::NonNullable),
&mut array_session().create_execution_ctx(),
))]
#[case::wider_i256_storage(DecimalByteParts::encode(
#[case::i256_storage(DecimalByteParts::encode(
&DecimalArray::new(
buffer![i256::from_i128(-99), i256::ZERO, i256::from_i128(99)],
DecimalDType::new(2, 0), Validity::NonNullable,
DecimalDType::new(76, 0), Validity::NonNullable,
),
&mut array_session().create_execution_ctx(),
))]
Expand Down
4 changes: 2 additions & 2 deletions encodings/sparse/src/canonical.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1135,7 +1135,7 @@ mod test {
let indices = buffer![0u32, 1u32, 7u32, 8u32].into_array();
let decimal_dtype = DecimalDType::new(3, 2);
let patch_values = DecimalArray::new(
buffer![100i128, 200i128, 300i128, 4000i128],
buffer![100i16, 200i16, 300i16, 4000i16],
decimal_dtype,
Validity::from_iter([true, true, true, false]),
)
Expand All @@ -1148,7 +1148,7 @@ mod test {
.arrow()
.execute_arrow(
DecimalArray::new(
buffer![100i128, 200, 123, 123, 123, 123, 123, 300, 4000, 123],
buffer![100i16, 200, 123, 123, 123, 123, 123, 300, 4000, 123],
decimal_dtype,
// NB: patch indices: [0, 1, 7, 8]; patch validity: [Valid, Valid, Valid, Invalid]; ergo 0, 1, 7 are valid.
Validity::from_mask(Mask::from_excluded_indices(10, vec![8]), Nullable),
Expand Down
8 changes: 3 additions & 5 deletions fuzz/src/array/mask.rs
Original file line number Diff line number Diff line change
Expand Up @@ -253,16 +253,14 @@ mod tests {
fn test_mask_decimal_array() {
let mut ctx = array_session().create_execution_ctx();
let dtype = DecimalDType::new(10, 2);
let array = DecimalArray::from_option_iter(
[Some(1i128), Some(2), Some(3), Some(4), Some(5)],
dtype,
);
let array =
DecimalArray::from_option_iter([Some(1i64), Some(2), Some(3), Some(4), Some(5)], dtype);
let mask = Mask::from_iter([true, true, false, true, true]);

let result = mask_canonical_array(canonical(array, &mut ctx), &mask, &mut ctx).unwrap();

let expected =
DecimalArray::from_option_iter([Some(1i128), Some(2), None, Some(4), Some(5)], dtype);
DecimalArray::from_option_iter([Some(1i64), Some(2), None, Some(4), Some(5)], dtype);
assert_arrays_eq!(result, expected, &mut ctx);
}

Expand Down
27 changes: 21 additions & 6 deletions fuzz/src/array/sum/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -141,22 +141,37 @@ fn test_sum_signed_unambiguous(
#[case(DecimalType::I128)]
#[case(DecimalType::I256)]
fn test_sum_decimal_storage(#[case] values_type: DecimalType) -> VortexResult<()> {
// Storage is bounded by precision, so each width is exercised at the widest precision it
// may back.
let precision = match values_type {
DecimalType::I8 => 2,
DecimalType::I16 => 4,
DecimalType::I32 => 9,
DecimalType::I64 => 18,
DecimalType::I128 => 38,
DecimalType::I256 => 76,
};
let array = match_each_decimal_value_type!(values_type, |D| {
let value = DecimalValue::I8(99)
.cast::<D>()
.ok_or_else(|| vortex_err!("99 fits in every decimal storage type"))?;
DecimalArray::new(
buffer![value, value, -value],
DecimalDType::new(2, 0),
DecimalDType::new(precision, 0),
Validity::NonNullable,
)
.into_array()
});
let expected = Scalar::decimal(
DecimalValue::I64(99),
DecimalDType::new(12, 0),
Nullability::Nullable,
);
let sum_dtype = DecimalDType::new(u8::min(76, precision + 10), 0);
let expected_value =
match_each_decimal_value_type!(DecimalType::smallest_decimal_value_type(&sum_dtype), |R| {
DecimalValue::from(
DecimalValue::I8(99)
.cast::<R>()
.ok_or_else(|| vortex_err!("99 fits in every decimal storage type"))?,
)
});
let expected = Scalar::decimal(expected_value, sum_dtype, Nullability::Nullable);
let mut ctx = SESSION.create_execution_ctx();
assert_eq!(
sum_canonical_array(&array, &mut ctx)?,
Expand Down
2 changes: 1 addition & 1 deletion vortex-array/src/aggregate_fn/fns/min_max/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -934,7 +934,7 @@ mod tests {
#[test]
fn test_decimal() -> VortexResult<()> {
let decimal = DecimalArray::new(
buffer![100i32, 2000i32, 200i32],
buffer![100i16, 2000i16, 200i16],
DecimalDType::new(4, 2),
Validity::from_iter([true, false, true]),
);
Expand Down
10 changes: 5 additions & 5 deletions vortex-array/src/aggregate_fn/fns/sum/decimal.rs
Original file line number Diff line number Diff line change
Expand Up @@ -147,7 +147,7 @@ mod tests {
#[test]
fn sum_decimal_basic() -> VortexResult<()> {
let decimal = DecimalArray::new(
buffer![100i32, 200i32, 300i32],
buffer![100i16, 200i16, 300i16],
DecimalDType::new(4, 2),
Validity::AllValid,
);
Expand All @@ -169,7 +169,7 @@ mod tests {
#[test]
fn sum_decimal_with_nulls() -> VortexResult<()> {
let decimal = DecimalArray::new(
buffer![100i32, 200i32, 300i32, 400i32],
buffer![100i16, 200i16, 300i16, 400i16],
DecimalDType::new(4, 2),
Validity::from_iter([true, false, true, true]),
);
Expand All @@ -191,7 +191,7 @@ mod tests {
#[test]
fn sum_decimal_negative_values() -> VortexResult<()> {
let decimal = DecimalArray::new(
buffer![100i32, -200i32, 300i32, -50i32],
buffer![100i16, -200i16, 300i16, -50i16],
DecimalDType::new(4, 2),
Validity::AllValid,
);
Expand Down Expand Up @@ -283,7 +283,7 @@ mod tests {
#[test]
fn sum_decimal_single_value() -> VortexResult<()> {
let decimal =
DecimalArray::new(buffer![42i32], DecimalDType::new(3, 1), Validity::AllValid);
DecimalArray::new(buffer![42i16], DecimalDType::new(3, 1), Validity::AllValid);

let result = sum(
&decimal.into_array(),
Expand All @@ -302,7 +302,7 @@ mod tests {
#[test]
fn sum_decimal_all_nulls_except_one() -> VortexResult<()> {
let decimal = DecimalArray::new(
buffer![100i32, 200i32, 300i32, 400i32],
buffer![100i16, 200i16, 300i16, 400i16],
DecimalDType::new(4, 2),
Validity::from_iter([false, false, true, false]),
);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -515,7 +515,7 @@ mod tests {
#[test]
fn decimal_matches_materialized_size() -> VortexResult<()> {
let array = DecimalArray::new(
buffer![12345i64, -123i64, 0i64],
buffer![12345i32, -123i32, 0i32],
DecimalDType::new(5, 2),
Validity::NonNullable,
)
Expand Down
17 changes: 6 additions & 11 deletions vortex-array/src/arrays/constant/vtable/canonical.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,11 +38,9 @@ use crate::builders::builder_with_capacity_in;
use crate::dtype::DType;
use crate::dtype::DecimalType;
use crate::dtype::Nullability;
use crate::match_each_decimal_value;
use crate::match_each_decimal_value_type;
use crate::match_each_native_ptype;
use crate::match_smallest_list_offset_type;
use crate::scalar::DecimalValue;
use crate::scalar::Scalar;
use crate::validity::Validity;

Expand Down Expand Up @@ -104,15 +102,12 @@ pub(crate) fn constant_canonicalize(
return Ok(Canonical::Decimal(all_null));
};

let decimal_array = match_each_decimal_value!(value, |value| {
// SAFETY: Constant decimal values with correct type and validity.
unsafe {
DecimalArray::new_unchecked(
Buffer::full(value, array.len()),
*decimal_type,
validity,
)
}
let storage = value.decimal_type().min(size);
let decimal_array = match_each_decimal_value_type!(storage, |D| {
let value = value
.cast::<D>()
.vortex_expect("decimal scalar fits its precision");
DecimalArray::new(Buffer::full(value, array.len()), *decimal_type, validity)
});
Canonical::Decimal(decimal_array)
}
Expand Down
Loading
Loading