From be5c547eab4335ae2ef54957a6258f21dd8750a4 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sun, 6 Sep 2026 00:33:05 +0800 Subject: [PATCH] feat: normalize marked Variant arrays at the native Parquet boundary --- native/core/src/parquet/cast_column.rs | 16 + .../core/src/parquet/cast_column/variant.rs | 570 ++++++++++++++++++ .../src/parquet/cast_column/variant/tests.rs | 367 +++++++++++ native/core/src/parquet/schema_adapter.rs | 99 ++- 4 files changed, 1049 insertions(+), 3 deletions(-) create mode 100644 native/core/src/parquet/cast_column/variant.rs create mode 100644 native/core/src/parquet/cast_column/variant/tests.rs diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index 231f88c5074..8c72455f4db 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -14,6 +14,9 @@ // KIND, either express or implied. See the License for the // specific language governing permissions and limitations // under the License. +mod variant; + +use self::variant::normalize_variant_array; use arrow::{ array::{make_array, Array, ArrayRef, LargeListArray, ListArray, MapArray, StructArray}, compute::CastOptions, @@ -26,6 +29,7 @@ use datafusion::common::format::DEFAULT_CAST_OPTIONS; use datafusion::common::{DataFusionError, Result as DataFusionResult}; use datafusion::logical_expr::ColumnarValue; use datafusion::physical_expr::PhysicalExpr; +use parquet::variant::VariantType; use std::{ fmt::{self, Display}, hash::Hash, @@ -243,6 +247,18 @@ impl PhysicalExpr for CometCastColumnExpr { fn evaluate(&self, batch: &RecordBatch) -> DataFusionResult { let value = self.expr.evaluate(batch)?; + if self.target_field.has_valid_extension_type::() { + return match value { + ColumnarValue::Array(array) => Ok(ColumnarValue::Array(normalize_variant_array( + &array, + &self.target_field, + )?)), + ColumnarValue::Scalar(_) => Err(DataFusionError::Execution( + "Variant Parquet projection requires an array".to_string(), + )), + }; + } + // Use == (PartialEq) instead of equals_datatype because equals_datatype // ignores field names in nested types (Struct, List, Map). We need to detect // when field names differ (e.g., Struct("a","b") vs Struct("c","d")) so that diff --git a/native/core/src/parquet/cast_column/variant.rs b/native/core/src/parquet/cast_column/variant.rs new file mode 100644 index 00000000000..51fbde70fd4 --- /dev/null +++ b/native/core/src/parquet/cast_column/variant.rs @@ -0,0 +1,570 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use arrow::{ + array::{ + make_array, Array, ArrayRef, AsArray, BinaryArray, BinaryBuilder, ListLikeArray, + StructArray, + }, + buffer::NullBuffer, + compute::cast, + datatypes::{DataType, FieldRef}, + error::ArrowError, +}; +use datafusion::common::{DataFusionError, Result as DataFusionResult}; +use parquet::variant::{ + unshred_variant, ListBuilder, MetadataBuilder, ObjectBuilder, ParentState, + ReadOnlyMetadataBuilder, ValueBuilder, Variant, VariantArray, VariantMetadata, +}; +use std::{ + panic::{catch_unwind, AssertUnwindSafe}, + sync::Arc, +}; + +pub(super) fn normalize_variant_array( + array: &ArrayRef, + target_field: &FieldRef, +) -> DataFusionResult { + let DataType::Struct(fields) = target_field.data_type() else { + return Err(DataFusionError::Execution( + "Variant extension field must use Struct storage".to_string(), + )); + }; + if fields.len() != 2 + || fields[0].name() != "value" + || fields[1].name() != "metadata" + || fields + .iter() + .any(|field| field.data_type() != &DataType::Binary) + { + return Err(DataFusionError::Execution( + "Variant output must contain Binary children [value, metadata]".to_string(), + )); + } + + // VariantArray resolves metadata/value/typed_value by name, so the reader's child order is + // irrelevant. Legacy Spark residuals must be put in Arrow order before the single upstream + // unshred call; the whole output is then put back in the order expected by released Spark 4. + let variant = VariantArray::try_new(array.as_ref())?; + let prepared = prepare_variant_for_unshredding(&variant)?; + let unshredded = unshred_variant(&prepared)?; + let value = unshredded.value_field().ok_or_else(|| { + DataFusionError::Execution("Unshredded Variant is missing its value field".to_string()) + })?; + let value = cast(value.as_ref(), &DataType::Binary)?; + let metadata = cast(unshredded.metadata_field().as_ref(), &DataType::Binary)?; + let value = reorder_variant_values(&value, &metadata, unshredded.inner().nulls())?; + + Ok(Arc::new(StructArray::try_new( + fields.clone(), + vec![value, metadata], + unshredded.inner().nulls().cloned(), + )?)) +} + +/// Arrow validates every residual `value` while unshredding. Spark versions before SPARK-58949 +/// wrote object keys in Java UTF-16 order, so rewrite every reachable legacy residual to Arrow's +/// UTF-8 order before unshredding. `metadata_rows` carries each root metadata row through nested +/// lists. +fn rewrite_shredding_state( + state: &StructArray, + metadata: &BinaryArray, + metadata_rows: &[Option], +) -> DataFusionResult<(ArrayRef, bool)> { + if state.len() != metadata_rows.len() { + return Err(DataFusionError::Execution( + "Variant shredding state and metadata row mapping have different lengths".to_string(), + )); + } + + let active_rows = metadata_rows + .iter() + .enumerate() + .map(|(index, row)| state.is_valid(index).then_some(*row).flatten()) + .collect::>(); + let mut fields = state.fields().iter().cloned().collect::>(); + let mut columns = state.columns().to_vec(); + let mut changed = false; + + if let Some(index) = fields.iter().position(|field| field.name() == "value") { + let (value, value_changed) = + rewrite_residual_values(&columns[index], metadata, &active_rows)?; + if value_changed { + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(value.data_type().clone()), + ); + columns[index] = value; + changed = true; + } + } + + if let Some(index) = fields + .iter() + .position(|field| field.name() == "typed_value") + { + let typed_rows = active_rows + .iter() + .enumerate() + .map(|(row, metadata)| columns[index].is_valid(row).then_some(*metadata).flatten()) + .collect::>(); + let (typed_value, typed_changed) = + rewrite_typed_value(&columns[index], metadata, &typed_rows)?; + if typed_changed { + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(typed_value.data_type().clone()), + ); + columns[index] = typed_value; + changed = true; + } + } + + if !changed { + return Ok((Arc::new(state.clone()), false)); + } + Ok(( + Arc::new(StructArray::try_new( + fields.into(), + columns, + state.nulls().cloned(), + )?), + true, + )) +} + +fn rewrite_residual_values( + value: &ArrayRef, + metadata: &BinaryArray, + metadata_rows: &[Option], +) -> DataFusionResult<(ArrayRef, bool)> { + let binary = cast(value.as_ref(), &DataType::Binary)?; + let binary = binary.as_binary::(); + let mut output = BinaryBuilder::new(); + let mut changed = false; + + for (index, metadata_row) in metadata_rows.iter().enumerate() { + if binary.is_null(index) { + output.append_null(); + continue; + } + let Some(metadata_row) = metadata_row else { + output.append_value(binary.value(index)); + continue; + }; + if metadata.is_null(*metadata_row) { + return Err(DataFusionError::Execution(format!( + "Variant metadata is null at row {metadata_row}" + ))); + } + + let rebuilt = catch_unwind(AssertUnwindSafe( + || -> Result>, ArrowError> { + let metadata = VariantMetadata::try_new(metadata.value(*metadata_row))?; + let variant = Variant::new_with_metadata(metadata.clone(), binary.value(index)); + if is_compatible_variant(&variant, VariantObjectKeyOrder::ArrowUtf8) { + return Ok(None); + } + if !is_compatible_variant(&variant, VariantObjectKeyOrder::SparkUtf16) { + return Err(ArrowError::InvalidArgumentError( + "Variant residual is neither UTF-8 nor Spark UTF-16 ordered".to_string(), + )); + } + Ok(Some(variant_bytes( + &metadata, + variant, + VariantObjectKeyOrder::ArrowUtf8, + )?)) + }, + )) + .map_err(|_| { + DataFusionError::Execution(format!("Invalid Variant residual at row {metadata_row}")) + })??; + changed |= rebuilt.is_some(); + output.append_value(rebuilt.as_deref().unwrap_or_else(|| binary.value(index))); + } + + if changed { + Ok((Arc::new(output.finish()), true)) + } else { + Ok((Arc::clone(value), false)) + } +} + +fn list_metadata_rows( + list: &L, + parent_rows: &[Option], +) -> DataFusionResult>> { + let mut child_rows = vec![None; list.values().len()]; + for (index, metadata_row) in parent_rows.iter().enumerate() { + let Some(metadata_row) = metadata_row else { + continue; + }; + for child_index in list.element_range(index) { + match child_rows[child_index] { + Some(existing) if existing != *metadata_row => { + return Err(DataFusionError::Execution( + "A shared Variant list child refers to different metadata rows".to_string(), + )); + } + _ => child_rows[child_index] = Some(*metadata_row), + } + } + } + Ok(child_rows) +} + +fn rewrite_list_typed_value( + array: &ArrayRef, + list: &L, + metadata: &BinaryArray, + metadata_rows: &[Option], +) -> DataFusionResult<(ArrayRef, bool)> { + let child_rows = list_metadata_rows(list, metadata_rows)?; + let values = list.values().as_struct_opt().ok_or_else(|| { + DataFusionError::Execution(format!( + "Invalid shredded Variant list values: expected Struct, got {}", + list.values().data_type() + )) + })?; + let (values, changed) = rewrite_shredding_state(values, metadata, &child_rows)?; + if !changed { + return Ok((Arc::clone(array), false)); + } + + let data_type = match array.data_type() { + DataType::List(field) => DataType::List(Arc::new( + field + .as_ref() + .clone() + .with_data_type(values.data_type().clone()), + )), + DataType::LargeList(field) => DataType::LargeList(Arc::new( + field + .as_ref() + .clone() + .with_data_type(values.data_type().clone()), + )), + DataType::ListView(field) => DataType::ListView(Arc::new( + field + .as_ref() + .clone() + .with_data_type(values.data_type().clone()), + )), + DataType::LargeListView(field) => DataType::LargeListView(Arc::new( + field + .as_ref() + .clone() + .with_data_type(values.data_type().clone()), + )), + data_type => { + return Err(DataFusionError::Execution(format!( + "Expected a Variant list, got {data_type}" + ))); + } + }; + let data = array + .to_data() + .into_builder() + .data_type(data_type) + .child_data(vec![values.to_data()]) + .build()?; + Ok((make_array(data), true)) +} + +fn rewrite_typed_value( + typed_value: &ArrayRef, + metadata: &BinaryArray, + metadata_rows: &[Option], +) -> DataFusionResult<(ArrayRef, bool)> { + match typed_value.data_type() { + DataType::Struct(_) => { + let object = typed_value.as_struct(); + let mut fields = object.fields().iter().cloned().collect::>(); + let mut columns = object.columns().to_vec(); + let mut changed = false; + for (index, column) in object.columns().iter().enumerate() { + let child = column.as_struct_opt().ok_or_else(|| { + DataFusionError::Execution(format!( + "Invalid shredded Variant object field '{}': expected Struct, got {}", + fields[index].name(), + column.data_type() + )) + })?; + let (child, child_changed) = + rewrite_shredding_state(child, metadata, metadata_rows)?; + if child_changed { + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(child.data_type().clone()), + ); + columns[index] = child; + changed = true; + } + } + if !changed { + return Ok((Arc::clone(typed_value), false)); + } + Ok(( + Arc::new(StructArray::try_new( + fields.into(), + columns, + object.nulls().cloned(), + )?), + true, + )) + } + DataType::List(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list::(), + metadata, + metadata_rows, + ), + DataType::LargeList(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list::(), + metadata, + metadata_rows, + ), + DataType::ListView(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list_view::(), + metadata, + metadata_rows, + ), + DataType::LargeListView(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list_view::(), + metadata, + metadata_rows, + ), + _ => Ok((Arc::clone(typed_value), false)), + } +} + +fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult { + if variant.typed_value_field().is_none() { + return Ok(variant.clone()); + } + + let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let metadata = metadata.as_binary::(); + let metadata_rows = (0..variant.len()) + .map(|index| variant.inner().is_valid(index).then_some(index)) + .collect::>(); + let (array, changed) = rewrite_shredding_state(variant.inner(), metadata, &metadata_rows)?; + if changed { + Ok(VariantArray::try_new(array.as_ref())?) + } else { + Ok(variant.clone()) + } +} + +/// Supplies sort-only field names whose Rust ordering matches Java `String.compareTo` ordering. +/// Field IDs still come from the original metadata dictionary. +#[derive(Debug)] +struct SparkMetadataBuilder<'a, 'm> { + metadata: &'a VariantMetadata<'m>, + sort_keys: Vec, +} + +impl<'a, 'm> SparkMetadataBuilder<'a, 'm> { + fn new(metadata: &'a VariantMetadata<'m>) -> Self { + let sort_keys = metadata + .iter() + .map(|field_name| { + field_name + .encode_utf16() + .map(|unit| char::from_u32(0x10000 + u32::from(unit)).unwrap()) + .collect() + }) + .collect(); + Self { + metadata, + sort_keys, + } + } +} + +impl MetadataBuilder for SparkMetadataBuilder<'_, '_> { + fn try_upsert_field_name(&mut self, field_name: &str) -> Result { + self.metadata + .get_entry(field_name) + .map(|(field_id, _)| field_id) + .ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Field name '{field_name}' not found in metadata dictionary" + )) + }) + } + + fn field_name(&self, field_id: usize) -> &str { + &self.sort_keys[field_id] + } + + fn num_field_names(&self) -> usize { + self.metadata.len() + } + + fn truncate_field_names(&mut self, new_size: usize) { + debug_assert_eq!(self.metadata.len(), new_size); + } + + fn finish(&mut self) -> usize { + self.metadata.size() + } +} + +#[derive(Clone, Copy)] +enum VariantObjectKeyOrder { + ArrowUtf8, + SparkUtf16, +} + +fn is_compatible_variant(variant: &Variant<'_, '_>, order: VariantObjectKeyOrder) -> bool { + match variant { + Variant::Object(object) => { + let mut previous = None; + object.iter().all(|(name, value)| { + let ordered = previous + .map(|previous: &str| match order { + VariantObjectKeyOrder::ArrowUtf8 => previous <= name, + VariantObjectKeyOrder::SparkUtf16 => { + previous.encode_utf16().cmp(name.encode_utf16()) + != std::cmp::Ordering::Greater + } + }) + .unwrap_or(true); + previous = Some(name); + ordered && is_compatible_variant(&value, order) + }) + } + Variant::List(list) => list + .iter() + .all(|value| is_compatible_variant(&value, order)), + _ => true, + } +} + +fn variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, + order: VariantObjectKeyOrder, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + match variant { + Variant::Object(object) => { + let mut metadata_builder: Box = match order { + VariantObjectKeyOrder::ArrowUtf8 => { + Box::new(ReadOnlyMetadataBuilder::new(metadata)) + } + VariantObjectKeyOrder::SparkUtf16 => Box::new(SparkMetadataBuilder::new(metadata)), + }; + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, metadata_builder.as_mut()), + matches!(order, VariantObjectKeyOrder::ArrowUtf8), + ); + for (name, value) in object.iter() { + let value = variant_bytes(metadata, value, order)?; + builder + .try_insert_bytes(name, Variant::new_with_metadata(metadata.clone(), &value))?; + } + builder.finish(); + } + Variant::List(list) => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + let mut builder = ListBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + matches!(order, VariantObjectKeyOrder::ArrowUtf8), + ); + for value in list.iter() { + let value = variant_bytes(metadata, value, order)?; + builder.append_value_bytes(Variant::new_with_metadata(metadata.clone(), &value)); + } + builder.finish(); + } + variant => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + ValueBuilder::try_append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + variant, + )?; + } + } + Ok(value_builder.into_inner()) +} + +/// Released Spark 4 profiles search object fields in Java UTF-16 order. Convert whole-value output +/// to that order until #5474 can remove this rewrite after every supported profile includes +/// SPARK-58949. Values already in the requested order remain byte-for-byte unchanged. +/// https://github.com/apache/datafusion-comet/issues/5474 +fn reorder_variant_values( + value: &ArrayRef, + metadata: &ArrayRef, + parent_nulls: Option<&NullBuffer>, +) -> DataFusionResult { + let value = value.as_binary::(); + let metadata = metadata.as_binary::(); + let mut output = BinaryBuilder::new(); + + for index in 0..value.len() { + if parent_nulls.is_some_and(|nulls| nulls.is_null(index)) { + output.append_null(); + continue; + } + if value.is_null(index) { + return Err(DataFusionError::Execution(format!( + "Variant value is null at row {index}" + ))); + } + if metadata.is_null(index) { + return Err(DataFusionError::Execution(format!( + "Variant metadata is null at row {index}" + ))); + } + + let rebuilt = catch_unwind(AssertUnwindSafe( + || -> Result>, ArrowError> { + let metadata = VariantMetadata::try_new(metadata.value(index))?; + let variant = Variant::new_with_metadata(metadata.clone(), value.value(index)); + if is_compatible_variant(&variant, VariantObjectKeyOrder::SparkUtf16) { + return Ok(None); + } + Ok(Some(variant_bytes( + &metadata, + variant, + VariantObjectKeyOrder::SparkUtf16, + )?)) + }, + )) + .map_err(|_| { + DataFusionError::Execution(format!("Invalid Variant value at row {index}")) + })??; + output.append_value(rebuilt.as_deref().unwrap_or_else(|| value.value(index))); + } + + Ok(Arc::new(output.finish())) +} + +#[cfg(test)] +mod tests; diff --git a/native/core/src/parquet/cast_column/variant/tests.rs b/native/core/src/parquet/cast_column/variant/tests.rs new file mode 100644 index 00000000000..fcad80c8e15 --- /dev/null +++ b/native/core/src/parquet/cast_column/variant/tests.rs @@ -0,0 +1,367 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use super::*; +use arrow::{ + array::{Int64Array, ListArray}, + buffer::OffsetBuffer, + datatypes::{Field, Fields}, +}; +use parquet::variant::{ + shred_variant, VariantArrayBuilder, VariantBuilder, VariantBuilderExt, VariantType, +}; + +fn target_field(nullable: bool) -> FieldRef { + Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + nullable, + ) + .with_extension_type(VariantType), + ) +} + +fn unicode_object_keys() -> Vec { + let mut keys = (0..30).map(|i| format!("k{i:02}")).collect::>(); + keys.push("\u{e000}".to_string()); + keys.push("๐Ÿ˜€".to_string()); + keys +} + +fn assert_spark_unicode_variant(variant: Variant<'_, '_>) { + let Variant::Object(object) = variant else { + panic!("expected object") + }; + let fields = object.iter().collect::>(); + assert_eq!(fields.len(), 32); + assert_eq!(fields[30].0, "๐Ÿ˜€"); + assert_eq!(fields[31].0, "\u{e000}"); + assert_eq!(object.get("๐Ÿ˜€").unwrap().as_int64(), Some(531)); +} + +fn assert_spark_unicode_output(output: &StructArray) { + let value = output.column(0).as_binary::(); + let metadata = output.column(1).as_binary::(); + assert_spark_unicode_variant(Variant::new(metadata.value(0), value.value(0))); +} + +#[test] +fn normalize_full_shredding_reorders_children_and_preserves_parent_nulls() { + let mut builder = VariantArrayBuilder::new(3); + builder.append_variant(Variant::from(1_i64)); + builder.append_null(); + builder.append_variant(Variant::from(3_i64)); + let base = builder.build(); + let metadata = Arc::clone(base.metadata_field()); + let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![Some(10), None, Some(30)])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("typed_value", DataType::Int64, true), + Field::new("metadata", metadata.data_type().clone(), false), + ]), + vec![typed_value, metadata], + base.inner().nulls().cloned(), + ) + .unwrap(), + ); + + let output = normalize_variant_array(&physical, &target_field(true)).unwrap(); + let output = output.as_struct(); + assert_eq!( + output + .fields() + .iter() + .map(|field| field.name().as_str()) + .collect::>(), + ["value", "metadata"] + ); + assert!(output + .columns() + .iter() + .all(|column| column.data_type() == &DataType::Binary)); + assert!(output.is_null(1)); + + let variant = VariantArray::try_new(output).unwrap(); + assert_eq!(variant.value(0), Variant::from(10_i64)); + assert_eq!(variant.value(2), Variant::from(30_i64)); +} + +#[test] +fn normalize_fully_shredded_object_orders_for_spark() { + let keys = unicode_object_keys(); + let (metadata, _) = VariantBuilder::new() + .with_field_names(keys.iter().map(String::as_str)) + .finish(); + let mut fields = Vec::with_capacity(keys.len()); + let mut columns = Vec::with_capacity(keys.len()); + for (index, key) in keys.iter().enumerate() { + let state = StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![if key == "๐Ÿ˜€" { + 531 + } else { + index as i64 + }]))], + None, + ) + .unwrap(); + fields.push(Field::new(key, state.data_type().clone(), false)); + columns.push(Arc::new(state) as ArrayRef); + } + let typed_value: ArrayRef = + Arc::new(StructArray::try_new(fields.into(), columns, None).unwrap()); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]), + vec![ + Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])), + typed_value, + ], + None, + ) + .unwrap(), + ); + + let output = normalize_variant_array(&physical, &target_field(false)).unwrap(); + assert_spark_unicode_output(output.as_struct()); +} + +#[test] +fn canonical_and_shredded_values_normalize_equally() { + let mut builder = VariantArrayBuilder::new(6); + builder.new_object().with_field("known", 1_i64).finish(); + builder + .new_object() + .with_field("known", 2_i64) + .with_field("extra", 3_i64) + .finish(); + builder + .new_list() + .with_value(4_i64) + .with_value(Variant::Null) + .finish(); + builder.append_variant(Variant::from(5_i64)); + builder.append_variant(Variant::Null); + builder.append_null(); + let canonical = builder.build(); + let shredded = shred_variant( + &canonical, + &DataType::Struct(Fields::from(vec![Field::new( + "known", + DataType::Int64, + true, + )])), + ) + .unwrap(); + + let normalize = |array: &VariantArray| { + let array: ArrayRef = Arc::new(array.inner().clone()); + let output = normalize_variant_array(&array, &target_field(true)).unwrap(); + VariantArray::try_new(output.as_ref()).unwrap() + }; + let canonical = normalize(&canonical); + let shredded = normalize(&shredded); + + for index in 0..canonical.len() { + assert_eq!(canonical.is_null(index), shredded.is_null(index)); + if canonical.is_valid(index) { + assert_eq!(canonical.value(index), shredded.value(index)); + } + } +} + +#[test] +fn normalize_unshredded_variant_orders_for_spark_and_is_idempotent() { + let keys = unicode_object_keys(); + let mut builder = VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); + let mut object = builder.new_object(); + for (index, key) in keys.iter().enumerate() { + object.insert(key, if key == "๐Ÿ˜€" { 531 } else { index as i64 }); + } + object.finish(); + let (metadata, value) = builder.finish(); + + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("value", DataType::Binary, false), + ]), + vec![ + Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])), + Arc::new(BinaryArray::from(vec![Some(value.as_slice())])), + ], + None, + ) + .unwrap(), + ); + + let first = normalize_variant_array(&physical, &target_field(false)).unwrap(); + assert_spark_unicode_output(first.as_struct()); + let first_value = first.as_struct().column(0).as_binary::().value(0); + + let second = normalize_variant_array(&first, &target_field(false)).unwrap(); + assert_spark_unicode_output(second.as_struct()); + assert_eq!( + second.as_struct().column(0).as_binary::().value(0), + first_value + ); +} + +#[test] +fn normalize_partially_shredded_legacy_residual() { + let keys = unicode_object_keys(); + let mut builder = VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); + let mut object = builder.new_object(); + for (index, key) in keys.iter().enumerate().skip(1) { + object.insert(key, if key == "๐Ÿ˜€" { 531 } else { index as i64 }); + } + object.finish(); + let (metadata, value) = builder.finish(); + let metadata_array: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])); + let value_array: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value.as_slice())])); + let legacy_value = reorder_variant_values(&value_array, &metadata_array, None).unwrap(); + + let shredded: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![0]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("k00", shredded.data_type().clone(), false)]), + vec![shredded], + None, + ) + .unwrap(), + ); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("value", DataType::Binary, true), + Field::new("typed_value", typed_value.data_type().clone(), true), + ]), + vec![metadata_array, legacy_value, typed_value], + None, + ) + .unwrap(), + ); + + let output = normalize_variant_array(&physical, &target_field(false)).unwrap(); + assert_spark_unicode_output(output.as_struct()); +} + +#[test] +fn normalize_nested_list_residuals_use_their_root_metadata() { + fn legacy_row(keys: &[&str]) -> (Vec, Vec) { + let mut builder = VariantBuilder::new().with_field_names(keys.iter().copied()); + let mut object = builder.new_object(); + for (index, key) in keys.iter().enumerate() { + object.insert(key, index as i64); + } + object.finish(); + let (metadata, value) = builder.finish(); + let metadata_array: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])); + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value.as_slice())])); + let value = reorder_variant_values(&value, &metadata_array, None).unwrap(); + (metadata, value.as_binary::().value(0).to_vec()) + } + + let (metadata0, value0) = legacy_row(&["a", "\u{e000}", "๐Ÿ˜€"]); + let (metadata1, value1) = legacy_row(&["b", "zz", "\u{ffff}", "๐€€"]); + let states: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("value", DataType::Binary, true)]), + vec![Arc::new(BinaryArray::from(vec![ + Some(value0.as_slice()), + Some(value1.as_slice()), + ]))], + None, + ) + .unwrap(), + ); + let list: ArrayRef = Arc::new( + ListArray::try_new( + Arc::new(Field::new("element", states.data_type().clone(), false)), + OffsetBuffer::new(vec![0, 1, 2].into()), + states, + None, + ) + .unwrap(), + ); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![ + Some(metadata0.as_slice()), + Some(metadata1.as_slice()), + ])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", list.data_type().clone(), false), + ]), + vec![metadata, list], + None, + ) + .unwrap(), + ); + + let output = normalize_variant_array(&physical, &target_field(false)).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + for (index, key) in ["๐Ÿ˜€", "๐€€"].into_iter().enumerate() { + let Variant::List(list) = output.value(index) else { + panic!("expected list") + }; + let Variant::Object(object) = list.get(0).unwrap() else { + panic!("expected object") + }; + assert_eq!(object.get(key).unwrap().as_int64(), Some(index as i64 + 2)); + } +} + +#[test] +fn normalize_null_parent_ignores_empty_children() { + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]), + vec![ + Arc::new(BinaryArray::from(vec![Some(&b""[..])])), + Arc::new(BinaryArray::from(vec![Some(&b""[..])])), + ], + Some(NullBuffer::from(vec![false])), + ) + .unwrap(), + ); + + let output = normalize_variant_array(&physical, &target_field(true)).unwrap(); + assert!(output.is_null(0)); + assert!(output.as_struct().column(0).is_null(0)); +} diff --git a/native/core/src/parquet/schema_adapter.rs b/native/core/src/parquet/schema_adapter.rs index 8b362618221..b5d83a5864a 100644 --- a/native/core/src/parquet/schema_adapter.rs +++ b/native/core/src/parquet/schema_adapter.rs @@ -33,7 +33,7 @@ use datafusion_physical_expr_adapter::{ replace_columns_with_literals, DefaultPhysicalExprAdapterFactory, PhysicalExprAdapter, PhysicalExprAdapterFactory, }; -use parquet::arrow::PARQUET_FIELD_ID_META_KEY; +use parquet::{arrow::PARQUET_FIELD_ID_META_KEY, variant::VariantType}; use std::collections::{HashMap, HashSet}; use std::fmt::{self, Display}; use std::hash::{Hash, Hasher}; @@ -577,6 +577,7 @@ impl PhysicalExprAdapter for SparkPhysicalExprAdapter { self.wrap_all_type_mismatches(expr)? } }; + let expr = self.wrap_direct_variant_column(expr)?; // For case-insensitive mode: remap column names from logical back to // original physical names. The default adapter was given a remapped @@ -605,6 +606,37 @@ impl PhysicalExprAdapter for SparkPhysicalExprAdapter { } impl SparkPhysicalExprAdapter { + /// The default adapter leaves an identical physical/logical Field as a bare Column. Variant + /// still needs normalization because a canonical unshredded value may already have that exact + /// marked layout. + fn wrap_direct_variant_column( + &self, + expr: Arc, + ) -> DataFusionResult> { + let Some(column) = expr.downcast_ref::() else { + return Ok(expr); + }; + let Ok(logical_field) = self.logical_file_schema.field_with_name(column.name()) else { + return Ok(expr); + }; + if !logical_field.has_valid_extension_type::() { + return Ok(expr); + } + let Some(physical_field) = self.physical_file_schema.fields().get(column.index()) else { + return Ok(expr); + }; + + Ok(Arc::new( + CometCastColumnExpr::try_new( + expr, + Arc::clone(physical_field), + Arc::new(logical_field.clone()), + None, + )? + .with_parquet_options(self.parquet_options.clone()), + )) + } + /// Wrap ALL Column expressions that have type mismatches with CometCastColumnExpr. /// This is the fallback path when the default adapter fails (e.g., for complex /// nested type casts like List or Map). Uses `spark_parquet_convert` @@ -650,7 +682,9 @@ impl SparkPhysicalExprAdapter { Arc::clone(&e) }; - if logical_field.data_type() != physical_field.data_type() { + if logical_field.has_valid_extension_type::() + || logical_field.data_type() != physical_field.data_type() + { // Mirror the same string/binary -> non-string/binary rejection in // `replace_with_spark_cast`; this branch is reached when the default // adapter rejected the cast and we'd otherwise build a CometCastColumnExpr @@ -709,6 +743,22 @@ impl SparkPhysicalExprAdapter { }; let physical_type = input_field.data_type(); + if cast + .target_field() + .has_valid_extension_type::() + { + let comet_cast: Arc = Arc::new( + CometCastColumnExpr::try_new( + child, + input_field, + Arc::clone(cast.target_field()), + None, + )? + .with_parquet_options(self.parquet_options.clone()), + ); + return Ok(Transformed::yes(comet_cast)); + } + // Identity cast: DataFusion's default adapter inserts a CastExpr // whenever the logical and physical Arrow Fields differ in any // attribute (data type, nullability, or metadata), so with identical @@ -1161,6 +1211,7 @@ impl PhysicalExpr for RejectOnNonEmpty { #[cfg(test)] mod test { + use crate::parquet::cast_column::CometCastColumnExpr; use crate::parquet::parquet_support::SparkParquetOptions; use crate::parquet::schema_adapter::SparkPhysicalExprAdapterFactory; use arrow::array::UInt32Array; @@ -1169,7 +1220,7 @@ mod test { Int64Array, StringArray, TimestampMicrosecondArray, }; use arrow::datatypes::SchemaRef; - use arrow::datatypes::{DataType, Field, Schema}; + use arrow::datatypes::{DataType, Field, Fields, Schema}; use arrow::record_batch::RecordBatch; use datafusion::common::DataFusionError; use datafusion::datasource::listing::PartitionedFile; @@ -1177,6 +1228,8 @@ mod test { use datafusion::datasource::source::DataSourceExec; use datafusion::execution::object_store::ObjectStoreUrl; use datafusion::execution::TaskContext; + use datafusion::physical_expr::expressions::Column; + use datafusion::physical_expr::PhysicalExpr; use datafusion::physical_plan::ExecutionPlan; use datafusion_comet_spark_expr::test_common::file_util::get_temp_filename; use datafusion_comet_spark_expr::EvalMode; @@ -1184,6 +1237,7 @@ mod test { use futures::StreamExt; use parquet::arrow::ArrowWriter; use parquet::arrow::PARQUET_FIELD_ID_META_KEY; + use parquet::variant::VariantType; use std::collections::HashMap; use std::fs::File; use std::sync::Arc; @@ -1935,4 +1989,43 @@ mod test { rewritten.err() ); } + + #[test] + fn marked_variant_column_is_wrapped_after_unicode_name_remap() { + let storage = DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])); + let logical = Arc::new(Schema::new(vec![Field::new( + "mรผnchen", + storage.clone(), + true, + ) + .with_extension_type(VariantType)])); + let physical = Arc::new(Schema::new(vec![ + Field::new("MรœNCHEN", storage, true).with_extension_type(VariantType) + ])); + let mut options = SparkParquetOptions::new(EvalMode::Legacy, "UTC", false); + options.case_sensitive = false; + let adapter = SparkPhysicalExprAdapterFactory::new(options, None) + .create(Arc::clone(&logical), Arc::clone(&physical)) + .unwrap(); + + let expr: Arc = Arc::new(Column::new("mรผnchen", 0)); + let rewritten = adapter.rewrite(expr).unwrap(); + let cast = rewritten + .downcast_ref::() + .expect("marked Variant must retain its normalization wrapper"); + assert_eq!( + cast.children()[0] + .downcast_ref::() + .expect("normalization input must remain a column") + .name(), + "MรœNCHEN" + ); + assert!(rewritten + .return_field(&physical) + .unwrap() + .has_valid_extension_type::()); + } }