sdf-jkl commented on code in PR #10157:
URL: https://github.com/apache/arrow-rs/pull/10157#discussion_r3574577053
##########
parquet-variant-compute/src/shred_variant.rs:
##########
@@ -2829,4 +2818,207 @@ mod tests {
let shredding_type = ShreddedSchemaBuilder::default().build();
assert_eq!(shredding_type, DataType::Null);
}
+
+ // This test wants to cover that the variant can/can't be shredded to the
given data type.
+ #[test]
+ fn test_variant_type_shredded_correctly() {
+ // array contains all variant types
+ let mut array_builder = VariantArrayBuilder::new(30);
+ array_builder.append_value(Variant::Null);
+ array_builder.append_value(Variant::Int8(1));
+ array_builder.append_value(Variant::Int16(2));
+ array_builder.append_value(Variant::Int32(3));
+ array_builder.append_value(Variant::Int64(4));
+
array_builder.append_value(Variant::Date(NaiveDate::from_epoch_days(12345).unwrap()));
+ array_builder.append_value(Variant::TimestampMicros(
+ DateTime::from_timestamp_micros(123456789).unwrap(),
+ ));
+ array_builder.append_value(Variant::TimestampNtzMicros(
+ DateTime::from_timestamp_micros(123456789)
+ .unwrap()
+ .naive_utc(),
+ ));
+
array_builder.append_value(Variant::TimestampNanos(DateTime::from_timestamp_nanos(
+ 1234567890123,
+ )));
+ array_builder.append_value(Variant::TimestampNtzNanos(
+ DateTime::from_timestamp_nanos(1234567890123).naive_utc(),
+ ));
+ array_builder.append_value(VariantDecimal4::try_new(123, 2).unwrap());
+ array_builder.append_value(VariantDecimal8::try_new(123, 2).unwrap());
+ array_builder.append_value(VariantDecimal16::try_new(123, 2).unwrap());
+ array_builder.append_value(Variant::Float(5.0));
+ array_builder.append_value(Variant::Double(6f64));
+ array_builder.append_value(Variant::BooleanTrue);
+ array_builder.append_value(Variant::BooleanFalse);
+ array_builder.append_value(Variant::Binary("helow".as_bytes()));
+ array_builder.append_value(Variant::String("hello"));
+ array_builder.append_value(Variant::ShortString(
+ ShortString::try_from("world").unwrap(),
+ ));
+ array_builder.append_value(Variant::Time(
+ NaiveTime::from_num_seconds_from_midnight_opt(12345, 123).unwrap(),
+ ));
+
+ let array = array_builder.build();
+
+ fn can_shred_to(v: &Variant, dt: &DataType) -> bool {
+ matches!(
+ (v, dt),
+ (Variant::Int8(_), DataType::Int8)
+ | (Variant::Int8(_), DataType::Int16)
+ | (Variant::Int8(_), DataType::Int32)
+ | (Variant::Int8(_), DataType::Int64)
+ | (Variant::Int16(_), DataType::Int8)
+ | (Variant::Int16(_), DataType::Int16)
+ | (Variant::Int16(_), DataType::Int32)
+ | (Variant::Int16(_), DataType::Int64)
+ | (Variant::Int32(_), DataType::Int8)
+ | (Variant::Int32(_), DataType::Int16)
+ | (Variant::Int32(_), DataType::Int32)
+ | (Variant::Int32(_), DataType::Int64)
+ | (Variant::Int64(_), DataType::Int8)
+ | (Variant::Int64(_), DataType::Int16)
+ | (Variant::Int64(_), DataType::Int32)
+ | (Variant::Int64(_), DataType::Int64)
+ | (Variant::Date(_), DataType::Date32)
+ | (
+ Variant::TimestampMicros(_),
+ DataType::Timestamp(TimeUnit::Microsecond, Some(_)),
+ )
+ | (
+ Variant::TimestampMicros(_),
+ DataType::Timestamp(TimeUnit::Nanosecond, Some(_))
+ )
+ | (
+ Variant::TimestampNtzMicros(_),
+ DataType::Timestamp(TimeUnit::Microsecond, None),
+ )
+ | (
+ Variant::TimestampNtzMicros(_),
+ DataType::Timestamp(TimeUnit::Nanosecond, None)
+ )
+ | (
+ Variant::TimestampNanos(_),
+ DataType::Timestamp(TimeUnit::Nanosecond, Some(_)),
+ )
+ | (
+ Variant::TimestampNtzNanos(_),
+ DataType::Timestamp(TimeUnit::Nanosecond, None),
+ )
+ | (Variant::Decimal4(_), DataType::Decimal32(_, _))
+ | (Variant::Decimal4(_), DataType::Decimal64(_, _))
+ | (Variant::Decimal4(_), DataType::Decimal128(_, _))
+ | (Variant::Decimal8(_), DataType::Decimal32(_, _))
+ | (Variant::Decimal8(_), DataType::Decimal64(_, _))
+ | (Variant::Decimal8(_), DataType::Decimal128(_, _))
+ | (Variant::Decimal16(_), DataType::Decimal32(_, _))
+ | (Variant::Decimal16(_), DataType::Decimal64(_, _))
+ | (Variant::Decimal16(_), DataType::Decimal128(_, _))
+ | (Variant::Float(_), DataType::Float32)
+ | (Variant::Float(_), DataType::Float64)
+ | (Variant::Double(_), DataType::Float32)
+ | (Variant::Double(_), DataType::Float64)
Review Comment:
Per equivalence class, decimals are in the same family as Int, so they
should be able to cast between each other.
Float and Double are not in the same family, so they should not be cast-able.
##########
parquet-variant-compute/src/type_conversion.rs:
##########
@@ -287,6 +444,143 @@ where
}
}
+/// Return the unscaled integer representation for Arrow decimal type `O` from
a `Variant`.
+///
+/// This function is unlike `variant_to_unscaled_decim`, it would never
rescale the decimal value,
+/// and only return the unscaled integer representation for the specific
decimal variants.
+pub(crate) fn shred_variant_to_unscaled_decimal<O>(
+ variant: &Variant<'_, '_>,
+ precision: u8,
+ scale: i8,
+) -> Option<O::Native>
+where
+ O: ShredDecimalVariant,
+ O::Native: DecimalCast,
+{
+ match variant {
+ Variant::Decimal4(_) | Variant::Decimal8(_) | Variant::Decimal16(_) =>
{
+ O::shred_variant(variant, precision, scale)
+ }
+ _ => None,
+ }
+}
+pub(crate) trait ShredDecimalVariant: DecimalType {
+ fn shred_variant(value: &Variant<'_, '_>, precision: u8, scale: i8) ->
Option<Self::Native>;
+}
+
+impl ShredDecimalVariant for Decimal32Type {
+ fn shred_variant(value: &Variant<'_, '_>, precision: u8, scale: i8) ->
Option<Self::Native> {
+ match *value {
+ Variant::Decimal4(d) => rescale_decimal::<Decimal32Type,
Decimal32Type>(
Review Comment:
`rescale_decimal` rounds when input_scale > output_scale.
1.23 shredded to Decimal32(9,0) gives typed_value=1, value=NULL — the 0.23
is unrecoverable. Same class, but shredding also has to be exact, and it isn't.
Spark rescales then rejects if inexact.
##########
parquet-variant-compute/src/shred_variant.rs:
##########
@@ -2829,4 +2818,207 @@ mod tests {
let shredding_type = ShreddedSchemaBuilder::default().build();
assert_eq!(shredding_type, DataType::Null);
}
+
+ // This test wants to cover that the variant can/can't be shredded to the
given data type.
+ #[test]
+ fn test_variant_type_shredded_correctly() {
+ // array contains all variant types
+ let mut array_builder = VariantArrayBuilder::new(30);
+ array_builder.append_value(Variant::Null);
+ array_builder.append_value(Variant::Int8(1));
+ array_builder.append_value(Variant::Int16(2));
+ array_builder.append_value(Variant::Int32(3));
+ array_builder.append_value(Variant::Int64(4));
+
array_builder.append_value(Variant::Date(NaiveDate::from_epoch_days(12345).unwrap()));
+ array_builder.append_value(Variant::TimestampMicros(
+ DateTime::from_timestamp_micros(123456789).unwrap(),
+ ));
+ array_builder.append_value(Variant::TimestampNtzMicros(
+ DateTime::from_timestamp_micros(123456789)
+ .unwrap()
+ .naive_utc(),
+ ));
+
array_builder.append_value(Variant::TimestampNanos(DateTime::from_timestamp_nanos(
+ 1234567890123,
+ )));
+ array_builder.append_value(Variant::TimestampNtzNanos(
+ DateTime::from_timestamp_nanos(1234567890123).naive_utc(),
+ ));
+ array_builder.append_value(VariantDecimal4::try_new(123, 2).unwrap());
+ array_builder.append_value(VariantDecimal8::try_new(123, 2).unwrap());
+ array_builder.append_value(VariantDecimal16::try_new(123, 2).unwrap());
+ array_builder.append_value(Variant::Float(5.0));
+ array_builder.append_value(Variant::Double(6f64));
+ array_builder.append_value(Variant::BooleanTrue);
+ array_builder.append_value(Variant::BooleanFalse);
+ array_builder.append_value(Variant::Binary("helow".as_bytes()));
+ array_builder.append_value(Variant::String("hello"));
+ array_builder.append_value(Variant::ShortString(
+ ShortString::try_from("world").unwrap(),
+ ));
+ array_builder.append_value(Variant::Time(
+ NaiveTime::from_num_seconds_from_midnight_opt(12345, 123).unwrap(),
+ ));
+
+ let array = array_builder.build();
+
+ fn can_shred_to(v: &Variant, dt: &DataType) -> bool {
+ matches!(
+ (v, dt),
+ (Variant::Int8(_), DataType::Int8)
+ | (Variant::Int8(_), DataType::Int16)
+ | (Variant::Int8(_), DataType::Int32)
+ | (Variant::Int8(_), DataType::Int64)
+ | (Variant::Int16(_), DataType::Int8)
+ | (Variant::Int16(_), DataType::Int16)
+ | (Variant::Int16(_), DataType::Int32)
+ | (Variant::Int16(_), DataType::Int64)
+ | (Variant::Int32(_), DataType::Int8)
+ | (Variant::Int32(_), DataType::Int16)
+ | (Variant::Int32(_), DataType::Int32)
+ | (Variant::Int32(_), DataType::Int64)
+ | (Variant::Int64(_), DataType::Int8)
+ | (Variant::Int64(_), DataType::Int16)
+ | (Variant::Int64(_), DataType::Int32)
+ | (Variant::Int64(_), DataType::Int64)
+ | (Variant::Date(_), DataType::Date32)
+ | (
+ Variant::TimestampMicros(_),
+ DataType::Timestamp(TimeUnit::Microsecond, Some(_)),
+ )
+ | (
+ Variant::TimestampMicros(_),
+ DataType::Timestamp(TimeUnit::Nanosecond, Some(_))
+ )
+ | (
+ Variant::TimestampNtzMicros(_),
+ DataType::Timestamp(TimeUnit::Microsecond, None),
+ )
+ | (
+ Variant::TimestampNtzMicros(_),
+ DataType::Timestamp(TimeUnit::Nanosecond, None)
+ )
+ | (
+ Variant::TimestampNanos(_),
+ DataType::Timestamp(TimeUnit::Nanosecond, Some(_)),
+ )
+ | (
+ Variant::TimestampNtzNanos(_),
+ DataType::Timestamp(TimeUnit::Nanosecond, None),
+ )
Review Comment:
We have micros to nanos here, but the backwards widening case is missing.
--
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]