Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 105 additions & 3 deletions parquet-variant-compute/src/shred_variant.rs
Original file line number Diff line number Diff line change
Expand Up @@ -169,13 +169,16 @@ pub(crate) fn make_variant_to_shredded_variant_arrow_row_builder<'a>(
)?;
VariantToShreddedVariantRowBuilder::Array(typed_value_builder)
}
// Supported shredded primitive types, see Variant shredding spec:
// https://github.com/apache/parquet-format/blob/master/VariantShredding.md#shredded-value-types
// Supported Arrow canonical Variant primitive types:
// https://arrow.apache.org/docs/format/CanonicalExtensions.html#primitive-type-mappings
DataType::Boolean
| DataType::Int8
| DataType::Int16
| DataType::Int32
| DataType::Int64
| DataType::UInt8
| DataType::UInt16
| DataType::UInt32
| DataType::Float32
| DataType::Float64
| DataType::Decimal32(..)
Expand Down Expand Up @@ -1459,7 +1462,7 @@ mod tests {
let input = VariantArray::from_iter([Variant::from(42)]);

let invalid_types = vec![
DataType::UInt8,
DataType::UInt64,
DataType::Float16,
DataType::Decimal256(38, 10),
DataType::Date64,
Expand Down Expand Up @@ -1503,6 +1506,105 @@ mod tests {
}
}

#[test]
fn test_unsigned_shredding() {
let fallback = Variant::Int8(-1);
let cases = [
(
vec![
Variant::Int16(0),
Variant::Int16(i16::from(u8::MAX)),
fallback.clone(),
],
DataType::UInt8,
),
(
vec![
Variant::Int32(0),
Variant::Int32(i32::from(u16::MAX)),
fallback.clone(),
],
DataType::UInt16,
),
(
vec![
Variant::Int64(0),
Variant::Int64(i64::from(u32::MAX)),
fallback.clone(),
],
DataType::UInt32,
),
];

for (expected, data_type) in cases {
let input = VariantArray::from_iter(expected.iter().cloned());
let shredded = shred_variant(&input, &data_type).unwrap();

assert_eq!(
shredded.typed_value_column().unwrap().data_type(),
&data_type
);
let array = ArrayRef::from(shredded);
let shredded = VariantArray::try_new(array.as_ref()).unwrap();
let unshredded = crate::unshred_variant(&shredded).unwrap();
let fallback_index = expected.len() - 1;
for (index, expected) in expected[..fallback_index].iter().cloned().enumerate() {
assert_eq!(shredded.value(index), expected);
assert_eq!(unshredded.value(index), expected);
}
assert!(
shredded
.typed_value_column()
.unwrap()
.is_null(fallback_index)
);
assert_eq!(shredded.value(fallback_index), fallback);
assert_eq!(unshredded.value(fallback_index), fallback);
}
}

#[test]
fn test_nested_unsigned_shredding() {
let input = build_variant_array(vec![VariantRow::Object(vec![(
"items",
VariantValue::List(vec![VariantValue::from(Variant::Int16(i16::from(u8::MAX)))]),
)])]);
let list_type = DataType::List(Arc::new(Field::new("item", DataType::UInt8, true)));
let target = ShreddedSchemaBuilder::default()
.with_path("items", list_type)
.unwrap()
.build();

let shredded = shred_variant(&input, &target).unwrap();
let object = shredded
.typed_value_column()
.unwrap()
.as_any()
.downcast_ref::<StructArray>()
.unwrap();
let items =
ShreddedVariantFieldArray::try_new(object.column_by_name("items").unwrap()).unwrap();
let items = items
.typed_value_column()
.unwrap()
.as_any()
.downcast_ref::<ListArray>()
.unwrap();
let item = ShreddedVariantFieldArray::try_new(items.values().as_ref()).unwrap();
assert_eq!(
item.typed_value_column().unwrap().data_type(),
&DataType::UInt8
);
let array = ArrayRef::from(shredded);
let shredded = VariantArray::try_new(&array).unwrap();
let unshredded = crate::unshred_variant(&shredded).unwrap();
let value = unshredded.value(0);
let object = value.as_object().unwrap();
let items = object.get("items").unwrap();
let items = items.as_list().unwrap();
assert_eq!(items.get(0), Some(Variant::Int16(i16::from(u8::MAX))));
}

#[test]
fn test_array_shredding_as_list() {
let input = build_variant_array(vec![
Expand Down
15 changes: 14 additions & 1 deletion parquet-variant-compute/src/unshred_variant.rs
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,8 @@ use arrow::array::{
use arrow::datatypes::{
ArrowPrimitiveType, DataType, Date32Type, Decimal32Type, Decimal64Type, Decimal128Type,
DecimalType, Float32Type, Float64Type, Int8Type, Int16Type, Int32Type, Int64Type,
Time64MicrosecondType, TimeUnit, TimestampMicrosecondType, TimestampNanosecondType,
Time64MicrosecondType, TimeUnit, TimestampMicrosecondType, TimestampNanosecondType, UInt8Type,
UInt16Type, UInt32Type,
};
use arrow::error::{ArrowError, Result};
use arrow::temporal_conversions::time64us_to_time;
Expand Down Expand Up @@ -101,6 +102,9 @@ enum UnshredVariantRowBuilder<'a> {
PrimitiveInt16(UnshredPrimitiveRowBuilder<'a, PrimitiveArray<Int16Type>>),
PrimitiveInt32(UnshredPrimitiveRowBuilder<'a, PrimitiveArray<Int32Type>>),
PrimitiveInt64(UnshredPrimitiveRowBuilder<'a, PrimitiveArray<Int64Type>>),
PrimitiveUInt8(UnshredPrimitiveRowBuilder<'a, PrimitiveArray<UInt8Type>>),
PrimitiveUInt16(UnshredPrimitiveRowBuilder<'a, PrimitiveArray<UInt16Type>>),
PrimitiveUInt32(UnshredPrimitiveRowBuilder<'a, PrimitiveArray<UInt32Type>>),
PrimitiveFloat32(UnshredPrimitiveRowBuilder<'a, PrimitiveArray<Float32Type>>),
PrimitiveFloat64(UnshredPrimitiveRowBuilder<'a, PrimitiveArray<Float64Type>>),
Decimal32(DecimalUnshredRowBuilder<'a, Decimal32Type, VariantDecimal4>),
Expand Down Expand Up @@ -146,6 +150,9 @@ impl<'a> UnshredVariantRowBuilder<'a> {
Self::PrimitiveInt16(b) => b.append_row(builder, metadata, index),
Self::PrimitiveInt32(b) => b.append_row(builder, metadata, index),
Self::PrimitiveInt64(b) => b.append_row(builder, metadata, index),
Self::PrimitiveUInt8(b) => b.append_row(builder, metadata, index),
Self::PrimitiveUInt16(b) => b.append_row(builder, metadata, index),
Self::PrimitiveUInt32(b) => b.append_row(builder, metadata, index),
Self::PrimitiveFloat32(b) => b.append_row(builder, metadata, index),
Self::PrimitiveFloat64(b) => b.append_row(builder, metadata, index),
Self::Decimal32(b) => b.append_row(builder, metadata, index),
Expand Down Expand Up @@ -204,6 +211,9 @@ impl<'a> UnshredVariantRowBuilder<'a> {
DataType::Int16 => primitive_builder!(PrimitiveInt16, as_primitive),
DataType::Int32 => primitive_builder!(PrimitiveInt32, as_primitive),
DataType::Int64 => primitive_builder!(PrimitiveInt64, as_primitive),
DataType::UInt8 => primitive_builder!(PrimitiveUInt8, as_primitive),
DataType::UInt16 => primitive_builder!(PrimitiveUInt16, as_primitive),
DataType::UInt32 => primitive_builder!(PrimitiveUInt32, as_primitive),
DataType::Float32 => primitive_builder!(PrimitiveFloat32, as_primitive),
DataType::Float64 => primitive_builder!(PrimitiveFloat64, as_primitive),
DataType::Decimal32(p, s) if VariantDecimal4::is_valid_precision_and_scale(p, s) => {
Expand Down Expand Up @@ -443,6 +453,9 @@ impl_append_to_variant_builder!(PrimitiveArray<Int8Type>);
impl_append_to_variant_builder!(PrimitiveArray<Int16Type>);
impl_append_to_variant_builder!(PrimitiveArray<Int32Type>);
impl_append_to_variant_builder!(PrimitiveArray<Int64Type>);
impl_append_to_variant_builder!(PrimitiveArray<UInt8Type>, |value| i16::from(value));
impl_append_to_variant_builder!(PrimitiveArray<UInt16Type>, |value| i32::from(value));
impl_append_to_variant_builder!(PrimitiveArray<UInt32Type>, |value| i64::from(value));
impl_append_to_variant_builder!(PrimitiveArray<Float32Type>);
impl_append_to_variant_builder!(PrimitiveArray<Float64Type>);

Expand Down
25 changes: 21 additions & 4 deletions parquet-variant-compute/src/variant_array.rs
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ use arrow::compute::cast;
use arrow::datatypes::{
Date32Type, Decimal32Type, Decimal64Type, Decimal128Type, Float16Type, Float32Type,
Float64Type, Int8Type, Int16Type, Int32Type, Int64Type, Time64MicrosecondType,
TimestampMicrosecondType, TimestampNanosecondType,
TimestampMicrosecondType, TimestampNanosecondType, UInt8Type, UInt16Type, UInt32Type,
};
use arrow::error::Result;
use arrow_schema::extension::{ExtensionType, Uuid as UuidExtension};
Expand Down Expand Up @@ -996,6 +996,23 @@ fn typed_value_to_variant<'a>(
DataType::Int64 => {
primitive_conversion_single_value!(Int64Type, typed_value, index)
}
DataType::UInt8 => {
generic_conversion_single_value!(UInt8Type, as_primitive, i16::from, typed_value, index)
}
DataType::UInt16 => generic_conversion_single_value!(
UInt16Type,
as_primitive,
i32::from,
typed_value,
index
),
DataType::UInt32 => generic_conversion_single_value!(
UInt32Type,
as_primitive,
i64::from,
typed_value,
index
),
DataType::Float16 => {
primitive_conversion_single_value!(Float16Type, typed_value, index)
}
Expand Down Expand Up @@ -1141,10 +1158,10 @@ fn canonicalize_and_verify_data_type(data_type: &DataType) -> Result<Cow<'_, Dat
let new_data_type = match data_type {
// Primitive arrow types that have a direct variant counterpart are allowed
Null | Boolean => borrow!(),
Int8 | Int16 | Int32 | Int64 | Float32 | Float64 => borrow!(),
Int8 | Int16 | Int32 | Int64 | UInt8 | UInt16 | UInt32 | Float32 | Float64 => borrow!(),

// Unsigned integers and half-float are not allowed
UInt8 | UInt16 | UInt32 | UInt64 | Float16 => fail!(),
// UInt64 and half-float have no canonical Variant mapping
UInt64 | Float16 => fail!(),

// Most decimal types are allowed, with restrictions on precision and scale
//
Expand Down
31 changes: 31 additions & 0 deletions parquet-variant-compute/src/variant_get.rs
Original file line number Diff line number Diff line change
Expand Up @@ -497,6 +497,7 @@ mod test {
LargeStringArray, ListArray, ListBuilder, ListViewArray, MapBuilder, NullArray,
NullBuilder, StringArray, StringBuilder, StringViewArray, StructArray,
Time32MillisecondArray, Time32SecondArray, Time64MicrosecondArray, Time64NanosecondArray,
UInt8Array, UInt16Array, UInt32Array,
};
use arrow::buffer::{NullBuffer, OffsetBuffer, ScalarBuffer};
use arrow::compute::{CastOptions, cast};
Expand Down Expand Up @@ -679,6 +680,21 @@ mod test {
Int64Array,
i64
);
numeric_partially_shredded_variant_array_fn!(
partially_shredded_uint8_variant_array,
UInt8Array,
u8
);
numeric_partially_shredded_variant_array_fn!(
partially_shredded_uint16_variant_array,
UInt16Array,
u16
);
numeric_partially_shredded_variant_array_fn!(
partially_shredded_uint32_variant_array,
UInt32Array,
u32
);
numeric_partially_shredded_variant_array_fn!(
partially_shredded_float32_variant_array,
Float32Array,
Expand Down Expand Up @@ -739,6 +755,21 @@ mod test {
numeric_partially_shredded_test!(f64, partially_shredded_float64_variant_array);
}

#[test]
fn get_variant_partially_shredded_uint8_as_variant() {
numeric_partially_shredded_test!(i16, partially_shredded_uint8_variant_array);
}

#[test]
fn get_variant_partially_shredded_uint16_as_variant() {
numeric_partially_shredded_test!(i32, partially_shredded_uint16_variant_array);
}

#[test]
fn get_variant_partially_shredded_uint32_as_variant() {
numeric_partially_shredded_test!(i64, partially_shredded_uint32_variant_array);
}

#[test]
fn get_variant_partially_shredded_bool_as_variant() {
let array = partially_shredded_bool_variant_array();
Expand Down
5 changes: 5 additions & 0 deletions parquet/src/arrow/arrow_writer/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1640,6 +1640,11 @@ fn write_leaf(
let array = column.as_primitive::<Int64Type>();
write_primitive(typed, array.values(), levels)
}
ArrowDataType::UInt32 => {
let array: arrow_array::Int64Array =
column.as_primitive::<UInt32Type>().unary(i64::from);
write_primitive(typed, array.values(), levels)
}
ArrowDataType::UInt64 => {
let values = column.as_primitive::<UInt64Type>().values();
// follow C++ implementation and use overflow/reinterpret cast from u64 to i64 which will map
Expand Down
9 changes: 9 additions & 0 deletions parquet/src/arrow/schema/complex.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@ use std::collections::HashMap;
use std::sync::Arc;

use crate::arrow::schema::extension::try_add_extension_type;
#[cfg(feature = "variant_experimental")]
use crate::arrow::schema::extension::validate_variant_type;
use crate::arrow::schema::primitive::convert_primitive;
use crate::arrow::schema::virtual_type::{RowGroupIndex, RowNumber};
use crate::arrow::{PARQUET_FIELD_ID_META_KEY, ProjectionMask};
Expand Down Expand Up @@ -747,6 +749,13 @@ fn convert_field(

match arrow_hint {
Some(hint) => {
#[cfg(feature = "variant_experimental")]
if matches!(
parquet_type.get_basic_info().logical_type_ref(),
Some(crate::basic::LogicalType::Variant(_))
) {
validate_variant_type(parquet_type)?;
}
// If the inferred type is a dictionary, preserve dictionary metadata
#[allow(deprecated)]
let field = match (&data_type, hint.dict_id(), hint.dict_is_ordered()) {
Expand Down
34 changes: 34 additions & 0 deletions parquet/src/arrow/schema/extension.rs
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ pub(crate) fn try_add_extension_type(
Ok(match parquet_logical_type {
#[cfg(feature = "variant_experimental")]
LogicalType::Variant(_) => {
validate_variant_type(parquet_type)?;
let mut arrow_field = arrow_field;
arrow_field.try_with_extension_type(parquet_variant_compute::VariantType)?;
arrow_field
Expand Down Expand Up @@ -99,6 +100,39 @@ pub(crate) fn try_add_extension_type(
})
}

#[cfg(feature = "variant_experimental")]
pub(crate) fn validate_variant_type(parquet_type: &Type) -> Result<(), ParquetError> {
use crate::basic::ConvertedType;

match parquet_type {
Type::PrimitiveType { basic_info, .. } => {
let bit_width = match basic_info.logical_type_ref() {
Some(LogicalType::Integer(integer)) if !integer.is_signed => {
Some(integer.bit_width)
}
_ => match basic_info.converted_type() {
ConvertedType::UINT_8 => Some(8),
ConvertedType::UINT_16 => Some(16),
ConvertedType::UINT_32 => Some(32),
ConvertedType::UINT_64 => Some(64),
_ => None,
},
};
if let Some(bit_width) = bit_width {
return Err(ParquetError::General(format!(
"Illegal shredded value type: UInt{bit_width}"
)));
}
}
Type::GroupType { fields, .. } => {
for field in fields {
validate_variant_type(field)?;
}
}
}
Ok(())
}

/// Returns true if [`try_add_extension_type`] would add an extension type
/// to the specified Parquet field.
///
Expand Down
Loading