From 45a0ed44ed9c58ede31410de077a0882e72fd4f8 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Fri, 21 Aug 2026 19:10:04 +0800 Subject: [PATCH 01/14] feat: project Spark 4 VARIANT columns in native Parquet scans --- native/core/src/execution/jni_api.rs | 6 +- native/core/src/execution/planner.rs | 10 +- native/core/src/execution/serde.rs | 57 +++++--- native/core/src/execution/utils.rs | 9 +- native/core/src/parquet/cast_column.rs | 121 ++++++++++++++++- native/core/src/parquet/mod.rs | 6 +- native/core/src/parquet/schema_adapter.rs | 18 ++- native/proto/src/proto/types.proto | 1 + .../apache/comet/rules/CometExecRule.scala | 12 +- .../apache/comet/rules/CometScanRule.scala | 6 +- .../apache/comet/serde/QueryPlanSerde.scala | 1 + .../apache/comet/serde/namedExpressions.scala | 9 +- .../serde/operator/CometNativeScan.scala | 15 +-- .../apache/spark/sql/comet/util/Utils.scala | 19 ++- .../apache/comet/shims/CometTypeShim.scala | 6 + .../apache/comet/shims/CometTypeShim.scala | 23 +++- .../sql-tests/expressions/misc/variant.sql | 34 +++-- .../comet/parquet/ParquetReadSuite.scala | 123 +++++++++++++++++- 18 files changed, 402 insertions(+), 74 deletions(-) diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index d80754736b7..9d0e9ea7c82 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -686,6 +686,7 @@ fn prepare_output( let schema_addrs = unsafe { schema_addrs.get_elements(env, ReleaseMode::NoCopyBack)? }; let schema_addrs = &*schema_addrs; + let output_schema = output_batch.schema(); let results = output_batch.columns(); let num_rows = output_batch.num_rows(); @@ -712,6 +713,7 @@ fn prepare_output( let mut i = 0; while i < results.len() { let array_ref = results.get(i).ok_or(CometError::IndexOutOfBounds(i))?; + let field = output_schema.field(i); if array_ref.offset() != 0 { // https://github.com/apache/datafusion-comet/issues/2051 @@ -728,11 +730,11 @@ fn prepare_output( new_array .to_data() - .move_to_spark(array_addrs[i], schema_addrs[i])?; + .move_to_spark(field, array_addrs[i], schema_addrs[i])?; } else { array_ref .to_data() - .move_to_spark(array_addrs[i], schema_addrs[i])?; + .move_to_spark(field, array_addrs[i], schema_addrs[i])?; } i += 1; } diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index d109627e825..6ef23ea163c 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -43,7 +43,7 @@ use crate::execution::{ }, planner::expression_registry::ExpressionRegistry, planner::operator_registry::OperatorRegistry, - serde::to_arrow_datatype, + serde::{to_arrow_datatype, to_arrow_field}, shuffle::{SchemaAlignExec, ShuffleWriterExec}, }; use crate::jvm_bridge::{jni_call, JVMClasses}; @@ -3772,15 +3772,17 @@ pub(crate) fn convert_spark_types_to_arrow_schema( let arrow_fields = spark_types .iter() .map(|spark_type| { - let field = Field::new( + let field = to_arrow_field( String::clone(&spark_type.name), - to_arrow_datatype(spark_type.data_type.as_ref().unwrap()), + spark_type.data_type.as_ref().unwrap(), spark_type.nullable, ); if spark_type.metadata.is_empty() { field } else { - field.with_metadata(spark_type.metadata.clone()) + let mut metadata = spark_type.metadata.clone(); + metadata.extend(field.metadata().clone()); + field.with_metadata(metadata) } }) .collect_vec(); diff --git a/native/core/src/execution/serde.rs b/native/core/src/execution/serde.rs index 89d23c9c3b0..84e8d97fa12 100644 --- a/native/core/src/execution/serde.rs +++ b/native/core/src/execution/serde.rs @@ -31,9 +31,9 @@ use datafusion_comet_proto::{ spark_expression::DataType, spark_operator, }; -use parquet::arrow::PARQUET_FIELD_ID_META_KEY; +use parquet::{arrow::PARQUET_FIELD_ID_META_KEY, variant::VariantType}; use prost::Message; -use std::{collections::HashMap, io::Cursor, sync::Arc}; +use std::{io::Cursor, sync::Arc}; /// Deserialize bytes to protobuf type of expression pub fn deserialize_expr(buf: &[u8]) -> Result { @@ -106,6 +106,10 @@ pub fn to_arrow_datatype(dt_value: &DataType) -> ArrowDataType { // Spark's CalendarIntervalType stores months, days, and microseconds. Arrow stores the // same components with nanosecond precision. DataTypeId::CalendarInterval => ArrowDataType::Interval(IntervalUnit::MonthDayNano), + DataTypeId::Variant => ArrowDataType::Struct(Fields::from(vec![ + Field::new("value", ArrowDataType::Binary, false), + Field::new("metadata", ArrowDataType::Binary, false), + ])), DataTypeId::Null => ArrowDataType::Null, DataTypeId::List => match dt_value .type_info @@ -117,9 +121,9 @@ pub fn to_arrow_datatype(dt_value: &DataType) -> ArrowDataType { { DatatypeStruct::List(info) => { let field = with_parquet_field_id( - Field::new( + to_arrow_field( "item", - to_arrow_datatype(info.element_type.as_ref().unwrap()), + info.element_type.as_ref().unwrap(), info.contains_null, ), info.element_field_id, @@ -138,17 +142,13 @@ pub fn to_arrow_datatype(dt_value: &DataType) -> ArrowDataType { { DatatypeStruct::Map(info) => { let key_field = with_parquet_field_id( - Field::new( - "key", - to_arrow_datatype(info.key_type.as_ref().unwrap()), - false, - ), + to_arrow_field("key", info.key_type.as_ref().unwrap(), false), info.key_field_id, ); let value_field = with_parquet_field_id( - Field::new( + to_arrow_field( "value", - to_arrow_datatype(info.value_type.as_ref().unwrap()), + info.value_type.as_ref().unwrap(), info.value_contains_null, ), info.value_field_id, @@ -176,16 +176,18 @@ pub fn to_arrow_datatype(dt_value: &DataType) -> ArrowDataType { .iter() .enumerate() .map(|(idx, name)| { - let field = Field::new( + let field = to_arrow_field( name, - to_arrow_datatype(&info.field_datatypes[idx]), + &info.field_datatypes[idx], info.field_nullable[idx], ); // Attach Spark field metadata (currently parquet.field.id) when present. // field_metadata is parallel to field_names; either empty or full length. if let Some(meta) = info.field_metadata.get(idx) { if !meta.metadata.is_empty() { - return field.with_metadata(meta.metadata.clone()); + let mut metadata = meta.metadata.clone(); + metadata.extend(field.metadata().clone()); + return field.with_metadata(metadata); } } field @@ -198,13 +200,32 @@ pub fn to_arrow_datatype(dt_value: &DataType) -> ArrowDataType { } } +/// Converts a protobuf type to an Arrow field, preserving logical extension identity. +pub fn to_arrow_field( + name: impl Into, + data_type: &DataType, + nullable: bool, +) -> Field { + let field = Field::new(name, to_arrow_datatype(data_type), nullable); + if DataTypeId::try_from(data_type.type_id).unwrap() == DataTypeId::Variant { + field.with_extension_type(VariantType) + } else { + field + } +} + +pub fn is_variant_field(field: &Field) -> bool { + field.has_valid_extension_type::() +} + /// Attach a Parquet field ID without changing synthetic fields when Catalyst did not supply one. fn with_parquet_field_id(field: Field, field_id: Option) -> Field { match field_id { - Some(id) => field.with_metadata(HashMap::from([( - PARQUET_FIELD_ID_META_KEY.to_string(), - id.to_string(), - )])), + Some(id) => { + let mut metadata = field.metadata().clone(); + metadata.insert(PARQUET_FIELD_ID_META_KEY.to_string(), id.to_string()); + field.with_metadata(metadata) + } None => field, } } diff --git a/native/core/src/execution/utils.rs b/native/core/src/execution/utils.rs index 6195e3f0aea..efc15826d6c 100644 --- a/native/core/src/execution/utils.rs +++ b/native/core/src/execution/utils.rs @@ -19,17 +19,18 @@ use crate::execution::operators::ExecutionError; use arrow::{ array::ArrayData, + datatypes::Field, ffi::{FFI_ArrowArray, FFI_ArrowSchema}, }; pub trait SparkArrowConvert { /// Move Arrow Arrays to C data interface. - fn move_to_spark(&self, array: i64, schema: i64) -> Result<(), ExecutionError>; + fn move_to_spark(&self, field: &Field, array: i64, schema: i64) -> Result<(), ExecutionError>; } impl SparkArrowConvert for ArrayData { /// Move this ArrowData to pointers of Arrow C data interface. - fn move_to_spark(&self, array: i64, schema: i64) -> Result<(), ExecutionError> { + fn move_to_spark(&self, field: &Field, array: i64, schema: i64) -> Result<(), ExecutionError> { let array_ptr = array as *mut FFI_ArrowArray; let schema_ptr = schema as *mut FFI_ArrowSchema; @@ -40,7 +41,7 @@ impl SparkArrowConvert for ArrayData { if array_ptr.align_offset(array_align) != 0 || schema_ptr.align_offset(schema_align) != 0 { unsafe { std::ptr::write_unaligned(array_ptr, FFI_ArrowArray::new(self)); - std::ptr::write_unaligned(schema_ptr, FFI_ArrowSchema::try_from(self.data_type())?); + std::ptr::write_unaligned(schema_ptr, FFI_ArrowSchema::try_from(field)?); } } else { // SAFETY: `array_ptr` and `schema_ptr` are aligned correctly. @@ -56,7 +57,7 @@ impl SparkArrowConvert for ArrayData { ); unsafe { std::ptr::write(array_ptr, FFI_ArrowArray::new(self)); - std::ptr::write(schema_ptr, FFI_ArrowSchema::try_from(self.data_type())?); + std::ptr::write(schema_ptr, FFI_ArrowSchema::try_from(field)?); } } diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index 1cc928d1d59..939be3a1050 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -19,17 +19,21 @@ use arrow::{ make_array, Array, ArrayRef, LargeListArray, ListArray, MapArray, StructArray, TimestampMicrosecondArray, TimestampMillisecondArray, }, - compute::CastOptions, + compute::{cast, CastOptions}, datatypes::{DataType, FieldRef, Schema, TimeUnit}, record_batch::RecordBatch, }; -use crate::parquet::parquet_support::{spark_parquet_convert, SparkParquetOptions}; +use crate::{ + execution::serde::is_variant_field, + parquet::parquet_support::{spark_parquet_convert, SparkParquetOptions}, +}; use datafusion::common::format::DEFAULT_CAST_OPTIONS; -use datafusion::common::Result as DataFusionResult; use datafusion::common::ScalarValue; +use datafusion::common::{DataFusionError, Result as DataFusionResult}; use datafusion::logical_expr::ColumnarValue; use datafusion::physical_expr::PhysicalExpr; +use parquet::variant::{unshred_variant, VariantArray}; use std::{ fmt::{self, Display}, hash::Hash, @@ -176,6 +180,42 @@ fn cast_timestamp_micros_to_millis_scalar( ScalarValue::TimestampMillisecond(new_val, target_tz) } +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(), + )); + } + + let variant = VariantArray::try_new(array.as_ref())?; + let unshredded = unshred_variant(&variant)?; + 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 output = StructArray::try_new( + fields.clone(), + vec![value, metadata], + unshredded.inner().nulls().cloned(), + )?; + Ok(Arc::new(output)) +} + #[derive(Debug, Clone, Eq)] pub struct CometCastColumnExpr { /// The physical expression producing the value to cast. @@ -260,6 +300,18 @@ impl PhysicalExpr for CometCastColumnExpr { fn evaluate(&self, batch: &RecordBatch) -> DataFusionResult { let value = self.expr.evaluate(batch)?; + if is_variant_field(&self.target_field) { + 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 @@ -349,9 +401,70 @@ impl PhysicalExpr for CometCastColumnExpr { #[cfg(test)] mod tests { use super::*; - use arrow::array::{Array, Int32Array, StringArray}; + use arrow::array::{Array, AsArray, Int32Array, Int64Array, StringArray}; use arrow::datatypes::{Field, Fields}; use datafusion::physical_expr::expressions::Column; + use parquet::variant::{Variant, VariantArrayBuilder, VariantType}; + + #[test] + fn test_normalize_shredded_variant_for_spark() { + 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_fields = Fields::from(vec![ + Field::new("typed_value", DataType::Int64, true), + Field::new("metadata", metadata.data_type().clone(), false), + ]); + let physical = StructArray::try_new( + physical_fields, + vec![typed_value, metadata], + base.inner().nulls().cloned(), + ) + .unwrap(); + + let input_field = Arc::new(Field::new("v", physical.data_type().clone(), true)); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), true).with_extension_type(VariantType), + ); + let schema = Schema::new(vec![Arc::clone(&input_field)]); + let batch = RecordBatch::try_new(Arc::new(schema), vec![Arc::new(physical)]).unwrap(); + let expr = CometCastColumnExpr::new( + Arc::new(Column::new("v", 0)), + input_field, + target_field, + None, + ); + + let ColumnarValue::Array(output) = expr.evaluate(&batch).unwrap() else { + panic!("expected array") + }; + let output = output.as_struct(); + assert_eq!( + output + .fields() + .iter() + .map(|field| field.name().as_str()) + .collect::>(), + vec!["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 test_cast_timestamp_micros_to_millis_array() { diff --git a/native/core/src/parquet/mod.rs b/native/core/src/parquet/mod.rs index cfa03220c10..ea61fe54ac7 100644 --- a/native/core/src/parquet/mod.rs +++ b/native/core/src/parquet/mod.rs @@ -308,8 +308,10 @@ pub extern "system" fn Java_org_apache_comet_parquet_Native_currentColumnBatch( .ok_or_else(|| CometError::Execution { source: ExecutionError::GeneralError("There is no more data to read".to_string()), }); - let data = batch_reader?.column(column_idx as usize).into_data(); - data.move_to_spark(array_addr, schema_addr) + let batch = batch_reader?; + let field = batch.schema().field(column_idx as usize).clone(); + let data = batch.column(column_idx as usize).into_data(); + data.move_to_spark(&field, array_addr, schema_addr) .map_err(|e| e.into()) }) } diff --git a/native/core/src/parquet/schema_adapter.rs b/native/core/src/parquet/schema_adapter.rs index c6586b4681e..74d8afd7bc3 100644 --- a/native/core/src/parquet/schema_adapter.rs +++ b/native/core/src/parquet/schema_adapter.rs @@ -15,6 +15,7 @@ // specific language governing permissions and limitations // under the License. +use crate::execution::serde::is_variant_field; use crate::parquet::cast_column::CometCastColumnExpr; use crate::parquet::parquet_support::{spark_parquet_convert, SparkParquetOptions}; use arrow::array::new_empty_array; @@ -586,7 +587,9 @@ impl SparkPhysicalExprAdapter { Arc::clone(&e) }; - if logical_field.data_type() != physical_field.data_type() { + if is_variant_field(logical_field) + || 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 @@ -645,6 +648,19 @@ impl SparkPhysicalExprAdapter { }; let physical_type = input_field.data_type(); + if is_variant_field(cast.target_field()) { + let comet_cast: Arc = Arc::new( + CometCastColumnExpr::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 diff --git a/native/proto/src/proto/types.proto b/native/proto/src/proto/types.proto index 643a9cbabb1..8557086e0e7 100644 --- a/native/proto/src/proto/types.proto +++ b/native/proto/src/proto/types.proto @@ -63,6 +63,7 @@ message DataType { YEAR_MONTH_INTERVAL = 18; DAY_TIME_INTERVAL = 19; CALENDAR_INTERVAL = 20; + VARIANT = 21; } DataTypeId type_id = 1; diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 3c2668cfe92..f615d664763 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -55,7 +55,7 @@ import org.apache.comet.CometSparkSessionExtensions._ import org.apache.comet.rules.CometExecRule.allExecs import org.apache.comet.serde._ import org.apache.comet.serde.operator._ -import org.apache.comet.shims.{ShimCometStreaming, ShimSubqueryBroadcast} +import org.apache.comet.shims.{CometTypeShim, ShimCometStreaming, ShimSubqueryBroadcast} object CometExecRule { @@ -122,6 +122,7 @@ object CometExecRule { */ case class CometExecRule(session: SparkSession) extends Rule[SparkPlan] + with CometTypeShim with ShimSubqueryBroadcast { private lazy val showTransformations = CometConf.COMET_EXPLAIN_TRANSFORMATIONS.get() @@ -731,6 +732,15 @@ case class CometExecRule(session: SparkSession) private def tryConvertToComet( op: SparkPlan, handler: CometOperatorSerde[_]): Option[SparkPlan] = { + if (!op.isInstanceOf[CometScanExec] && + (op.output ++ op.children.flatMap(_.output)).exists(attr => + containsVariantType(attr.dataType))) { + withFallbackReason( + op, + "Native operators do not support schemas containing type VariantType") + return None + } + val serde = handler.asInstanceOf[CometOperatorSerde[SparkPlan]] if (isOperatorEnabled(serde, op)) { // For operators that require native children (like writes), check if all data-producing diff --git a/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala index a524da3af92..41ae750daf1 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala @@ -963,8 +963,10 @@ case class CometScanRule(session: SparkSession) private def isSchemaSupported(scanExec: FileSourceScanExec, r: HadoopFsRelation): Boolean = { val fallbackReasons = new ListBuffer[String]() val typeChecker = CometScanTypeChecker() - val schemaSupported = - typeChecker.isSchemaSupported(scanExec.requiredSchema, fallbackReasons) + val schemaSupported = scanExec.requiredSchema.fields.forall { field => + isVariantType(field.dataType) || + typeChecker.isTypeSupported(field.dataType, field.name, fallbackReasons) + } if (!schemaSupported) { withFallbackReason( scanExec, diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index 6802dfaa646..d5f03a63ea2 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -588,6 +588,7 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { case _: YearMonthIntervalType => 18 case _: DayTimeIntervalType => 19 case CalendarIntervalType => 20 + case dt if isVariantType(dt) => 21 case dt => logWarning(s"Cannot serialize Spark data type: $dt") return None diff --git a/spark/src/main/scala/org/apache/comet/serde/namedExpressions.scala b/spark/src/main/scala/org/apache/comet/serde/namedExpressions.scala index edd083c282a..209634bad33 100644 --- a/spark/src/main/scala/org/apache/comet/serde/namedExpressions.scala +++ b/spark/src/main/scala/org/apache/comet/serde/namedExpressions.scala @@ -23,6 +23,7 @@ import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, AttributeRef import org.apache.comet.CometSparkSessionExtensions.withFallbackReason import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, serializeDataType} +import org.apache.comet.shims.CometTypeShim object CometAlias extends CometExpressionSerde[Alias] { override def convert( @@ -33,9 +34,13 @@ object CometAlias extends CometExpressionSerde[Alias] { } } -object CometAttributeReference extends CometExpressionSerde[AttributeReference] { +object CometAttributeReference + extends CometExpressionSerde[AttributeReference] + with CometTypeShim { override def getSupportLevel(attr: AttributeReference): SupportLevel = { - if (serializeDataType(attr.dataType).isDefined) { + if (isVariantType(attr.dataType)) { + Unsupported(Some(s"unsupported expression input of type ${attr.dataType}")) + } else if (serializeDataType(attr.dataType).isDefined) { Compatible() } else { Unsupported(Some(s"unsupported datatype: ${attr.dataType}")) diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala index e395ac6d9d3..c80fd59a39d 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala @@ -29,7 +29,7 @@ import org.apache.spark.sql.comet.{CometNativeExec, CometNativeScanExec, CometSc import org.apache.spark.sql.execution.{FileSourceScanExec, InSubqueryExec, SubqueryAdaptiveBroadcastExec} import org.apache.spark.sql.execution.datasources.parquet.ParquetUtils import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{ArrayType, DataType, MapType, StructField, StructType} +import org.apache.spark.sql.types.{StructField, StructType} import org.apache.comet.{CometConf, ConfigEntry} import org.apache.comet.CometConf.COMET_EXEC_ENABLED @@ -51,15 +51,6 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS // like "file_size" could collide with a real column of the same name. Prefix to avoid it. private val constantMetadataFieldPrefix = "_comet_metadata_" - private def containsVariantType(dataType: DataType): Boolean = dataType match { - case dt if isVariantType(dt) => true - case StructType(fields) => fields.exists(field => containsVariantType(field.dataType)) - case ArrayType(elementType, _) => containsVariantType(elementType) - case MapType(keyType, valueType, _) => - containsVariantType(keyType) || containsVariantType(valueType) - case _ => false - } - /** Determine whether the scan is supported and tag the Spark plan with any fallback reasons */ def isSupported(scanExec: FileSourceScanExec): Boolean = { @@ -198,8 +189,8 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS // unrequested struct. The complete relation schema still contains that unsupported type, // and serializing it would throw even though the native reader never needs those bytes. // Keep ordinary fields unchanged and replace a requested Variant-bearing root with its - // already-validated, pruned required field. A requested actual Variant never reaches this - // point because CometScanRule keeps those scans on Spark. + // already-validated, pruned required field. Direct top-level Variant fields are retained; + // unsupported nested Variant fields are rejected by CometScanRule. val nativeDataSchema = StructType(scan.relation.dataSchema.fields.flatMap { field => if (containsVariantType(field.dataType)) { scan.requiredSchema.fields.find(requiredField => diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala b/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala index 9d4b0bce881..92d1127377d 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala @@ -48,6 +48,9 @@ import org.apache.comet.shims.CometTypeShim import org.apache.comet.vector.CometVector object Utils extends CometTypeShim with Logging { + private val ArrowExtensionNameKey = "ARROW:extension:name" + private val VariantExtensionName = "arrow.parquet.variant" + def getConfPath(confFileName: String): String = { sys.env .get(COMET_CONF_DIR_ENV) @@ -78,11 +81,17 @@ object Utils extends CometTypeShim with Logging { val elementType = fromArrowField(elementField) ArrayType(elementType, containsNull = elementField.isNullable) case ArrowType.Struct.INSTANCE => - val fields = field.getChildren().asScala.map { child => - val dt = fromArrowField(child) - StructField(child.getName, dt, child.isNullable) - } - StructType(fields.toSeq) + Option(field.getMetadata) + .flatMap(metadata => Option(metadata.get(ArrowExtensionNameKey))) + .filter(_ == VariantExtensionName) + .flatMap(_ => variantType) + .getOrElse { + val fields = field.getChildren().asScala.map { child => + val dt = fromArrowField(child) + StructField(child.getName, dt, child.isNullable) + } + StructType(fields.toSeq) + } case arrowType => fromArrowType(arrowType) } } diff --git a/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala b/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala index b71476c3dd1..3fd97509fcf 100644 --- a/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala +++ b/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala @@ -39,6 +39,12 @@ trait CometTypeShim { @nowarn // Spark 4 feature; VariantType doesn't exist in Spark 3.x. def isVariantType(dt: DataType): Boolean = false + @nowarn // Spark 4 feature; VariantType doesn't exist in Spark 3.x. + def containsVariantType(dt: DataType): Boolean = false + + @nowarn // Spark 4 feature; VariantType doesn't exist in Spark 3.x. + def variantType: Option[DataType] = None + @nowarn // Spark 4.1 feature; TimeType doesn't exist in Spark 3.x. def isTimeType(dt: DataType): Boolean = false } diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala index f48955a7da5..e40023c6030 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala @@ -49,17 +49,26 @@ trait CometTypeShim { // Spark 4.0's `PushVariantIntoScan` rewrites `VariantType` columns into a `StructType` whose // fields each carry `__VARIANT_METADATA_KEY` metadata, then pushes `variant_get` paths down as - // ordinary struct field accesses. Comet's native scans don't understand the on-disk Parquet - // variant shredding layout, so reading such a struct natively returns nulls. Detect the marker - // and force scan fallback. + // ordinary struct field accesses. The direct whole-value scan path does not support that pushed + // VariantStruct representation. Detect the marker and force scan fallback. def isVariantStruct(s: StructType): Boolean = VariantMetadata.isVariantStruct(s) - // Comet has no native execution path for Spark 4's `VariantType` (introduced in - // SPARK-45827). Serdes call this to route casts/expressions touching the type back to Spark - // rather than serializing an unsupported datatype into the native plan. Stubbed to `false` in - // Spark 3.x where `VariantType` does not exist. + // Outside direct top-level Parquet projection, Comet has no native execution path for Spark 4's + // `VariantType` (introduced in SPARK-45827). Serdes call this to route casts and expressions + // touching the type back to Spark. Stubbed to `false` in Spark 3.x. def isVariantType(dt: DataType): Boolean = dt.isInstanceOf[VariantType] + def containsVariantType(dt: DataType): Boolean = dt match { + case dt if isVariantType(dt) => true + case StructType(fields) => fields.exists(field => containsVariantType(field.dataType)) + case ArrayType(elementType, _) => containsVariantType(elementType) + case MapType(keyType, valueType, _) => + containsVariantType(keyType) || containsVariantType(valueType) + case _ => false + } + + def variantType: Option[DataType] = Some(VariantType) + def isTimeType(dt: DataType): Boolean = dt.getClass.getSimpleName.startsWith("TimeType") diff --git a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql index 328254b340a..435c395a033 100644 --- a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql +++ b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql @@ -15,21 +15,23 @@ -- specific language governing permissions and limitations -- under the License. --- Confirms Comet falls back to Spark when a parquet scan's schema contains a --- VariantType column. VariantType is a Spark 4.0+ data type that Comet does --- not currently support, so any scan exposing it must be executed by Spark. +-- Confirms direct top-level VariantType projection through Comet's ordinary +-- native Parquet scan. Expressions, operators, nested Variant, and Iceberg +-- remain unsupported. -- MinSparkVersion: 4.0 +-- Config: spark.sql.variant.writeShredding.enabled=false statement CREATE TABLE test_variant(id INT, v VARIANT, tail STRING) USING parquet statement INSERT INTO test_variant VALUES - (1, parse_json('{"a": 1, "b": "hello"}'), 'first'), - (2, parse_json('{"a": 2, "b": "world"}'), NULL), - (3, parse_json('null'), 'variant-null'), - (4, NULL, 'sql-null') + (1, parse_json('{"a": 1, "b": "hello"}'), 'object'), + (2, parse_json('[1, true, "x"]'), 'array'), + (3, parse_json('42'), 'scalar'), + (4, parse_json('null'), 'json-null'), + (5, CAST(NULL AS VARIANT), 'sql-null') -- A plain Parquet scan can remain native when its required schema prunes the -- Variant column completely, including both SQL NULL and Variant null values. @@ -43,8 +45,21 @@ SELECT tail FROM test_variant ORDER BY id query SELECT id, tail FROM test_variant WHERE tail IS NOT NULL ORDER BY id +-- Full-value projection is scan-only: no native expression or pass-through operator carries v. +query +SELECT v FROM test_variant + +query +SELECT id, v, tail FROM test_variant + +query expect_fallback(type VariantType) +SELECT v FROM test_variant ORDER BY id + +query expect_fallback(type VariantType) +SELECT v FROM test_variant LIMIT 1 + query expect_fallback(type VariantType) -SELECT id, v FROM test_variant ORDER BY id +SELECT /*+ REPARTITION(2, id) */ id, v FROM test_variant query expect_fallback(type VariantType) SELECT variant_get(v, '$.a', 'int') AS a FROM test_variant ORDER BY id @@ -55,6 +70,9 @@ SELECT id FROM test_variant WHERE variant_get(v, '$.a', 'int') = 1 query expect_fallback(type VariantType) SELECT COUNT(*) FROM test_variant WHERE v IS NOT NULL +query expect_fallback(type VariantType) +SELECT CAST(v AS STRING) FROM test_variant + statement CREATE TABLE test_variant_struct(id INT, s STRUCT, tail STRING) USING parquet diff --git a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala index b256c917a1d..cb5d440452b 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala @@ -24,6 +24,7 @@ import java.math.{BigDecimal, BigInteger} import java.time.{ZoneId, ZoneOffset} import java.util.{Base64, Collections} +import scala.jdk.CollectionConverters._ import scala.reflect.ClassTag import scala.reflect.runtime.universe.TypeTag @@ -39,7 +40,8 @@ import org.apache.parquet.schema.MessageTypeParser import org.apache.spark.SparkException import org.apache.spark.sql.{CometTestBase, DataFrame, Row} import org.apache.spark.sql.catalyst.util.DateTimeUtils -import org.apache.spark.sql.comet.{CometNativeScanExec, CometScanExec} +import org.apache.spark.sql.comet.{CometColumnarToRowExec, CometNativeColumnarToRowExec, CometNativeScanExec, CometScanExec} +import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.execution.datasources.parquet.ParquetUtils import org.apache.spark.sql.internal.SQLConf @@ -47,7 +49,8 @@ import org.apache.spark.sql.types._ import com.google.common.primitives.UnsignedLong -import org.apache.comet.CometConf +import org.apache.comet.{CometConf, CometSparkSessionExtensions} +import org.apache.comet.vector.CometStructVector abstract class ParquetReadSuite extends CometTestBase { import testImplicits._ @@ -85,6 +88,122 @@ abstract class ParquetReadSuite extends CometTestBase { } } + test("native scan projects Variant through a Spark-compatible vector") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + def normalizedRows(df: DataFrame, variantOrdinal: Int): Seq[Seq[Any]] = { + df + .collect() + .map { row => + row.toSeq.updated( + variantOrdinal, + Option(row.get(variantOrdinal)).map(_.toString).orNull) + } + .sortBy(_.map(value => Option(value).fold("0")(v => "1" + v.toString)).mkString("\u0000")) + .toSeq + } + + Seq(false, true).foreach { shredded => + withTable("variant_projection") { + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + "spark.sql.variant.writeShredding.enabled" -> shredded.toString, + "spark.sql.variant.forceShreddingSchemaForTest" -> "a BIGINT") { + sql("CREATE TABLE variant_projection(id INT, v VARIANT, tail STRING) USING parquet") + sql("""INSERT INTO variant_projection VALUES + |(1, parse_json('{"a": 10, "b": "hello"}'), 'object'), + |(2, parse_json('[1, true, "x"]'), 'array'), + |(3, parse_json('42'), 'scalar'), + |(4, parse_json('null'), 'json-null'), + |(5, CAST(NULL AS VARIANT), 'sql-null')""".stripMargin) + } + + val queries = Seq( + "SELECT v FROM variant_projection" -> 0, + "SELECT id, v, tail FROM variant_projection" -> 1) + var expected = Seq.empty[Seq[Seq[Any]]] + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + "spark.sql.variant.allowReadingShredded" -> "true") { + expected = queries.map { case (query, variantOrdinal) => + normalizedRows(sql(query), variantOrdinal) + } + } + + // Phase A handles only whole values; Spark's pushed VariantStruct remains a fallback. + withSQLConf( + CometConf.COMET_NATIVE_COLUMNAR_TO_ROW_ENABLED.key -> "true", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val plans = queries.zip(expected).map { case ((query, variantOrdinal), expectedRows) => + val df = sql(query) + assert(normalizedRows(df, variantOrdinal) == expectedRows) + df.queryExecution.executedPlan + } + + if (!shredded) { + queries.foreach { case (query, _) => checkSparkAnswerAndOperator(sql(query)) } + } + + plans.foreach { cometPlan => + assert(collect(cometPlan) { case _: CometNativeScanExec => true }.size == 1) + assert(collect(cometPlan) { case _: CometNativeColumnarToRowExec => true }.isEmpty) + assert(collect(cometPlan) { case _: CometColumnarToRowExec => true }.nonEmpty) + } + + val scan = collect(plans.head) { case scan: CometNativeScanExec => scan }.head + + val summaries = scan + .executeColumnar() + .mapPartitions { batches => + batches.map { batch => + try { + val vector = batch.column(0) + val struct = vector.asInstanceOf[CometStructVector] + val field = struct.getValueVector.getField + val getVariant = struct.getClass.getMethod("getVariant", Integer.TYPE) + val values = (0 until batch.numRows()).map { rowId => + if (struct.isNullAt(rowId)) { + None + } else { + Some(getVariant.invoke(struct, Int.box(rowId)).toString) + } + } + ( + Utils.isVariantType(struct.dataType()), + field.getName, + field.isNullable, + field.getMetadata.get("ARROW:extension:name"), + field.getChildren.asScala.map(_.getName).toSeq, + Seq(struct.getChild(0).dataType(), struct.getChild(1).dataType()), + values) + } finally { + batch.close() + } + } + } + .collect() + + assert(summaries.nonEmpty) + summaries.foreach { + case (isVariant, name, nullable, extension, children, childTypes, _) => + assert(isVariant) + assert(name == "v") + assert(nullable) + assert(extension == "arrow.parquet.variant") + assert(children == Seq("value", "metadata")) + assert(childTypes == Seq(BinaryType, BinaryType)) + } + val values = summaries.flatMap(_._7) + assert(values.count(_.isEmpty) == 1) + assert( + values.flatten.toSet == + Set("{\"a\":10,\"b\":\"hello\"}", "[1,true,\"x\"]", "42", "null")) + } + } + } + } + // Spark ignores ARROW:schema during Parquet schema inference: // https://github.com/apache/spark/blob/v4.2.0/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/parquet/ParquetFileFormat.scala#L585-L599 // With binaryAsString, Spark maps unannotated BINARY to StringType: From 784c316cf4ccc884fcf1542a4b0beee6eb7d8f8d Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sun, 23 Aug 2026 02:01:59 +0800 Subject: [PATCH 02/14] review --- native/core/src/execution/planner.rs | 175 ++++-- native/core/src/execution/utils.rs | 50 +- native/core/src/parquet/cast_column.rs | 534 +++++++++++++++++- .../apache/comet/rules/CometExecRule.scala | 13 +- .../rules/EliminateRedundantTransitions.scala | 5 +- .../serde/operator/CometNativeScan.scala | 57 +- .../apache/comet/shims/CometTypeShim.scala | 4 + .../apache/comet/shims/CometTypeShim.scala | 18 + .../sql-tests/expressions/misc/variant.sql | 69 ++- .../parquet/CometParquetWriterSuite.scala | 31 +- .../comet/parquet/ParquetReadSuite.scala | 176 +++++- .../sql/comet/CometMapInBatchSuite.scala | 33 +- 12 files changed, 1067 insertions(+), 98 deletions(-) diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 6ef23ea163c..0d8d062f9a0 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -115,7 +115,7 @@ use crate::parquet::parquet_exec::init_datasource_exec; use arrow::array::{ new_empty_array, Array, ArrayRef, BinaryBuilder, BooleanArray, Date32Array, Decimal128Array, Float32Array, Float64Array, Int16Array, Int32Array, Int64Array, Int8Array, ListArray, - NullArray, StringBuilder, TimestampMicrosecondArray, + NullArray, RecordBatch, StringBuilder, TimestampMicrosecondArray, }; use arrow::buffer::{BooleanBuffer, NullBuffer, OffsetBuffer}; use arrow::row::{OwnedRow, RowConverter, SortField}; @@ -923,6 +923,25 @@ impl PhysicalPlanner { } } + /// Decode a Spark existence default into the scalar consumed by the Parquet schema adapter. + /// Ordinary defaults are literals. A Variant default is transported as a constant + /// `CreateNamedStruct(value, metadata)` because Variant has Struct storage in Arrow. + fn create_default_value( + &self, + spark_expr: &Expr, + input_schema: SchemaRef, + ) -> Result { + let expr = self.create_expr(spark_expr, Arc::clone(&input_schema))?; + if let Some(literal) = expr.downcast_ref::() { + return Ok(literal.value().clone()); + } + + let array = expr + .evaluate(&RecordBatch::new_empty(input_schema))? + .into_array_of_size(1)?; + Ok(ScalarValue::try_from_array(array.as_ref(), 0)?) + } + /// Create a DataFusion physical sort expression from Spark physical expression fn create_sort_expr<'a>( &'a self, @@ -1580,44 +1599,38 @@ impl PhysicalPlanner { .collect() }; - let default_values: Option> = if !common - .default_values - .is_empty() - { - // We have default values. Extract the two lists (same length) of values and - // indexes in the schema, and then create a HashMap to use in the SchemaMapper. - let default_values: Result, DataFusionError> = common - .default_values - .iter() - .map(|expr| { - let literal = self.create_expr(expr, Arc::clone(&required_schema))?; - let df_literal = - literal.downcast_ref::().ok_or_else(|| { - GeneralError("Expected literal of default value.".to_string()) + if common.default_values.len() != common.default_values_indexes.len() { + return Err(GeneralError(format!( + "NativeScan has {} default values but {} default indexes", + common.default_values.len(), + common.default_values_indexes.len() + ))); + } + let default_values: Option> = + if common.default_values.is_empty() { + None + } else { + let defaults = common + .default_values + .iter() + .zip(&common.default_values_indexes) + .map(|(expr, offset)| { + let idx = usize::try_from(*offset).map_err(|_| { + GeneralError(format!("Invalid default value index {offset}")) })?; - Ok(df_literal.value().clone()) - }) - .collect(); - let default_values = default_values?; - let default_values_indexes: Vec = common - .default_values_indexes - .iter() - .map(|offset| *offset as usize) - .collect(); - Some( - default_values_indexes - .into_iter() - .zip(default_values) - .map(|(idx, scalar_value)| { - let field = required_schema.field(idx); - let column = Column::new(field.name().as_str(), idx); - (column, scalar_value) + let field = required_schema.fields().get(idx).ok_or_else(|| { + GeneralError(format!( + "Default value index {idx} is outside schema with {} fields", + required_schema.fields().len() + )) + })?; + let value = + self.create_default_value(expr, Arc::clone(&required_schema))?; + Ok((Column::new(field.name(), idx), value)) }) - .collect(), - ) - } else { - None - }; + .collect::, ExecutionError>>()?; + Some(defaults) + }; // Get one file from this partition (we know it's not empty due to early return above) let one_file = partition_files @@ -4600,7 +4613,8 @@ mod tests { use std::{sync::Arc, task::Poll}; use arrow::array::{ - Array, DictionaryArray, Int32Array, Int8Array, ListArray, RecordBatch, StringArray, + Array, BinaryArray, DictionaryArray, Int32Array, Int8Array, ListArray, RecordBatch, + StringArray, }; use arrow::datatypes::{DataType, Field, FieldRef, Fields, Schema}; use datafusion::catalog::memory::DataSourceExec; @@ -4613,7 +4627,10 @@ mod tests { use datafusion::error::DataFusionError; use datafusion::logical_expr::ScalarUDF; use datafusion::physical_plan::ExecutionPlan; - use datafusion::{assert_batches_eq, physical_plan::common::collect, prelude::SessionContext}; + use datafusion::{ + assert_batches_eq, physical_plan::common::collect, prelude::SessionContext, + scalar::ScalarValue, + }; use datafusion_physical_expr_adapter::PhysicalExprAdapterFactory; use tempfile::TempDir; use tokio::sync::mpsc; @@ -4623,6 +4640,7 @@ mod tests { use crate::execution::operators::ExecutionError; use crate::execution::planner::literal_to_array_ref; use crate::execution::planner::parse_file_scan_tasks_from_common; + use crate::execution::serde::to_arrow_field; use crate::parquet::parquet_support::SparkParquetOptions; use crate::parquet::schema_adapter::SparkPhysicalExprAdapterFactory; use datafusion_comet_proto::spark_expression::expr::ExprStruct; @@ -4636,6 +4654,85 @@ mod tests { }; use datafusion_comet_spark_expr::EvalMode; + #[test] + fn test_create_default_value_from_literal_and_variant_struct() { + fn literal(value: literal::Value, type_id: i32) -> Expr { + Expr { + expr_struct: Some(ExprStruct::Literal(spark_expression::Literal { + value: Some(value), + datatype: Some(spark_expression::DataType { + type_id, + type_info: None, + }), + is_null: false, + })), + query_context: None, + expr_id: None, + } + } + + let planner = PhysicalPlanner::default(); + let int_field = Field::new("n", DataType::Int32, true); + let int_default = literal(literal::Value::IntVal(7), 3); + assert_eq!( + planner + .create_default_value(&int_default, Arc::new(Schema::new(vec![int_field.clone()])),) + .unwrap(), + ScalarValue::Int32(Some(7)) + ); + + let variant_field = to_arrow_field( + "v", + &spark_expression::DataType { + type_id: 21, + type_info: None, + }, + true, + ); + let variant_default = Expr { + expr_struct: Some(ExprStruct::CreateNamedStruct( + spark_expression::CreateNamedStruct { + values: vec![ + literal(literal::Value::BytesVal(vec![1, 2]), 8), + literal(literal::Value::BytesVal(vec![3, 4]), 8), + ], + names: vec!["value".to_string(), "metadata".to_string()], + }, + )), + query_context: None, + expr_id: None, + }; + let value = planner + .create_default_value( + &variant_default, + Arc::new(Schema::new(vec![variant_field.clone()])), + ) + .unwrap(); + let ScalarValue::Struct(value) = value else { + panic!("expected Variant default to use Struct storage") + }; + assert_eq!(value.fields()[0].name(), "value"); + assert_eq!(value.fields()[1].name(), "metadata"); + assert_eq!( + value + .column(0) + .as_any() + .downcast_ref::() + .unwrap() + .value(0), + &[1, 2] + ); + assert_eq!( + value + .column(1) + .as_any() + .downcast_ref::() + .unwrap() + .value(0), + &[3, 4] + ); + } + #[test] fn test_unpack_dictionary_primitive() { let op_scan = Operator { diff --git a/native/core/src/execution/utils.rs b/native/core/src/execution/utils.rs index efc15826d6c..435da3fe030 100644 --- a/native/core/src/execution/utils.rs +++ b/native/core/src/execution/utils.rs @@ -20,9 +20,26 @@ use crate::execution::operators::ExecutionError; use arrow::{ array::ArrayData, datatypes::Field, + error::ArrowError, ffi::{FFI_ArrowArray, FFI_ArrowSchema}, }; +fn ffi_schema_for_field(field: &Field) -> Result { + if field.name().contains('\0') { + // Spark keeps Parquet field names as strings, while ArrowSchema exports names as C strings: + // https://github.com/apache/spark/blob/v4.1.3/sql/api/src/main/scala/org/apache/spark/sql/types/StructField.scala#L32-L51 + // https://github.com/apache/spark/blob/v4.1.3/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/parquet/ParquetSchemaConverter.scala#L576-L647 + // https://github.com/apache/arrow-rs/blob/58.4.0/arrow-schema/src/ffi.rs#L168-L175 + // The logical output name is owned by Spark's plan, so substitute only at this boundary. + let field = field + .clone() + .with_name(field.name().replace('\0', "\u{fffd}")); + FFI_ArrowSchema::try_from(&field) + } else { + FFI_ArrowSchema::try_from(field) + } +} + pub trait SparkArrowConvert { /// Move Arrow Arrays to C data interface. fn move_to_spark(&self, field: &Field, array: i64, schema: i64) -> Result<(), ExecutionError>; @@ -36,12 +53,14 @@ impl SparkArrowConvert for ArrayData { let array_align = std::mem::align_of::(); let schema_align = std::mem::align_of::(); + let ffi_array = FFI_ArrowArray::new(self); + let ffi_schema = ffi_schema_for_field(field)?; // Check if the pointer alignment is correct. if array_ptr.align_offset(array_align) != 0 || schema_ptr.align_offset(schema_align) != 0 { unsafe { - std::ptr::write_unaligned(array_ptr, FFI_ArrowArray::new(self)); - std::ptr::write_unaligned(schema_ptr, FFI_ArrowSchema::try_from(field)?); + std::ptr::write_unaligned(array_ptr, ffi_array); + std::ptr::write_unaligned(schema_ptr, ffi_schema); } } else { // SAFETY: `array_ptr` and `schema_ptr` are aligned correctly. @@ -56,8 +75,8 @@ impl SparkArrowConvert for ArrayData { "move_to_spark: schema_ptr not aligned" ); unsafe { - std::ptr::write(array_ptr, FFI_ArrowArray::new(self)); - std::ptr::write(schema_ptr, FFI_ArrowSchema::try_from(field)?); + std::ptr::write(array_ptr, ffi_array); + std::ptr::write(schema_ptr, ffi_schema); } } @@ -66,3 +85,26 @@ impl SparkArrowConvert for ArrayData { } pub use datafusion_comet_common::bytes_to_i128; + +#[cfg(test)] +mod tests { + use super::*; + use arrow::datatypes::DataType; + use std::collections::HashMap; + + #[test] + fn test_ffi_schema_sanitizes_nul_name_and_preserves_metadata() { + let field = Field::new("v\0tail", DataType::Int32, true).with_metadata(HashMap::from([( + "ARROW:extension:name".to_string(), + "arrow.parquet.variant".to_string(), + )])); + + let ffi_schema = ffi_schema_for_field(&field).unwrap(); + let exported = Field::try_from(&ffi_schema).unwrap(); + + assert_eq!(exported.name(), "v\u{fffd}tail"); + assert_eq!(exported.data_type(), field.data_type()); + assert_eq!(exported.is_nullable(), field.is_nullable()); + assert_eq!(exported.metadata(), field.metadata()); + } +} diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index 939be3a1050..a16626d2d1c 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -16,11 +16,13 @@ // under the License. use arrow::{ array::{ - make_array, Array, ArrayRef, LargeListArray, ListArray, MapArray, StructArray, - TimestampMicrosecondArray, TimestampMillisecondArray, + make_array, Array, ArrayRef, BinaryArray, BinaryBuilder, LargeListArray, ListArray, + MapArray, StructArray, TimestampMicrosecondArray, TimestampMillisecondArray, }, + buffer::NullBuffer, compute::{cast, CastOptions}, datatypes::{DataType, FieldRef, Schema, TimeUnit}, + error::ArrowError, record_batch::RecordBatch, }; @@ -33,10 +35,14 @@ use datafusion::common::ScalarValue; use datafusion::common::{DataFusionError, Result as DataFusionResult}; use datafusion::logical_expr::ColumnarValue; use datafusion::physical_expr::PhysicalExpr; -use parquet::variant::{unshred_variant, VariantArray}; +use parquet::variant::{ + unshred_variant, MetadataBuilder, ParentState, ReadOnlyMetadataBuilder, ValueBuilder, Variant, + VariantArray, VariantMetadata, +}; use std::{ fmt::{self, Display}, hash::Hash, + panic::{catch_unwind, AssertUnwindSafe}, sync::Arc, }; @@ -201,13 +207,21 @@ fn normalize_variant_array( )); } - let variant = VariantArray::try_new(array.as_ref())?; + let array = decode_variant_metadata_dictionary(array)?; + let variant = prepare_variant_for_unshredding(&VariantArray::try_new(array.as_ref())?)?; let unshredded = unshred_variant(&variant)?; 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(), + VariantObjectKeyOrder::SparkUtf16, + false, + )?; let output = StructArray::try_new( fields.clone(), vec![value, metadata], @@ -216,6 +230,231 @@ fn normalize_variant_array( Ok(Arc::new(output)) } +/// Arrow's unshredder fully validates any residual `value` in a partially shredded object. Spark +/// writes object keys in Java UTF-16 order, so put that residual value in Arrow UTF-8 order only +/// while it passes through the upstream unshredder. +fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult { + let (Some(value), Some(_)) = (variant.value_field(), variant.typed_value_field()) else { + return Ok(variant.clone()); + }; + + let value = cast(value.as_ref(), &DataType::Binary)?; + let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let value = reorder_variant_values( + &value, + &metadata, + variant.inner().nulls(), + VariantObjectKeyOrder::ArrowUtf8, + true, + )?; + + let value_index = variant + .inner() + .fields() + .iter() + .position(|field| field.name() == "value") + .unwrap(); + let mut fields = variant.inner().fields().iter().cloned().collect::>(); + fields[value_index] = Arc::new( + fields[value_index] + .as_ref() + .clone() + .with_data_type(DataType::Binary), + ); + let mut columns = variant.inner().columns().to_vec(); + columns[value_index] = value; + let array = StructArray::try_new(fields.into(), columns, variant.inner().nulls().cloned())?; + Ok(VariantArray::try_new(&array)?) +} + +/// Arrow-rs parquet-variant-compute allows dictionary-encoded metadata in its contract, but 58.4's +/// `VariantArray::try_new` validates only Binary, LargeBinary, and BinaryView. Decode just that +/// child and keep the physical struct otherwise unchanged. +/// https://github.com/apache/arrow-rs/blob/0ff81c1215cc026a1de93ce3d2078df1ecba6f09/parquet-variant-compute/src/variant_array.rs#L276-L310 +fn decode_variant_metadata_dictionary(array: &ArrayRef) -> DataFusionResult { + let Some(struct_array) = array.as_any().downcast_ref::() else { + return Ok(Arc::clone(array)); + }; + let Some((metadata_index, metadata_field)) = struct_array + .fields() + .iter() + .enumerate() + .find(|(_, field)| field.name() == "metadata") + else { + return Ok(Arc::clone(array)); + }; + let DataType::Dictionary(_, value_type) = metadata_field.data_type() else { + return Ok(Arc::clone(array)); + }; + + let decoded = cast(struct_array.column(metadata_index).as_ref(), value_type)?; + let mut fields = struct_array.fields().iter().cloned().collect::>(); + fields[metadata_index] = Arc::new( + metadata_field + .as_ref() + .clone() + .with_data_type(decoded.data_type().clone()), + ); + let mut columns = struct_array.columns().to_vec(); + columns[metadata_index] = decoded; + Ok(Arc::new(StructArray::try_new( + fields.into(), + columns, + struct_array.nulls().cloned(), + )?)) +} + +/// Supplies sort-only field names whose Rust ordering matches Java `String.compareTo` ordering. +/// The original metadata dictionary still supplies the field IDs written to the Variant value. +#[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, + } +} + +/// Reorder object keys for either Arrow's UTF-8 order or Spark's Java UTF-16 order. Preserve +/// already-compatible values byte-for-byte and retain the original metadata dictionary. +fn reorder_variant_values( + value: &ArrayRef, + metadata: &ArrayRef, + parent_nulls: Option<&NullBuffer>, + order: VariantObjectKeyOrder, + allow_null_value: bool, +) -> DataFusionResult { + let value = value.as_any().downcast_ref::().unwrap(); + let metadata = metadata.as_any().downcast_ref::().unwrap(); + 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) { + if allow_null_value { + output.append_null(); + continue; + } + 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 metadata = VariantMetadata::try_new(metadata.value(index))?; + let rebuilt = catch_unwind(AssertUnwindSafe(|| { + let variant = Variant::new_with_metadata(metadata.clone(), value.value(index)); + if is_compatible_variant(&variant, order) { + return None; + } + let mut value_builder = ValueBuilder::new(); + match order { + VariantObjectKeyOrder::ArrowUtf8 => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(&metadata); + ValueBuilder::append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + variant, + ); + } + VariantObjectKeyOrder::SparkUtf16 => { + let mut metadata_builder = SparkMetadataBuilder::new(&metadata); + ValueBuilder::append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + variant, + ); + } + } + Some(value_builder.into_inner()) + })) + .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())) +} + #[derive(Debug, Clone, Eq)] pub struct CometCastColumnExpr { /// The physical expression producing the value to cast. @@ -401,19 +640,59 @@ impl PhysicalExpr for CometCastColumnExpr { #[cfg(test)] mod tests { use super::*; - use arrow::array::{Array, AsArray, Int32Array, Int64Array, StringArray}; - use arrow::datatypes::{Field, Fields}; + use arrow::array::{ + Array, AsArray, BinaryArray, DictionaryArray, Int32Array, Int64Array, StringArray, + }; + use arrow::datatypes::{Field, Fields, Int32Type}; use datafusion::physical_expr::expressions::Column; - use parquet::variant::{Variant, VariantArrayBuilder, VariantType}; + use parquet::variant::{VariantArrayBuilder, VariantBuilder, 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_object(output: &StructArray) { + let value = output.column(0).as_binary::(); + let metadata = output.column(1).as_binary::(); + let Variant::Object(object) = Variant::new(metadata.value(0), value.value(0)) 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}"); + let emoji = fields + .binary_search_by(|(name, _)| name.encode_utf16().cmp("😀".encode_utf16())) + .unwrap(); + assert_eq!(fields[emoji].1, Variant::from(531_i64)); + let private_use = fields + .binary_search_by(|(name, _)| name.encode_utf16().cmp("\u{e000}".encode_utf16())) + .unwrap(); + assert_eq!(fields[private_use].1, Variant::from(30_i64)); + } #[test] - fn test_normalize_shredded_variant_for_spark() { + fn test_normalize_shredded_variant_with_dictionary_metadata_for_spark() { 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 metadata = cast(base.metadata_field().as_ref(), &DataType::Binary).unwrap(); + let metadata_bytes = metadata.as_binary::().value(0).to_vec(); + let metadata_values: ArrayRef = + Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new( + DictionaryArray::::try_new( + Int32Array::from(vec![Some(0), Some(0), Some(0)]), + metadata_values, + ) + .unwrap(), + ); let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![Some(10), None, Some(30)])); let physical_fields = Fields::from(vec![ Field::new("typed_value", DataType::Int64, true), @@ -466,6 +745,243 @@ mod tests { assert_eq!(variant.value(2), Variant::from(30_i64)); } + #[test] + fn test_normalize_shredded_variant_uses_spark_object_key_order() { + let keys = unicode_object_keys(); + + let metadata_builder = + VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); + let (metadata_bytes, _) = metadata_builder.finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + + let mut object_fields = Vec::with_capacity(keys.len()); + let mut object_columns = Vec::with_capacity(keys.len()); + for (index, key) in keys.iter().enumerate() { + let value = if key == "😀" { 531 } else { index as i64 }; + let field = Field::new("typed_value", DataType::Int64, false); + let column: ArrayRef = Arc::new(Int64Array::from(vec![value])); + let shredded_field = + StructArray::try_new(Fields::from(vec![field]), vec![column], None).unwrap(); + object_fields.push(Field::new(key, shredded_field.data_type().clone(), false)); + object_columns.push(Arc::new(shredded_field) as ArrayRef); + } + let typed_value: ArrayRef = + Arc::new(StructArray::try_new(object_fields.into(), object_columns, None).unwrap()); + let physical_fields = Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]); + let physical: ArrayRef = Arc::new( + StructArray::try_new(physical_fields, vec![metadata, typed_value], None).unwrap(), + ); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_spark_unicode_object(output); + } + + #[test] + fn test_normalize_partially_shredded_spark_object_key_order() { + 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_bytes, value_bytes) = builder.finish(); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let spark_value = reorder_variant_values( + &value, + &metadata, + None, + VariantObjectKeyOrder::SparkUtf16, + false, + ) + .unwrap(); + + let shredded_k00: 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_k00.data_type().clone(), + false, + )]), + vec![shredded_k00], + 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, spark_value, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + assert_spark_unicode_object(output.as_struct()); + } + + #[test] + fn test_normalize_unshredded_variant_uses_spark_object_key_order() { + 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_bytes, value_bytes) = builder.finish(); + + let Variant::Object(canonical) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + let canonical_fields = canonical.iter().collect::>(); + assert_eq!(canonical_fields[30].0, "\u{e000}"); + assert_eq!(canonical_fields[31].0, "😀"); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let physical: ArrayRef = + Arc::new(StructArray::try_new(physical_fields, vec![value, metadata], None).unwrap()); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false) + .with_extension_type(VariantType), + ); + + let first = normalize_variant_array(&physical, &target_field).unwrap(); + let first_value = first + .as_struct() + .column(0) + .as_binary::() + .value(0) + .to_vec(); + assert_spark_unicode_object(first.as_struct()); + + // Spark-produced already-unshredded input is UTF-16 ordered. Normalizing it again must + // remain valid without Arrow's UTF-8-order full validation. + let second = normalize_variant_array(&first, &target_field).unwrap(); + assert_spark_unicode_object(second.as_struct()); + assert_eq!( + second.as_struct().column(0).as_binary::().value(0), + first_value + ); + } + + #[test] + fn test_normalize_spark_ordered_variant_preserves_value_bytes() { + let mut builder = VariantBuilder::new(); + let mut object = builder.new_object(); + object.insert("b", 1_i64); + object.insert("a", 2_i64); + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + + let Variant::Object(object) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + assert_eq!( + object.iter().map(|(name, _)| name).collect::>(), + vec!["a", "b"] + ); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let physical: ArrayRef = + Arc::new(StructArray::try_new(physical_fields, vec![value, metadata], None).unwrap()); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + assert_eq!( + output.as_struct().column(0).as_binary::().value(0), + value_bytes + ); + } + + #[test] + fn test_normalize_variant_skips_empty_children_of_null_parent() { + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&b""[..])])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&b""[..])])); + let physical_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + physical_fields, + vec![value, metadata], + Some(NullBuffer::from(vec![false])), + ) + .unwrap(), + ); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), true).with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + assert!(output.is_null(0)); + assert!(output.as_struct().column(0).is_null(0)); + } + #[test] fn test_cast_timestamp_micros_to_millis_array() { // Create a TimestampMicrosecond array with some values diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index f615d664763..91ea6225e6a 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -732,8 +732,14 @@ case class CometExecRule(session: SparkSession) private def tryConvertToComet( op: SparkPlan, handler: CometOperatorSerde[_]): Option[SparkPlan] = { + // Get the actual data-producing children (unwrap WriteFilesExec if present). + val dataProducingChildren = op.children.flatMap { + case writeFiles: WriteFilesExec => Seq(writeFiles.child) + case other => Seq(other) + } + if (!op.isInstanceOf[CometScanExec] && - (op.output ++ op.children.flatMap(_.output)).exists(attr => + (op.output ++ dataProducingChildren.flatMap(_.output)).exists(attr => containsVariantType(attr.dataType))) { withFallbackReason( op, @@ -747,11 +753,6 @@ case class CometExecRule(session: SparkSession) // children are CometNativeExec. This prevents runtime failures when the native operator // expects Arrow arrays but receives non-Arrow data (e.g., OnHeapColumnVector). if (serde.requiresNativeChildren && op.children.nonEmpty) { - // Get the actual data-producing children (unwrap WriteFilesExec if present) - val dataProducingChildren = op.children.flatMap { - case writeFiles: WriteFilesExec => Seq(writeFiles.child) - case other => Seq(other) - } if (!dataProducingChildren.forall(_.isInstanceOf[CometNativeExec])) { withFallbackReason( op, diff --git a/spark/src/main/scala/org/apache/comet/rules/EliminateRedundantTransitions.scala b/spark/src/main/scala/org/apache/comet/rules/EliminateRedundantTransitions.scala index ec277cfc7bc..0176c8db1f5 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EliminateRedundantTransitions.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EliminateRedundantTransitions.scala @@ -32,7 +32,7 @@ import org.apache.spark.sql.execution.exchange.ReusedExchangeExec import org.apache.comet.CometConf import org.apache.comet.CometSparkSessionExtensions.withInfo import org.apache.comet.serde.NativeOptIn -import org.apache.comet.shims.ShimSQLConf +import org.apache.comet.shims.{CometTypeShim, ShimSQLConf} // This rule is responsible for eliminating redundant transitions between row-based and // columnar-based operators for Comet. Currently, three potential redundant transitions are: @@ -58,6 +58,7 @@ import org.apache.comet.shims.ShimSQLConf case class EliminateRedundantTransitions(session: SparkSession) extends Rule[SparkPlan] with ShimCometMapInBatch + with CometTypeShim with ShimSQLConf { private lazy val showTransformations = CometConf.COMET_EXPLAIN_TRANSFORMATIONS.get() @@ -206,6 +207,8 @@ case class EliminateRedundantTransitions(session: SparkSession) } else { matchMapInArrow(plan) .orElse(matchMapInPandas(plan)) + .filterNot(info => + (info.output ++ info.child.output).exists(attr => containsVariantType(attr.dataType))) .flatMap(info => extractColumnarChild(info.child).map(child => (info, child))) } } diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala index c80fd59a39d..0a5197762f5 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala @@ -23,7 +23,7 @@ import scala.collection.mutable.ListBuffer import scala.jdk.CollectionConverters._ import org.apache.spark.internal.Logging -import org.apache.spark.sql.catalyst.expressions.{Expression, Literal} +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, Literal} import org.apache.spark.sql.catalyst.util.ResolveDefaultColumns.getExistenceDefaultValues import org.apache.spark.sql.comet.{CometNativeExec, CometNativeScanExec, CometScanExec} import org.apache.spark.sql.execution.{FileSourceScanExec, InSubqueryExec, SubqueryAdaptiveBroadcastExec} @@ -50,6 +50,31 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS // DataFusion's table_partition_cols literal substitution matches by name, so a bare name // like "file_size" could collide with a real column of the same name. Prefix to avoid it. private val constantMetadataFieldPrefix = "_comet_metadata_" + private val unsupportedDefaultReason = + "Full native scan disabled because one or more column default values are not supported" + + private def serializeExistenceDefaultValues( + schema: StructType, + output: Seq[Attribute]): Option[(Seq[Expr], Seq[java.lang.Long])] = { + val serialized = getExistenceDefaultValues(schema).iterator + .zip(schema.fields.iterator) + .zipWithIndex + .collect { + case ((value, field), index) if value != null => + // Variant expressions remain unsupported generally. Only scan defaults use the existing + // physical Arrow storage struct so the native schema adapter can fill a missing column. + val proto = + if (isVariantType(field.dataType)) { + variantDefaultExpression(value).flatMap(exprToProto(_, output)) + } else { + exprToProto(Literal(value), output) + } + proto.map(_ -> java.lang.Long.valueOf(index.toLong)) + } + .toSeq + + if (serialized.forall(_.isDefined)) Some(serialized.flatten.unzip) else None + } /** Determine whether the scan is supported and tag the Spark plan with any fallback reasons */ def isSupported(scanExec: FileSourceScanExec): Boolean = { @@ -93,6 +118,10 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS withFallbackReason(scanExec, "Full native scan disabled because ignoreMissingFiles enabled") } + if (serializeExistenceDefaultValues(scanExec.requiredSchema, scanExec.output).isEmpty) { + withFallbackReason(scanExec, unsupportedDefaultReason) + } + // the scan is supported if no fallback reasons were added to the node !hasFallbackReason(scanExec) } @@ -144,23 +173,15 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS commonBuilder.addAllDataFilters(dataFilters.asJava) } - val possibleDefaultValues = getExistenceDefaultValues(scan.requiredSchema) - if (possibleDefaultValues.exists(_ != null)) { - // Our schema has default values. Serialize two lists, one with the default values - // and another with the indexes in the schema so the native side can map missing - // columns to these default values. - val (defaultValues, indexes) = possibleDefaultValues.iterator.zipWithIndex - .filter { case (expr, _) => expr != null } - .map { case (expr, index) => - // ResolveDefaultColumnsUtil.getExistenceDefaultValues has evaluated these - // expressions and they should now just be literals. - (Literal(expr), index.toLong.asInstanceOf[java.lang.Long]) - } - .toList - .unzip - commonBuilder.addAllDefaultValues( - defaultValues.flatMap(exprToProto(_, scan.output)).asJava) - commonBuilder.addAllDefaultValuesIndexes(indexes.asJava) + serializeExistenceDefaultValues(scan.requiredSchema, scan.output) match { + case Some((defaultValues, indexes)) => + // Keep each value paired with its original required-schema index. Dropping an + // unsupported value while retaining its index would shift every later default. + commonBuilder.addAllDefaultValues(defaultValues.asJava) + commonBuilder.addAllDefaultValuesIndexes(indexes.asJava) + case None => + withFallbackReason(scan, unsupportedDefaultReason) + return None } // Extract object store options from first file (S3 configs apply to all files in scan). diff --git a/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala b/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala index 3fd97509fcf..16d4904a26a 100644 --- a/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala +++ b/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala @@ -21,6 +21,7 @@ package org.apache.comet.shims import scala.annotation.nowarn +import org.apache.spark.sql.catalyst.expressions.Expression import org.apache.spark.sql.types.{DataType, StructType} trait CometTypeShim { @@ -45,6 +46,9 @@ trait CometTypeShim { @nowarn // Spark 4 feature; VariantType doesn't exist in Spark 3.x. def variantType: Option[DataType] = None + @nowarn // Spark 4 feature; VariantVal doesn't exist in Spark 3.x. + def variantDefaultExpression(value: Any): Option[Expression] = None + @nowarn // Spark 4.1 feature; TimeType doesn't exist in Spark 3.x. def isTimeType(dt: DataType): Boolean = false } diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala index e40023c6030..b6d5925b333 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala @@ -19,8 +19,10 @@ package org.apache.comet.shims +import org.apache.spark.sql.catalyst.expressions.{CreateNamedStruct, Expression, Literal} import org.apache.spark.sql.execution.datasources.VariantMetadata import org.apache.spark.sql.types.{ArrayType, DataType, MapType, StringType, StructType, VariantType} +import org.apache.spark.unsafe.types.VariantVal trait CometTypeShim { // A `StringType` carries collation metadata in Spark 4.0. Only non-default (non-UTF8_BINARY) @@ -69,6 +71,22 @@ trait CometTypeShim { def variantType: Option[DataType] = Some(VariantType) + // Expose Variant defaults to the native scan as their Arrow storage struct without enabling + // Variant literals in Comet's general expression serde. + def variantDefaultExpression(value: Any): Option[Expression] = value match { + case variant: VariantVal => + val variantValue = variant.getValue + val metadata = variant.getMetadata + if (variantValue == null || metadata == null) { + None + } else { + Some( + CreateNamedStruct( + Seq(Literal("value"), Literal(variantValue), Literal("metadata"), Literal(metadata)))) + } + case _ => None + } + def isTimeType(dt: DataType): Boolean = dt.getClass.getSimpleName.startsWith("TimeType") diff --git a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql index 435c395a033..33fc0092e66 100644 --- a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql +++ b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql @@ -21,13 +21,14 @@ -- MinSparkVersion: 4.0 -- Config: spark.sql.variant.writeShredding.enabled=false +-- Config: spark.sql.variant.pushVariantIntoScan=false statement CREATE TABLE test_variant(id INT, v VARIANT, tail STRING) USING parquet statement INSERT INTO test_variant VALUES - (1, parse_json('{"a": 1, "b": "hello"}'), 'object'), + (1, parse_json('{"b": "hello", "a": 1}'), 'object'), (2, parse_json('[1, true, "x"]'), 'array'), (3, parse_json('42'), 'scalar'), (4, parse_json('null'), 'json-null'), @@ -52,6 +53,16 @@ SELECT v FROM test_variant query SELECT id, v, tail FROM test_variant +-- Spark's pushed VariantStruct remains an explicit fallback in Phase A. +statement +SET spark.sql.variant.pushVariantIntoScan=true + +query expect_fallback(shredded; not supported by native scan) +SELECT v FROM test_variant + +statement +SET spark.sql.variant.pushVariantIntoScan=false + query expect_fallback(type VariantType) SELECT v FROM test_variant ORDER BY id @@ -73,6 +84,62 @@ SELECT COUNT(*) FROM test_variant WHERE v IS NOT NULL query expect_fallback(type VariantType) SELECT CAST(v AS STRING) FROM test_variant +-- A Variant existence default is read from Spark's table schema and applied only when an old +-- Parquet file does not contain the column. variant_get remains a Spark expression, while the +-- ordinary Parquet scan and missing-column substitution stay native. +statement +CREATE TABLE test_variant_defaults_sql(id INT) USING parquet + +statement +INSERT INTO test_variant_defaults_sql VALUES (1) + +statement +ALTER TABLE test_variant_defaults_sql ADD COLUMNS( + v VARIANT DEFAULT parse_json('{"a":1}'), n INT DEFAULT 7) + +statement +INSERT INTO test_variant_defaults_sql VALUES (2, parse_json('{"a":2}'), 8) + +statement +SET spark.sql.parquet.enableVectorizedReader=false + +statement +SET spark.comet.scan.allowDisabledParquetVectorizedReader=true + +query expect_fallback(type VariantType) +SELECT id, variant_get(v, '$.a', 'int') AS a, n +FROM test_variant_defaults_sql ORDER BY id + +statement +SET spark.sql.parquet.enableVectorizedReader=true + +statement +SET spark.comet.scan.allowDisabledParquetVectorizedReader=false + +-- Arrow and Spark order supplementary Unicode object keys differently. Force a shredded field so +-- the native scan reconstructs the whole 32-field value before Spark's variant_get binary search. +statement +SET spark.sql.variant.writeShredding.enabled=true + +statement +SET spark.sql.variant.forceShreddingSchemaForTest=k00 BIGINT + +statement +CREATE TABLE test_variant_unicode(v VARIANT) USING parquet + +statement +INSERT INTO test_variant_unicode VALUES (parse_json( + '{"k00":0,"k01":1,"k02":2,"k03":3,"k04":4,"k05":5,"k06":6,"k07":7,"k08":8,"k09":9,"k10":10,"k11":11,"k12":12,"k13":13,"k14":14,"k15":15,"k16":16,"k17":17,"k18":18,"k19":19,"k20":20,"k21":21,"k22":22,"k23":23,"k24":24,"k25":25,"k26":26,"k27":27,"k28":28,"k29":29,"\uE000":30,"😀":531}')) + +statement +SET spark.sql.variant.writeShredding.enabled=false + +statement +SET spark.sql.variant.allowReadingShredded=true + +query expect_fallback(type VariantType) +SELECT variant_get(v, '$.😀', 'bigint') FROM test_variant_unicode + statement CREATE TABLE test_variant_struct(id INT, s STRUCT, tail STRING) USING parquet diff --git a/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala b/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala index eef77d88246..9fbe38ac97e 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala @@ -38,13 +38,42 @@ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{ArrayType, LongType, MapType, Metadata, MetadataBuilder, StringType, StructField, StructType} import org.apache.comet.CometConf -import org.apache.comet.CometSparkSessionExtensions.isSpark35Plus +import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus} import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator, SchemaGenOptions} class CometParquetWriterSuite extends CometTestBase { import testImplicits._ + test("parquet write with Variant input falls back to Spark") { + assume(isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempPath { dir => + val inputPath = new File(dir, "input.parquet").getAbsolutePath + val outputPath = new File(dir, "output.parquet").getAbsolutePath + + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + "spark.sql.variant.writeShredding.enabled" -> "false") { + sql("SELECT parse_json('42') AS v").write.parquet(inputPath) + } + + val input = spark.read.parquet(inputPath) + withSQLConf( + CometConf.COMET_NATIVE_PARQUET_WRITE_ENABLED.key -> "true", + CometConf.COMET_OPERATOR_DATA_WRITING_COMMAND_ALLOW_INCOMPAT.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val plan = captureWritePlan(path => input.write.parquet(path), outputPath) + assertNoCometNativeWriteExec(plan) + } + + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + assert(spark.read.parquet(outputPath).collect().map(_.get(0).toString).toSeq == Seq("42")) + } + } + } + test("partitioned write with empty string partition value") { withTempPath { path => Seq(("", 1), ("a", 2)) diff --git a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala index cb5d440452b..8b8e56bcd19 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala @@ -55,6 +55,16 @@ import org.apache.comet.vector.CometStructVector abstract class ParquetReadSuite extends CometTestBase { import testImplicits._ + private def normalizedVariantRows(df: DataFrame, variantOrdinal: Int): Seq[Seq[Any]] = { + df + .collect() + .map { row => + row.toSeq.updated(variantOrdinal, Option(row.get(variantOrdinal)).map(_.toString).orNull) + } + .sortBy(_.map(value => Option(value).fold("0")(v => "1" + v.toString)).mkString("\u0000")) + .toSeq + } + testStandardAndLegacyModes("decimals") { Seq(16, 1024).foreach { batchSize => withSQLConf( @@ -91,18 +101,6 @@ abstract class ParquetReadSuite extends CometTestBase { test("native scan projects Variant through a Spark-compatible vector") { assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") - def normalizedRows(df: DataFrame, variantOrdinal: Int): Seq[Seq[Any]] = { - df - .collect() - .map { row => - row.toSeq.updated( - variantOrdinal, - Option(row.get(variantOrdinal)).map(_.toString).orNull) - } - .sortBy(_.map(value => Option(value).fold("0")(v => "1" + v.toString)).mkString("\u0000")) - .toSeq - } - Seq(false, true).foreach { shredded => withTable("variant_projection") { withSQLConf( @@ -111,7 +109,7 @@ abstract class ParquetReadSuite extends CometTestBase { "spark.sql.variant.forceShreddingSchemaForTest" -> "a BIGINT") { sql("CREATE TABLE variant_projection(id INT, v VARIANT, tail STRING) USING parquet") sql("""INSERT INTO variant_projection VALUES - |(1, parse_json('{"a": 10, "b": "hello"}'), 'object'), + |(1, parse_json('{"b": "hello", "a": 10}'), 'object'), |(2, parse_json('[1, true, "x"]'), 'array'), |(3, parse_json('42'), 'scalar'), |(4, parse_json('null'), 'json-null'), @@ -126,7 +124,7 @@ abstract class ParquetReadSuite extends CometTestBase { CometConf.COMET_ENABLED.key -> "false", "spark.sql.variant.allowReadingShredded" -> "true") { expected = queries.map { case (query, variantOrdinal) => - normalizedRows(sql(query), variantOrdinal) + normalizedVariantRows(sql(query), variantOrdinal) } } @@ -137,7 +135,7 @@ abstract class ParquetReadSuite extends CometTestBase { "spark.sql.variant.pushVariantIntoScan" -> "false") { val plans = queries.zip(expected).map { case ((query, variantOrdinal), expectedRows) => val df = sql(query) - assert(normalizedRows(df, variantOrdinal) == expectedRows) + assert(normalizedVariantRows(df, variantOrdinal) == expectedRows) df.queryExecution.executedPlan } @@ -204,6 +202,154 @@ abstract class ParquetReadSuite extends CometTestBase { } } + test("native scan preserves Variant existence default pairing") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTable("variant_defaults") { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + sql("CREATE TABLE variant_defaults(v VARIANT DEFAULT parse_json('1')) USING parquet") + sql("INSERT INTO variant_defaults VALUES (parse_json('42'))") + sql("ALTER TABLE variant_defaults ADD COLUMNS(n INT DEFAULT 7)") + } + + withSQLConf("spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql("SELECT v, n FROM variant_defaults") + assert(normalizedVariantRows(df, 0) == Seq(Seq("42", 7))) + val cometPlan = df.queryExecution.executedPlan + assert(collect(cometPlan) { case _: CometNativeScanExec => true }.size == 1) + } + } + } + + test("native scan fills a Variant existence default for an old Parquet file") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTable("variant_defaults") { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + sql("CREATE TABLE variant_defaults(id INT) USING parquet") + sql("INSERT INTO variant_defaults VALUES (1)") + sql("""ALTER TABLE variant_defaults ADD COLUMNS( + | v VARIANT DEFAULT parse_json('{"b":2,"a":1}'), n INT DEFAULT 7)""".stripMargin) + sql("""INSERT INTO variant_defaults VALUES + |(2, parse_json('42'), 8), + |(3, CAST(NULL AS VARIANT), 9), + |(4, parse_json('null'), 10)""".stripMargin) + } + + val query = "SELECT id, v, n FROM variant_defaults ORDER BY id" + // Spark's vectorized Parquet reader cannot append a VariantVal when a column is missing. + // The row reader implements the intended existence-default semantics and is the reference. + var expected = Seq.empty[Seq[Any]] + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + SQLConf.PARQUET_VECTORIZED_READER_ENABLED.key -> "false", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + expected = normalizedVariantRows(sql(query), 1) + } + assert( + expected == Seq( + Seq(1, "{\"a\":1,\"b\":2}", 7), + Seq(2, "42", 8), + Seq(3, null, 9), + Seq(4, "null", 10))) + + withSQLConf("spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql(query) + assert(normalizedVariantRows(df, 1) == expected) + val cometPlan = df.queryExecution.executedPlan + assert(collect(cometPlan) { case _: CometNativeScanExec => true }.size == 1) + } + } + } + + test("native scan exports a NUL-containing Parquet field name") { + val fieldName = "v" + 0.toChar + "suffix" + withTempPath { path => + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark.range(3).toDF(fieldName).write.parquet(path.getCanonicalPath) + } + withParquetTable(path.getCanonicalPath, "nul_field") { + val (_, cometPlan) = + checkSparkAnswerAndOperator(sql(s"SELECT `$fieldName` FROM nul_field")) + assert(collect(cometPlan) { case _: CometNativeScanExec => true }.size == 1) + } + } + } + + test("native scan decodes dictionary-encoded Variant metadata") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempDir { dir => + val path = new Path(dir.toURI.toString, "dictionary-variant.parquet") + val parquetSchema = MessageTypeParser.parseMessageType("""message root { + | optional group v { + | required binary value; + | required binary metadata; + | } + |} + |""".stripMargin) + val valueField = new ArrowField( + "value", + FieldType.notNullable(ArrowType.Binary.INSTANCE), + Collections.emptyList[ArrowField]()) + val metadataField = new ArrowField( + "metadata", + new FieldType( + false, + ArrowType.Binary.INSTANCE, + new DictionaryEncoding(0L, false, new ArrowType.Int(32, true))), + Collections.emptyList[ArrowField]()) + val variantField = new ArrowField( + "v", + new FieldType( + true, + ArrowType.Struct.INSTANCE, + null, + Collections.singletonMap("ARROW:extension:name", "arrow.parquet.variant")), + Seq(valueField, metadataField).asJava) + val arrowSchema = new ArrowSchema(Collections.singletonList(variantField)) + val footer = Collections.singletonMap( + "ARROW:schema", + Base64.getEncoder.encodeToString(arrowSchema.serializeAsMessage())) + + val variant = sql("SELECT parse_json('42')").head().get(0) + val value = variant.getClass.getMethod("getValue").invoke(variant).asInstanceOf[Array[Byte]] + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + val writer = ExampleParquetWriter + .builder(path) + .withType(parquetSchema) + .withDictionaryEncoding(true) + .withExtraMetaData(footer) + .withConf(spark.sessionState.newHadoopConf()) + .build() + + try { + (0 until 3).foreach { _ => + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(0, Binary.fromConstantByteArray(value)) + group.add(1, Binary.fromConstantByteArray(metadata)) + writer.write(row) + } + } finally { + writer.close() + } + + withTable("dictionary_variant") { + sql(s"""CREATE TABLE dictionary_variant(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + withSQLConf("spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql("SELECT v FROM dictionary_variant") + assert(normalizedVariantRows(df, 0) == Seq.fill(3)(Seq("42"))) + assert(collect(df.queryExecution.executedPlan) { case _: CometNativeScanExec => + true + }.size == 1) + } + } + } + } + // Spark ignores ARROW:schema during Parquet schema inference: // https://github.com/apache/spark/blob/v4.2.0/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/parquet/ParquetFileFormat.scala#L585-L599 // With binaryAsString, Spark maps unannotated BINARY to StringType: diff --git a/spark/src/test/spark-4.x/org/apache/spark/sql/comet/CometMapInBatchSuite.scala b/spark/src/test/spark-4.x/org/apache/spark/sql/comet/CometMapInBatchSuite.scala index 09802dddaa4..9af3814501a 100644 --- a/spark/src/test/spark-4.x/org/apache/spark/sql/comet/CometMapInBatchSuite.scala +++ b/spark/src/test/spark-4.x/org/apache/spark/sql/comet/CometMapInBatchSuite.scala @@ -27,7 +27,7 @@ import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, ExprId, PythonUDF} import org.apache.spark.sql.execution.{ColumnarToRowExec, LeafExecNode} import org.apache.spark.sql.execution.python.MapInArrowExec -import org.apache.spark.sql.types.{LongType, StructField, StructType} +import org.apache.spark.sql.types.{ArrayType, DataType, LongType, StructField, StructType, VariantType} import org.apache.spark.sql.vectorized.ColumnarBatch import org.apache.comet.{CometConf, ExtendedExplainInfo} @@ -78,10 +78,18 @@ class CometMapInBatchSuite extends CometTestBase { } private def buildPlan(): MapInArrowExec = { - val cometChild = StubCometLeaf(Seq(AttributeReference("id", LongType)(ExprId(0L)))) + buildPlan(LongType, LongType) + } + + private def buildPlan(inputType: DataType, outputType: DataType): MapInArrowExec = { + val input = AttributeReference("id", inputType)(ExprId(0L)) + val output = Seq(AttributeReference("id", outputType)(ExprId(1L))) + val cometChild = StubCometLeaf(Seq(input)) MapInArrowExec( - stubPythonUDF, - cometChild.output, + stubPythonUDF.copy( + dataType = StructType(Seq(StructField("id", outputType))), + children = Seq(input)), + output, ColumnarToRowExec(cometChild), isBarrier = false, profile = None) @@ -96,6 +104,23 @@ class CometMapInBatchSuite extends CometTestBase { } } + test("rule does not rewrite MapInArrowExec with Variant-bearing input or output") { + val nestedVariant = StructType(Seq(StructField("v", VariantType))) + val plans = Seq( + buildPlan(VariantType, LongType), + buildPlan(nestedVariant, LongType), + buildPlan(LongType, ArrayType(VariantType, containsNull = true))) + + withSQLConf(CometConf.COMET_PYARROW_UDF_ENABLED.key -> "true") { + plans.foreach { plan => + val rewritten = EliminateRedundantTransitions(spark).apply(plan) + assert( + !rewritten.exists(_.isInstanceOf[CometMapInBatchExec]), + s"unexpected CometMapInBatchExec for Variant-bearing schema:\n$rewritten") + } + } + } + test("rule does not rewrite when feature is disabled") { withSQLConf(CometConf.COMET_PYARROW_UDF_ENABLED.key -> "false") { val rewritten = EliminateRedundantTransitions(spark).apply(buildPlan()) From c355fefd9c0b7e96d86523a7214bb2cdd47e1a55 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sun, 23 Aug 2026 12:26:52 +0800 Subject: [PATCH 03/14] review --- native/core/src/parquet/cast_column.rs | 60 ++++++++++++++++++- .../sql-tests/expressions/misc/variant.sql | 7 +-- 2 files changed, 62 insertions(+), 5 deletions(-) diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index a16626d2d1c..2137ac8ef3e 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -423,8 +423,11 @@ fn reorder_variant_values( ))); } - let metadata = VariantMetadata::try_new(metadata.value(index))?; let rebuilt = catch_unwind(AssertUnwindSafe(|| { + // Spark encodes empty object keys with equal metadata offsets, which Arrow 58.4's + // full validator rejects. Keep shallow parsing and all accesses inside this boundary. + // https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant/src/variant/metadata.rs#L307-L317 + let metadata = VariantMetadata::new(metadata.value(index)); let variant = Variant::new_with_metadata(metadata.clone(), value.value(index)); if is_compatible_variant(&variant, order) { return None; @@ -953,6 +956,61 @@ mod tests { ); } + #[test] + fn test_normalize_variant_preserves_empty_object_keys() { + let mut builder = VariantBuilder::new(); + let mut object = builder.new_object(); + object.insert("", 1_i64); + let mut nested = object.new_object("nested"); + nested.insert("", 2_i64); + nested.finish(); + object.finish(); + let (mut metadata_bytes, value_bytes) = builder.finish(); + + // Spark leaves the metadata dictionary unsorted. Equal offsets encode the empty key. + metadata_bytes[0] &= !0x10; + assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]), + vec![value, metadata], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_eq!(output.column(0).as_binary::().value(0), value_bytes); + assert_eq!(output.column(1).as_binary::().value(0), metadata_bytes); + + let Variant::Object(object) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + assert_eq!(object.get(""), Some(Variant::from(1_i64))); + let Variant::Object(nested) = object.get("nested").unwrap() else { + panic!("expected nested object") + }; + assert_eq!(nested.get(""), Some(Variant::from(2_i64))); + } + #[test] fn test_normalize_variant_skips_empty_children_of_null_parent() { let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&b""[..])])); diff --git a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql index 33fc0092e66..37650c26182 100644 --- a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql +++ b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql @@ -22,6 +22,7 @@ -- MinSparkVersion: 4.0 -- Config: spark.sql.variant.writeShredding.enabled=false -- Config: spark.sql.variant.pushVariantIntoScan=false +-- Config: spark.sql.variant.forceShreddingSchemaForTest=k00 BIGINT statement CREATE TABLE test_variant(id INT, v VARIANT, tail STRING) USING parquet @@ -32,7 +33,8 @@ INSERT INTO test_variant VALUES (2, parse_json('[1, true, "x"]'), 'array'), (3, parse_json('42'), 'scalar'), (4, parse_json('null'), 'json-null'), - (5, CAST(NULL AS VARIANT), 'sql-null') + (5, CAST(NULL AS VARIANT), 'sql-null'), + (6, parse_json('{"":1,"nested":{"":2}}'), 'empty-key') -- A plain Parquet scan can remain native when its required schema prunes the -- Variant column completely, including both SQL NULL and Variant null values. @@ -121,9 +123,6 @@ SET spark.comet.scan.allowDisabledParquetVectorizedReader=false statement SET spark.sql.variant.writeShredding.enabled=true -statement -SET spark.sql.variant.forceShreddingSchemaForTest=k00 BIGINT - statement CREATE TABLE test_variant_unicode(v VARIANT) USING parquet From c556a522f1953075f2f8a8fd4c2235edf5e4fdb9 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sun, 23 Aug 2026 22:09:23 +0800 Subject: [PATCH 04/14] add link about spark and arrow-rs issue that need to be fixed so we can get cleaner code --- native/core/src/parquet/cast_column.rs | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index 2137ac8ef3e..6e9838e9889 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -270,7 +270,8 @@ fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult DataFusionResult { let Some(struct_array) = array.as_any().downcast_ref::() else { return Ok(Arc::clone(array)); @@ -392,6 +393,11 @@ fn is_compatible_variant(variant: &Variant<'_, '_>, order: VariantObjectKeyOrder /// Reorder object keys for either Arrow's UTF-8 order or Spark's Java UTF-16 order. Preserve /// already-compatible values byte-for-byte and retain the original metadata dictionary. +/// SPARK-56637 tracks this mismatch. The metadata dictionary's sorted flag affects dictionary +/// lookup, not object-entry ordering; Spark's builder and lookup must agree while continuing to +/// read Variant values already written by Spark 4.x in UTF-16 order. +/// https://issues.apache.org/jira/browse/SPARK-56637 +/// https://github.com/apache/spark/pull/55928 fn reorder_variant_values( value: &ArrayRef, metadata: &ArrayRef, @@ -427,6 +433,7 @@ fn reorder_variant_values( // Spark encodes empty object keys with equal metadata offsets, which Arrow 58.4's // full validator rejects. Keep shallow parsing and all accesses inside this boundary. // https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant/src/variant/metadata.rs#L307-L317 + // Upstream fix: https://github.com/apache/arrow-rs/pull/10352 let metadata = VariantMetadata::new(metadata.value(index)); let variant = Variant::new_with_metadata(metadata.clone(), value.value(index)); if is_compatible_variant(&variant, order) { From 9742a7b7d18f23fd7d834162e085601c4e13cfcb Mon Sep 17 00:00:00 2001 From: peterxcli Date: Sun, 23 Aug 2026 23:23:30 +0800 Subject: [PATCH 05/14] update spark issue link --- native/core/src/parquet/cast_column.rs | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index 6e9838e9889..c4a77427d8e 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -393,11 +393,11 @@ fn is_compatible_variant(variant: &Variant<'_, '_>, order: VariantObjectKeyOrder /// Reorder object keys for either Arrow's UTF-8 order or Spark's Java UTF-16 order. Preserve /// already-compatible values byte-for-byte and retain the original metadata dictionary. -/// SPARK-56637 tracks this mismatch. The metadata dictionary's sorted flag affects dictionary -/// lookup, not object-entry ordering; Spark's builder and lookup must agree while continuing to -/// read Variant values already written by Spark 4.x in UTF-16 order. -/// https://issues.apache.org/jira/browse/SPARK-56637 -/// https://github.com/apache/spark/pull/55928 +/// SPARK-58949 tracks this mismatch and legacy compatibility. The metadata dictionary's sorted +/// flag affects dictionary lookup, not object-entry ordering; Spark's builder and lookup must +/// agree while continuing to read Variant values already written by Spark 4.x in UTF-16 order. +/// https://issues.apache.org/jira/browse/SPARK-58949 +/// https://github.com/apache/parquet-java/issues/3735 fn reorder_variant_values( value: &ArrayRef, metadata: &ArrayRef, From 33e513cf24ad043728152d97596a5d25e2514cc0 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Mon, 24 Aug 2026 22:00:05 +0800 Subject: [PATCH 06/14] widen unsigned shredded field --- native/core/src/parquet/cast_column.rs | 146 ++++++++++++++++++ .../sql-tests/expressions/misc/variant.sql | 21 +++ 2 files changed, 167 insertions(+) diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index c4a77427d8e..3933b5b0ea4 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -208,6 +208,7 @@ fn normalize_variant_array( } let array = decode_variant_metadata_dictionary(array)?; + let array = widen_unsigned_variant_typed_value(&array)?; let variant = prepare_variant_for_unshredding(&VariantArray::try_new(array.as_ref())?)?; let unshredded = unshred_variant(&variant)?; let value = unshredded.value_field().ok_or_else(|| { @@ -230,6 +231,81 @@ fn normalize_variant_array( Ok(Arc::new(output)) } +fn widen_unsigned_variant_type(data_type: &DataType) -> Option { + fn widen_field(field: &FieldRef) -> Option { + widen_unsigned_variant_type(field.data_type()) + .map(|data_type| Arc::new(field.as_ref().clone().with_data_type(data_type))) + } + + match data_type { + DataType::UInt8 => Some(DataType::Int16), + DataType::UInt16 => Some(DataType::Int32), + DataType::UInt32 => Some(DataType::Int64), + DataType::List(field) => widen_field(field).map(DataType::List), + DataType::LargeList(field) => widen_field(field).map(DataType::LargeList), + DataType::ListView(field) => widen_field(field).map(DataType::ListView), + DataType::LargeListView(field) => widen_field(field).map(DataType::LargeListView), + DataType::Struct(fields) => { + let mut changed = false; + let fields = fields + .iter() + .map(|field| match widen_field(field) { + Some(field) => { + changed = true; + field + } + None => Arc::clone(field), + }) + .collect::>(); + changed.then(|| DataType::Struct(fields.into())) + } + _ => None, + } +} + +/// Parquet restores unsigned integer annotations as Arrow unsigned arrays, while Spark widens +/// those values to the next signed width. Arrow Variant accepts only the latter representation. +/// arrow-rs #10416/#10417 would move this widening into `VariantArray`/`unshred_variant`; remove +/// both local `widen_unsigned_variant_*` helpers after that ships and Comet upgrades: +/// https://github.com/apache/arrow-rs/issues/10416 +/// https://github.com/apache/arrow-rs/pull/10417 +/// Arrow #50622/#50810 instead proposes removing unsigned `typed_value` mappings because the +/// Parquet Variant shredding table permits only signed integer fields. Until upstream resolves +/// that choice, keep this compatibility path for unsigned files Spark already reads: +/// https://github.com/apache/arrow/issues/50622 +/// https://github.com/apache/arrow/pull/50810 +fn widen_unsigned_variant_typed_value(array: &ArrayRef) -> DataFusionResult { + let Some(struct_array) = array.as_any().downcast_ref::() else { + return Ok(Arc::clone(array)); + }; + let Some((typed_value_index, typed_value_field)) = struct_array + .fields() + .iter() + .enumerate() + .find(|(_, field)| field.name() == "typed_value") + else { + return Ok(Arc::clone(array)); + }; + let Some(data_type) = widen_unsigned_variant_type(typed_value_field.data_type()) else { + return Ok(Arc::clone(array)); + }; + + let mut fields = struct_array.fields().iter().cloned().collect::>(); + fields[typed_value_index] = Arc::new( + typed_value_field + .as_ref() + .clone() + .with_data_type(data_type.clone()), + ); + let mut columns = struct_array.columns().to_vec(); + columns[typed_value_index] = cast(columns[typed_value_index].as_ref(), &data_type)?; + Ok(Arc::new(StructArray::try_new( + fields.into(), + columns, + struct_array.nulls().cloned(), + )?)) +} + /// Arrow's unshredder fully validates any residual `value` in a partially shredded object. Spark /// writes object keys in Java UTF-16 order, so put that residual value in Arrow UTF-8 order only /// while it passes through the upstream unshredder. @@ -652,6 +728,7 @@ mod tests { use super::*; use arrow::array::{ Array, AsArray, BinaryArray, DictionaryArray, Int32Array, Int64Array, StringArray, + UInt16Array, UInt32Array, UInt8Array, }; use arrow::datatypes::{Field, Fields, Int32Type}; use datafusion::physical_expr::expressions::Column; @@ -755,6 +832,75 @@ mod tests { assert_eq!(variant.value(2), Variant::from(30_i64)); } + #[test] + fn test_normalize_shredded_variant_widens_unsigned_values() { + let metadata_builder = VariantBuilder::new().with_field_names(["u8", "u16", "u32"]); + let (metadata_bytes, _) = metadata_builder.finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + + let fields = [ + ("u8", Arc::new(UInt8Array::from(vec![u8::MAX])) as ArrayRef), + ( + "u16", + Arc::new(UInt16Array::from(vec![u16::MAX])) as ArrayRef, + ), + ( + "u32", + Arc::new(UInt32Array::from(vec![u32::MAX])) as ArrayRef, + ), + ]; + let mut object_fields = Vec::with_capacity(fields.len()); + let mut object_columns = Vec::with_capacity(fields.len()); + for (name, value) in fields { + let shredded = StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + value.data_type().clone(), + false, + )]), + vec![value], + None, + ) + .unwrap(); + object_fields.push(Field::new(name, shredded.data_type().clone(), false)); + object_columns.push(Arc::new(shredded) as ArrayRef); + } + let typed_value: ArrayRef = + Arc::new(StructArray::try_new(object_fields.into(), object_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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let variant = output.value(0); + let Variant::Object(object) = variant else { + panic!("expected object") + }; + assert_eq!(object.get("u8"), Some(Variant::from(255_i16))); + assert_eq!(object.get("u16"), Some(Variant::from(65_535_i32))); + assert_eq!(object.get("u32"), Some(Variant::from(4_294_967_295_i64))); + } + #[test] fn test_normalize_shredded_variant_uses_spark_object_key_order() { let keys = unicode_object_keys(); diff --git a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql index 37650c26182..50cafbfebdc 100644 --- a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql +++ b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql @@ -139,6 +139,27 @@ SET spark.sql.variant.allowReadingShredded=true query expect_fallback(type VariantType) SELECT variant_get(v, '$.😀', 'bigint') FROM test_variant_unicode +-- Spark SQL cannot write unsigned Parquet integer annotations. Generate the signed representation +-- produced after widening here; the unsigned-to-signed conversion itself is covered in Rust. +statement +SET spark.sql.variant.forceShreddingSchemaForTest=u8 SMALLINT, u16 INT, u32 BIGINT + +statement +SET spark.sql.variant.writeShredding.enabled=true + +statement +CREATE TABLE test_variant_widened(v VARIANT) USING parquet + +statement +INSERT INTO test_variant_widened VALUES + (parse_json('{"u8":255,"u16":65535,"u32":4294967295}')) + +statement +SET spark.sql.variant.writeShredding.enabled=false + +query +SELECT v FROM test_variant_widened + statement CREATE TABLE test_variant_struct(id INT, s STRUCT, tail STRING) USING parquet From ed39c4c028d2b5e993a1a8b6ee68e915f8f4acef Mon Sep 17 00:00:00 2001 From: peterxcli Date: Wed, 26 Aug 2026 01:57:35 +0800 Subject: [PATCH 07/14] fix: preserve Spark Variant compatibility when unshredding --- native/core/src/parquet/cast_column.rs | 1226 ++++++++++++++++- native/core/src/parquet/schema_adapter.rs | 68 +- .../sql-tests/expressions/misc/variant.sql | 26 +- 3 files changed, 1266 insertions(+), 54 deletions(-) diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index 3933b5b0ea4..d452184b8d4 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -16,8 +16,9 @@ // under the License. use arrow::{ array::{ - make_array, Array, ArrayRef, BinaryArray, BinaryBuilder, LargeListArray, ListArray, - MapArray, StructArray, TimestampMicrosecondArray, TimestampMillisecondArray, + make_array, Array, ArrayRef, AsArray, BinaryArray, BinaryBuilder, LargeListArray, + ListArray, ListLikeArray, MapArray, StructArray, TimestampMicrosecondArray, + TimestampMillisecondArray, }, buffer::NullBuffer, compute::{cast, CastOptions}, @@ -36,10 +37,12 @@ use datafusion::common::{DataFusionError, Result as DataFusionResult}; use datafusion::logical_expr::ColumnarValue; use datafusion::physical_expr::PhysicalExpr; use parquet::variant::{ - unshred_variant, MetadataBuilder, ParentState, ReadOnlyMetadataBuilder, ValueBuilder, Variant, - VariantArray, VariantMetadata, + unshred_variant, BorrowedShreddingState, ListBuilder, MetadataBuilder, ObjectBuilder, + ParentState, ReadOnlyMetadataBuilder, ValueBuilder, Variant, VariantArray, VariantBuilder, + VariantDecimal4, VariantDecimal8, VariantMetadata, WritableMetadataBuilder, }; use std::{ + collections::HashSet, fmt::{self, Display}, hash::Hash, panic::{catch_unwind, AssertUnwindSafe}, @@ -209,20 +212,26 @@ fn normalize_variant_array( let array = decode_variant_metadata_dictionary(array)?; let array = widen_unsigned_variant_typed_value(&array)?; - let variant = prepare_variant_for_unshredding(&VariantArray::try_new(array.as_ref())?)?; - let unshredded = unshred_variant(&variant)?; + let variant = VariantArray::try_new(array.as_ref())?; + let was_shredded = variant.typed_value_field().is_some(); + let unshredded = unshred_variant_for_spark(&variant)?; 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(), - VariantObjectKeyOrder::SparkUtf16, - false, - )?; + let (value, metadata) = if was_shredded { + rebuild_shredded_variant_for_spark(&variant, &value, &metadata, unshredded.inner().nulls())? + } else { + let value = reorder_variant_values( + &value, + &metadata, + unshredded.inner().nulls(), + VariantObjectKeyOrder::SparkUtf16, + false, + )?; + (value, metadata) + }; let output = StructArray::try_new( fields.clone(), vec![value, metadata], @@ -231,6 +240,20 @@ fn normalize_variant_array( Ok(Arc::new(output)) } +fn unshred_variant_for_spark(variant: &VariantArray) -> DataFusionResult { + let first = + prepare_variant_for_unshredding(variant).and_then(|array| Ok(unshred_variant(&array)?)); + let first_error = match first { + Ok(array) => return Ok(array), + Err(error) => error, + }; + let Some(variant) = canonicalize_spark_empty_key_metadata(variant)? else { + return Err(first_error); + }; + let variant = prepare_variant_for_unshredding(&variant)?; + Ok(unshred_variant(&variant)?) +} + fn widen_unsigned_variant_type(data_type: &DataType) -> Option { fn widen_field(field: &FieldRef) -> Option { widen_unsigned_variant_type(field.data_type()) @@ -343,6 +366,135 @@ fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult DataFusionResult> { + type Replacement = Option<(Vec, Option>)>; + + let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let metadata = metadata.as_binary::(); + let value = variant + .value_field() + .map(|value| cast(value.as_ref(), &DataType::Binary)) + .transpose()?; + let value = value.as_ref().map(|value| value.as_binary::()); + let mut replacements = Vec::with_capacity(variant.len()); + let mut changed = false; + + for index in 0..variant.len() { + if variant.inner().is_null(index) || metadata.is_null(index) { + replacements.push(None); + continue; + } + let metadata_bytes = metadata.value(index); + if VariantMetadata::try_new(metadata_bytes).is_ok() { + replacements.push(None); + continue; + } + + let replacement = catch_unwind(AssertUnwindSafe(|| -> Result { + let old_metadata = VariantMetadata::new(metadata_bytes); + let mut names = old_metadata + .iter_try() + .map(|name| name.map(str::to_string)) + .collect::, _>>()?; + if !names.iter().any(String::is_empty) { + return Ok(None); + } + names.sort_unstable(); + if names.windows(2).any(|names| names[0] == names[1]) { + return Ok(None); + } + + let mut builder = + VariantBuilder::new().with_field_names(names.iter().map(String::as_str)); + match value { + Some(value) if !value.is_null(index) => { + builder + .append_value(Variant::new_with_metadata(old_metadata, value.value(index))); + let (metadata, value) = builder.finish(); + Ok(Some((metadata, Some(value)))) + } + _ => Ok(Some((builder.finish().0, None))), + } + })) + .map_err(|_| { + DataFusionError::Execution(format!( + "Invalid Variant metadata with an empty key at row {index}" + )) + })??; + changed |= replacement.is_some(); + replacements.push(replacement); + } + + if !changed { + return Ok(None); + } + + let mut metadata_builder = BinaryBuilder::new(); + let mut value_builder = value.map(|_| BinaryBuilder::new()); + for (index, replacement) in replacements.iter().enumerate() { + match replacement { + Some((metadata, value)) => { + metadata_builder.append_value(metadata); + if let Some(builder) = &mut value_builder { + match value { + Some(value) => builder.append_value(value), + None => builder.append_null(), + } + } + } + None => { + if metadata.is_null(index) { + metadata_builder.append_null(); + } else { + metadata_builder.append_value(metadata.value(index)); + } + if let (Some(value), Some(builder)) = (value, &mut value_builder) { + if value.is_null(index) { + builder.append_null(); + } else { + builder.append_value(value.value(index)); + } + } + } + } + } + + let mut fields = variant.inner().fields().iter().cloned().collect::>(); + let mut columns = variant.inner().columns().to_vec(); + let metadata_index = fields + .iter() + .position(|field| field.name() == "metadata") + .unwrap(); + fields[metadata_index] = Arc::new( + fields[metadata_index] + .as_ref() + .clone() + .with_data_type(DataType::Binary), + ); + columns[metadata_index] = Arc::new(metadata_builder.finish()); + if let Some(mut value_builder) = value_builder { + let value_index = fields + .iter() + .position(|field| field.name() == "value") + .unwrap(); + fields[value_index] = Arc::new( + fields[value_index] + .as_ref() + .clone() + .with_data_type(DataType::Binary), + ); + columns[value_index] = Arc::new(value_builder.finish()); + } + let array = StructArray::try_new(fields.into(), columns, variant.inner().nulls().cloned())?; + Ok(Some(VariantArray::try_new(&array)?)) +} + /// Arrow-rs parquet-variant-compute allows dictionary-encoded metadata in its contract, but 58.4's /// `VariantArray::try_new` validates only Binary, LargeBinary, and BinaryView. Decode just that /// child and keep the physical struct otherwise unchanged. @@ -467,6 +619,590 @@ fn is_compatible_variant(variant: &Variant<'_, '_>, order: VariantObjectKeyOrder } } +fn compact_spark_integer(value: i64) -> Variant<'static, 'static> { + if let Ok(value) = i8::try_from(value) { + Variant::Int8(value) + } else if let Ok(value) = i16::try_from(value) { + Variant::Int16(value) + } else if let Ok(value) = i32::try_from(value) { + Variant::Int32(value) + } else { + Variant::Int64(value) + } +} + +fn compact_spark_typed_variant<'m, 'v>(variant: Variant<'m, 'v>) -> Variant<'m, 'v> { + match variant { + Variant::Int16(value) => compact_spark_integer(value.into()), + Variant::Int32(value) => compact_spark_integer(value.into()), + Variant::Int64(value) => compact_spark_integer(value), + Variant::Decimal8(value) => i32::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal4::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal4) + .unwrap_or(Variant::Decimal8(value)), + Variant::Decimal16(value) => i32::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal4::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal4) + .or_else(|| { + i64::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal8::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal8) + }) + .unwrap_or(Variant::Decimal16(value)), + Variant::Float(value) if value.is_nan() => Variant::Float(f32::from_bits(0x7fc0_0000)), + Variant::Double(value) if value.is_nan() => { + Variant::Double(f64::from_bits(0x7ff8_0000_0000_0000)) + } + variant => variant, + } +} + +/// Re-encode a residual Variant against `metadata`. Scalar widths are intentionally preserved, +/// matching Spark's `VariantBuilder.appendVariant` behavior. +fn spark_variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + match variant { + Variant::Object(object) => { + let mut metadata_builder = SparkMetadataBuilder::new(metadata); + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for (name, value) in object.iter() { + let value = spark_variant_bytes(metadata, value)?; + 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), + false, + ); + for value in list.iter() { + let value = spark_variant_bytes(metadata, value)?; + 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()) +} + +fn variant_binary_value(array: &ArrayRef, index: usize) -> Result, ArrowError> { + if array.is_null(index) { + return Ok(None); + } + let value = match array.data_type() { + DataType::Binary => array.as_binary::().value(index), + DataType::LargeBinary => array.as_binary::().value(index), + DataType::BinaryView => array.as_binary_view().value(index), + data_type => { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant value must be binary-like, got {data_type}" + ))) + } + }; + Ok(Some(value)) +} + +fn shredding_state_has_value(state: &BorrowedShreddingState<'_>, index: usize) -> bool { + state + .typed_value_field() + .is_some_and(|array| array.is_valid(index)) + || state + .value_field() + .is_some_and(|array| array.is_valid(index)) +} + +fn collect_spark_field_name(name: &str, field_names: &mut Vec, seen: &mut HashSet) { + if seen.insert(name.to_string()) { + field_names.push(name.to_string()); + } +} + +fn collect_residual_field_names( + variant: Variant<'_, '_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + match variant { + Variant::Object(object) => { + for (name, value) in object.iter() { + collect_spark_field_name(name, field_names, seen); + collect_residual_field_names(value, field_names, seen)?; + } + } + Variant::List(list) => { + for value in list.iter() { + collect_residual_field_names(value, field_names, seen)?; + } + } + _ => {} + } + Ok(()) +} + +fn collect_list_field_names( + list: &L, + index: usize, + source_metadata: &VariantMetadata<'_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + let values = list.values().as_struct(); + let state = BorrowedShreddingState::try_from(values)?; + for element_index in list.element_range(index) { + collect_shredded_field_names( + state.clone(), + element_index, + source_metadata, + field_names, + seen, + )?; + } + Ok(()) +} + +fn collect_shredded_field_names( + state: BorrowedShreddingState<'_>, + index: usize, + source_metadata: &VariantMetadata<'_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + let Some(typed_value) = state + .typed_value_field() + .filter(|array| array.is_valid(index)) + else { + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + collect_residual_field_names( + Variant::new_with_metadata(source_metadata.clone(), value), + field_names, + seen, + )?; + } + } + return Ok(()); + }; + + match typed_value.data_type() { + DataType::Struct(_) => { + let object = typed_value.as_struct(); + for (field, column) in object.fields().iter().zip(object.columns()) { + let child = column.as_struct_opt().ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Invalid shredded Variant object field '{}': expected Struct, got {}", + field.name(), + column.data_type() + )) + })?; + if child.is_null(index) { + return Err(ArrowError::InvalidArgumentError(format!( + "Shredded Variant object field '{}' is null", + field.name() + ))); + } + let child_state = BorrowedShreddingState::try_from(child)?; + if shredding_state_has_value(&child_state, index) { + collect_spark_field_name(field.name(), field_names, seen); + collect_shredded_field_names( + child_state, + index, + source_metadata, + field_names, + seen, + )?; + } + } + + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + let Variant::Object(residual) = + Variant::new_with_metadata(source_metadata.clone(), value) + else { + return Err(ArrowError::InvalidArgumentError( + "Partially shredded Variant object has a non-object value".to_string(), + )); + }; + for (name, value) in residual.iter() { + if object.fields().iter().any(|field| field.name() == name) { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant field '{name}' appears in both value and typed_value" + ))); + } + collect_spark_field_name(name, field_names, seen); + collect_residual_field_names(value, field_names, seen)?; + } + } + } + } + DataType::List(_) => collect_list_field_names( + typed_value.as_list::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::LargeList(_) => collect_list_field_names( + typed_value.as_list::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::ListView(_) => collect_list_field_names( + typed_value.as_list_view::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::LargeListView(_) => collect_list_field_names( + typed_value.as_list_view::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::FixedSizeList(_, _) => collect_list_field_names( + typed_value.as_fixed_size_list(), + index, + source_metadata, + field_names, + seen, + )?, + _ => {} + } + Ok(()) +} + +fn spark_typed_variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + ValueBuilder::try_append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + compact_spark_typed_variant(variant), + )?; + Ok(value_builder.into_inner()) +} + +fn spark_list_bytes( + list: &L, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Variant::List(semantic) = semantic else { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant list did not unshred to a list".to_string(), + )); + }; + let semantic = semantic.iter_try().collect::, _>>()?; + let element_range = list.element_range(index); + if element_range.len() != semantic.len() { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant list length changed while unshredding".to_string(), + )); + } + + let values = list.values().as_struct(); + let state = BorrowedShreddingState::try_from(values)?; + let mut elements = Vec::with_capacity(semantic.len()); + for (element_index, semantic) in element_range.zip(semantic) { + elements.push(spark_shredded_variant_bytes( + state.clone(), + element_index, + semantic, + source_metadata, + target_metadata, + )?); + } + + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = ReadOnlyMetadataBuilder::new(target_metadata); + let mut builder = ListBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for element in elements { + builder.append_value_bytes(Variant::new_with_metadata( + target_metadata.clone(), + &element, + )); + } + builder.finish(); + Ok(value_builder.into_inner()) +} + +fn spark_object_bytes( + state: BorrowedShreddingState<'_>, + object: &StructArray, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Variant::Object(semantic) = semantic else { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant object did not unshred to an object".to_string(), + )); + }; + let mut entries = Vec::new(); + for (field, column) in object.fields().iter().zip(object.columns()) { + let child = column.as_struct_opt().ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Invalid shredded Variant object field '{}': expected Struct, got {}", + field.name(), + column.data_type() + )) + })?; + if child.is_null(index) { + return Err(ArrowError::InvalidArgumentError(format!( + "Shredded Variant object field '{}' is null", + field.name() + ))); + } + let child_state = BorrowedShreddingState::try_from(child)?; + if shredding_state_has_value(&child_state, index) { + let value = semantic.get(field.name()).ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Unshredded Variant is missing field '{}'", + field.name() + )) + })?; + entries.push(( + field.name().to_string(), + spark_shredded_variant_bytes( + child_state, + index, + value, + source_metadata, + target_metadata, + )?, + )); + } + } + + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + let Variant::Object(residual) = + Variant::new_with_metadata(source_metadata.clone(), value) + else { + return Err(ArrowError::InvalidArgumentError( + "Partially shredded Variant object has a non-object value".to_string(), + )); + }; + for (name, value) in residual.iter() { + if object.fields().iter().any(|field| field.name() == name) { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant field '{name}' appears in both value and typed_value" + ))); + } + entries.push(( + name.to_string(), + spark_variant_bytes(target_metadata, value)?, + )); + } + } + } + + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = SparkMetadataBuilder::new(target_metadata); + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for (name, value) in entries { + builder.try_insert_bytes( + &name, + Variant::new_with_metadata(target_metadata.clone(), &value), + )?; + } + builder.finish(); + Ok(value_builder.into_inner()) +} + +// ponytail: recursive child buffers can be O(depth²); use a streaming encoder only if deeply +// nested Variant profiles show this compatibility path is a bottleneck. +fn spark_shredded_variant_bytes( + state: BorrowedShreddingState<'_>, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Some(typed_value) = state + .typed_value_field() + .filter(|array| array.is_valid(index)) + else { + return match state.value_field() { + Some(value) => match variant_binary_value(value, index)? { + Some(value) => spark_variant_bytes( + target_metadata, + Variant::new_with_metadata(source_metadata.clone(), value), + ), + None => Err(ArrowError::InvalidArgumentError( + "Shredded Variant has neither value nor typed_value".to_string(), + )), + }, + None => Err(ArrowError::InvalidArgumentError( + "Shredded Variant has neither value nor typed_value".to_string(), + )), + }; + }; + + match typed_value.data_type() { + DataType::Struct(_) => spark_object_bytes( + state, + typed_value.as_struct(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::List(_) => spark_list_bytes( + typed_value.as_list::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::LargeList(_) => spark_list_bytes( + typed_value.as_list::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::ListView(_) => spark_list_bytes( + typed_value.as_list_view::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::LargeListView(_) => spark_list_bytes( + typed_value.as_list_view::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::FixedSizeList(_, _) => spark_list_bytes( + typed_value.as_fixed_size_list(), + index, + semantic, + source_metadata, + target_metadata, + ), + _ => spark_typed_variant_bytes(target_metadata, semantic), + } +} + +/// Arrow's unshredder preserves the source metadata but rebuilds object slots in UTF-8 order. +/// Rebuild from the physical shredding state so Spark's metadata insertion order, UTF-16 object +/// headers, typed scalar widths, and residual scalar bytes all remain compatible. +fn rebuild_shredded_variant_for_spark( + source: &VariantArray, + value: &ArrayRef, + metadata: &ArrayRef, + parent_nulls: Option<&NullBuffer>, +) -> DataFusionResult<(ArrayRef, ArrayRef)> { + let source_metadata = cast(source.metadata_field().as_ref(), &DataType::Binary)?; + let source_metadata = source_metadata.as_binary::(); + let source_state = source.shredding_state().borrow(); + let value = value.as_binary::(); + let metadata = metadata.as_binary::(); + let mut value_output = BinaryBuilder::new(); + let mut metadata_output = BinaryBuilder::new(); + + for index in 0..value.len() { + if parent_nulls.is_some_and(|nulls| nulls.is_null(index)) { + value_output.append_null(); + metadata_output.append_null(); + continue; + } + if value.is_null(index) || metadata.is_null(index) { + return Err(DataFusionError::Execution(format!( + "Variant value or metadata is null at row {index}" + ))); + } + + let (rebuilt_value, rebuilt_metadata) = catch_unwind(AssertUnwindSafe( + || -> Result<(Vec, Vec), ArrowError> { + let source_metadata = VariantMetadata::new(source_metadata.value(index)); + let semantic_metadata = VariantMetadata::try_new(metadata.value(index))?; + let semantic = Variant::new_with_metadata(semantic_metadata, value.value(index)); + let mut field_names = Vec::new(); + collect_shredded_field_names( + source_state.clone(), + index, + &source_metadata, + &mut field_names, + &mut HashSet::new(), + ) + .map_err(|error| { + ArrowError::InvalidArgumentError(format!( + "Failed to collect Spark Variant metadata: {error}" + )) + })?; + + let mut metadata_builder = + WritableMetadataBuilder::from_iter(field_names.iter().map(String::as_str)); + metadata_builder.finish(); + let mut rebuilt_metadata = metadata_builder.into_inner(); + // Spark's VariantBuilder never marks its insertion-ordered dictionary as sorted. + rebuilt_metadata[0] &= !0x10; + let target = VariantMetadata::new(&rebuilt_metadata); + let rebuilt_value = spark_shredded_variant_bytes( + source_state.clone(), + index, + semantic, + &source_metadata, + &target, + ) + .map_err(|error| { + ArrowError::InvalidArgumentError(format!( + "Failed to rebuild Spark Variant value: {error}" + )) + })?; + Ok((rebuilt_value, rebuilt_metadata)) + }, + )) + .map_err(|_| { + DataFusionError::Execution(format!("Invalid shredded Variant at row {index}")) + })??; + value_output.append_value(rebuilt_value); + metadata_output.append_value(rebuilt_metadata); + } + + Ok(( + Arc::new(value_output.finish()), + Arc::new(metadata_output.finish()), + )) +} + /// Reorder object keys for either Arrow's UTF-8 order or Spark's Java UTF-16 order. Preserve /// already-compatible values byte-for-byte and retain the original metadata dictionary. /// SPARK-58949 tracks this mismatch and legacy compatibility. The metadata dictionary's sorted @@ -505,7 +1241,7 @@ fn reorder_variant_values( ))); } - let rebuilt = catch_unwind(AssertUnwindSafe(|| { + let rebuilt = catch_unwind(AssertUnwindSafe(|| -> DataFusionResult>> { // Spark encodes empty object keys with equal metadata offsets, which Arrow 58.4's // full validator rejects. Keep shallow parsing and all accesses inside this boundary. // https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant/src/variant/metadata.rs#L307-L317 @@ -513,28 +1249,25 @@ fn reorder_variant_values( let metadata = VariantMetadata::new(metadata.value(index)); let variant = Variant::new_with_metadata(metadata.clone(), value.value(index)); if is_compatible_variant(&variant, order) { - return None; + return Ok(None); } - let mut value_builder = ValueBuilder::new(); - match order { + let value = match order { VariantObjectKeyOrder::ArrowUtf8 => { + let mut value_builder = ValueBuilder::new(); let mut metadata_builder = ReadOnlyMetadataBuilder::new(&metadata); - ValueBuilder::append_variant( - ParentState::variant(&mut value_builder, &mut metadata_builder), - variant, - ); - } - VariantObjectKeyOrder::SparkUtf16 => { - let mut metadata_builder = SparkMetadataBuilder::new(&metadata); - ValueBuilder::append_variant( + ValueBuilder::try_append_variant( ParentState::variant(&mut value_builder, &mut metadata_builder), variant, - ); + )?; + value_builder.into_inner() } - } - Some(value_builder.into_inner()) + VariantObjectKeyOrder::SparkUtf16 => spark_variant_bytes(&metadata, variant)?, + }; + Ok(Some(value)) })) - .map_err(|_| DataFusionError::Execution(format!("Invalid Variant value at row {index}")))?; + .map_err(|_| { + DataFusionError::Execution(format!("Invalid Variant value at row {index}")) + })??; output.append_value(rebuilt.as_deref().unwrap_or_else(|| value.value(index))); } @@ -727,8 +1460,8 @@ impl PhysicalExpr for CometCastColumnExpr { mod tests { use super::*; use arrow::array::{ - Array, AsArray, BinaryArray, DictionaryArray, Int32Array, Int64Array, StringArray, - UInt16Array, UInt32Array, UInt8Array, + Array, AsArray, BinaryArray, Decimal128Array, DictionaryArray, Int32Array, Int64Array, + StringArray, UInt16Array, UInt32Array, UInt8Array, }; use arrow::datatypes::{Field, Fields, Int32Type}; use datafusion::physical_expr::expressions::Column; @@ -741,10 +1474,8 @@ mod tests { keys } - fn assert_spark_unicode_object(output: &StructArray) { - let value = output.column(0).as_binary::(); - let metadata = output.column(1).as_binary::(); - let Variant::Object(object) = Variant::new(metadata.value(0), value.value(0)) else { + fn assert_spark_unicode_variant(variant: Variant<'_, '_>) { + let Variant::Object(object) = variant else { panic!("expected object") }; let fields = object.iter().collect::>(); @@ -755,11 +1486,17 @@ mod tests { let emoji = fields .binary_search_by(|(name, _)| name.encode_utf16().cmp("😀".encode_utf16())) .unwrap(); - assert_eq!(fields[emoji].1, Variant::from(531_i64)); + assert_eq!(fields[emoji].1.as_int64(), Some(531)); let private_use = fields .binary_search_by(|(name, _)| name.encode_utf16().cmp("\u{e000}".encode_utf16())) .unwrap(); - assert_eq!(fields[private_use].1, Variant::from(30_i64)); + assert_eq!(fields[private_use].1.as_int64(), Some(30)); + } + + fn assert_spark_unicode_object(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] @@ -828,8 +1565,8 @@ mod tests { 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)); + assert_eq!(variant.value(0), Variant::from(10_i8)); + assert_eq!(variant.value(2), Variant::from(30_i8)); } #[test] @@ -901,6 +1638,264 @@ mod tests { assert_eq!(object.get("u32"), Some(Variant::from(4_294_967_295_i64))); } + #[test] + fn test_normalize_shredded_variant_compacts_spark_integer_widths() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![ + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + ])); + let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![ + 1, + i64::from(i8::MAX) + 1, + i64::from(i16::MAX) + 1, + i64::from(i32::MAX) + 1, + ])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", DataType::Int64, false), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!(output.value(0), Variant::Int8(1)); + assert_eq!(output.value(1), Variant::Int16(128)); + assert_eq!(output.value(2), Variant::Int32(32_768)); + assert_eq!(output.value(3), Variant::Int64(2_147_483_648)); + } + + #[test] + fn test_compact_spark_typed_variant_canonicalizes_nan() { + let Variant::Float(float) = + compact_spark_typed_variant(Variant::Float(f32::from_bits(0x7fc0_0001))) + else { + panic!("expected float") + }; + assert_eq!(float.to_bits(), 0x7fc0_0000); + + let Variant::Double(double) = + compact_spark_typed_variant(Variant::Double(f64::from_bits(0x7ff8_0000_0000_0001))) + else { + panic!("expected double") + }; + assert_eq!(double.to_bits(), 0x7ff8_0000_0000_0000); + + let (metadata, _) = VariantBuilder::new().finish(); + let metadata = VariantMetadata::new(&metadata); + let residual = f32::from_bits(0x7fc0_0001); + let bytes = spark_variant_bytes(&metadata, Variant::Float(residual)).unwrap(); + let Variant::Float(output) = Variant::new_with_metadata(metadata, &bytes) else { + panic!("expected residual float") + }; + assert_eq!(output.to_bits(), residual.to_bits()); + } + + #[test] + fn test_normalize_shredded_variant_rejects_missing_required_value() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![None])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", DataType::Int64, true), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + assert!(normalize_variant_array(&physical, &target_field).is_err()); + } + + #[test] + fn test_normalize_shredded_variant_uses_physical_metadata_order() { + let (metadata_bytes, _) = VariantBuilder::new().with_field_names(["a", "b"]).finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let mut fields = Vec::new(); + let mut columns = Vec::new(); + for (name, value) in [("b", 2_i64), ("a", 1_i64)] { + let child = StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![value]))], + None, + ) + .unwrap(); + fields.push(Field::new(name, child.data_type().clone(), false)); + columns.push(Arc::new(child) 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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_eq!( + output.column(1).as_binary::().value(0), + &[0x01, 2, 0, 1, 2, b'b', b'a'] + ); + + let mut expected = VariantBuilder::new().with_field_names(["b", "a"]); + let mut object = expected.new_object(); + object.insert("b", 2_i8); + object.insert("a", 1_i8); + object.finish(); + let (_, expected_value) = expected.finish(); + assert_eq!(output.column(0).as_binary::().value(0), expected_value); + } + + #[test] + fn test_normalize_shredded_variant_preserves_residual_scalar_width() { + let mut builder = VariantBuilder::new().with_field_names(["known", "residual"]); + let mut object = builder.new_object(); + object.insert("residual", 1_i64); + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let known: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![2]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), + vec![known], + 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, value, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + assert_eq!(object.get("known"), Some(Variant::Int8(2))); + assert_eq!(object.get("residual"), Some(Variant::Int64(1))); + } + + #[test] + fn test_normalize_shredded_variant_compacts_spark_decimal_width() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let typed_value: ArrayRef = Arc::new( + Decimal128Array::from(vec![123_i128]) + .with_precision_and_scale(38, 2) + .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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!( + output.value(0), + Variant::Decimal4(VariantDecimal4::try_new(123, 2).unwrap()) + ); + } + #[test] fn test_normalize_shredded_variant_uses_spark_object_key_order() { let keys = unicode_object_keys(); @@ -944,6 +1939,85 @@ mod tests { assert_spark_unicode_object(output); } + #[test] + fn test_normalize_nested_shredded_variant_uses_spark_object_key_order() { + let keys = unicode_object_keys(); + let field_names = std::iter::once("nested") + .chain(keys.iter().map(String::as_str)) + .collect::>(); + let (metadata_bytes, _) = VariantBuilder::new().with_field_names(field_names).finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + + let mut nested_fields = Vec::with_capacity(keys.len()); + let mut nested_columns = Vec::with_capacity(keys.len()); + for (index, key) in keys.iter().enumerate() { + let value = if key == "😀" { 531 } else { index as i64 }; + let state = StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![value]))], + None, + ) + .unwrap(); + nested_fields.push(Field::new(key, state.data_type().clone(), false)); + nested_columns.push(Arc::new(state) as ArrayRef); + } + let nested_value: ArrayRef = + Arc::new(StructArray::try_new(nested_fields.into(), nested_columns, None).unwrap()); + let nested_state: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + nested_value.data_type().clone(), + false, + )]), + vec![nested_value], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "nested", + nested_state.data_type().clone(), + false, + )]), + vec![nested_state], + 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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let Variant::Object(object) = output.value(0) else { + panic!("expected outer object") + }; + assert_spark_unicode_variant(object.get("nested").expect("nested field")); + } + #[test] fn test_normalize_partially_shredded_spark_object_key_order() { let keys = unicode_object_keys(); @@ -1164,6 +2238,80 @@ mod tests { assert_eq!(nested.get(""), Some(Variant::from(2_i64))); } + #[test] + fn test_normalize_partially_shredded_nested_unicode_and_empty_keys() { + let keys = unicode_object_keys(); + let mut builder = VariantBuilder::new().with_field_names(["known"]); + let mut object = builder.new_object(); + object.insert("", 1_i64); + let mut nested = object.new_object("nested"); + for (index, key) in keys.iter().enumerate() { + nested.insert(key, if key == "😀" { 531_i64 } else { index as i64 }); + } + nested.finish(); + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let known: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![3]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), + vec![known], + 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, value, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + assert_eq!(object.get("").unwrap().as_int64(), Some(1)); + assert_eq!(object.get("known").unwrap().as_int64(), Some(3)); + let Variant::Object(nested) = object.get("nested").unwrap() else { + panic!("expected nested object") + }; + let fields = nested.iter().collect::>(); + assert_eq!(fields.len(), 32); + assert_eq!(fields[30].0, "😀"); + assert_eq!(fields[31].0, "\u{e000}"); + assert_eq!(nested.get("😀").unwrap().as_int64(), Some(531)); + } + #[test] fn test_normalize_variant_skips_empty_children_of_null_parent() { let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&b""[..])])); diff --git a/native/core/src/parquet/schema_adapter.rs b/native/core/src/parquet/schema_adapter.rs index 74d8afd7bc3..7e8ed2bde15 100644 --- a/native/core/src/parquet/schema_adapter.rs +++ b/native/core/src/parquet/schema_adapter.rs @@ -1119,13 +1119,12 @@ impl PhysicalExpr for RejectOnNonEmpty { mod test { use crate::parquet::parquet_support::SparkParquetOptions; use crate::parquet::schema_adapter::SparkPhysicalExprAdapterFactory; - use arrow::array::UInt32Array; use arrow::array::{ - BinaryArray, Date32Array, Decimal128Array, Float32Array, Float64Array, Int32Array, - Int64Array, StringArray, TimestampMicrosecondArray, + Array, BinaryArray, Date32Array, Decimal128Array, Float32Array, Float64Array, Int32Array, + Int64Array, StringArray, StructArray, TimestampMicrosecondArray, UInt16Array, UInt32Array, + UInt8Array, }; - use arrow::datatypes::SchemaRef; - use arrow::datatypes::{DataType, Field, Schema}; + use arrow::datatypes::{DataType, Field, Fields, Schema, SchemaRef}; use arrow::record_batch::RecordBatch; use datafusion::common::DataFusionError; use datafusion::datasource::listing::PartitionedFile; @@ -1139,6 +1138,7 @@ mod test { use datafusion_physical_expr_adapter::PhysicalExprAdapterFactory; use futures::StreamExt; use parquet::arrow::ArrowWriter; + use parquet::variant::{Variant, VariantArray, VariantBuilder, VariantType}; use std::fs::File; use std::sync::Arc; @@ -1695,6 +1695,64 @@ mod test { Ok(()) } + #[tokio::test] + async fn parquet_roundtrip_shredded_variant_unsigned_values() -> Result<(), DataFusionError> { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let values = [ + ( + "u8", + Arc::new(UInt8Array::from(vec![u8::MAX])) as Arc, + ), + ("u16", Arc::new(UInt16Array::from(vec![u16::MAX]))), + ("u32", Arc::new(UInt32Array::from(vec![u32::MAX]))), + ]; + let mut file_fields = Vec::with_capacity(values.len()); + let mut columns = Vec::with_capacity(values.len()); + let mut required_fields = Vec::with_capacity(values.len()); + for (name, value) in values { + let metadata = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical = StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", value.data_type().clone(), false), + ]), + vec![metadata, value], + None, + )?; + file_fields.push( + Field::new(name, physical.data_type().clone(), false) + .with_extension_type(VariantType), + ); + columns.push(Arc::new(physical) as Arc); + required_fields.push( + Field::new( + name, + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + } + + let file_schema = Arc::new(Schema::new(file_fields)); + let batch = RecordBatch::try_new(Arc::clone(&file_schema), columns)?; + let output = roundtrip(&batch, Arc::new(Schema::new(required_fields))).await?; + + let expected = [ + Variant::from(255_i16), + Variant::from(65_535_i32), + Variant::from(4_294_967_295_i64), + ]; + for (column, expected) in output.columns().iter().zip(expected) { + let variant = VariantArray::try_new(column.as_ref())?; + assert_eq!(variant.value(0), expected); + } + Ok(()) + } + /// Create a Parquet file containing a single batch and then read the batch back using /// the specified required_schema. This will cause the PhysicalExprAdapter code to be used. async fn roundtrip( diff --git a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql index 50cafbfebdc..2674970ec94 100644 --- a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql +++ b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql @@ -127,8 +127,11 @@ statement CREATE TABLE test_variant_unicode(v VARIANT) USING parquet statement -INSERT INTO test_variant_unicode VALUES (parse_json( - '{"k00":0,"k01":1,"k02":2,"k03":3,"k04":4,"k05":5,"k06":6,"k07":7,"k08":8,"k09":9,"k10":10,"k11":11,"k12":12,"k13":13,"k14":14,"k15":15,"k16":16,"k17":17,"k18":18,"k19":19,"k20":20,"k21":21,"k22":22,"k23":23,"k24":24,"k25":25,"k26":26,"k27":27,"k28":28,"k29":29,"\uE000":30,"😀":531}')) +INSERT INTO test_variant_unicode VALUES + (parse_json( + '{"k00":0,"k01":1,"k02":2,"k03":3,"k04":4,"k05":5,"k06":6,"k07":7,"k08":8,"k09":9,"k10":10,"k11":11,"k12":12,"k13":13,"k14":14,"k15":15,"k16":16,"k17":17,"k18":18,"k19":19,"k20":20,"k21":21,"k22":22,"k23":23,"k24":24,"k25":25,"k26":26,"k27":27,"k28":28,"k29":29,"\uE000":30,"😀":531}')), + (parse_json( + '{"":-2,"k00":0,"nested":{"":-1,"k00":0,"k01":1,"k02":2,"k03":3,"k04":4,"k05":5,"k06":6,"k07":7,"k08":8,"k09":9,"k10":10,"k11":11,"k12":12,"k13":13,"k14":14,"k15":15,"k16":16,"k17":17,"k18":18,"k19":19,"k20":20,"k21":21,"k22":22,"k23":23,"k24":24,"k25":25,"k26":26,"k27":27,"k28":28,"k29":29,"\uE000":30,"😀":532}}')) statement SET spark.sql.variant.writeShredding.enabled=false @@ -136,29 +139,32 @@ SET spark.sql.variant.writeShredding.enabled=false statement SET spark.sql.variant.allowReadingShredded=true +query +SELECT v FROM test_variant_unicode + query expect_fallback(type VariantType) -SELECT variant_get(v, '$.😀', 'bigint') FROM test_variant_unicode +SELECT variant_get(v, '$.😀', 'bigint'), variant_get(v, '$.nested.😀', 'bigint') +FROM test_variant_unicode --- Spark SQL cannot write unsigned Parquet integer annotations. Generate the signed representation --- produced after widening here; the unsigned-to-signed conversion itself is covered in Rust. +-- Spark rebuilds typed fields in physical shredding-schema order and chooses integer/decimal +-- widths from the runtime value, independently of the Parquet physical width. statement -SET spark.sql.variant.forceShreddingSchemaForTest=u8 SMALLINT, u16 INT, u32 BIGINT +SET spark.sql.variant.forceShreddingSchemaForTest=b BIGINT, a BIGINT, d DECIMAL(38,2) statement SET spark.sql.variant.writeShredding.enabled=true statement -CREATE TABLE test_variant_widened(v VARIANT) USING parquet +CREATE TABLE test_variant_typed_bytes(v VARIANT) USING parquet statement -INSERT INTO test_variant_widened VALUES - (parse_json('{"u8":255,"u16":65535,"u32":4294967295}')) +INSERT INTO test_variant_typed_bytes VALUES (parse_json('{"a":1,"b":2,"d":1.23}')) statement SET spark.sql.variant.writeShredding.enabled=false query -SELECT v FROM test_variant_widened +SELECT v FROM test_variant_typed_bytes statement CREATE TABLE test_variant_struct(id INT, s STRUCT, tail STRING) From 360a653c0f0444b589ab4ac01c8ed0c1f70f415f Mon Sep 17 00:00:00 2001 From: peterxcli Date: Wed, 26 Aug 2026 15:48:40 +0800 Subject: [PATCH 08/14] fix: normalize shredded Variant types before unshredding --- native/core/src/parquet/cast_column.rs | 176 +++++++++++++++++---- native/core/src/parquet/schema_adapter.rs | 177 ++++++++++++++++++++-- 2 files changed, 305 insertions(+), 48 deletions(-) diff --git a/native/core/src/parquet/cast_column.rs b/native/core/src/parquet/cast_column.rs index d452184b8d4..eafa4ed0661 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -21,7 +21,7 @@ use arrow::{ TimestampMillisecondArray, }, buffer::NullBuffer, - compute::{cast, CastOptions}, + compute::{cast, cast_with_options, CastOptions}, datatypes::{DataType, FieldRef, Schema, TimeUnit}, error::ArrowError, record_batch::RecordBatch, @@ -211,7 +211,7 @@ fn normalize_variant_array( } let array = decode_variant_metadata_dictionary(array)?; - let array = widen_unsigned_variant_typed_value(&array)?; + let array = normalize_variant_typed_value(&array)?; let variant = VariantArray::try_new(array.as_ref())?; let was_shredded = variant.typed_value_field().is_some(); let unshredded = unshred_variant_for_spark(&variant)?; @@ -254,9 +254,9 @@ fn unshred_variant_for_spark(variant: &VariantArray) -> DataFusionResult Option { - fn widen_field(field: &FieldRef) -> Option { - widen_unsigned_variant_type(field.data_type()) +fn normalize_variant_type(data_type: &DataType) -> Option { + fn normalize_field(field: &FieldRef) -> Option { + normalize_variant_type(field.data_type()) .map(|data_type| Arc::new(field.as_ref().clone().with_data_type(data_type))) } @@ -264,15 +264,21 @@ fn widen_unsigned_variant_type(data_type: &DataType) -> Option { DataType::UInt8 => Some(DataType::Int16), DataType::UInt16 => Some(DataType::Int32), DataType::UInt32 => Some(DataType::Int64), - DataType::List(field) => widen_field(field).map(DataType::List), - DataType::LargeList(field) => widen_field(field).map(DataType::LargeList), - DataType::ListView(field) => widen_field(field).map(DataType::ListView), - DataType::LargeListView(field) => widen_field(field).map(DataType::LargeListView), + DataType::Timestamp(TimeUnit::Millisecond, timezone) => { + Some(DataType::Timestamp(TimeUnit::Microsecond, timezone.clone())) + } + DataType::FixedSizeList(field, _) => Some(DataType::List( + normalize_field(field).unwrap_or_else(|| Arc::clone(field)), + )), + DataType::List(field) => normalize_field(field).map(DataType::List), + DataType::LargeList(field) => normalize_field(field).map(DataType::LargeList), + DataType::ListView(field) => normalize_field(field).map(DataType::ListView), + DataType::LargeListView(field) => normalize_field(field).map(DataType::LargeListView), DataType::Struct(fields) => { let mut changed = false; let fields = fields .iter() - .map(|field| match widen_field(field) { + .map(|field| match normalize_field(field) { Some(field) => { changed = true; field @@ -286,10 +292,12 @@ fn widen_unsigned_variant_type(data_type: &DataType) -> Option { } } -/// Parquet restores unsigned integer annotations as Arrow unsigned arrays, while Spark widens -/// those values to the next signed width. Arrow Variant accepts only the latter representation. +/// Normalize Arrow types that Spark's Parquet reader accepts but `VariantArray` 58.4 rejects. +/// Parquet restores unsigned integers to Arrow unsigned arrays and millisecond timestamps at their +/// annotated unit; embedded Arrow schemas may also restore fixed-size lists. Spark widens the +/// integers and timestamps and treats fixed-size lists as ordinary Variant arrays. /// arrow-rs #10416/#10417 would move this widening into `VariantArray`/`unshred_variant`; remove -/// both local `widen_unsigned_variant_*` helpers after that ships and Comet upgrades: +/// the unsigned arms after that ships and Comet upgrades: /// https://github.com/apache/arrow-rs/issues/10416 /// https://github.com/apache/arrow-rs/pull/10417 /// Arrow #50622/#50810 instead proposes removing unsigned `typed_value` mappings because the @@ -297,7 +305,7 @@ fn widen_unsigned_variant_type(data_type: &DataType) -> Option { /// that choice, keep this compatibility path for unsigned files Spark already reads: /// https://github.com/apache/arrow/issues/50622 /// https://github.com/apache/arrow/pull/50810 -fn widen_unsigned_variant_typed_value(array: &ArrayRef) -> DataFusionResult { +fn normalize_variant_typed_value(array: &ArrayRef) -> DataFusionResult { let Some(struct_array) = array.as_any().downcast_ref::() else { return Ok(Arc::clone(array)); }; @@ -309,7 +317,7 @@ fn widen_unsigned_variant_typed_value(array: &ArrayRef) -> DataFusionResult DataFusionResult collect_list_field_names( - typed_value.as_fixed_size_list(), - index, - source_metadata, - field_names, - seen, - )?, _ => {} } Ok(()) @@ -1109,13 +1114,6 @@ fn spark_shredded_variant_bytes( source_metadata, target_metadata, ), - DataType::FixedSizeList(_, _) => spark_list_bytes( - typed_value.as_fixed_size_list(), - index, - semantic, - source_metadata, - target_metadata, - ), _ => spark_typed_variant_bytes(target_metadata, semantic), } } @@ -1460,8 +1458,9 @@ impl PhysicalExpr for CometCastColumnExpr { mod tests { use super::*; use arrow::array::{ - Array, AsArray, BinaryArray, Decimal128Array, DictionaryArray, Int32Array, Int64Array, - StringArray, UInt16Array, UInt32Array, UInt8Array, + Array, AsArray, BinaryArray, Decimal128Array, DictionaryArray, FixedSizeListArray, + Int32Array, Int64Array, StringArray, TimestampMillisecondArray, UInt16Array, UInt32Array, + UInt8Array, }; use arrow::datatypes::{Field, Fields, Int32Type}; use datafusion::physical_expr::expressions::Column; @@ -1499,6 +1498,38 @@ mod tests { assert_spark_unicode_variant(Variant::new(metadata.value(0), value.value(0))); } + fn normalize_typed_value(typed_value: ArrayRef, field_names: &[&str]) -> VariantArray { + let (metadata_bytes, _) = VariantBuilder::new() + .with_field_names(field_names.iter().copied()) + .finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + 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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + VariantArray::try_new(output.as_ref()).unwrap() + } + #[test] fn test_normalize_shredded_variant_with_dictionary_metadata_for_spark() { let mut builder = VariantArrayBuilder::new(3); @@ -1638,6 +1669,87 @@ mod tests { assert_eq!(object.get("u32"), Some(Variant::from(4_294_967_295_i64))); } + #[test] + fn test_normalize_shredded_variant_widens_millisecond_timestamps() { + let millis = 1_704_067_200_123_i64; + let ltz: ArrayRef = + Arc::new(TimestampMillisecondArray::from(vec![millis]).with_timezone("UTC")); + let ntz: ArrayRef = Arc::new(TimestampMillisecondArray::from(vec![millis])); + let shredded = |value: ArrayRef| -> ArrayRef { + Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + value.data_type().clone(), + false, + )]), + vec![value], + None, + ) + .unwrap(), + ) + }; + let ltz = shredded(ltz); + let ntz = shredded(ntz); + let object: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("ltz", ltz.data_type().clone(), false), + Field::new("ntz", ntz.data_type().clone(), false), + ]), + vec![ltz, ntz], + None, + ) + .unwrap(), + ); + let output = normalize_typed_value(object, &["ltz", "ntz"]); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + + let Some(Variant::TimestampMicros(ltz)) = object.get("ltz") else { + panic!("expected timestamp") + }; + assert_eq!(ltz.timestamp_micros(), millis * 1_000); + + let Some(Variant::TimestampNtzMicros(ntz)) = object.get("ntz") else { + panic!("expected timestamp_ntz") + }; + assert_eq!(ntz.and_utc().timestamp_micros(), millis * 1_000); + } + + #[test] + fn test_normalize_shredded_variant_converts_fixed_size_list() { + let elements: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![42, 43]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + FixedSizeListArray::try_new( + Arc::new(Field::new("element", elements.data_type().clone(), false)), + 2, + elements, + None, + ) + .unwrap(), + ); + + let output = normalize_typed_value(typed_value, &[]); + let Variant::List(list) = output.value(0) else { + panic!("expected list") + }; + assert_eq!( + list.iter() + .map(|value| value.as_int64()) + .collect::>(), + vec![Some(42), Some(43)] + ); + } + #[test] fn test_normalize_shredded_variant_compacts_spark_integer_widths() { let (metadata_bytes, _) = VariantBuilder::new().finish(); diff --git a/native/core/src/parquet/schema_adapter.rs b/native/core/src/parquet/schema_adapter.rs index 7e8ed2bde15..d395ae95007 100644 --- a/native/core/src/parquet/schema_adapter.rs +++ b/native/core/src/parquet/schema_adapter.rs @@ -1120,10 +1120,11 @@ mod test { use crate::parquet::parquet_support::SparkParquetOptions; use crate::parquet::schema_adapter::SparkPhysicalExprAdapterFactory; use arrow::array::{ - Array, BinaryArray, Date32Array, Decimal128Array, Float32Array, Float64Array, Int32Array, - Int64Array, StringArray, StructArray, TimestampMicrosecondArray, UInt16Array, UInt32Array, - UInt8Array, + Array, ArrayRef, BinaryArray, Date32Array, Decimal128Array, FixedSizeListArray, + Float32Array, Float64Array, Int32Array, Int64Array, StringArray, StructArray, + TimestampMicrosecondArray, TimestampMillisecondArray, UInt16Array, UInt32Array, UInt8Array, }; + use arrow::buffer::NullBuffer; use arrow::datatypes::{DataType, Field, Fields, Schema, SchemaRef}; use arrow::record_batch::RecordBatch; use datafusion::common::DataFusionError; @@ -1137,7 +1138,7 @@ mod test { use datafusion_comet_spark_expr::EvalMode; use datafusion_physical_expr_adapter::PhysicalExprAdapterFactory; use futures::StreamExt; - use parquet::arrow::ArrowWriter; + use parquet::arrow::{arrow_writer::ArrowWriterOptions, ArrowWriter}; use parquet::variant::{Variant, VariantArray, VariantBuilder, VariantType}; use std::fs::File; use std::sync::Arc; @@ -1695,6 +1696,18 @@ mod test { Ok(()) } + fn required_variant_field(name: &str, nullable: bool) -> Field { + Field::new( + name, + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + nullable, + ) + .with_extension_type(VariantType) + } + #[tokio::test] async fn parquet_roundtrip_shredded_variant_unsigned_values() -> Result<(), DataFusionError> { let (metadata_bytes, _) = VariantBuilder::new().finish(); @@ -1724,17 +1737,7 @@ mod test { .with_extension_type(VariantType), ); columns.push(Arc::new(physical) as Arc); - required_fields.push( - Field::new( - name, - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); + required_fields.push(required_variant_field(name, false)); } let file_schema = Arc::new(Schema::new(file_fields)); @@ -1753,16 +1756,158 @@ mod test { Ok(()) } + #[tokio::test] + async fn parquet_roundtrip_shredded_variant_millisecond_timestamps( + ) -> Result<(), DataFusionError> { + let millis = 1_704_067_200_123_i64; + let micros = 1_704_067_200_123_456_i64; + let shredded = |value: ArrayRef| -> ArrayRef { + Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + value.data_type().clone(), + false, + )]), + vec![value], + None, + ) + .unwrap(), + ) + }; + let ltz = shredded(Arc::new( + TimestampMillisecondArray::from(vec![millis]).with_timezone("UTC"), + )); + let ntz = shredded(Arc::new(TimestampMillisecondArray::from(vec![millis]))); + let micros_control = shredded(Arc::new( + TimestampMicrosecondArray::from(vec![micros]).with_timezone("UTC"), + )); + let typed_value: ArrayRef = Arc::new(StructArray::try_new( + Fields::from(vec![ + Field::new("ltz", ltz.data_type().clone(), false), + Field::new("ntz", ntz.data_type().clone(), false), + Field::new("micros", micros_control.data_type().clone(), false), + ]), + vec![ltz, ntz, micros_control], + None, + )?); + let (metadata_bytes, _) = VariantBuilder::new() + .with_field_names(["ltz", "ntz", "micros"]) + .finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical = StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]), + vec![metadata, typed_value], + None, + )?; + let file_schema = Arc::new(Schema::new(vec![Field::new( + "v", + physical.data_type().clone(), + false, + ) + .with_extension_type(VariantType)])); + let batch = RecordBatch::try_new(Arc::clone(&file_schema), vec![Arc::new(physical)])?; + + let output = roundtrip_with_options( + &batch, + Arc::new(Schema::new(vec![required_variant_field("v", false)])), + ArrowWriterOptions::new().with_skip_arrow_metadata(true), + ) + .await?; + let variant = VariantArray::try_new(output.column(0).as_ref())?; + let Variant::Object(object) = variant.value(0) else { + panic!("expected object") + }; + let Some(Variant::TimestampMicros(ltz)) = object.get("ltz") else { + panic!("expected timestamp") + }; + assert_eq!(ltz.timestamp_micros(), millis * 1_000); + let Some(Variant::TimestampNtzMicros(ntz)) = object.get("ntz") else { + panic!("expected timestamp_ntz") + }; + assert_eq!(ntz.and_utc().timestamp_micros(), millis * 1_000); + let Some(Variant::TimestampMicros(control)) = object.get("micros") else { + panic!("expected timestamp") + }; + assert_eq!(control.timestamp_micros(), micros); + Ok(()) + } + + #[tokio::test] + async fn parquet_roundtrip_shredded_variant_fixed_size_list() -> Result<(), DataFusionError> { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![ + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + ])); + let elements: ArrayRef = Arc::new(StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![42, 43, 0, 0]))], + None, + )?); + let typed_value: ArrayRef = Arc::new(FixedSizeListArray::try_new( + Arc::new(Field::new("element", elements.data_type().clone(), false)), + 2, + elements, + None, + )?); + let physical = StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]), + vec![metadata, typed_value], + Some(NullBuffer::from(vec![true, false])), + )?; + let file_schema = Arc::new(Schema::new(vec![Field::new( + "v", + physical.data_type().clone(), + true, + ) + .with_extension_type(VariantType)])); + let batch = RecordBatch::try_new(Arc::clone(&file_schema), vec![Arc::new(physical)])?; + + let output = roundtrip( + &batch, + Arc::new(Schema::new(vec![required_variant_field("v", true)])), + ) + .await?; + let variant = VariantArray::try_new(output.column(0).as_ref())?; + let Variant::List(list) = variant.value(0) else { + panic!("expected list") + }; + assert_eq!( + list.iter() + .map(|value| value.as_int64()) + .collect::>(), + vec![Some(42), Some(43)] + ); + assert!(variant.inner().is_null(1)); + Ok(()) + } + /// Create a Parquet file containing a single batch and then read the batch back using /// the specified required_schema. This will cause the PhysicalExprAdapter code to be used. async fn roundtrip( batch: &RecordBatch, required_schema: SchemaRef, + ) -> Result { + roundtrip_with_options(batch, required_schema, ArrowWriterOptions::new()).await + } + + async fn roundtrip_with_options( + batch: &RecordBatch, + required_schema: SchemaRef, + writer_options: ArrowWriterOptions, ) -> Result { let filename = get_temp_filename(); let filename = filename.as_path().as_os_str().to_str().unwrap().to_string(); let file = File::create(&filename)?; - let mut writer = ArrowWriter::try_new(file, Arc::clone(&batch.schema()), None)?; + let mut writer = + ArrowWriter::try_new_with_options(file, Arc::clone(&batch.schema()), writer_options)?; writer.write(batch)?; writer.close()?; From a800bc854be020a9d365be993d047e1d3a2cacc8 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Thu, 27 Aug 2026 00:44:53 +0800 Subject: [PATCH 09/14] refactor: isolate Parquet Variant normalization --- native/core/src/parquet/cast_column.rs | 2047 +---------------- .../core/src/parquet/cast_column/variant.rs | 1122 +++++++++ .../src/parquet/cast_column/variant/tests.rs | 935 ++++++++ 3 files changed, 2078 insertions(+), 2026 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 eafa4ed0661..502622c880f 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -14,41 +14,34 @@ // 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, AsArray, BinaryArray, BinaryBuilder, LargeListArray, - ListArray, ListLikeArray, MapArray, StructArray, TimestampMicrosecondArray, - TimestampMillisecondArray, + make_array, Array, ArrayRef, LargeListArray, ListArray, MapArray, StructArray, + TimestampMicrosecondArray, TimestampMillisecondArray, }, - buffer::NullBuffer, - compute::{cast, cast_with_options, CastOptions}, + compute::CastOptions, datatypes::{DataType, FieldRef, Schema, TimeUnit}, - error::ArrowError, record_batch::RecordBatch, }; - -use crate::{ - execution::serde::is_variant_field, - parquet::parquet_support::{spark_parquet_convert, SparkParquetOptions}, +use datafusion::common::{ + format::DEFAULT_CAST_OPTIONS, DataFusionError, Result as DataFusionResult, ScalarValue, }; -use datafusion::common::format::DEFAULT_CAST_OPTIONS; -use datafusion::common::ScalarValue; -use datafusion::common::{DataFusionError, Result as DataFusionResult}; use datafusion::logical_expr::ColumnarValue; use datafusion::physical_expr::PhysicalExpr; -use parquet::variant::{ - unshred_variant, BorrowedShreddingState, ListBuilder, MetadataBuilder, ObjectBuilder, - ParentState, ReadOnlyMetadataBuilder, ValueBuilder, Variant, VariantArray, VariantBuilder, - VariantDecimal4, VariantDecimal8, VariantMetadata, WritableMetadataBuilder, -}; use std::{ - collections::HashSet, fmt::{self, Display}, hash::Hash, - panic::{catch_unwind, AssertUnwindSafe}, sync::Arc, }; +use crate::{ + execution::serde::is_variant_field, + parquet::parquet_support::{spark_parquet_convert, SparkParquetOptions}, +}; + /// Returns true if two DataTypes are structurally equivalent (same data layout) /// but may differ in field names within nested types. fn types_differ_only_in_field_names(physical: &DataType, logical: &DataType) -> bool { @@ -189,1089 +182,6 @@ fn cast_timestamp_micros_to_millis_scalar( ScalarValue::TimestampMillisecond(new_val, target_tz) } -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(), - )); - } - - let array = decode_variant_metadata_dictionary(array)?; - let array = normalize_variant_typed_value(&array)?; - let variant = VariantArray::try_new(array.as_ref())?; - let was_shredded = variant.typed_value_field().is_some(); - let unshredded = unshred_variant_for_spark(&variant)?; - 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, metadata) = if was_shredded { - rebuild_shredded_variant_for_spark(&variant, &value, &metadata, unshredded.inner().nulls())? - } else { - let value = reorder_variant_values( - &value, - &metadata, - unshredded.inner().nulls(), - VariantObjectKeyOrder::SparkUtf16, - false, - )?; - (value, metadata) - }; - let output = StructArray::try_new( - fields.clone(), - vec![value, metadata], - unshredded.inner().nulls().cloned(), - )?; - Ok(Arc::new(output)) -} - -fn unshred_variant_for_spark(variant: &VariantArray) -> DataFusionResult { - let first = - prepare_variant_for_unshredding(variant).and_then(|array| Ok(unshred_variant(&array)?)); - let first_error = match first { - Ok(array) => return Ok(array), - Err(error) => error, - }; - let Some(variant) = canonicalize_spark_empty_key_metadata(variant)? else { - return Err(first_error); - }; - let variant = prepare_variant_for_unshredding(&variant)?; - Ok(unshred_variant(&variant)?) -} - -fn normalize_variant_type(data_type: &DataType) -> Option { - fn normalize_field(field: &FieldRef) -> Option { - normalize_variant_type(field.data_type()) - .map(|data_type| Arc::new(field.as_ref().clone().with_data_type(data_type))) - } - - match data_type { - DataType::UInt8 => Some(DataType::Int16), - DataType::UInt16 => Some(DataType::Int32), - DataType::UInt32 => Some(DataType::Int64), - DataType::Timestamp(TimeUnit::Millisecond, timezone) => { - Some(DataType::Timestamp(TimeUnit::Microsecond, timezone.clone())) - } - DataType::FixedSizeList(field, _) => Some(DataType::List( - normalize_field(field).unwrap_or_else(|| Arc::clone(field)), - )), - DataType::List(field) => normalize_field(field).map(DataType::List), - DataType::LargeList(field) => normalize_field(field).map(DataType::LargeList), - DataType::ListView(field) => normalize_field(field).map(DataType::ListView), - DataType::LargeListView(field) => normalize_field(field).map(DataType::LargeListView), - DataType::Struct(fields) => { - let mut changed = false; - let fields = fields - .iter() - .map(|field| match normalize_field(field) { - Some(field) => { - changed = true; - field - } - None => Arc::clone(field), - }) - .collect::>(); - changed.then(|| DataType::Struct(fields.into())) - } - _ => None, - } -} - -/// Normalize Arrow types that Spark's Parquet reader accepts but `VariantArray` 58.4 rejects. -/// Parquet restores unsigned integers to Arrow unsigned arrays and millisecond timestamps at their -/// annotated unit; embedded Arrow schemas may also restore fixed-size lists. Spark widens the -/// integers and timestamps and treats fixed-size lists as ordinary Variant arrays. -/// arrow-rs #10416/#10417 would move this widening into `VariantArray`/`unshred_variant`; remove -/// the unsigned arms after that ships and Comet upgrades: -/// https://github.com/apache/arrow-rs/issues/10416 -/// https://github.com/apache/arrow-rs/pull/10417 -/// Arrow #50622/#50810 instead proposes removing unsigned `typed_value` mappings because the -/// Parquet Variant shredding table permits only signed integer fields. Until upstream resolves -/// that choice, keep this compatibility path for unsigned files Spark already reads: -/// https://github.com/apache/arrow/issues/50622 -/// https://github.com/apache/arrow/pull/50810 -fn normalize_variant_typed_value(array: &ArrayRef) -> DataFusionResult { - let Some(struct_array) = array.as_any().downcast_ref::() else { - return Ok(Arc::clone(array)); - }; - let Some((typed_value_index, typed_value_field)) = struct_array - .fields() - .iter() - .enumerate() - .find(|(_, field)| field.name() == "typed_value") - else { - return Ok(Arc::clone(array)); - }; - let Some(data_type) = normalize_variant_type(typed_value_field.data_type()) else { - return Ok(Arc::clone(array)); - }; - - let mut fields = struct_array.fields().iter().cloned().collect::>(); - fields[typed_value_index] = Arc::new( - typed_value_field - .as_ref() - .clone() - .with_data_type(data_type.clone()), - ); - let mut columns = struct_array.columns().to_vec(); - columns[typed_value_index] = cast_with_options( - columns[typed_value_index].as_ref(), - &data_type, - &DEFAULT_CAST_OPTIONS, - )?; - Ok(Arc::new(StructArray::try_new( - fields.into(), - columns, - struct_array.nulls().cloned(), - )?)) -} - -/// Arrow's unshredder fully validates any residual `value` in a partially shredded object. Spark -/// writes object keys in Java UTF-16 order, so put that residual value in Arrow UTF-8 order only -/// while it passes through the upstream unshredder. -fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult { - let (Some(value), Some(_)) = (variant.value_field(), variant.typed_value_field()) else { - return Ok(variant.clone()); - }; - - let value = cast(value.as_ref(), &DataType::Binary)?; - let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; - let value = reorder_variant_values( - &value, - &metadata, - variant.inner().nulls(), - VariantObjectKeyOrder::ArrowUtf8, - true, - )?; - - let value_index = variant - .inner() - .fields() - .iter() - .position(|field| field.name() == "value") - .unwrap(); - let mut fields = variant.inner().fields().iter().cloned().collect::>(); - fields[value_index] = Arc::new( - fields[value_index] - .as_ref() - .clone() - .with_data_type(DataType::Binary), - ); - let mut columns = variant.inner().columns().to_vec(); - columns[value_index] = value; - let array = StructArray::try_new(fields.into(), columns, variant.inner().nulls().cloned())?; - Ok(VariantArray::try_new(&array)?) -} - -/// Arrow 58.4 rejects Spark metadata dictionaries containing an empty object key because their -/// unsorted offsets can be equal. Rebuild only those rows before any residual value is traversed, -/// sorting the dictionary and remapping the residual value's field IDs at the same time. -/// https://github.com/apache/arrow-rs/pull/10352 -fn canonicalize_spark_empty_key_metadata( - variant: &VariantArray, -) -> DataFusionResult> { - type Replacement = Option<(Vec, Option>)>; - - let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; - let metadata = metadata.as_binary::(); - let value = variant - .value_field() - .map(|value| cast(value.as_ref(), &DataType::Binary)) - .transpose()?; - let value = value.as_ref().map(|value| value.as_binary::()); - let mut replacements = Vec::with_capacity(variant.len()); - let mut changed = false; - - for index in 0..variant.len() { - if variant.inner().is_null(index) || metadata.is_null(index) { - replacements.push(None); - continue; - } - let metadata_bytes = metadata.value(index); - if VariantMetadata::try_new(metadata_bytes).is_ok() { - replacements.push(None); - continue; - } - - let replacement = catch_unwind(AssertUnwindSafe(|| -> Result { - let old_metadata = VariantMetadata::new(metadata_bytes); - let mut names = old_metadata - .iter_try() - .map(|name| name.map(str::to_string)) - .collect::, _>>()?; - if !names.iter().any(String::is_empty) { - return Ok(None); - } - names.sort_unstable(); - if names.windows(2).any(|names| names[0] == names[1]) { - return Ok(None); - } - - let mut builder = - VariantBuilder::new().with_field_names(names.iter().map(String::as_str)); - match value { - Some(value) if !value.is_null(index) => { - builder - .append_value(Variant::new_with_metadata(old_metadata, value.value(index))); - let (metadata, value) = builder.finish(); - Ok(Some((metadata, Some(value)))) - } - _ => Ok(Some((builder.finish().0, None))), - } - })) - .map_err(|_| { - DataFusionError::Execution(format!( - "Invalid Variant metadata with an empty key at row {index}" - )) - })??; - changed |= replacement.is_some(); - replacements.push(replacement); - } - - if !changed { - return Ok(None); - } - - let mut metadata_builder = BinaryBuilder::new(); - let mut value_builder = value.map(|_| BinaryBuilder::new()); - for (index, replacement) in replacements.iter().enumerate() { - match replacement { - Some((metadata, value)) => { - metadata_builder.append_value(metadata); - if let Some(builder) = &mut value_builder { - match value { - Some(value) => builder.append_value(value), - None => builder.append_null(), - } - } - } - None => { - if metadata.is_null(index) { - metadata_builder.append_null(); - } else { - metadata_builder.append_value(metadata.value(index)); - } - if let (Some(value), Some(builder)) = (value, &mut value_builder) { - if value.is_null(index) { - builder.append_null(); - } else { - builder.append_value(value.value(index)); - } - } - } - } - } - - let mut fields = variant.inner().fields().iter().cloned().collect::>(); - let mut columns = variant.inner().columns().to_vec(); - let metadata_index = fields - .iter() - .position(|field| field.name() == "metadata") - .unwrap(); - fields[metadata_index] = Arc::new( - fields[metadata_index] - .as_ref() - .clone() - .with_data_type(DataType::Binary), - ); - columns[metadata_index] = Arc::new(metadata_builder.finish()); - if let Some(mut value_builder) = value_builder { - let value_index = fields - .iter() - .position(|field| field.name() == "value") - .unwrap(); - fields[value_index] = Arc::new( - fields[value_index] - .as_ref() - .clone() - .with_data_type(DataType::Binary), - ); - columns[value_index] = Arc::new(value_builder.finish()); - } - let array = StructArray::try_new(fields.into(), columns, variant.inner().nulls().cloned())?; - Ok(Some(VariantArray::try_new(&array)?)) -} - -/// Arrow-rs parquet-variant-compute allows dictionary-encoded metadata in its contract, but 58.4's -/// `VariantArray::try_new` validates only Binary, LargeBinary, and BinaryView. Decode just that -/// child and keep the physical struct otherwise unchanged. -/// https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant-compute/src/variant_array.rs#L276-L310 -/// Upstream issue: https://github.com/apache/arrow-rs/issues/10802 -fn decode_variant_metadata_dictionary(array: &ArrayRef) -> DataFusionResult { - let Some(struct_array) = array.as_any().downcast_ref::() else { - return Ok(Arc::clone(array)); - }; - let Some((metadata_index, metadata_field)) = struct_array - .fields() - .iter() - .enumerate() - .find(|(_, field)| field.name() == "metadata") - else { - return Ok(Arc::clone(array)); - }; - let DataType::Dictionary(_, value_type) = metadata_field.data_type() else { - return Ok(Arc::clone(array)); - }; - - let decoded = cast(struct_array.column(metadata_index).as_ref(), value_type)?; - let mut fields = struct_array.fields().iter().cloned().collect::>(); - fields[metadata_index] = Arc::new( - metadata_field - .as_ref() - .clone() - .with_data_type(decoded.data_type().clone()), - ); - let mut columns = struct_array.columns().to_vec(); - columns[metadata_index] = decoded; - Ok(Arc::new(StructArray::try_new( - fields.into(), - columns, - struct_array.nulls().cloned(), - )?)) -} - -/// Supplies sort-only field names whose Rust ordering matches Java `String.compareTo` ordering. -/// The original metadata dictionary still supplies the field IDs written to the Variant value. -#[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 compact_spark_integer(value: i64) -> Variant<'static, 'static> { - if let Ok(value) = i8::try_from(value) { - Variant::Int8(value) - } else if let Ok(value) = i16::try_from(value) { - Variant::Int16(value) - } else if let Ok(value) = i32::try_from(value) { - Variant::Int32(value) - } else { - Variant::Int64(value) - } -} - -fn compact_spark_typed_variant<'m, 'v>(variant: Variant<'m, 'v>) -> Variant<'m, 'v> { - match variant { - Variant::Int16(value) => compact_spark_integer(value.into()), - Variant::Int32(value) => compact_spark_integer(value.into()), - Variant::Int64(value) => compact_spark_integer(value), - Variant::Decimal8(value) => i32::try_from(value.integer()) - .ok() - .and_then(|integer| VariantDecimal4::try_new(integer, value.scale()).ok()) - .map(Variant::Decimal4) - .unwrap_or(Variant::Decimal8(value)), - Variant::Decimal16(value) => i32::try_from(value.integer()) - .ok() - .and_then(|integer| VariantDecimal4::try_new(integer, value.scale()).ok()) - .map(Variant::Decimal4) - .or_else(|| { - i64::try_from(value.integer()) - .ok() - .and_then(|integer| VariantDecimal8::try_new(integer, value.scale()).ok()) - .map(Variant::Decimal8) - }) - .unwrap_or(Variant::Decimal16(value)), - Variant::Float(value) if value.is_nan() => Variant::Float(f32::from_bits(0x7fc0_0000)), - Variant::Double(value) if value.is_nan() => { - Variant::Double(f64::from_bits(0x7ff8_0000_0000_0000)) - } - variant => variant, - } -} - -/// Re-encode a residual Variant against `metadata`. Scalar widths are intentionally preserved, -/// matching Spark's `VariantBuilder.appendVariant` behavior. -fn spark_variant_bytes( - metadata: &VariantMetadata<'_>, - variant: Variant<'_, '_>, -) -> Result, ArrowError> { - let mut value_builder = ValueBuilder::new(); - match variant { - Variant::Object(object) => { - let mut metadata_builder = SparkMetadataBuilder::new(metadata); - let mut builder = ObjectBuilder::new( - ParentState::variant(&mut value_builder, &mut metadata_builder), - false, - ); - for (name, value) in object.iter() { - let value = spark_variant_bytes(metadata, value)?; - 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), - false, - ); - for value in list.iter() { - let value = spark_variant_bytes(metadata, value)?; - 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()) -} - -fn variant_binary_value(array: &ArrayRef, index: usize) -> Result, ArrowError> { - if array.is_null(index) { - return Ok(None); - } - let value = match array.data_type() { - DataType::Binary => array.as_binary::().value(index), - DataType::LargeBinary => array.as_binary::().value(index), - DataType::BinaryView => array.as_binary_view().value(index), - data_type => { - return Err(ArrowError::InvalidArgumentError(format!( - "Variant value must be binary-like, got {data_type}" - ))) - } - }; - Ok(Some(value)) -} - -fn shredding_state_has_value(state: &BorrowedShreddingState<'_>, index: usize) -> bool { - state - .typed_value_field() - .is_some_and(|array| array.is_valid(index)) - || state - .value_field() - .is_some_and(|array| array.is_valid(index)) -} - -fn collect_spark_field_name(name: &str, field_names: &mut Vec, seen: &mut HashSet) { - if seen.insert(name.to_string()) { - field_names.push(name.to_string()); - } -} - -fn collect_residual_field_names( - variant: Variant<'_, '_>, - field_names: &mut Vec, - seen: &mut HashSet, -) -> Result<(), ArrowError> { - match variant { - Variant::Object(object) => { - for (name, value) in object.iter() { - collect_spark_field_name(name, field_names, seen); - collect_residual_field_names(value, field_names, seen)?; - } - } - Variant::List(list) => { - for value in list.iter() { - collect_residual_field_names(value, field_names, seen)?; - } - } - _ => {} - } - Ok(()) -} - -fn collect_list_field_names( - list: &L, - index: usize, - source_metadata: &VariantMetadata<'_>, - field_names: &mut Vec, - seen: &mut HashSet, -) -> Result<(), ArrowError> { - let values = list.values().as_struct(); - let state = BorrowedShreddingState::try_from(values)?; - for element_index in list.element_range(index) { - collect_shredded_field_names( - state.clone(), - element_index, - source_metadata, - field_names, - seen, - )?; - } - Ok(()) -} - -fn collect_shredded_field_names( - state: BorrowedShreddingState<'_>, - index: usize, - source_metadata: &VariantMetadata<'_>, - field_names: &mut Vec, - seen: &mut HashSet, -) -> Result<(), ArrowError> { - let Some(typed_value) = state - .typed_value_field() - .filter(|array| array.is_valid(index)) - else { - if let Some(value) = state.value_field() { - if let Some(value) = variant_binary_value(value, index)? { - collect_residual_field_names( - Variant::new_with_metadata(source_metadata.clone(), value), - field_names, - seen, - )?; - } - } - return Ok(()); - }; - - match typed_value.data_type() { - DataType::Struct(_) => { - let object = typed_value.as_struct(); - for (field, column) in object.fields().iter().zip(object.columns()) { - let child = column.as_struct_opt().ok_or_else(|| { - ArrowError::InvalidArgumentError(format!( - "Invalid shredded Variant object field '{}': expected Struct, got {}", - field.name(), - column.data_type() - )) - })?; - if child.is_null(index) { - return Err(ArrowError::InvalidArgumentError(format!( - "Shredded Variant object field '{}' is null", - field.name() - ))); - } - let child_state = BorrowedShreddingState::try_from(child)?; - if shredding_state_has_value(&child_state, index) { - collect_spark_field_name(field.name(), field_names, seen); - collect_shredded_field_names( - child_state, - index, - source_metadata, - field_names, - seen, - )?; - } - } - - if let Some(value) = state.value_field() { - if let Some(value) = variant_binary_value(value, index)? { - let Variant::Object(residual) = - Variant::new_with_metadata(source_metadata.clone(), value) - else { - return Err(ArrowError::InvalidArgumentError( - "Partially shredded Variant object has a non-object value".to_string(), - )); - }; - for (name, value) in residual.iter() { - if object.fields().iter().any(|field| field.name() == name) { - return Err(ArrowError::InvalidArgumentError(format!( - "Variant field '{name}' appears in both value and typed_value" - ))); - } - collect_spark_field_name(name, field_names, seen); - collect_residual_field_names(value, field_names, seen)?; - } - } - } - } - DataType::List(_) => collect_list_field_names( - typed_value.as_list::(), - index, - source_metadata, - field_names, - seen, - )?, - DataType::LargeList(_) => collect_list_field_names( - typed_value.as_list::(), - index, - source_metadata, - field_names, - seen, - )?, - DataType::ListView(_) => collect_list_field_names( - typed_value.as_list_view::(), - index, - source_metadata, - field_names, - seen, - )?, - DataType::LargeListView(_) => collect_list_field_names( - typed_value.as_list_view::(), - index, - source_metadata, - field_names, - seen, - )?, - _ => {} - } - Ok(()) -} - -fn spark_typed_variant_bytes( - metadata: &VariantMetadata<'_>, - variant: Variant<'_, '_>, -) -> Result, ArrowError> { - let mut value_builder = ValueBuilder::new(); - let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); - ValueBuilder::try_append_variant( - ParentState::variant(&mut value_builder, &mut metadata_builder), - compact_spark_typed_variant(variant), - )?; - Ok(value_builder.into_inner()) -} - -fn spark_list_bytes( - list: &L, - index: usize, - semantic: Variant<'_, '_>, - source_metadata: &VariantMetadata<'_>, - target_metadata: &VariantMetadata<'_>, -) -> Result, ArrowError> { - let Variant::List(semantic) = semantic else { - return Err(ArrowError::InvalidArgumentError( - "Shredded Variant list did not unshred to a list".to_string(), - )); - }; - let semantic = semantic.iter_try().collect::, _>>()?; - let element_range = list.element_range(index); - if element_range.len() != semantic.len() { - return Err(ArrowError::InvalidArgumentError( - "Shredded Variant list length changed while unshredding".to_string(), - )); - } - - let values = list.values().as_struct(); - let state = BorrowedShreddingState::try_from(values)?; - let mut elements = Vec::with_capacity(semantic.len()); - for (element_index, semantic) in element_range.zip(semantic) { - elements.push(spark_shredded_variant_bytes( - state.clone(), - element_index, - semantic, - source_metadata, - target_metadata, - )?); - } - - let mut value_builder = ValueBuilder::new(); - let mut metadata_builder = ReadOnlyMetadataBuilder::new(target_metadata); - let mut builder = ListBuilder::new( - ParentState::variant(&mut value_builder, &mut metadata_builder), - false, - ); - for element in elements { - builder.append_value_bytes(Variant::new_with_metadata( - target_metadata.clone(), - &element, - )); - } - builder.finish(); - Ok(value_builder.into_inner()) -} - -fn spark_object_bytes( - state: BorrowedShreddingState<'_>, - object: &StructArray, - index: usize, - semantic: Variant<'_, '_>, - source_metadata: &VariantMetadata<'_>, - target_metadata: &VariantMetadata<'_>, -) -> Result, ArrowError> { - let Variant::Object(semantic) = semantic else { - return Err(ArrowError::InvalidArgumentError( - "Shredded Variant object did not unshred to an object".to_string(), - )); - }; - let mut entries = Vec::new(); - for (field, column) in object.fields().iter().zip(object.columns()) { - let child = column.as_struct_opt().ok_or_else(|| { - ArrowError::InvalidArgumentError(format!( - "Invalid shredded Variant object field '{}': expected Struct, got {}", - field.name(), - column.data_type() - )) - })?; - if child.is_null(index) { - return Err(ArrowError::InvalidArgumentError(format!( - "Shredded Variant object field '{}' is null", - field.name() - ))); - } - let child_state = BorrowedShreddingState::try_from(child)?; - if shredding_state_has_value(&child_state, index) { - let value = semantic.get(field.name()).ok_or_else(|| { - ArrowError::InvalidArgumentError(format!( - "Unshredded Variant is missing field '{}'", - field.name() - )) - })?; - entries.push(( - field.name().to_string(), - spark_shredded_variant_bytes( - child_state, - index, - value, - source_metadata, - target_metadata, - )?, - )); - } - } - - if let Some(value) = state.value_field() { - if let Some(value) = variant_binary_value(value, index)? { - let Variant::Object(residual) = - Variant::new_with_metadata(source_metadata.clone(), value) - else { - return Err(ArrowError::InvalidArgumentError( - "Partially shredded Variant object has a non-object value".to_string(), - )); - }; - for (name, value) in residual.iter() { - if object.fields().iter().any(|field| field.name() == name) { - return Err(ArrowError::InvalidArgumentError(format!( - "Variant field '{name}' appears in both value and typed_value" - ))); - } - entries.push(( - name.to_string(), - spark_variant_bytes(target_metadata, value)?, - )); - } - } - } - - let mut value_builder = ValueBuilder::new(); - let mut metadata_builder = SparkMetadataBuilder::new(target_metadata); - let mut builder = ObjectBuilder::new( - ParentState::variant(&mut value_builder, &mut metadata_builder), - false, - ); - for (name, value) in entries { - builder.try_insert_bytes( - &name, - Variant::new_with_metadata(target_metadata.clone(), &value), - )?; - } - builder.finish(); - Ok(value_builder.into_inner()) -} - -// ponytail: recursive child buffers can be O(depth²); use a streaming encoder only if deeply -// nested Variant profiles show this compatibility path is a bottleneck. -fn spark_shredded_variant_bytes( - state: BorrowedShreddingState<'_>, - index: usize, - semantic: Variant<'_, '_>, - source_metadata: &VariantMetadata<'_>, - target_metadata: &VariantMetadata<'_>, -) -> Result, ArrowError> { - let Some(typed_value) = state - .typed_value_field() - .filter(|array| array.is_valid(index)) - else { - return match state.value_field() { - Some(value) => match variant_binary_value(value, index)? { - Some(value) => spark_variant_bytes( - target_metadata, - Variant::new_with_metadata(source_metadata.clone(), value), - ), - None => Err(ArrowError::InvalidArgumentError( - "Shredded Variant has neither value nor typed_value".to_string(), - )), - }, - None => Err(ArrowError::InvalidArgumentError( - "Shredded Variant has neither value nor typed_value".to_string(), - )), - }; - }; - - match typed_value.data_type() { - DataType::Struct(_) => spark_object_bytes( - state, - typed_value.as_struct(), - index, - semantic, - source_metadata, - target_metadata, - ), - DataType::List(_) => spark_list_bytes( - typed_value.as_list::(), - index, - semantic, - source_metadata, - target_metadata, - ), - DataType::LargeList(_) => spark_list_bytes( - typed_value.as_list::(), - index, - semantic, - source_metadata, - target_metadata, - ), - DataType::ListView(_) => spark_list_bytes( - typed_value.as_list_view::(), - index, - semantic, - source_metadata, - target_metadata, - ), - DataType::LargeListView(_) => spark_list_bytes( - typed_value.as_list_view::(), - index, - semantic, - source_metadata, - target_metadata, - ), - _ => spark_typed_variant_bytes(target_metadata, semantic), - } -} - -/// Arrow's unshredder preserves the source metadata but rebuilds object slots in UTF-8 order. -/// Rebuild from the physical shredding state so Spark's metadata insertion order, UTF-16 object -/// headers, typed scalar widths, and residual scalar bytes all remain compatible. -fn rebuild_shredded_variant_for_spark( - source: &VariantArray, - value: &ArrayRef, - metadata: &ArrayRef, - parent_nulls: Option<&NullBuffer>, -) -> DataFusionResult<(ArrayRef, ArrayRef)> { - let source_metadata = cast(source.metadata_field().as_ref(), &DataType::Binary)?; - let source_metadata = source_metadata.as_binary::(); - let source_state = source.shredding_state().borrow(); - let value = value.as_binary::(); - let metadata = metadata.as_binary::(); - let mut value_output = BinaryBuilder::new(); - let mut metadata_output = BinaryBuilder::new(); - - for index in 0..value.len() { - if parent_nulls.is_some_and(|nulls| nulls.is_null(index)) { - value_output.append_null(); - metadata_output.append_null(); - continue; - } - if value.is_null(index) || metadata.is_null(index) { - return Err(DataFusionError::Execution(format!( - "Variant value or metadata is null at row {index}" - ))); - } - - let (rebuilt_value, rebuilt_metadata) = catch_unwind(AssertUnwindSafe( - || -> Result<(Vec, Vec), ArrowError> { - let source_metadata = VariantMetadata::new(source_metadata.value(index)); - let semantic_metadata = VariantMetadata::try_new(metadata.value(index))?; - let semantic = Variant::new_with_metadata(semantic_metadata, value.value(index)); - let mut field_names = Vec::new(); - collect_shredded_field_names( - source_state.clone(), - index, - &source_metadata, - &mut field_names, - &mut HashSet::new(), - ) - .map_err(|error| { - ArrowError::InvalidArgumentError(format!( - "Failed to collect Spark Variant metadata: {error}" - )) - })?; - - let mut metadata_builder = - WritableMetadataBuilder::from_iter(field_names.iter().map(String::as_str)); - metadata_builder.finish(); - let mut rebuilt_metadata = metadata_builder.into_inner(); - // Spark's VariantBuilder never marks its insertion-ordered dictionary as sorted. - rebuilt_metadata[0] &= !0x10; - let target = VariantMetadata::new(&rebuilt_metadata); - let rebuilt_value = spark_shredded_variant_bytes( - source_state.clone(), - index, - semantic, - &source_metadata, - &target, - ) - .map_err(|error| { - ArrowError::InvalidArgumentError(format!( - "Failed to rebuild Spark Variant value: {error}" - )) - })?; - Ok((rebuilt_value, rebuilt_metadata)) - }, - )) - .map_err(|_| { - DataFusionError::Execution(format!("Invalid shredded Variant at row {index}")) - })??; - value_output.append_value(rebuilt_value); - metadata_output.append_value(rebuilt_metadata); - } - - Ok(( - Arc::new(value_output.finish()), - Arc::new(metadata_output.finish()), - )) -} - -/// Reorder object keys for either Arrow's UTF-8 order or Spark's Java UTF-16 order. Preserve -/// already-compatible values byte-for-byte and retain the original metadata dictionary. -/// SPARK-58949 tracks this mismatch and legacy compatibility. The metadata dictionary's sorted -/// flag affects dictionary lookup, not object-entry ordering; Spark's builder and lookup must -/// agree while continuing to read Variant values already written by Spark 4.x in UTF-16 order. -/// https://issues.apache.org/jira/browse/SPARK-58949 -/// https://github.com/apache/parquet-java/issues/3735 -fn reorder_variant_values( - value: &ArrayRef, - metadata: &ArrayRef, - parent_nulls: Option<&NullBuffer>, - order: VariantObjectKeyOrder, - allow_null_value: bool, -) -> DataFusionResult { - let value = value.as_any().downcast_ref::().unwrap(); - let metadata = metadata.as_any().downcast_ref::().unwrap(); - 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) { - if allow_null_value { - output.append_null(); - continue; - } - 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(|| -> DataFusionResult>> { - // Spark encodes empty object keys with equal metadata offsets, which Arrow 58.4's - // full validator rejects. Keep shallow parsing and all accesses inside this boundary. - // https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant/src/variant/metadata.rs#L307-L317 - // Upstream fix: https://github.com/apache/arrow-rs/pull/10352 - let metadata = VariantMetadata::new(metadata.value(index)); - let variant = Variant::new_with_metadata(metadata.clone(), value.value(index)); - if is_compatible_variant(&variant, order) { - return Ok(None); - } - let value = match order { - VariantObjectKeyOrder::ArrowUtf8 => { - let mut value_builder = ValueBuilder::new(); - let mut metadata_builder = ReadOnlyMetadataBuilder::new(&metadata); - ValueBuilder::try_append_variant( - ParentState::variant(&mut value_builder, &mut metadata_builder), - variant, - )?; - value_builder.into_inner() - } - VariantObjectKeyOrder::SparkUtf16 => spark_variant_bytes(&metadata, variant)?, - }; - Ok(Some(value)) - })) - .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())) -} - #[derive(Debug, Clone, Eq)] pub struct CometCastColumnExpr { /// The physical expression producing the value to cast. @@ -1457,78 +367,16 @@ impl PhysicalExpr for CometCastColumnExpr { #[cfg(test)] mod tests { use super::*; - use arrow::array::{ - Array, AsArray, BinaryArray, Decimal128Array, DictionaryArray, FixedSizeListArray, - Int32Array, Int64Array, StringArray, TimestampMillisecondArray, UInt16Array, UInt32Array, - UInt8Array, + use arrow::{ + array::{ + Array, AsArray, BinaryArray, DictionaryArray, Int32Array, Int64Array, StringArray, + TimestampMillisecondArray, + }, + compute::cast, + datatypes::{Field, Fields, Int32Type}, }; - use arrow::datatypes::{Field, Fields, Int32Type}; use datafusion::physical_expr::expressions::Column; - use parquet::variant::{VariantArrayBuilder, VariantBuilder, 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}"); - let emoji = fields - .binary_search_by(|(name, _)| name.encode_utf16().cmp("😀".encode_utf16())) - .unwrap(); - assert_eq!(fields[emoji].1.as_int64(), Some(531)); - let private_use = fields - .binary_search_by(|(name, _)| name.encode_utf16().cmp("\u{e000}".encode_utf16())) - .unwrap(); - assert_eq!(fields[private_use].1.as_int64(), Some(30)); - } - - fn assert_spark_unicode_object(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))); - } - - fn normalize_typed_value(typed_value: ArrayRef, field_names: &[&str]) -> VariantArray { - let (metadata_bytes, _) = VariantBuilder::new() - .with_field_names(field_names.iter().copied()) - .finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - 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![metadata, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - VariantArray::try_new(output.as_ref()).unwrap() - } + use parquet::variant::{Variant, VariantArray, VariantArrayBuilder, VariantType}; #[test] fn test_normalize_shredded_variant_with_dictionary_metadata_for_spark() { @@ -1600,859 +448,6 @@ mod tests { assert_eq!(variant.value(2), Variant::from(30_i8)); } - #[test] - fn test_normalize_shredded_variant_widens_unsigned_values() { - let metadata_builder = VariantBuilder::new().with_field_names(["u8", "u16", "u32"]); - let (metadata_bytes, _) = metadata_builder.finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - - let fields = [ - ("u8", Arc::new(UInt8Array::from(vec![u8::MAX])) as ArrayRef), - ( - "u16", - Arc::new(UInt16Array::from(vec![u16::MAX])) as ArrayRef, - ), - ( - "u32", - Arc::new(UInt32Array::from(vec![u32::MAX])) as ArrayRef, - ), - ]; - let mut object_fields = Vec::with_capacity(fields.len()); - let mut object_columns = Vec::with_capacity(fields.len()); - for (name, value) in fields { - let shredded = StructArray::try_new( - Fields::from(vec![Field::new( - "typed_value", - value.data_type().clone(), - false, - )]), - vec![value], - None, - ) - .unwrap(); - object_fields.push(Field::new(name, shredded.data_type().clone(), false)); - object_columns.push(Arc::new(shredded) as ArrayRef); - } - let typed_value: ArrayRef = - Arc::new(StructArray::try_new(object_fields.into(), object_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![metadata, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = VariantArray::try_new(output.as_ref()).unwrap(); - let variant = output.value(0); - let Variant::Object(object) = variant else { - panic!("expected object") - }; - assert_eq!(object.get("u8"), Some(Variant::from(255_i16))); - assert_eq!(object.get("u16"), Some(Variant::from(65_535_i32))); - assert_eq!(object.get("u32"), Some(Variant::from(4_294_967_295_i64))); - } - - #[test] - fn test_normalize_shredded_variant_widens_millisecond_timestamps() { - let millis = 1_704_067_200_123_i64; - let ltz: ArrayRef = - Arc::new(TimestampMillisecondArray::from(vec![millis]).with_timezone("UTC")); - let ntz: ArrayRef = Arc::new(TimestampMillisecondArray::from(vec![millis])); - let shredded = |value: ArrayRef| -> ArrayRef { - Arc::new( - StructArray::try_new( - Fields::from(vec![Field::new( - "typed_value", - value.data_type().clone(), - false, - )]), - vec![value], - None, - ) - .unwrap(), - ) - }; - let ltz = shredded(ltz); - let ntz = shredded(ntz); - let object: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![ - Field::new("ltz", ltz.data_type().clone(), false), - Field::new("ntz", ntz.data_type().clone(), false), - ]), - vec![ltz, ntz], - None, - ) - .unwrap(), - ); - let output = normalize_typed_value(object, &["ltz", "ntz"]); - let Variant::Object(object) = output.value(0) else { - panic!("expected object") - }; - - let Some(Variant::TimestampMicros(ltz)) = object.get("ltz") else { - panic!("expected timestamp") - }; - assert_eq!(ltz.timestamp_micros(), millis * 1_000); - - let Some(Variant::TimestampNtzMicros(ntz)) = object.get("ntz") else { - panic!("expected timestamp_ntz") - }; - assert_eq!(ntz.and_utc().timestamp_micros(), millis * 1_000); - } - - #[test] - fn test_normalize_shredded_variant_converts_fixed_size_list() { - let elements: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), - vec![Arc::new(Int64Array::from(vec![42, 43]))], - None, - ) - .unwrap(), - ); - let typed_value: ArrayRef = Arc::new( - FixedSizeListArray::try_new( - Arc::new(Field::new("element", elements.data_type().clone(), false)), - 2, - elements, - None, - ) - .unwrap(), - ); - - let output = normalize_typed_value(typed_value, &[]); - let Variant::List(list) = output.value(0) else { - panic!("expected list") - }; - assert_eq!( - list.iter() - .map(|value| value.as_int64()) - .collect::>(), - vec![Some(42), Some(43)] - ); - } - - #[test] - fn test_normalize_shredded_variant_compacts_spark_integer_widths() { - let (metadata_bytes, _) = VariantBuilder::new().finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![ - Some(metadata_bytes.as_slice()), - Some(metadata_bytes.as_slice()), - Some(metadata_bytes.as_slice()), - Some(metadata_bytes.as_slice()), - ])); - let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![ - 1, - i64::from(i8::MAX) + 1, - i64::from(i16::MAX) + 1, - i64::from(i32::MAX) + 1, - ])); - let physical: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![ - Field::new("metadata", DataType::Binary, false), - Field::new("typed_value", DataType::Int64, false), - ]), - vec![metadata, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = VariantArray::try_new(output.as_ref()).unwrap(); - assert_eq!(output.value(0), Variant::Int8(1)); - assert_eq!(output.value(1), Variant::Int16(128)); - assert_eq!(output.value(2), Variant::Int32(32_768)); - assert_eq!(output.value(3), Variant::Int64(2_147_483_648)); - } - - #[test] - fn test_compact_spark_typed_variant_canonicalizes_nan() { - let Variant::Float(float) = - compact_spark_typed_variant(Variant::Float(f32::from_bits(0x7fc0_0001))) - else { - panic!("expected float") - }; - assert_eq!(float.to_bits(), 0x7fc0_0000); - - let Variant::Double(double) = - compact_spark_typed_variant(Variant::Double(f64::from_bits(0x7ff8_0000_0000_0001))) - else { - panic!("expected double") - }; - assert_eq!(double.to_bits(), 0x7ff8_0000_0000_0000); - - let (metadata, _) = VariantBuilder::new().finish(); - let metadata = VariantMetadata::new(&metadata); - let residual = f32::from_bits(0x7fc0_0001); - let bytes = spark_variant_bytes(&metadata, Variant::Float(residual)).unwrap(); - let Variant::Float(output) = Variant::new_with_metadata(metadata, &bytes) else { - panic!("expected residual float") - }; - assert_eq!(output.to_bits(), residual.to_bits()); - } - - #[test] - fn test_normalize_shredded_variant_rejects_missing_required_value() { - let (metadata_bytes, _) = VariantBuilder::new().finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![None])); - let physical: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![ - Field::new("metadata", DataType::Binary, false), - Field::new("typed_value", DataType::Int64, true), - ]), - vec![metadata, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - assert!(normalize_variant_array(&physical, &target_field).is_err()); - } - - #[test] - fn test_normalize_shredded_variant_uses_physical_metadata_order() { - let (metadata_bytes, _) = VariantBuilder::new().with_field_names(["a", "b"]).finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let mut fields = Vec::new(); - let mut columns = Vec::new(); - for (name, value) in [("b", 2_i64), ("a", 1_i64)] { - let child = StructArray::try_new( - Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), - vec![Arc::new(Int64Array::from(vec![value]))], - None, - ) - .unwrap(); - fields.push(Field::new(name, child.data_type().clone(), false)); - columns.push(Arc::new(child) 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![metadata, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = output.as_struct(); - assert_eq!( - output.column(1).as_binary::().value(0), - &[0x01, 2, 0, 1, 2, b'b', b'a'] - ); - - let mut expected = VariantBuilder::new().with_field_names(["b", "a"]); - let mut object = expected.new_object(); - object.insert("b", 2_i8); - object.insert("a", 1_i8); - object.finish(); - let (_, expected_value) = expected.finish(); - assert_eq!(output.column(0).as_binary::().value(0), expected_value); - } - - #[test] - fn test_normalize_shredded_variant_preserves_residual_scalar_width() { - let mut builder = VariantBuilder::new().with_field_names(["known", "residual"]); - let mut object = builder.new_object(); - object.insert("residual", 1_i64); - object.finish(); - let (metadata_bytes, value_bytes) = builder.finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); - let known: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), - vec![Arc::new(Int64Array::from(vec![2]))], - None, - ) - .unwrap(), - ); - let typed_value: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), - vec![known], - 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, value, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = VariantArray::try_new(output.as_ref()).unwrap(); - let Variant::Object(object) = output.value(0) else { - panic!("expected object") - }; - assert_eq!(object.get("known"), Some(Variant::Int8(2))); - assert_eq!(object.get("residual"), Some(Variant::Int64(1))); - } - - #[test] - fn test_normalize_shredded_variant_compacts_spark_decimal_width() { - let (metadata_bytes, _) = VariantBuilder::new().finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let typed_value: ArrayRef = Arc::new( - Decimal128Array::from(vec![123_i128]) - .with_precision_and_scale(38, 2) - .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![metadata, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = VariantArray::try_new(output.as_ref()).unwrap(); - assert_eq!( - output.value(0), - Variant::Decimal4(VariantDecimal4::try_new(123, 2).unwrap()) - ); - } - - #[test] - fn test_normalize_shredded_variant_uses_spark_object_key_order() { - let keys = unicode_object_keys(); - - let metadata_builder = - VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); - let (metadata_bytes, _) = metadata_builder.finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - - let mut object_fields = Vec::with_capacity(keys.len()); - let mut object_columns = Vec::with_capacity(keys.len()); - for (index, key) in keys.iter().enumerate() { - let value = if key == "😀" { 531 } else { index as i64 }; - let field = Field::new("typed_value", DataType::Int64, false); - let column: ArrayRef = Arc::new(Int64Array::from(vec![value])); - let shredded_field = - StructArray::try_new(Fields::from(vec![field]), vec![column], None).unwrap(); - object_fields.push(Field::new(key, shredded_field.data_type().clone(), false)); - object_columns.push(Arc::new(shredded_field) as ArrayRef); - } - let typed_value: ArrayRef = - Arc::new(StructArray::try_new(object_fields.into(), object_columns, None).unwrap()); - let physical_fields = Fields::from(vec![ - Field::new("metadata", DataType::Binary, false), - Field::new("typed_value", typed_value.data_type().clone(), false), - ]); - let physical: ArrayRef = Arc::new( - StructArray::try_new(physical_fields, vec![metadata, typed_value], None).unwrap(), - ); - let target_fields = Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ]); - let target_field = Arc::new( - Field::new("v", DataType::Struct(target_fields), false) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = output.as_struct(); - assert_spark_unicode_object(output); - } - - #[test] - fn test_normalize_nested_shredded_variant_uses_spark_object_key_order() { - let keys = unicode_object_keys(); - let field_names = std::iter::once("nested") - .chain(keys.iter().map(String::as_str)) - .collect::>(); - let (metadata_bytes, _) = VariantBuilder::new().with_field_names(field_names).finish(); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - - let mut nested_fields = Vec::with_capacity(keys.len()); - let mut nested_columns = Vec::with_capacity(keys.len()); - for (index, key) in keys.iter().enumerate() { - let value = if key == "😀" { 531 } else { index as i64 }; - let state = StructArray::try_new( - Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), - vec![Arc::new(Int64Array::from(vec![value]))], - None, - ) - .unwrap(); - nested_fields.push(Field::new(key, state.data_type().clone(), false)); - nested_columns.push(Arc::new(state) as ArrayRef); - } - let nested_value: ArrayRef = - Arc::new(StructArray::try_new(nested_fields.into(), nested_columns, None).unwrap()); - let nested_state: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![Field::new( - "typed_value", - nested_value.data_type().clone(), - false, - )]), - vec![nested_value], - None, - ) - .unwrap(), - ); - let typed_value: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![Field::new( - "nested", - nested_state.data_type().clone(), - false, - )]), - vec![nested_state], - 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![metadata, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = VariantArray::try_new(output.as_ref()).unwrap(); - let Variant::Object(object) = output.value(0) else { - panic!("expected outer object") - }; - assert_spark_unicode_variant(object.get("nested").expect("nested field")); - } - - #[test] - fn test_normalize_partially_shredded_spark_object_key_order() { - 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_bytes, value_bytes) = builder.finish(); - - let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let spark_value = reorder_variant_values( - &value, - &metadata, - None, - VariantObjectKeyOrder::SparkUtf16, - false, - ) - .unwrap(); - - let shredded_k00: 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_k00.data_type().clone(), - false, - )]), - vec![shredded_k00], - 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, spark_value, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - assert_spark_unicode_object(output.as_struct()); - } - - #[test] - fn test_normalize_unshredded_variant_uses_spark_object_key_order() { - 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_bytes, value_bytes) = builder.finish(); - - let Variant::Object(canonical) = Variant::new(&metadata_bytes, &value_bytes) else { - panic!("expected object") - }; - let canonical_fields = canonical.iter().collect::>(); - assert_eq!(canonical_fields[30].0, "\u{e000}"); - assert_eq!(canonical_fields[31].0, "😀"); - - let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let physical_fields = Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ]); - let physical: ArrayRef = - Arc::new(StructArray::try_new(physical_fields, vec![value, metadata], None).unwrap()); - let target_fields = Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ]); - let target_field = Arc::new( - Field::new("v", DataType::Struct(target_fields), false) - .with_extension_type(VariantType), - ); - - let first = normalize_variant_array(&physical, &target_field).unwrap(); - let first_value = first - .as_struct() - .column(0) - .as_binary::() - .value(0) - .to_vec(); - assert_spark_unicode_object(first.as_struct()); - - // Spark-produced already-unshredded input is UTF-16 ordered. Normalizing it again must - // remain valid without Arrow's UTF-8-order full validation. - let second = normalize_variant_array(&first, &target_field).unwrap(); - assert_spark_unicode_object(second.as_struct()); - assert_eq!( - second.as_struct().column(0).as_binary::().value(0), - first_value - ); - } - - #[test] - fn test_normalize_spark_ordered_variant_preserves_value_bytes() { - let mut builder = VariantBuilder::new(); - let mut object = builder.new_object(); - object.insert("b", 1_i64); - object.insert("a", 2_i64); - object.finish(); - let (metadata_bytes, value_bytes) = builder.finish(); - - let Variant::Object(object) = Variant::new(&metadata_bytes, &value_bytes) else { - panic!("expected object") - }; - assert_eq!( - object.iter().map(|(name, _)| name).collect::>(), - vec!["a", "b"] - ); - - let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let physical_fields = Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ]); - let physical: ArrayRef = - Arc::new(StructArray::try_new(physical_fields, vec![value, metadata], None).unwrap()); - let target_fields = Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ]); - let target_field = Arc::new( - Field::new("v", DataType::Struct(target_fields), false) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - assert_eq!( - output.as_struct().column(0).as_binary::().value(0), - value_bytes - ); - } - - #[test] - fn test_normalize_variant_preserves_empty_object_keys() { - let mut builder = VariantBuilder::new(); - let mut object = builder.new_object(); - object.insert("", 1_i64); - let mut nested = object.new_object("nested"); - nested.insert("", 2_i64); - nested.finish(); - object.finish(); - let (mut metadata_bytes, value_bytes) = builder.finish(); - - // Spark leaves the metadata dictionary unsorted. Equal offsets encode the empty key. - metadata_bytes[0] &= !0x10; - assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); - - let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let physical: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ]), - vec![value, metadata], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = output.as_struct(); - assert_eq!(output.column(0).as_binary::().value(0), value_bytes); - assert_eq!(output.column(1).as_binary::().value(0), metadata_bytes); - - let Variant::Object(object) = Variant::new(&metadata_bytes, &value_bytes) else { - panic!("expected object") - }; - assert_eq!(object.get(""), Some(Variant::from(1_i64))); - let Variant::Object(nested) = object.get("nested").unwrap() else { - panic!("expected nested object") - }; - assert_eq!(nested.get(""), Some(Variant::from(2_i64))); - } - - #[test] - fn test_normalize_partially_shredded_nested_unicode_and_empty_keys() { - let keys = unicode_object_keys(); - let mut builder = VariantBuilder::new().with_field_names(["known"]); - let mut object = builder.new_object(); - object.insert("", 1_i64); - let mut nested = object.new_object("nested"); - for (index, key) in keys.iter().enumerate() { - nested.insert(key, if key == "😀" { 531_i64 } else { index as i64 }); - } - nested.finish(); - object.finish(); - let (metadata_bytes, value_bytes) = builder.finish(); - assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); - - let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); - let known: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), - vec![Arc::new(Int64Array::from(vec![3]))], - None, - ) - .unwrap(), - ); - let typed_value: ArrayRef = Arc::new( - StructArray::try_new( - Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), - vec![known], - 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, value, typed_value], - None, - ) - .unwrap(), - ); - let target_field = Arc::new( - Field::new( - "v", - DataType::Struct(Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ])), - false, - ) - .with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - let output = VariantArray::try_new(output.as_ref()).unwrap(); - let Variant::Object(object) = output.value(0) else { - panic!("expected object") - }; - assert_eq!(object.get("").unwrap().as_int64(), Some(1)); - assert_eq!(object.get("known").unwrap().as_int64(), Some(3)); - let Variant::Object(nested) = object.get("nested").unwrap() else { - panic!("expected nested object") - }; - let fields = nested.iter().collect::>(); - assert_eq!(fields.len(), 32); - assert_eq!(fields[30].0, "😀"); - assert_eq!(fields[31].0, "\u{e000}"); - assert_eq!(nested.get("😀").unwrap().as_int64(), Some(531)); - } - - #[test] - fn test_normalize_variant_skips_empty_children_of_null_parent() { - let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&b""[..])])); - let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&b""[..])])); - let physical_fields = Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ]); - let physical: ArrayRef = Arc::new( - StructArray::try_new( - physical_fields, - vec![value, metadata], - Some(NullBuffer::from(vec![false])), - ) - .unwrap(), - ); - let target_fields = Fields::from(vec![ - Field::new("value", DataType::Binary, false), - Field::new("metadata", DataType::Binary, false), - ]); - let target_field = Arc::new( - Field::new("v", DataType::Struct(target_fields), true).with_extension_type(VariantType), - ); - - let output = normalize_variant_array(&physical, &target_field).unwrap(); - assert!(output.is_null(0)); - assert!(output.as_struct().column(0).is_null(0)); - } - #[test] fn test_cast_timestamp_micros_to_millis_array() { // Create a TimestampMicrosecond array with some values 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..994de41de42 --- /dev/null +++ b/native/core/src/parquet/cast_column/variant.rs @@ -0,0 +1,1122 @@ +// 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::{Array, ArrayRef, AsArray, BinaryArray, BinaryBuilder, ListLikeArray, StructArray}, + buffer::NullBuffer, + compute::{cast, cast_with_options}, + datatypes::{DataType, FieldRef, TimeUnit}, + error::ArrowError, +}; +use datafusion::common::{ + format::DEFAULT_CAST_OPTIONS, DataFusionError, Result as DataFusionResult, +}; +use parquet::variant::{ + unshred_variant, BorrowedShreddingState, ListBuilder, MetadataBuilder, ObjectBuilder, + ParentState, ReadOnlyMetadataBuilder, ValueBuilder, Variant, VariantArray, VariantBuilder, + VariantDecimal4, VariantDecimal8, VariantMetadata, WritableMetadataBuilder, +}; +use std::{ + collections::HashSet, + 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(), + )); + } + + let array = decode_variant_metadata_dictionary(array)?; + let array = normalize_variant_typed_value(&array)?; + let variant = VariantArray::try_new(array.as_ref())?; + let was_shredded = variant.typed_value_field().is_some(); + let unshredded = unshred_variant_for_spark(&variant)?; + 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, metadata) = if was_shredded { + rebuild_shredded_variant_for_spark(&variant, &value, &metadata, unshredded.inner().nulls())? + } else { + let value = reorder_variant_values( + &value, + &metadata, + unshredded.inner().nulls(), + VariantObjectKeyOrder::SparkUtf16, + false, + )?; + (value, metadata) + }; + let output = StructArray::try_new( + fields.clone(), + vec![value, metadata], + unshredded.inner().nulls().cloned(), + )?; + Ok(Arc::new(output)) +} + +fn unshred_variant_for_spark(variant: &VariantArray) -> DataFusionResult { + let first = + prepare_variant_for_unshredding(variant).and_then(|array| Ok(unshred_variant(&array)?)); + let first_error = match first { + Ok(array) => return Ok(array), + Err(error) => error, + }; + let Some(variant) = canonicalize_spark_empty_key_metadata(variant)? else { + return Err(first_error); + }; + let variant = prepare_variant_for_unshredding(&variant)?; + Ok(unshred_variant(&variant)?) +} + +fn normalize_variant_type(data_type: &DataType) -> Option { + fn normalize_field(field: &FieldRef) -> Option { + normalize_variant_type(field.data_type()) + .map(|data_type| Arc::new(field.as_ref().clone().with_data_type(data_type))) + } + + match data_type { + DataType::UInt8 => Some(DataType::Int16), + DataType::UInt16 => Some(DataType::Int32), + DataType::UInt32 => Some(DataType::Int64), + DataType::Timestamp(TimeUnit::Millisecond, timezone) => { + Some(DataType::Timestamp(TimeUnit::Microsecond, timezone.clone())) + } + DataType::FixedSizeList(field, _) => Some(DataType::List( + normalize_field(field).unwrap_or_else(|| Arc::clone(field)), + )), + DataType::List(field) => normalize_field(field).map(DataType::List), + DataType::LargeList(field) => normalize_field(field).map(DataType::LargeList), + DataType::ListView(field) => normalize_field(field).map(DataType::ListView), + DataType::LargeListView(field) => normalize_field(field).map(DataType::LargeListView), + DataType::Struct(fields) => { + let mut changed = false; + let fields = fields + .iter() + .map(|field| match normalize_field(field) { + Some(field) => { + changed = true; + field + } + None => Arc::clone(field), + }) + .collect::>(); + changed.then(|| DataType::Struct(fields.into())) + } + _ => None, + } +} + +/// Normalize Arrow types that Spark's Parquet reader accepts but `VariantArray` 58.4 rejects. +/// Parquet restores unsigned integers to Arrow unsigned arrays and millisecond timestamps at their +/// annotated unit; embedded Arrow schemas may also restore fixed-size lists. Spark widens the +/// integers and timestamps and treats fixed-size lists as ordinary Variant arrays. +/// arrow-rs #10416/#10417 would move this widening into `VariantArray`/`unshred_variant`; remove +/// the unsigned arms after that ships and Comet upgrades: +/// https://github.com/apache/arrow-rs/issues/10416 +/// https://github.com/apache/arrow-rs/pull/10417 +/// Arrow #50622/#50810 instead proposes removing unsigned `typed_value` mappings because the +/// Parquet Variant shredding table permits only signed integer fields. Until upstream resolves +/// that choice, keep this compatibility path for unsigned files Spark already reads: +/// https://github.com/apache/arrow/issues/50622 +/// https://github.com/apache/arrow/pull/50810 +fn normalize_variant_typed_value(array: &ArrayRef) -> DataFusionResult { + let Some(struct_array) = array.as_any().downcast_ref::() else { + return Ok(Arc::clone(array)); + }; + let Some((typed_value_index, typed_value_field)) = struct_array + .fields() + .iter() + .enumerate() + .find(|(_, field)| field.name() == "typed_value") + else { + return Ok(Arc::clone(array)); + }; + let Some(data_type) = normalize_variant_type(typed_value_field.data_type()) else { + return Ok(Arc::clone(array)); + }; + + let mut fields = struct_array.fields().iter().cloned().collect::>(); + fields[typed_value_index] = Arc::new( + typed_value_field + .as_ref() + .clone() + .with_data_type(data_type.clone()), + ); + let mut columns = struct_array.columns().to_vec(); + columns[typed_value_index] = cast_with_options( + columns[typed_value_index].as_ref(), + &data_type, + &DEFAULT_CAST_OPTIONS, + )?; + Ok(Arc::new(StructArray::try_new( + fields.into(), + columns, + struct_array.nulls().cloned(), + )?)) +} + +/// Arrow's unshredder fully validates any residual `value` in a partially shredded object. Spark +/// writes object keys in Java UTF-16 order, so put that residual value in Arrow UTF-8 order only +/// while it passes through the upstream unshredder. +fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult { + let (Some(value), Some(_)) = (variant.value_field(), variant.typed_value_field()) else { + return Ok(variant.clone()); + }; + + let value = cast(value.as_ref(), &DataType::Binary)?; + let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let value = reorder_variant_values( + &value, + &metadata, + variant.inner().nulls(), + VariantObjectKeyOrder::ArrowUtf8, + true, + )?; + + let value_index = variant + .inner() + .fields() + .iter() + .position(|field| field.name() == "value") + .unwrap(); + let mut fields = variant.inner().fields().iter().cloned().collect::>(); + fields[value_index] = Arc::new( + fields[value_index] + .as_ref() + .clone() + .with_data_type(DataType::Binary), + ); + let mut columns = variant.inner().columns().to_vec(); + columns[value_index] = value; + let array = StructArray::try_new(fields.into(), columns, variant.inner().nulls().cloned())?; + Ok(VariantArray::try_new(&array)?) +} + +/// Arrow 58.4 rejects Spark metadata dictionaries containing an empty object key because their +/// unsorted offsets can be equal. Rebuild only those rows before any residual value is traversed, +/// sorting the dictionary and remapping the residual value's field IDs at the same time. +/// https://github.com/apache/arrow-rs/pull/10352 +fn canonicalize_spark_empty_key_metadata( + variant: &VariantArray, +) -> DataFusionResult> { + type Replacement = Option<(Vec, Option>)>; + + let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let metadata = metadata.as_binary::(); + let value = variant + .value_field() + .map(|value| cast(value.as_ref(), &DataType::Binary)) + .transpose()?; + let value = value.as_ref().map(|value| value.as_binary::()); + let mut replacements = Vec::with_capacity(variant.len()); + let mut changed = false; + + for index in 0..variant.len() { + if variant.inner().is_null(index) || metadata.is_null(index) { + replacements.push(None); + continue; + } + let metadata_bytes = metadata.value(index); + if VariantMetadata::try_new(metadata_bytes).is_ok() { + replacements.push(None); + continue; + } + + let replacement = catch_unwind(AssertUnwindSafe(|| -> Result { + let old_metadata = VariantMetadata::new(metadata_bytes); + let mut names = old_metadata + .iter_try() + .map(|name| name.map(str::to_string)) + .collect::, _>>()?; + if !names.iter().any(String::is_empty) { + return Ok(None); + } + names.sort_unstable(); + if names.windows(2).any(|names| names[0] == names[1]) { + return Ok(None); + } + + let mut builder = + VariantBuilder::new().with_field_names(names.iter().map(String::as_str)); + match value { + Some(value) if !value.is_null(index) => { + builder + .append_value(Variant::new_with_metadata(old_metadata, value.value(index))); + let (metadata, value) = builder.finish(); + Ok(Some((metadata, Some(value)))) + } + _ => Ok(Some((builder.finish().0, None))), + } + })) + .map_err(|_| { + DataFusionError::Execution(format!( + "Invalid Variant metadata with an empty key at row {index}" + )) + })??; + changed |= replacement.is_some(); + replacements.push(replacement); + } + + if !changed { + return Ok(None); + } + + let mut metadata_builder = BinaryBuilder::new(); + let mut value_builder = value.map(|_| BinaryBuilder::new()); + for (index, replacement) in replacements.iter().enumerate() { + match replacement { + Some((metadata, value)) => { + metadata_builder.append_value(metadata); + if let Some(builder) = &mut value_builder { + match value { + Some(value) => builder.append_value(value), + None => builder.append_null(), + } + } + } + None => { + if metadata.is_null(index) { + metadata_builder.append_null(); + } else { + metadata_builder.append_value(metadata.value(index)); + } + if let (Some(value), Some(builder)) = (value, &mut value_builder) { + if value.is_null(index) { + builder.append_null(); + } else { + builder.append_value(value.value(index)); + } + } + } + } + } + + let mut fields = variant.inner().fields().iter().cloned().collect::>(); + let mut columns = variant.inner().columns().to_vec(); + let metadata_index = fields + .iter() + .position(|field| field.name() == "metadata") + .unwrap(); + fields[metadata_index] = Arc::new( + fields[metadata_index] + .as_ref() + .clone() + .with_data_type(DataType::Binary), + ); + columns[metadata_index] = Arc::new(metadata_builder.finish()); + if let Some(mut value_builder) = value_builder { + let value_index = fields + .iter() + .position(|field| field.name() == "value") + .unwrap(); + fields[value_index] = Arc::new( + fields[value_index] + .as_ref() + .clone() + .with_data_type(DataType::Binary), + ); + columns[value_index] = Arc::new(value_builder.finish()); + } + let array = StructArray::try_new(fields.into(), columns, variant.inner().nulls().cloned())?; + Ok(Some(VariantArray::try_new(&array)?)) +} + +/// Arrow-rs parquet-variant-compute allows dictionary-encoded metadata in its contract, but 58.4's +/// `VariantArray::try_new` validates only Binary, LargeBinary, and BinaryView. Decode just that +/// child and keep the physical struct otherwise unchanged. +/// https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant-compute/src/variant_array.rs#L276-L310 +/// Upstream issue: https://github.com/apache/arrow-rs/issues/10802 +fn decode_variant_metadata_dictionary(array: &ArrayRef) -> DataFusionResult { + let Some(struct_array) = array.as_any().downcast_ref::() else { + return Ok(Arc::clone(array)); + }; + let Some((metadata_index, metadata_field)) = struct_array + .fields() + .iter() + .enumerate() + .find(|(_, field)| field.name() == "metadata") + else { + return Ok(Arc::clone(array)); + }; + let DataType::Dictionary(_, value_type) = metadata_field.data_type() else { + return Ok(Arc::clone(array)); + }; + + let decoded = cast(struct_array.column(metadata_index).as_ref(), value_type)?; + let mut fields = struct_array.fields().iter().cloned().collect::>(); + fields[metadata_index] = Arc::new( + metadata_field + .as_ref() + .clone() + .with_data_type(decoded.data_type().clone()), + ); + let mut columns = struct_array.columns().to_vec(); + columns[metadata_index] = decoded; + Ok(Arc::new(StructArray::try_new( + fields.into(), + columns, + struct_array.nulls().cloned(), + )?)) +} + +/// Supplies sort-only field names whose Rust ordering matches Java `String.compareTo` ordering. +/// The original metadata dictionary still supplies the field IDs written to the Variant value. +#[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 compact_spark_integer(value: i64) -> Variant<'static, 'static> { + if let Ok(value) = i8::try_from(value) { + Variant::Int8(value) + } else if let Ok(value) = i16::try_from(value) { + Variant::Int16(value) + } else if let Ok(value) = i32::try_from(value) { + Variant::Int32(value) + } else { + Variant::Int64(value) + } +} + +fn compact_spark_typed_variant<'m, 'v>(variant: Variant<'m, 'v>) -> Variant<'m, 'v> { + match variant { + Variant::Int16(value) => compact_spark_integer(value.into()), + Variant::Int32(value) => compact_spark_integer(value.into()), + Variant::Int64(value) => compact_spark_integer(value), + Variant::Decimal8(value) => i32::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal4::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal4) + .unwrap_or(Variant::Decimal8(value)), + Variant::Decimal16(value) => i32::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal4::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal4) + .or_else(|| { + i64::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal8::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal8) + }) + .unwrap_or(Variant::Decimal16(value)), + Variant::Float(value) if value.is_nan() => Variant::Float(f32::from_bits(0x7fc0_0000)), + Variant::Double(value) if value.is_nan() => { + Variant::Double(f64::from_bits(0x7ff8_0000_0000_0000)) + } + variant => variant, + } +} + +/// Re-encode a residual Variant against `metadata`. Scalar widths are intentionally preserved, +/// matching Spark's `VariantBuilder.appendVariant` behavior. +fn spark_variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + match variant { + Variant::Object(object) => { + let mut metadata_builder = SparkMetadataBuilder::new(metadata); + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for (name, value) in object.iter() { + let value = spark_variant_bytes(metadata, value)?; + 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), + false, + ); + for value in list.iter() { + let value = spark_variant_bytes(metadata, value)?; + 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()) +} + +fn variant_binary_value(array: &ArrayRef, index: usize) -> Result, ArrowError> { + if array.is_null(index) { + return Ok(None); + } + let value = match array.data_type() { + DataType::Binary => array.as_binary::().value(index), + DataType::LargeBinary => array.as_binary::().value(index), + DataType::BinaryView => array.as_binary_view().value(index), + data_type => { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant value must be binary-like, got {data_type}" + ))) + } + }; + Ok(Some(value)) +} + +fn shredding_state_has_value(state: &BorrowedShreddingState<'_>, index: usize) -> bool { + state + .typed_value_field() + .is_some_and(|array| array.is_valid(index)) + || state + .value_field() + .is_some_and(|array| array.is_valid(index)) +} + +fn collect_spark_field_name(name: &str, field_names: &mut Vec, seen: &mut HashSet) { + if seen.insert(name.to_string()) { + field_names.push(name.to_string()); + } +} + +fn collect_residual_field_names( + variant: Variant<'_, '_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + match variant { + Variant::Object(object) => { + for (name, value) in object.iter() { + collect_spark_field_name(name, field_names, seen); + collect_residual_field_names(value, field_names, seen)?; + } + } + Variant::List(list) => { + for value in list.iter() { + collect_residual_field_names(value, field_names, seen)?; + } + } + _ => {} + } + Ok(()) +} + +fn collect_list_field_names( + list: &L, + index: usize, + source_metadata: &VariantMetadata<'_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + let values = list.values().as_struct(); + let state = BorrowedShreddingState::try_from(values)?; + for element_index in list.element_range(index) { + collect_shredded_field_names( + state.clone(), + element_index, + source_metadata, + field_names, + seen, + )?; + } + Ok(()) +} + +fn collect_shredded_field_names( + state: BorrowedShreddingState<'_>, + index: usize, + source_metadata: &VariantMetadata<'_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + let Some(typed_value) = state + .typed_value_field() + .filter(|array| array.is_valid(index)) + else { + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + collect_residual_field_names( + Variant::new_with_metadata(source_metadata.clone(), value), + field_names, + seen, + )?; + } + } + return Ok(()); + }; + + match typed_value.data_type() { + DataType::Struct(_) => { + let object = typed_value.as_struct(); + for (field, column) in object.fields().iter().zip(object.columns()) { + let child = column.as_struct_opt().ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Invalid shredded Variant object field '{}': expected Struct, got {}", + field.name(), + column.data_type() + )) + })?; + if child.is_null(index) { + return Err(ArrowError::InvalidArgumentError(format!( + "Shredded Variant object field '{}' is null", + field.name() + ))); + } + let child_state = BorrowedShreddingState::try_from(child)?; + if shredding_state_has_value(&child_state, index) { + collect_spark_field_name(field.name(), field_names, seen); + collect_shredded_field_names( + child_state, + index, + source_metadata, + field_names, + seen, + )?; + } + } + + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + let Variant::Object(residual) = + Variant::new_with_metadata(source_metadata.clone(), value) + else { + return Err(ArrowError::InvalidArgumentError( + "Partially shredded Variant object has a non-object value".to_string(), + )); + }; + for (name, value) in residual.iter() { + if object.fields().iter().any(|field| field.name() == name) { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant field '{name}' appears in both value and typed_value" + ))); + } + collect_spark_field_name(name, field_names, seen); + collect_residual_field_names(value, field_names, seen)?; + } + } + } + } + DataType::List(_) => collect_list_field_names( + typed_value.as_list::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::LargeList(_) => collect_list_field_names( + typed_value.as_list::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::ListView(_) => collect_list_field_names( + typed_value.as_list_view::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::LargeListView(_) => collect_list_field_names( + typed_value.as_list_view::(), + index, + source_metadata, + field_names, + seen, + )?, + _ => {} + } + Ok(()) +} + +fn spark_typed_variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + ValueBuilder::try_append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + compact_spark_typed_variant(variant), + )?; + Ok(value_builder.into_inner()) +} + +fn spark_list_bytes( + list: &L, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Variant::List(semantic) = semantic else { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant list did not unshred to a list".to_string(), + )); + }; + let semantic = semantic.iter_try().collect::, _>>()?; + let element_range = list.element_range(index); + if element_range.len() != semantic.len() { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant list length changed while unshredding".to_string(), + )); + } + + let values = list.values().as_struct(); + let state = BorrowedShreddingState::try_from(values)?; + let mut elements = Vec::with_capacity(semantic.len()); + for (element_index, semantic) in element_range.zip(semantic) { + elements.push(spark_shredded_variant_bytes( + state.clone(), + element_index, + semantic, + source_metadata, + target_metadata, + )?); + } + + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = ReadOnlyMetadataBuilder::new(target_metadata); + let mut builder = ListBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for element in elements { + builder.append_value_bytes(Variant::new_with_metadata( + target_metadata.clone(), + &element, + )); + } + builder.finish(); + Ok(value_builder.into_inner()) +} + +fn spark_object_bytes( + state: BorrowedShreddingState<'_>, + object: &StructArray, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Variant::Object(semantic) = semantic else { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant object did not unshred to an object".to_string(), + )); + }; + let mut entries = Vec::new(); + for (field, column) in object.fields().iter().zip(object.columns()) { + let child = column.as_struct_opt().ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Invalid shredded Variant object field '{}': expected Struct, got {}", + field.name(), + column.data_type() + )) + })?; + if child.is_null(index) { + return Err(ArrowError::InvalidArgumentError(format!( + "Shredded Variant object field '{}' is null", + field.name() + ))); + } + let child_state = BorrowedShreddingState::try_from(child)?; + if shredding_state_has_value(&child_state, index) { + let value = semantic.get(field.name()).ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Unshredded Variant is missing field '{}'", + field.name() + )) + })?; + entries.push(( + field.name().to_string(), + spark_shredded_variant_bytes( + child_state, + index, + value, + source_metadata, + target_metadata, + )?, + )); + } + } + + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + let Variant::Object(residual) = + Variant::new_with_metadata(source_metadata.clone(), value) + else { + return Err(ArrowError::InvalidArgumentError( + "Partially shredded Variant object has a non-object value".to_string(), + )); + }; + for (name, value) in residual.iter() { + if object.fields().iter().any(|field| field.name() == name) { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant field '{name}' appears in both value and typed_value" + ))); + } + entries.push(( + name.to_string(), + spark_variant_bytes(target_metadata, value)?, + )); + } + } + } + + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = SparkMetadataBuilder::new(target_metadata); + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for (name, value) in entries { + builder.try_insert_bytes( + &name, + Variant::new_with_metadata(target_metadata.clone(), &value), + )?; + } + builder.finish(); + Ok(value_builder.into_inner()) +} + +// ponytail: recursive child buffers can be O(depth²); use a streaming encoder only if deeply +// nested Variant profiles show this compatibility path is a bottleneck. +fn spark_shredded_variant_bytes( + state: BorrowedShreddingState<'_>, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Some(typed_value) = state + .typed_value_field() + .filter(|array| array.is_valid(index)) + else { + return match state.value_field() { + Some(value) => match variant_binary_value(value, index)? { + Some(value) => spark_variant_bytes( + target_metadata, + Variant::new_with_metadata(source_metadata.clone(), value), + ), + None => Err(ArrowError::InvalidArgumentError( + "Shredded Variant has neither value nor typed_value".to_string(), + )), + }, + None => Err(ArrowError::InvalidArgumentError( + "Shredded Variant has neither value nor typed_value".to_string(), + )), + }; + }; + + match typed_value.data_type() { + DataType::Struct(_) => spark_object_bytes( + state, + typed_value.as_struct(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::List(_) => spark_list_bytes( + typed_value.as_list::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::LargeList(_) => spark_list_bytes( + typed_value.as_list::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::ListView(_) => spark_list_bytes( + typed_value.as_list_view::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::LargeListView(_) => spark_list_bytes( + typed_value.as_list_view::(), + index, + semantic, + source_metadata, + target_metadata, + ), + _ => spark_typed_variant_bytes(target_metadata, semantic), + } +} + +/// Arrow's unshredder preserves the source metadata but rebuilds object slots in UTF-8 order. +/// Rebuild from the physical shredding state so Spark's metadata insertion order, UTF-16 object +/// headers, typed scalar widths, and residual scalar bytes all remain compatible. +fn rebuild_shredded_variant_for_spark( + source: &VariantArray, + value: &ArrayRef, + metadata: &ArrayRef, + parent_nulls: Option<&NullBuffer>, +) -> DataFusionResult<(ArrayRef, ArrayRef)> { + let source_metadata = cast(source.metadata_field().as_ref(), &DataType::Binary)?; + let source_metadata = source_metadata.as_binary::(); + let source_state = source.shredding_state().borrow(); + let value = value.as_binary::(); + let metadata = metadata.as_binary::(); + let mut value_output = BinaryBuilder::new(); + let mut metadata_output = BinaryBuilder::new(); + + for index in 0..value.len() { + if parent_nulls.is_some_and(|nulls| nulls.is_null(index)) { + value_output.append_null(); + metadata_output.append_null(); + continue; + } + if value.is_null(index) || metadata.is_null(index) { + return Err(DataFusionError::Execution(format!( + "Variant value or metadata is null at row {index}" + ))); + } + + let (rebuilt_value, rebuilt_metadata) = catch_unwind(AssertUnwindSafe( + || -> Result<(Vec, Vec), ArrowError> { + let source_metadata = VariantMetadata::new(source_metadata.value(index)); + let semantic_metadata = VariantMetadata::try_new(metadata.value(index))?; + let semantic = Variant::new_with_metadata(semantic_metadata, value.value(index)); + let mut field_names = Vec::new(); + collect_shredded_field_names( + source_state.clone(), + index, + &source_metadata, + &mut field_names, + &mut HashSet::new(), + ) + .map_err(|error| { + ArrowError::InvalidArgumentError(format!( + "Failed to collect Spark Variant metadata: {error}" + )) + })?; + + let mut metadata_builder = + WritableMetadataBuilder::from_iter(field_names.iter().map(String::as_str)); + metadata_builder.finish(); + let mut rebuilt_metadata = metadata_builder.into_inner(); + // Spark's VariantBuilder never marks its insertion-ordered dictionary as sorted. + rebuilt_metadata[0] &= !0x10; + let target = VariantMetadata::new(&rebuilt_metadata); + let rebuilt_value = spark_shredded_variant_bytes( + source_state.clone(), + index, + semantic, + &source_metadata, + &target, + ) + .map_err(|error| { + ArrowError::InvalidArgumentError(format!( + "Failed to rebuild Spark Variant value: {error}" + )) + })?; + Ok((rebuilt_value, rebuilt_metadata)) + }, + )) + .map_err(|_| { + DataFusionError::Execution(format!("Invalid shredded Variant at row {index}")) + })??; + value_output.append_value(rebuilt_value); + metadata_output.append_value(rebuilt_metadata); + } + + Ok(( + Arc::new(value_output.finish()), + Arc::new(metadata_output.finish()), + )) +} + +/// Reorder object keys for either Arrow's UTF-8 order or Spark's Java UTF-16 order. Preserve +/// already-compatible values byte-for-byte and retain the original metadata dictionary. +/// SPARK-58949 tracks this mismatch and legacy compatibility. The metadata dictionary's sorted +/// flag affects dictionary lookup, not object-entry ordering; Spark's builder and lookup must +/// agree while continuing to read Variant values already written by Spark 4.x in UTF-16 order. +/// https://issues.apache.org/jira/browse/SPARK-58949 +/// https://github.com/apache/parquet-java/issues/3735 +fn reorder_variant_values( + value: &ArrayRef, + metadata: &ArrayRef, + parent_nulls: Option<&NullBuffer>, + order: VariantObjectKeyOrder, + allow_null_value: bool, +) -> DataFusionResult { + let value = value.as_any().downcast_ref::().unwrap(); + let metadata = metadata.as_any().downcast_ref::().unwrap(); + 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) { + if allow_null_value { + output.append_null(); + continue; + } + 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(|| -> DataFusionResult>> { + // Spark encodes empty object keys with equal metadata offsets, which Arrow 58.4's + // full validator rejects. Keep shallow parsing and all accesses inside this boundary. + // https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant/src/variant/metadata.rs#L307-L317 + // Upstream fix: https://github.com/apache/arrow-rs/pull/10352 + let metadata = VariantMetadata::new(metadata.value(index)); + let variant = Variant::new_with_metadata(metadata.clone(), value.value(index)); + if is_compatible_variant(&variant, order) { + return Ok(None); + } + let value = match order { + VariantObjectKeyOrder::ArrowUtf8 => { + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = ReadOnlyMetadataBuilder::new(&metadata); + ValueBuilder::try_append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + variant, + )?; + value_builder.into_inner() + } + VariantObjectKeyOrder::SparkUtf16 => spark_variant_bytes(&metadata, variant)?, + }; + Ok(Some(value)) + })) + .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..5f72435f642 --- /dev/null +++ b/native/core/src/parquet/cast_column/variant/tests.rs @@ -0,0 +1,935 @@ +// 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::{ + Decimal128Array, FixedSizeListArray, Int64Array, TimestampMillisecondArray, UInt16Array, + UInt32Array, UInt8Array, +}; +use arrow::datatypes::{Field, Fields}; +use parquet::variant::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}"); + let emoji = fields + .binary_search_by(|(name, _)| name.encode_utf16().cmp("😀".encode_utf16())) + .unwrap(); + assert_eq!(fields[emoji].1.as_int64(), Some(531)); + let private_use = fields + .binary_search_by(|(name, _)| name.encode_utf16().cmp("\u{e000}".encode_utf16())) + .unwrap(); + assert_eq!(fields[private_use].1.as_int64(), Some(30)); +} + +fn assert_spark_unicode_object(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))); +} + +fn normalize_typed_value(typed_value: ArrayRef, field_names: &[&str]) -> VariantArray { + let (metadata_bytes, _) = VariantBuilder::new() + .with_field_names(field_names.iter().copied()) + .finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + 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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + VariantArray::try_new(output.as_ref()).unwrap() +} + +#[test] +fn test_normalize_shredded_variant_widens_unsigned_values() { + let metadata_builder = VariantBuilder::new().with_field_names(["u8", "u16", "u32"]); + let (metadata_bytes, _) = metadata_builder.finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + + let fields = [ + ("u8", Arc::new(UInt8Array::from(vec![u8::MAX])) as ArrayRef), + ( + "u16", + Arc::new(UInt16Array::from(vec![u16::MAX])) as ArrayRef, + ), + ( + "u32", + Arc::new(UInt32Array::from(vec![u32::MAX])) as ArrayRef, + ), + ]; + let mut object_fields = Vec::with_capacity(fields.len()); + let mut object_columns = Vec::with_capacity(fields.len()); + for (name, value) in fields { + let shredded = StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + value.data_type().clone(), + false, + )]), + vec![value], + None, + ) + .unwrap(); + object_fields.push(Field::new(name, shredded.data_type().clone(), false)); + object_columns.push(Arc::new(shredded) as ArrayRef); + } + let typed_value: ArrayRef = + Arc::new(StructArray::try_new(object_fields.into(), object_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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let variant = output.value(0); + let Variant::Object(object) = variant else { + panic!("expected object") + }; + assert_eq!(object.get("u8"), Some(Variant::from(255_i16))); + assert_eq!(object.get("u16"), Some(Variant::from(65_535_i32))); + assert_eq!(object.get("u32"), Some(Variant::from(4_294_967_295_i64))); +} + +#[test] +fn test_normalize_shredded_variant_widens_millisecond_timestamps() { + let millis = 1_704_067_200_123_i64; + let ltz: ArrayRef = + Arc::new(TimestampMillisecondArray::from(vec![millis]).with_timezone("UTC")); + let ntz: ArrayRef = Arc::new(TimestampMillisecondArray::from(vec![millis])); + let shredded = |value: ArrayRef| -> ArrayRef { + Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + value.data_type().clone(), + false, + )]), + vec![value], + None, + ) + .unwrap(), + ) + }; + let ltz = shredded(ltz); + let ntz = shredded(ntz); + let object: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("ltz", ltz.data_type().clone(), false), + Field::new("ntz", ntz.data_type().clone(), false), + ]), + vec![ltz, ntz], + None, + ) + .unwrap(), + ); + let output = normalize_typed_value(object, &["ltz", "ntz"]); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + + let Some(Variant::TimestampMicros(ltz)) = object.get("ltz") else { + panic!("expected timestamp") + }; + assert_eq!(ltz.timestamp_micros(), millis * 1_000); + + let Some(Variant::TimestampNtzMicros(ntz)) = object.get("ntz") else { + panic!("expected timestamp_ntz") + }; + assert_eq!(ntz.and_utc().timestamp_micros(), millis * 1_000); +} + +#[test] +fn test_normalize_shredded_variant_converts_fixed_size_list() { + let elements: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![42, 43]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + FixedSizeListArray::try_new( + Arc::new(Field::new("element", elements.data_type().clone(), false)), + 2, + elements, + None, + ) + .unwrap(), + ); + + let output = normalize_typed_value(typed_value, &[]); + let Variant::List(list) = output.value(0) else { + panic!("expected list") + }; + assert_eq!( + list.iter() + .map(|value| value.as_int64()) + .collect::>(), + vec![Some(42), Some(43)] + ); +} + +#[test] +fn test_normalize_shredded_variant_compacts_spark_integer_widths() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![ + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + ])); + let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![ + 1, + i64::from(i8::MAX) + 1, + i64::from(i16::MAX) + 1, + i64::from(i32::MAX) + 1, + ])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", DataType::Int64, false), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!(output.value(0), Variant::Int8(1)); + assert_eq!(output.value(1), Variant::Int16(128)); + assert_eq!(output.value(2), Variant::Int32(32_768)); + assert_eq!(output.value(3), Variant::Int64(2_147_483_648)); +} + +#[test] +fn test_compact_spark_typed_variant_canonicalizes_nan() { + let Variant::Float(float) = + compact_spark_typed_variant(Variant::Float(f32::from_bits(0x7fc0_0001))) + else { + panic!("expected float") + }; + assert_eq!(float.to_bits(), 0x7fc0_0000); + + let Variant::Double(double) = + compact_spark_typed_variant(Variant::Double(f64::from_bits(0x7ff8_0000_0000_0001))) + else { + panic!("expected double") + }; + assert_eq!(double.to_bits(), 0x7ff8_0000_0000_0000); + + let (metadata, _) = VariantBuilder::new().finish(); + let metadata = VariantMetadata::new(&metadata); + let residual = f32::from_bits(0x7fc0_0001); + let bytes = spark_variant_bytes(&metadata, Variant::Float(residual)).unwrap(); + let Variant::Float(output) = Variant::new_with_metadata(metadata, &bytes) else { + panic!("expected residual float") + }; + assert_eq!(output.to_bits(), residual.to_bits()); +} + +#[test] +fn test_normalize_shredded_variant_rejects_missing_required_value() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![None])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", DataType::Int64, true), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + assert!(normalize_variant_array(&physical, &target_field).is_err()); +} + +#[test] +fn test_normalize_shredded_variant_uses_physical_metadata_order() { + let (metadata_bytes, _) = VariantBuilder::new().with_field_names(["a", "b"]).finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let mut fields = Vec::new(); + let mut columns = Vec::new(); + for (name, value) in [("b", 2_i64), ("a", 1_i64)] { + let child = StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![value]))], + None, + ) + .unwrap(); + fields.push(Field::new(name, child.data_type().clone(), false)); + columns.push(Arc::new(child) 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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_eq!( + output.column(1).as_binary::().value(0), + &[0x01, 2, 0, 1, 2, b'b', b'a'] + ); + + let mut expected = VariantBuilder::new().with_field_names(["b", "a"]); + let mut object = expected.new_object(); + object.insert("b", 2_i8); + object.insert("a", 1_i8); + object.finish(); + let (_, expected_value) = expected.finish(); + assert_eq!(output.column(0).as_binary::().value(0), expected_value); +} + +#[test] +fn test_normalize_shredded_variant_preserves_residual_scalar_width() { + let mut builder = VariantBuilder::new().with_field_names(["known", "residual"]); + let mut object = builder.new_object(); + object.insert("residual", 1_i64); + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let known: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![2]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), + vec![known], + 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, value, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + assert_eq!(object.get("known"), Some(Variant::Int8(2))); + assert_eq!(object.get("residual"), Some(Variant::Int64(1))); +} + +#[test] +fn test_normalize_shredded_variant_compacts_spark_decimal_width() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let typed_value: ArrayRef = Arc::new( + Decimal128Array::from(vec![123_i128]) + .with_precision_and_scale(38, 2) + .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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!( + output.value(0), + Variant::Decimal4(VariantDecimal4::try_new(123, 2).unwrap()) + ); +} + +#[test] +fn test_normalize_shredded_variant_uses_spark_object_key_order() { + let keys = unicode_object_keys(); + + let metadata_builder = VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); + let (metadata_bytes, _) = metadata_builder.finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + + let mut object_fields = Vec::with_capacity(keys.len()); + let mut object_columns = Vec::with_capacity(keys.len()); + for (index, key) in keys.iter().enumerate() { + let value = if key == "😀" { 531 } else { index as i64 }; + let field = Field::new("typed_value", DataType::Int64, false); + let column: ArrayRef = Arc::new(Int64Array::from(vec![value])); + let shredded_field = + StructArray::try_new(Fields::from(vec![field]), vec![column], None).unwrap(); + object_fields.push(Field::new(key, shredded_field.data_type().clone(), false)); + object_columns.push(Arc::new(shredded_field) as ArrayRef); + } + let typed_value: ArrayRef = + Arc::new(StructArray::try_new(object_fields.into(), object_columns, None).unwrap()); + let physical_fields = Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]); + let physical: ArrayRef = + Arc::new(StructArray::try_new(physical_fields, vec![metadata, typed_value], None).unwrap()); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false).with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_spark_unicode_object(output); +} + +#[test] +fn test_normalize_nested_shredded_variant_uses_spark_object_key_order() { + let keys = unicode_object_keys(); + let field_names = std::iter::once("nested") + .chain(keys.iter().map(String::as_str)) + .collect::>(); + let (metadata_bytes, _) = VariantBuilder::new().with_field_names(field_names).finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + + let mut nested_fields = Vec::with_capacity(keys.len()); + let mut nested_columns = Vec::with_capacity(keys.len()); + for (index, key) in keys.iter().enumerate() { + let value = if key == "😀" { 531 } else { index as i64 }; + let state = StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![value]))], + None, + ) + .unwrap(); + nested_fields.push(Field::new(key, state.data_type().clone(), false)); + nested_columns.push(Arc::new(state) as ArrayRef); + } + let nested_value: ArrayRef = + Arc::new(StructArray::try_new(nested_fields.into(), nested_columns, None).unwrap()); + let nested_state: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + nested_value.data_type().clone(), + false, + )]), + vec![nested_value], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "nested", + nested_state.data_type().clone(), + false, + )]), + vec![nested_state], + 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![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let Variant::Object(object) = output.value(0) else { + panic!("expected outer object") + }; + assert_spark_unicode_variant(object.get("nested").expect("nested field")); +} + +#[test] +fn test_normalize_partially_shredded_spark_object_key_order() { + 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_bytes, value_bytes) = builder.finish(); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let spark_value = reorder_variant_values( + &value, + &metadata, + None, + VariantObjectKeyOrder::SparkUtf16, + false, + ) + .unwrap(); + + let shredded_k00: 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_k00.data_type().clone(), + false, + )]), + vec![shredded_k00], + 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, spark_value, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + assert_spark_unicode_object(output.as_struct()); +} + +#[test] +fn test_normalize_unshredded_variant_uses_spark_object_key_order() { + 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_bytes, value_bytes) = builder.finish(); + + let Variant::Object(canonical) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + let canonical_fields = canonical.iter().collect::>(); + assert_eq!(canonical_fields[30].0, "\u{e000}"); + assert_eq!(canonical_fields[31].0, "😀"); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let physical: ArrayRef = + Arc::new(StructArray::try_new(physical_fields, vec![value, metadata], None).unwrap()); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false).with_extension_type(VariantType), + ); + + let first = normalize_variant_array(&physical, &target_field).unwrap(); + let first_value = first + .as_struct() + .column(0) + .as_binary::() + .value(0) + .to_vec(); + assert_spark_unicode_object(first.as_struct()); + + // Spark-produced already-unshredded input is UTF-16 ordered. Normalizing it again must + // remain valid without Arrow's UTF-8-order full validation. + let second = normalize_variant_array(&first, &target_field).unwrap(); + assert_spark_unicode_object(second.as_struct()); + assert_eq!( + second.as_struct().column(0).as_binary::().value(0), + first_value + ); +} + +#[test] +fn test_normalize_spark_ordered_variant_preserves_value_bytes() { + let mut builder = VariantBuilder::new(); + let mut object = builder.new_object(); + object.insert("b", 1_i64); + object.insert("a", 2_i64); + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + + let Variant::Object(object) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + assert_eq!( + object.iter().map(|(name, _)| name).collect::>(), + vec!["a", "b"] + ); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let physical: ArrayRef = + Arc::new(StructArray::try_new(physical_fields, vec![value, metadata], None).unwrap()); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false).with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + assert_eq!( + output.as_struct().column(0).as_binary::().value(0), + value_bytes + ); +} + +#[test] +fn test_normalize_variant_preserves_empty_object_keys() { + let mut builder = VariantBuilder::new(); + let mut object = builder.new_object(); + object.insert("", 1_i64); + let mut nested = object.new_object("nested"); + nested.insert("", 2_i64); + nested.finish(); + object.finish(); + let (mut metadata_bytes, value_bytes) = builder.finish(); + + // Spark leaves the metadata dictionary unsorted. Equal offsets encode the empty key. + metadata_bytes[0] &= !0x10; + assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]), + vec![value, metadata], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_eq!(output.column(0).as_binary::().value(0), value_bytes); + assert_eq!(output.column(1).as_binary::().value(0), metadata_bytes); + + let Variant::Object(object) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + assert_eq!(object.get(""), Some(Variant::from(1_i64))); + let Variant::Object(nested) = object.get("nested").unwrap() else { + panic!("expected nested object") + }; + assert_eq!(nested.get(""), Some(Variant::from(2_i64))); +} + +#[test] +fn test_normalize_partially_shredded_nested_unicode_and_empty_keys() { + let keys = unicode_object_keys(); + let mut builder = VariantBuilder::new().with_field_names(["known"]); + let mut object = builder.new_object(); + object.insert("", 1_i64); + let mut nested = object.new_object("nested"); + for (index, key) in keys.iter().enumerate() { + nested.insert(key, if key == "😀" { 531_i64 } else { index as i64 }); + } + nested.finish(); + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let known: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![3]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), + vec![known], + 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, value, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + assert_eq!(object.get("").unwrap().as_int64(), Some(1)); + assert_eq!(object.get("known").unwrap().as_int64(), Some(3)); + let Variant::Object(nested) = object.get("nested").unwrap() else { + panic!("expected nested object") + }; + let fields = nested.iter().collect::>(); + assert_eq!(fields.len(), 32); + assert_eq!(fields[30].0, "😀"); + assert_eq!(fields[31].0, "\u{e000}"); + assert_eq!(nested.get("😀").unwrap().as_int64(), Some(531)); +} + +#[test] +fn test_normalize_variant_skips_empty_children_of_null_parent() { + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&b""[..])])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&b""[..])])); + let physical_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + physical_fields, + vec![value, metadata], + Some(NullBuffer::from(vec![false])), + ) + .unwrap(), + ); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), true).with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + assert!(output.is_null(0)); + assert!(output.as_struct().column(0).is_null(0)); +} From ba37a688e5aaa1d9e4a17c55d42ea0211486f7a3 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Thu, 27 Aug 2026 00:45:03 +0800 Subject: [PATCH 10/14] fix: preserve strict Variant Parquet validation --- .../serde/operator/CometNativeScan.scala | 11 +++++++ .../comet/parquet/ParquetReadSuite.scala | 33 ++++++++++++++++++- 2 files changed, 43 insertions(+), 1 deletion(-) diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala index 0a5197762f5..02a4b21d185 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala @@ -122,6 +122,17 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS withFallbackReason(scanExec, unsupportedDefaultReason) } + // Spark's strict mode validates the legacy two-field Variant layout and reports SPARK-47546 + // errors itself: https://issues.apache.org/jira/browse/SPARK-47546 + if (scanExec.requiredSchema.fields.exists(field => isVariantType(field.dataType)) && + !SQLConf.get + .getConfString("spark.sql.variant.allowReadingShredded", "true") + .toBoolean) { + withFallbackReason( + scanExec, + "Full native scan disabled because Spark's strict unshredded Variant reader is enabled") + } + // the scan is supported if no fallback reasons were added to the node !hasFallbackReason(scanExec) } diff --git a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala index d8b9f367477..93a25f87e98 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala @@ -38,7 +38,7 @@ import org.apache.parquet.hadoop.example.ExampleParquetWriter import org.apache.parquet.io.api.Binary import org.apache.parquet.schema.MessageTypeParser import org.apache.spark.SparkException -import org.apache.spark.sql.{CometTestBase, DataFrame, Row} +import org.apache.spark.sql.{AnalysisException, CometTestBase, DataFrame, Row} import org.apache.spark.sql.catalyst.util.DateTimeUtils import org.apache.spark.sql.comet.{CometColumnarToRowExec, CometNativeColumnarToRowExec, CometNativeScanExec, CometScanExec} import org.apache.spark.sql.comet.util.Utils @@ -202,6 +202,37 @@ abstract class ParquetReadSuite extends CometTestBase { } } + test("strict unshredded Variant validation falls back to Spark") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempDir { dir => + val path = new File(dir, "data").getCanonicalPath + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + sql("SELECT named_struct('value', X'08') AS v").write + .parquet(path) + } + + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "false", + "spark.sql.variant.pushVariantIntoScan" -> "false", + "spark.sql.variant.inferShreddingSchema" -> "false") { + val result = spark.read + .schema("v VARIANT") + .parquet(path) + .selectExpr("to_json(v)") + assert(collect(result.queryExecution.executedPlan) { case _: CometNativeScanExec => + true + }.isEmpty) + + val error = intercept[SparkException](result.collect()).getCause + assert(error.isInstanceOf[AnalysisException]) + assert( + error.asInstanceOf[AnalysisException].getErrorClass == + "INVALID_VARIANT_FROM_PARQUET.WRONG_NUM_FIELDS") + } + } + } + test("native scan preserves Variant existence default pairing") { assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") From 1e832c153428120d4b24fe08779c7740e5eb4462 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Thu, 27 Aug 2026 02:35:25 +0800 Subject: [PATCH 11/14] review --- .../serde/operator/CometNativeScan.scala | 18 ++- .../comet/parquet/ParquetReadSuite.scala | 132 +++++++++++++++--- 2 files changed, 123 insertions(+), 27 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala index 02a4b21d185..d3225d9c333 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala @@ -122,17 +122,27 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS withFallbackReason(scanExec, unsupportedDefaultReason) } + val hasVariant = + scanExec.requiredSchema.fields.exists(field => isVariantType(field.dataType)) + // Spark's strict mode validates the legacy two-field Variant layout and reports SPARK-47546 // errors itself: https://issues.apache.org/jira/browse/SPARK-47546 - if (scanExec.requiredSchema.fields.exists(field => isVariantType(field.dataType)) && - !SQLConf.get - .getConfString("spark.sql.variant.allowReadingShredded", "true") - .toBoolean) { + if (hasVariant && + !SQLConf.get.getConfString("spark.sql.variant.allowReadingShredded").toBoolean) { withFallbackReason( scanExec, "Full native scan disabled because Spark's strict unshredded Variant reader is enabled") } + // Spark interprets TIMESTAMP(NANOS) leaves as raw longs in this legacy mode. The logical + // Variant schema does not expose whether a file contains such a shredded child, so preserve + // Spark semantics by falling back before reading any requested Variant value. + if (hasVariant && SQLConf.get.legacyParquetNanosAsLong) { + withFallbackReason( + scanExec, + "Full native scan disabled because spark.sql.legacy.parquet.nanosAsLong is enabled") + } + // the scan is supported if no fallback reasons were added to the node !hasFallbackReason(scanExec) } diff --git a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala index 93a25f87e98..f72936650ed 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala @@ -49,7 +49,7 @@ import org.apache.spark.sql.types._ import com.google.common.primitives.UnsignedLong -import org.apache.comet.{CometConf, CometSparkSessionExtensions} +import org.apache.comet.{CometConf, CometSparkSessionExtensions, ExtendedExplainInfo} import org.apache.comet.vector.CometStructVector abstract class ParquetReadSuite extends CometTestBase { @@ -202,33 +202,119 @@ abstract class ParquetReadSuite extends CometTestBase { } } - test("strict unshredded Variant validation falls back to Spark") { + test("native scan honors Spark's shredded Variant reader configuration") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTable("variant_reader_mode") { + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.writeShredding.enabled" -> "true", + "spark.sql.variant.forceShreddingSchemaForTest" -> "a BIGINT") { + sql("CREATE TABLE variant_reader_mode(v VARIANT) USING parquet") + sql("INSERT INTO variant_reader_mode VALUES (parse_json('{\"a\":1}'))") + } + + val key = "spark.sql.variant.allowReadingShredded" + val conf = SQLConf.get + val original = conf.getAllConfs.get(key) + try { + Seq(None, Some(false), Some(true)).foreach { configured => + configured match { + case Some(value) => conf.setConfString(key, value.toString) + case None => conf.unsetConf(key) + } + val allowShredded = conf.getConfString(key).toBoolean + + withSQLConf( + "spark.sql.variant.pushVariantIntoScan" -> "false", + "spark.sql.variant.inferShreddingSchema" -> "false") { + val result = sql("SELECT v FROM variant_reader_mode") + val nativeScans = collect(result.queryExecution.executedPlan) { + case _: CometNativeScanExec => true + } + + if (allowShredded) { + assert(normalizedVariantRows(result, 0) == Seq(Seq("{\"a\":1}"))) + assert(nativeScans.size == 1) + } else { + assert(nativeScans.isEmpty) + val error = intercept[SparkException](result.collect()).getCause + assert(error.isInstanceOf[AnalysisException]) + assert( + error.asInstanceOf[AnalysisException].getErrorClass == + "INVALID_VARIANT_FROM_PARQUET.WRONG_NUM_FIELDS") + } + } + } + } finally { + original match { + case Some(value) => conf.setConfString(key, value) + case None => conf.unsetConf(key) + } + } + } + } + + test("nanosAsLong Variant children fall back to Spark") { assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") withTempDir { dir => - val path = new File(dir, "data").getCanonicalPath - withSQLConf(CometConf.COMET_ENABLED.key -> "false") { - sql("SELECT named_struct('value', X'08') AS v").write - .parquet(path) + val rawNanos = Seq(1704067200123000000L, 1704067200123456789L) + val variant = sql("SELECT parse_json('0')").head().get(0) + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + + Seq(true, false).foreach { adjustedToUtc => + val path = new Path(dir.toURI.toString, s"nanos-$adjustedToUtc.parquet") + val parquetSchema = MessageTypeParser.parseMessageType(s"""message root { + | optional group v { + | optional binary value; + | required binary metadata; + | optional int64 typed_value (TIMESTAMP(NANOS,$adjustedToUtc)); + | } + |} + |""".stripMargin) + val writer = ExampleParquetWriter + .builder(path) + .withType(parquetSchema) + .withConf(spark.sessionState.newHadoopConf()) + .build() + + try { + rawNanos.foreach { value => + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(1, Binary.fromConstantByteArray(metadata)) + group.add(2, value) + writer.write(row) + } + } finally { + writer.close() + } } - withSQLConf( - "spark.sql.variant.allowReadingShredded" -> "false", - "spark.sql.variant.pushVariantIntoScan" -> "false", - "spark.sql.variant.inferShreddingSchema" -> "false") { - val result = spark.read - .schema("v VARIANT") - .parquet(path) - .selectExpr("to_json(v)") - assert(collect(result.queryExecution.executedPlan) { case _: CometNativeScanExec => - true - }.isEmpty) - - val error = intercept[SparkException](result.collect()).getCause - assert(error.isInstanceOf[AnalysisException]) - assert( - error.asInstanceOf[AnalysisException].getErrorClass == - "INVALID_VARIANT_FROM_PARQUET.WRONG_NUM_FIELDS") + withTable("variant_nanos_as_long") { + sql(s"""CREATE TABLE variant_nanos_as_long(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + withSQLConf( + SQLConf.PARQUET_VECTORIZED_READER_ENABLED.key -> "true", + "spark.sql.legacy.parquet.nanosAsLong" -> "true", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false", + "spark.sql.variant.inferShreddingSchema" -> "false") { + val result = sql("SELECT v FROM variant_nanos_as_long") + assert( + normalizedVariantRows(result, 0) == + rawNanos.flatMap(value => Seq.fill(2)(Seq(value.toString)))) + + val plan = result.queryExecution.executedPlan + assert(collect(plan) { case _: CometNativeScanExec => true }.isEmpty) + val fallbackReasons = new ExtendedExplainInfo().getFallbackReasons(plan) + assert( + fallbackReasons.exists(_.contains("spark.sql.legacy.parquet.nanosAsLong")), + s"Expected nanosAsLong fallback, found: ${fallbackReasons.mkString(", ")}") + } } } } From a784bea3911d3ee6a96253fc1b025739d01f1df2 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Thu, 27 Aug 2026 03:23:33 +0800 Subject: [PATCH 12/14] fix: normalize remaining Variant scan encodings --- .../core/src/parquet/cast_column/variant.rs | 708 ++++++++++++------ .../src/parquet/cast_column/variant/tests.rs | 293 +++++++- .../sql-tests/expressions/misc/variant.sql | 10 +- .../comet/parquet/ParquetReadSuite.scala | 78 +- 4 files changed, 833 insertions(+), 256 deletions(-) diff --git a/native/core/src/parquet/cast_column/variant.rs b/native/core/src/parquet/cast_column/variant.rs index 994de41de42..d2e098664a8 100644 --- a/native/core/src/parquet/cast_column/variant.rs +++ b/native/core/src/parquet/cast_column/variant.rs @@ -15,7 +15,10 @@ // specific language governing permissions and limitations // under the License. use arrow::{ - array::{Array, ArrayRef, AsArray, BinaryArray, BinaryBuilder, ListLikeArray, StructArray}, + array::{ + make_array, Array, ArrayRef, AsArray, BinaryArray, BinaryBuilder, ListLikeArray, + StructArray, + }, buffer::NullBuffer, compute::{cast, cast_with_options}, datatypes::{DataType, FieldRef, TimeUnit}, @@ -56,8 +59,7 @@ pub(super) fn normalize_variant_array( )); } - let array = decode_variant_metadata_dictionary(array)?; - let array = normalize_variant_typed_value(&array)?; + let array = normalize_variant_storage(array)?; let variant = VariantArray::try_new(array.as_ref())?; let was_shredded = variant.typed_value_field().is_some(); let unshredded = unshred_variant_for_spark(&variant)?; @@ -87,17 +89,27 @@ pub(super) fn normalize_variant_array( } fn unshred_variant_for_spark(variant: &VariantArray) -> DataFusionResult { - let first = - prepare_variant_for_unshredding(variant).and_then(|array| Ok(unshred_variant(&array)?)); - let first_error = match first { + let first_error = match unshred_variant(variant) { Ok(array) => return Ok(array), - Err(error) => error, + Err(error) => DataFusionError::from(error), }; - let Some(variant) = canonicalize_spark_empty_key_metadata(variant)? else { + + if let Ok(prepared) = prepare_variant_for_unshredding(variant, None) { + if let Ok(array) = unshred_variant(&prepared) { + return Ok(array); + } + } + let Some(metadata) = canonicalize_spark_empty_key_metadata(variant)? else { + return Err(first_error); + }; + let Ok(prepared) = prepare_variant_for_unshredding(variant, Some(metadata.as_binary::())) + else { return Err(first_error); }; - let variant = prepare_variant_for_unshredding(&variant)?; - Ok(unshred_variant(&variant)?) + match unshred_variant(&prepared) { + Ok(array) => Ok(array), + Err(_) => Err(first_error), + } } fn normalize_variant_type(data_type: &DataType) -> Option { @@ -107,9 +119,15 @@ fn normalize_variant_type(data_type: &DataType) -> Option { } match data_type { + DataType::Dictionary(_, value_type) => { + Some(normalize_variant_type(value_type).unwrap_or_else(|| value_type.as_ref().clone())) + } DataType::UInt8 => Some(DataType::Int16), DataType::UInt16 => Some(DataType::Int32), DataType::UInt32 => Some(DataType::Int64), + // Spark reads Parquet UINT_64 as Decimal(20, 0). This is lossless for the full range and + // lets the existing Spark-compatible rebuild choose the Variant decimal width per value. + DataType::UInt64 => Some(DataType::Decimal128(20, 0)), DataType::Timestamp(TimeUnit::Millisecond, timezone) => { Some(DataType::Timestamp(TimeUnit::Microsecond, timezone.clone())) } @@ -140,191 +158,382 @@ fn normalize_variant_type(data_type: &DataType) -> Option { /// Normalize Arrow types that Spark's Parquet reader accepts but `VariantArray` 58.4 rejects. /// Parquet restores unsigned integers to Arrow unsigned arrays and millisecond timestamps at their -/// annotated unit; embedded Arrow schemas may also restore fixed-size lists. Spark widens the -/// integers and timestamps and treats fixed-size lists as ordinary Variant arrays. -/// arrow-rs #10416/#10417 would move this widening into `VariantArray`/`unshred_variant`; remove -/// the unsigned arms after that ships and Comet upgrades: +/// annotated unit; embedded Arrow schemas may also restore dictionary arrays and fixed-size lists. +/// Spark decodes dictionaries, widens integers and timestamps, and treats fixed-size lists as +/// ordinary Variant arrays. +/// arrow-rs #10416/#10417 would move the UInt8/16/32 widening into +/// `VariantArray`/`unshred_variant`; remove those three arms after that ships and Comet upgrades: /// https://github.com/apache/arrow-rs/issues/10416 /// https://github.com/apache/arrow-rs/pull/10417 -/// Arrow #50622/#50810 instead proposes removing unsigned `typed_value` mappings because the -/// Parquet Variant shredding table permits only signed integer fields. Until upstream resolves -/// that choice, keep this compatibility path for unsigned files Spark already reads: +/// Arrow #50622/#50810 instead proposes removing unsigned Parquet `typed_value` mappings; keep +/// compatibility until upstream settles that schema boundary: /// https://github.com/apache/arrow/issues/50622 /// https://github.com/apache/arrow/pull/50810 -fn normalize_variant_typed_value(array: &ArrayRef) -> DataFusionResult { - let Some(struct_array) = array.as_any().downcast_ref::() else { +/// Spark maps Parquet UINT_64 to Decimal(20,0), so its UInt64 arm remains Spark-specific: +/// https://github.com/apache/spark/blob/v4.0.4/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/parquet/ParquetSchemaConverter.scala#L281-L296 +/// arrow-rs #10810 covers encoded canonical metadata only. Embedded Arrow schemas can also restore +/// dictionary `value` and `typed_value` children, so keep decoding those for Spark compatibility: +/// https://github.com/apache/arrow-rs/pull/10810 +fn normalize_variant_storage(array: &ArrayRef) -> DataFusionResult { + let Some(data_type) = normalize_variant_type(array.data_type()) else { return Ok(Arc::clone(array)); }; - let Some((typed_value_index, typed_value_field)) = struct_array - .fields() - .iter() - .enumerate() - .find(|(_, field)| field.name() == "typed_value") - else { - return Ok(Arc::clone(array)); - }; - let Some(data_type) = normalize_variant_type(typed_value_field.data_type()) else { - return Ok(Arc::clone(array)); - }; - - let mut fields = struct_array.fields().iter().cloned().collect::>(); - fields[typed_value_index] = Arc::new( - typed_value_field - .as_ref() - .clone() - .with_data_type(data_type.clone()), - ); - let mut columns = struct_array.columns().to_vec(); - columns[typed_value_index] = cast_with_options( - columns[typed_value_index].as_ref(), + Ok(cast_with_options( + array.as_ref(), &data_type, &DEFAULT_CAST_OPTIONS, - )?; - Ok(Arc::new(StructArray::try_new( - fields.into(), - columns, - struct_array.nulls().cloned(), - )?)) + )?) } -/// Arrow's unshredder fully validates any residual `value` in a partially shredded object. Spark -/// writes object keys in Java UTF-16 order, so put that residual value in Arrow UTF-8 order only -/// while it passes through the upstream unshredder. -fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult { - let (Some(value), Some(_)) = (variant.value_field(), variant.typed_value_field()) else { - return Ok(variant.clone()); - }; +/// Arrow fully validates residual `value` fields while unshredding. Spark versions before +/// SPARK-58949 wrote their object keys in Java UTF-16 order, so rewrite every reachable legacy +/// residual to Arrow's UTF-8 order before calling the upstream unshredder. Keep this input-side +/// compatibility for historical files after #5474 removes the output-side UTF-16 rewrite: +/// https://github.com/apache/datafusion-comet/issues/5474 +/// `metadata_rows` carries the root metadata row through nested lists. +fn rewrite_shredding_state( + state: &StructArray, + source_metadata: &BinaryArray, + target_metadata: &BinaryArray, + metadata_rows: &[Option], + remap_metadata: bool, +) -> 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; - let value = cast(value.as_ref(), &DataType::Binary)?; - let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; - let value = reorder_variant_values( - &value, - &metadata, - variant.inner().nulls(), - VariantObjectKeyOrder::ArrowUtf8, - true, - )?; + if let Some(index) = fields.iter().position(|field| field.name() == "value") { + let (value, value_changed) = rewrite_residual_values( + &columns[index], + source_metadata, + target_metadata, + &active_rows, + remap_metadata, + )?; + if value_changed { + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(value.data_type().clone()), + ); + columns[index] = value; + changed = true; + } + } - let value_index = variant - .inner() - .fields() + if let Some(index) = fields .iter() - .position(|field| field.name() == "value") - .unwrap(); - let mut fields = variant.inner().fields().iter().cloned().collect::>(); - fields[value_index] = Arc::new( - fields[value_index] - .as_ref() - .clone() - .with_data_type(DataType::Binary), - ); - let mut columns = variant.inner().columns().to_vec(); - columns[value_index] = value; - let array = StructArray::try_new(fields.into(), columns, variant.inner().nulls().cloned())?; - Ok(VariantArray::try_new(&array)?) -} + .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], + source_metadata, + target_metadata, + &typed_rows, + remap_metadata, + )?; + 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; + } + } -/// Arrow 58.4 rejects Spark metadata dictionaries containing an empty object key because their -/// unsorted offsets can be equal. Rebuild only those rows before any residual value is traversed, -/// sorting the dictionary and remapping the residual value's field IDs at the same time. -/// https://github.com/apache/arrow-rs/pull/10352 -fn canonicalize_spark_empty_key_metadata( - variant: &VariantArray, -) -> DataFusionResult> { - type Replacement = Option<(Vec, Option>)>; + if !changed { + return Ok((Arc::new(state.clone()), false)); + } + Ok(( + Arc::new(StructArray::try_new( + fields.into(), + columns, + state.nulls().cloned(), + )?), + true, + )) +} - let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; - let metadata = metadata.as_binary::(); - let value = variant - .value_field() - .map(|value| cast(value.as_ref(), &DataType::Binary)) - .transpose()?; - let value = value.as_ref().map(|value| value.as_binary::()); - let mut replacements = Vec::with_capacity(variant.len()); +fn rewrite_residual_values( + value: &ArrayRef, + source_metadata: &BinaryArray, + target_metadata: &BinaryArray, + metadata_rows: &[Option], + remap_metadata: bool, +) -> 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 in 0..variant.len() { - if variant.inner().is_null(index) || metadata.is_null(index) { - replacements.push(None); + for (index, metadata_row) in metadata_rows.iter().enumerate() { + if binary.is_null(index) { + output.append_null(); continue; } - let metadata_bytes = metadata.value(index); - if VariantMetadata::try_new(metadata_bytes).is_ok() { - replacements.push(None); + let Some(metadata_row) = metadata_row else { + output.append_value(binary.value(index)); continue; + }; + if source_metadata.is_null(*metadata_row) || target_metadata.is_null(*metadata_row) { + return Err(DataFusionError::Execution(format!( + "Variant metadata is null at row {metadata_row}" + ))); } - let replacement = catch_unwind(AssertUnwindSafe(|| -> Result { - let old_metadata = VariantMetadata::new(metadata_bytes); - let mut names = old_metadata - .iter_try() - .map(|name| name.map(str::to_string)) - .collect::, _>>()?; - if !names.iter().any(String::is_empty) { - return Ok(None); - } - names.sort_unstable(); - if names.windows(2).any(|names| names[0] == names[1]) { - return Ok(None); - } - - let mut builder = - VariantBuilder::new().with_field_names(names.iter().map(String::as_str)); - match value { - Some(value) if !value.is_null(index) => { - builder - .append_value(Variant::new_with_metadata(old_metadata, value.value(index))); - let (metadata, value) = builder.finish(); - Ok(Some((metadata, Some(value)))) + let rebuilt = catch_unwind(AssertUnwindSafe( + || -> Result>, ArrowError> { + let source = VariantMetadata::new(source_metadata.value(*metadata_row)); + let target = VariantMetadata::new(target_metadata.value(*metadata_row)); + let variant = Variant::new_with_metadata(source, binary.value(index)); + let arrow_ordered = + is_compatible_variant(&variant, VariantObjectKeyOrder::ArrowUtf8); + if !remap_metadata && arrow_ordered { + return Ok(None); } - _ => Ok(Some((builder.finish().0, None))), - } - })) + if !arrow_ordered + && !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(arrow_variant_bytes(&target, variant)?)) + }, + )) .map_err(|_| { - DataFusionError::Execution(format!( - "Invalid Variant metadata with an empty key at row {index}" - )) + DataFusionError::Execution(format!("Invalid Variant residual at row {metadata_row}")) })??; - changed |= replacement.is_some(); - replacements.push(replacement); + changed |= rebuilt.is_some(); + output.append_value(rebuilt.as_deref().unwrap_or_else(|| binary.value(index))); } - if !changed { - return Ok(None); + if changed { + Ok((Arc::new(output.finish()), true)) + } else { + Ok((Arc::clone(value), false)) } +} - let mut metadata_builder = BinaryBuilder::new(); - let mut value_builder = value.map(|_| BinaryBuilder::new()); - for (index, replacement) in replacements.iter().enumerate() { - match replacement { - Some((metadata, value)) => { - metadata_builder.append_value(metadata); - if let Some(builder) = &mut value_builder { - match value { - Some(value) => builder.append_value(value), - None => builder.append_null(), - } +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), } - None => { - if metadata.is_null(index) { - metadata_builder.append_null(); - } else { - metadata_builder.append_value(metadata.value(index)); - } - if let (Some(value), Some(builder)) = (value, &mut value_builder) { - if value.is_null(index) { - builder.append_null(); - } else { - builder.append_value(value.value(index)); - } + } + } + Ok(child_rows) +} + +fn rewrite_list_typed_value( + array: &ArrayRef, + list: &L, + source_metadata: &BinaryArray, + target_metadata: &BinaryArray, + metadata_rows: &[Option], + remap_metadata: bool, +) -> 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, + source_metadata, + target_metadata, + &child_rows, + remap_metadata, + )?; + 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, + source_metadata: &BinaryArray, + target_metadata: &BinaryArray, + metadata_rows: &[Option], + remap_metadata: bool, +) -> 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, + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + )?; + 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::(), + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + ), + DataType::LargeList(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list::(), + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + ), + DataType::ListView(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list_view::(), + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + ), + DataType::LargeListView(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list_view::(), + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + ), + _ => Ok((Arc::clone(typed_value), false)), } +} - let mut fields = variant.inner().fields().iter().cloned().collect::>(); - let mut columns = variant.inner().columns().to_vec(); +fn prepare_variant_for_unshredding( + variant: &VariantArray, + target_metadata: Option<&BinaryArray>, +) -> DataFusionResult { + if variant.typed_value_field().is_none() { + return Ok(variant.clone()); + } + let source_metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let source_metadata = source_metadata.as_binary::(); + let remap_metadata = target_metadata.is_some(); + let target_metadata = target_metadata.unwrap_or(source_metadata); + let metadata_rows = (0..variant.len()) + .map(|index| variant.inner().is_valid(index).then_some(index)) + .collect::>(); + let (array, _) = rewrite_shredding_state( + variant.inner(), + source_metadata, + target_metadata, + &metadata_rows, + remap_metadata, + )?; + let array = array.as_struct(); + let mut fields = array.fields().iter().cloned().collect::>(); + let mut columns = array.columns().to_vec(); let metadata_index = fields .iter() .position(|field| field.name() == "metadata") @@ -335,60 +544,80 @@ fn canonicalize_spark_empty_key_metadata( .clone() .with_data_type(DataType::Binary), ); - columns[metadata_index] = Arc::new(metadata_builder.finish()); - if let Some(mut value_builder) = value_builder { - let value_index = fields - .iter() - .position(|field| field.name() == "value") - .unwrap(); - fields[value_index] = Arc::new( - fields[value_index] - .as_ref() - .clone() - .with_data_type(DataType::Binary), - ); - columns[value_index] = Arc::new(value_builder.finish()); - } - let array = StructArray::try_new(fields.into(), columns, variant.inner().nulls().cloned())?; - Ok(Some(VariantArray::try_new(&array)?)) + columns[metadata_index] = Arc::new(target_metadata.clone()); + let array = StructArray::try_new(fields.into(), columns, array.nulls().cloned())?; + Ok(VariantArray::try_new(&array)?) } -/// Arrow-rs parquet-variant-compute allows dictionary-encoded metadata in its contract, but 58.4's -/// `VariantArray::try_new` validates only Binary, LargeBinary, and BinaryView. Decode just that -/// child and keep the physical struct otherwise unchanged. -/// https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant-compute/src/variant_array.rs#L276-L310 -/// Upstream issue: https://github.com/apache/arrow-rs/issues/10802 -fn decode_variant_metadata_dictionary(array: &ArrayRef) -> DataFusionResult { - let Some(struct_array) = array.as_any().downcast_ref::() else { - return Ok(Arc::clone(array)); - }; - let Some((metadata_index, metadata_field)) = struct_array - .fields() - .iter() - .enumerate() - .find(|(_, field)| field.name() == "metadata") - else { - return Ok(Arc::clone(array)); - }; - let DataType::Dictionary(_, value_type) = metadata_field.data_type() else { - return Ok(Arc::clone(array)); - }; +/// Arrow 58.4 rejects Spark metadata dictionaries containing an empty object key when Spark's +/// insertion-order encoding leaves equal offsets. On an upstream validation failure, canonicalize +/// only that exact case; the recursive rewrite above remaps every residual to the new field IDs. +/// https://github.com/apache/arrow-rs/pull/10352 +fn canonicalize_spark_empty_key_metadata( + variant: &VariantArray, +) -> DataFusionResult> { + let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let metadata = metadata.as_binary::(); + let mut output = BinaryBuilder::new(); + let mut changed = false; - let decoded = cast(struct_array.column(metadata_index).as_ref(), value_type)?; - let mut fields = struct_array.fields().iter().cloned().collect::>(); - fields[metadata_index] = Arc::new( - metadata_field - .as_ref() - .clone() - .with_data_type(decoded.data_type().clone()), - ); - let mut columns = struct_array.columns().to_vec(); - columns[metadata_index] = decoded; - Ok(Arc::new(StructArray::try_new( - fields.into(), - columns, - struct_array.nulls().cloned(), - )?)) + for index in 0..variant.len() { + if variant.inner().is_null(index) || metadata.is_null(index) { + if metadata.is_null(index) { + output.append_null(); + } else { + output.append_value(metadata.value(index)); + } + continue; + } + let metadata_bytes = metadata.value(index); + if VariantMetadata::try_new(metadata_bytes).is_ok() { + output.append_value(metadata_bytes); + continue; + } + + let replacement = catch_unwind(AssertUnwindSafe( + || -> Result>, ArrowError> { + let old_metadata = VariantMetadata::new(metadata_bytes); + let mut names = old_metadata + .iter_try() + .map(|name| name.map(str::to_string)) + .collect::, _>>()?; + if !names.iter().any(String::is_empty) { + return Ok(None); + } + let mut source = + WritableMetadataBuilder::from_iter(names.iter().map(String::as_str)); + source.finish(); + let mut source = source.into_inner(); + source[0] &= !0x10; + if source != metadata_bytes { + return Ok(None); + } + names.sort_unstable(); + if names.windows(2).any(|names| names[0] == names[1]) { + return Ok(None); + } + let metadata = VariantBuilder::new() + .with_field_names(names.iter().map(String::as_str)) + .finish() + .0; + VariantMetadata::try_new(&metadata)?; + Ok(Some(metadata)) + }, + )); + let Ok(Ok(Some(replacement))) = replacement else { + return Ok(None); + }; + output.append_value(replacement); + changed = true; + } + + if changed { + Ok(Some(Arc::new(output.finish()))) + } else { + Ok(None) + } } /// Supplies sort-only field names whose Rust ordering matches Java `String.compareTo` ordering. @@ -518,6 +747,48 @@ fn compact_spark_typed_variant<'m, 'v>(variant: Variant<'m, 'v>) -> Variant<'m, } } +fn arrow_variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + match variant { + Variant::Object(object) => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + true, + ); + for (name, value) in object.iter() { + let value = arrow_variant_bytes(metadata, value)?; + 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), + true, + ); + for value in list.iter() { + let value = arrow_variant_bytes(metadata, value)?; + 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()) +} + /// Re-encode a residual Variant against `metadata`. Scalar widths are intentionally preserved, /// matching Spark's `VariantBuilder.appendVariant` behavior. fn spark_variant_bytes( @@ -1049,10 +1320,11 @@ fn rebuild_shredded_variant_for_spark( /// Reorder object keys for either Arrow's UTF-8 order or Spark's Java UTF-16 order. Preserve /// already-compatible values byte-for-byte and retain the original metadata dictionary. -/// SPARK-58949 tracks this mismatch and legacy compatibility. The metadata dictionary's sorted -/// flag affects dictionary lookup, not object-entry ordering; Spark's builder and lookup must -/// agree while continuing to read Variant values already written by Spark 4.x in UTF-16 order. +/// SPARK-58949 fixes this mismatch and retains legacy lookup compatibility. Once every Spark +/// profile supported by Comet contains that fix, #5474 can remove this output-side UTF-16 rewrite; +/// the input-side residual rewrite above remains necessary for historical Spark-written files. /// https://issues.apache.org/jira/browse/SPARK-58949 +/// https://github.com/apache/datafusion-comet/issues/5474 /// https://github.com/apache/parquet-java/issues/3735 fn reorder_variant_values( value: &ArrayRef, @@ -1096,15 +1368,7 @@ fn reorder_variant_values( return Ok(None); } let value = match order { - VariantObjectKeyOrder::ArrowUtf8 => { - let mut value_builder = ValueBuilder::new(); - let mut metadata_builder = ReadOnlyMetadataBuilder::new(&metadata); - ValueBuilder::try_append_variant( - ParentState::variant(&mut value_builder, &mut metadata_builder), - variant, - )?; - value_builder.into_inner() - } + VariantObjectKeyOrder::ArrowUtf8 => arrow_variant_bytes(&metadata, variant)?, VariantObjectKeyOrder::SparkUtf16 => spark_variant_bytes(&metadata, variant)?, }; Ok(Some(value)) diff --git a/native/core/src/parquet/cast_column/variant/tests.rs b/native/core/src/parquet/cast_column/variant/tests.rs index 5f72435f642..2756177384e 100644 --- a/native/core/src/parquet/cast_column/variant/tests.rs +++ b/native/core/src/parquet/cast_column/variant/tests.rs @@ -16,11 +16,11 @@ // under the License. use super::*; use arrow::array::{ - Decimal128Array, FixedSizeListArray, Int64Array, TimestampMillisecondArray, UInt16Array, - UInt32Array, UInt8Array, + Decimal128Array, DictionaryArray, FixedSizeListArray, Int32Array, Int64Array, + TimestampMillisecondArray, UInt16Array, UInt32Array, UInt64Array, UInt8Array, }; -use arrow::datatypes::{Field, Fields}; -use parquet::variant::VariantType; +use arrow::datatypes::{Field, Fields, Int32Type}; +use parquet::variant::{VariantDecimal16, VariantType}; fn unicode_object_keys() -> Vec { let mut keys = (0..30).map(|i| format!("k{i:02}")).collect::>(); @@ -88,7 +88,7 @@ fn normalize_typed_value(typed_value: ArrayRef, field_names: &[&str]) -> Variant #[test] fn test_normalize_shredded_variant_widens_unsigned_values() { - let metadata_builder = VariantBuilder::new().with_field_names(["u8", "u16", "u32"]); + let metadata_builder = VariantBuilder::new().with_field_names(["u8", "u16", "u32", "u64"]); let (metadata_bytes, _) = metadata_builder.finish(); let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); @@ -102,6 +102,10 @@ fn test_normalize_shredded_variant_widens_unsigned_values() { "u32", Arc::new(UInt32Array::from(vec![u32::MAX])) as ArrayRef, ), + ( + "u64", + Arc::new(UInt64Array::from(vec![u64::MAX])) as ArrayRef, + ), ]; let mut object_fields = Vec::with_capacity(fields.len()); let mut object_columns = Vec::with_capacity(fields.len()); @@ -153,6 +157,77 @@ fn test_normalize_shredded_variant_widens_unsigned_values() { assert_eq!(object.get("u8"), Some(Variant::from(255_i16))); assert_eq!(object.get("u16"), Some(Variant::from(65_535_i32))); assert_eq!(object.get("u32"), Some(Variant::from(4_294_967_295_i64))); + assert_eq!( + object.get("u64"), + Some(Variant::Decimal16( + VariantDecimal16::try_new(18_446_744_073_709_551_615_i128, 0).unwrap() + )) + ); +} + +#[test] +fn test_normalize_variant_decodes_dictionary_storage() { + let mut builder = VariantBuilder::new(); + builder.append_value(42_i64); + let (metadata_bytes, value_bytes) = builder.finish(); + let values: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let value: ArrayRef = + Arc::new(DictionaryArray::::try_new(Int32Array::from(vec![0]), values).unwrap()); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", value.data_type().clone(), false), + Field::new("metadata", DataType::Binary, false), + ]), + vec![value, metadata], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!(output.value(0), Variant::Int64(42)); + + let values: ArrayRef = Arc::new(Int64Array::from(vec![42])); + let dictionary: ArrayRef = + Arc::new(DictionaryArray::::try_new(Int32Array::from(vec![0]), values).unwrap()); + let child: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + dictionary.data_type().clone(), + false, + )]), + vec![dictionary], + None, + ) + .unwrap(), + ); + let object: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("d", child.data_type().clone(), false)]), + vec![child], + None, + ) + .unwrap(), + ); + let output = normalize_typed_value(object, &["d"]); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + assert_eq!(object.get("d"), Some(Variant::Int8(42))); } #[test] @@ -831,32 +906,97 @@ fn test_normalize_variant_preserves_empty_object_keys() { assert_eq!(nested.get(""), Some(Variant::from(2_i64))); } +#[test] +fn test_compatibility_retry_rejects_non_spark_residual_order() { + let mut value_builder = VariantBuilder::new().with_field_names(["a", "b", "c"]); + let mut object = value_builder.new_object(); + object.insert("a", 1_i64); + object.insert("b", 2_i64); + object.insert("c", 3_i64); + object.finish(); + let (_, value) = value_builder.finish(); + + // Reassign the field IDs so the encoded slots read as b, a, c. That is neither canonical + // UTF-8 order nor Spark's UTF-16 order and must not be repaired by the compatibility retry. + let mut metadata_builder = WritableMetadataBuilder::from_iter(["b", "a", "c"]); + metadata_builder.finish(); + let metadata = metadata_builder.into_inner(); + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])); + + let error = rewrite_residual_values( + &value, + metadata.as_binary::(), + metadata.as_binary::(), + &[Some(0)], + false, + ) + .unwrap_err(); + assert!(error.to_string().contains("neither UTF-8 nor Spark UTF-16")); +} + +#[test] +fn test_empty_key_retry_rejects_non_spark_metadata() { + let mut metadata_builder = WritableMetadataBuilder::from_iter(["", "b", "a"]); + metadata_builder.finish(); + let mut metadata = metadata_builder.into_inner(); + metadata[0] |= 0x10; + + let mut value_builder = VariantBuilder::new(); + value_builder.append_value(1_i64); + let (_, value) = value_builder.finish(); + 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(value.as_slice())])), + Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])), + ], + None, + ) + .unwrap(), + ); + let variant = VariantArray::try_new(physical.as_ref()).unwrap(); + assert!(canonicalize_spark_empty_key_metadata(&variant) + .unwrap() + .is_none()); +} + #[test] fn test_normalize_partially_shredded_nested_unicode_and_empty_keys() { let keys = unicode_object_keys(); - let mut builder = VariantBuilder::new().with_field_names(["known"]); - let mut object = builder.new_object(); - object.insert("", 1_i64); - let mut nested = object.new_object("nested"); + let mut builder = VariantBuilder::new().with_field_names(["nested", "known"]); + let mut residual = builder.new_object(); + residual.insert("", -1_i64); for (index, key) in keys.iter().enumerate() { - nested.insert(key, if key == "😀" { 531_i64 } else { index as i64 }); + residual.insert(key, if key == "😀" { 531_i64 } else { index as i64 }); } - nested.finish(); - object.finish(); + residual.finish(); let (metadata_bytes, value_bytes) = builder.finish(); assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let value = reorder_variant_values( + &value, + &metadata, + None, + VariantObjectKeyOrder::SparkUtf16, + false, + ) + .unwrap(); let known: ArrayRef = Arc::new( StructArray::try_new( Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), - vec![Arc::new(Int64Array::from(vec![3]))], + vec![Arc::new(Int64Array::from(vec![99]))], None, ) .unwrap(), ); - let typed_value: ArrayRef = Arc::new( + let nested_typed_value: ArrayRef = Arc::new( StructArray::try_new( Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), vec![known], @@ -864,14 +1004,36 @@ fn test_normalize_partially_shredded_nested_unicode_and_empty_keys() { ) .unwrap(), ); + let nested: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", DataType::Binary, true), + Field::new("typed_value", nested_typed_value.data_type().clone(), true), + ]), + vec![value, nested_typed_value], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "nested", + nested.data_type().clone(), + false, + )]), + vec![nested], + 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, value, typed_value], + vec![metadata, typed_value], None, ) .unwrap(), @@ -893,16 +1055,103 @@ fn test_normalize_partially_shredded_nested_unicode_and_empty_keys() { let Variant::Object(object) = output.value(0) else { panic!("expected object") }; - assert_eq!(object.get("").unwrap().as_int64(), Some(1)); - assert_eq!(object.get("known").unwrap().as_int64(), Some(3)); let Variant::Object(nested) = object.get("nested").unwrap() else { panic!("expected nested object") }; let fields = nested.iter().collect::>(); - assert_eq!(fields.len(), 32); - assert_eq!(fields[30].0, "😀"); - assert_eq!(fields[31].0, "\u{e000}"); - assert_eq!(nested.get("😀").unwrap().as_int64(), Some(531)); + assert_eq!(fields.len(), 34); + assert_eq!(fields[0].0, ""); + assert_eq!(fields[32].0, "😀"); + assert_eq!(fields[33].0, "\u{e000}"); + assert_eq!(nested.get("known").unwrap().as_int64(), Some(99)); + let emoji = fields + .binary_search_by(|(name, _)| name.encode_utf16().cmp("😀".encode_utf16())) + .unwrap(); + assert_eq!(fields[emoji].1.as_int64(), Some(531)); +} + +#[test] +fn test_normalize_nested_list_residuals_tracks_root_metadata() { + fn row(keys: &[&str]) -> DataFusionResult<(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, + VariantObjectKeyOrder::SparkUtf16, + false, + )?; + Ok((metadata, value.as_binary::().value(0).to_vec())) + } + + let (metadata0, value0) = row(&["a", "\u{e000}", "😀"]).unwrap(); + let (metadata1, value1) = row(&["b", "zz", "\u{ffff}", "𐀀"]).unwrap(); + 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( + arrow::array::ListArray::try_new( + Arc::new(Field::new("element", states.data_type().clone(), false)), + arrow::buffer::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 target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + let output = normalize_variant_array(&physical, &target_field).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] diff --git a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql index 2674970ec94..69a98502d8f 100644 --- a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql +++ b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql @@ -22,6 +22,7 @@ -- MinSparkVersion: 4.0 -- Config: spark.sql.variant.writeShredding.enabled=false -- Config: spark.sql.variant.pushVariantIntoScan=false +-- Config: spark.sql.variant.allowReadingShredded=true -- Config: spark.sql.variant.forceShreddingSchemaForTest=k00 BIGINT statement @@ -118,8 +119,11 @@ SET spark.sql.parquet.enableVectorizedReader=true statement SET spark.comet.scan.allowDisabledParquetVectorizedReader=false --- Arrow and Spark order supplementary Unicode object keys differently. Force a shredded field so --- the native scan reconstructs the whole 32-field value before Spark's variant_get binary search. +-- Arrow and Spark order supplementary Unicode object keys differently. Force top-level and nested +-- shredded fields so the native scan reconstructs their residual values before Spark's lookup. +statement +SET spark.sql.variant.forceShreddingSchemaForTest=k00 BIGINT, nested STRUCT + statement SET spark.sql.variant.writeShredding.enabled=true @@ -131,7 +135,7 @@ INSERT INTO test_variant_unicode VALUES (parse_json( '{"k00":0,"k01":1,"k02":2,"k03":3,"k04":4,"k05":5,"k06":6,"k07":7,"k08":8,"k09":9,"k10":10,"k11":11,"k12":12,"k13":13,"k14":14,"k15":15,"k16":16,"k17":17,"k18":18,"k19":19,"k20":20,"k21":21,"k22":22,"k23":23,"k24":24,"k25":25,"k26":26,"k27":27,"k28":28,"k29":29,"\uE000":30,"😀":531}')), (parse_json( - '{"":-2,"k00":0,"nested":{"":-1,"k00":0,"k01":1,"k02":2,"k03":3,"k04":4,"k05":5,"k06":6,"k07":7,"k08":8,"k09":9,"k10":10,"k11":11,"k12":12,"k13":13,"k14":14,"k15":15,"k16":16,"k17":17,"k18":18,"k19":19,"k20":20,"k21":21,"k22":22,"k23":23,"k24":24,"k25":25,"k26":26,"k27":27,"k28":28,"k29":29,"\uE000":30,"😀":532}}')) + '{"":-2,"k00":0,"nested":{"known":99,"":-1,"k00":0,"k01":1,"k02":2,"k03":3,"k04":4,"k05":5,"k06":6,"k07":7,"k08":8,"k09":9,"k10":10,"k11":11,"k12":12,"k13":13,"k14":14,"k15":15,"k16":16,"k17":17,"k18":18,"k19":19,"k20":20,"k21":21,"k22":22,"k23":23,"k24":24,"k25":25,"k26":26,"k27":27,"k28":28,"k29":29,"\uE000":30,"😀":532}}')) statement SET spark.sql.variant.writeShredding.enabled=false diff --git a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala index f72936650ed..742ec883792 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala @@ -393,28 +393,39 @@ abstract class ParquetReadSuite extends CometTestBase { } } - test("native scan decodes dictionary-encoded Variant metadata") { + test("native scan decodes dictionary-encoded Variant storage") { assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") withTempDir { dir => val path = new Path(dir.toURI.toString, "dictionary-variant.parquet") val parquetSchema = MessageTypeParser.parseMessageType("""message root { | optional group v { - | required binary value; + | optional binary value; | required binary metadata; + | optional int64 typed_value; | } |} |""".stripMargin) val valueField = new ArrowField( "value", - FieldType.notNullable(ArrowType.Binary.INSTANCE), + new FieldType( + true, + ArrowType.Binary.INSTANCE, + new DictionaryEncoding(0L, false, new ArrowType.Int(32, true))), Collections.emptyList[ArrowField]()) val metadataField = new ArrowField( "metadata", new FieldType( false, ArrowType.Binary.INSTANCE, - new DictionaryEncoding(0L, false, new ArrowType.Int(32, true))), + new DictionaryEncoding(1L, false, new ArrowType.Int(32, true))), + Collections.emptyList[ArrowField]()) + val typedValueField = new ArrowField( + "typed_value", + new FieldType( + true, + new ArrowType.Int(64, true), + new DictionaryEncoding(2L, false, new ArrowType.Int(32, true))), Collections.emptyList[ArrowField]()) val variantField = new ArrowField( "v", @@ -423,7 +434,7 @@ abstract class ParquetReadSuite extends CometTestBase { ArrowType.Struct.INSTANCE, null, Collections.singletonMap("ARROW:extension:name", "arrow.parquet.variant")), - Seq(valueField, metadataField).asJava) + Seq(valueField, metadataField, typedValueField).asJava) val arrowSchema = new ArrowSchema(Collections.singletonList(variantField)) val footer = Collections.singletonMap( "ARROW:schema", @@ -442,11 +453,16 @@ abstract class ParquetReadSuite extends CometTestBase { .build() try { - (0 until 3).foreach { _ => + (0 until 4).foreach { index => val row = new SimpleGroup(parquetSchema) val group = row.addGroup(0) - group.add(0, Binary.fromConstantByteArray(value)) + if (index % 2 == 0) { + group.add(0, Binary.fromConstantByteArray(value)) + } group.add(1, Binary.fromConstantByteArray(metadata)) + if (index % 2 != 0) { + group.add(2, 42L) + } writer.write(row) } } finally { @@ -456,9 +472,53 @@ abstract class ParquetReadSuite extends CometTestBase { withTable("dictionary_variant") { sql(s"""CREATE TABLE dictionary_variant(v VARIANT) |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) - withSQLConf("spark.sql.variant.pushVariantIntoScan" -> "false") { + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { val df = sql("SELECT v FROM dictionary_variant") - assert(normalizedVariantRows(df, 0) == Seq.fill(3)(Seq("42"))) + assert(normalizedVariantRows(df, 0) == Seq.fill(4)(Seq("42"))) + assert(collect(df.queryExecution.executedPlan) { case _: CometNativeScanExec => + true + }.size == 1) + } + } + } + } + + test("native scan widens shredded Variant UINT_64") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempDir { dir => + val path = new Path(dir.toURI.toString, "uint64-variant.parquet") + val parquetSchema = MessageTypeParser.parseMessageType("""message root { + | optional group v { + | required binary metadata; + | required int64 typed_value (UINT_64); + | } + |} + |""".stripMargin) + val variant = sql("SELECT parse_json('0')").head().get(0) + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + val writer = createParquetWriter(parquetSchema, path, dictionaryEnabled = false) + try { + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(0, Binary.fromConstantByteArray(metadata)) + group.add(1, -1L) + writer.write(row) + } finally { + writer.close() + } + + withTable("uint64_variant") { + sql(s"""CREATE TABLE uint64_variant(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql("SELECT v FROM uint64_variant") + assert(normalizedVariantRows(df, 0) == Seq(Seq("18446744073709551615"))) assert(collect(df.queryExecution.executedPlan) { case _: CometNativeScanExec => true }.size == 1) From 389b134b89256a9ff50cd5623315fd8689295e8e Mon Sep 17 00:00:00 2001 From: peterxcli Date: Thu, 27 Aug 2026 14:54:45 +0800 Subject: [PATCH 13/14] fix: align native Variant Parquet scan semantics --- native/core/Cargo.toml | 2 +- .../core/src/parquet/cast_column/variant.rs | 34 +++++- .../src/parquet/cast_column/variant/tests.rs | 48 +++++++- .../serde/operator/CometNativeScan.scala | 10 ++ .../comet/parquet/ParquetReadSuite.scala | 108 +++++++++++++++++- 5 files changed, 193 insertions(+), 9 deletions(-) diff --git a/native/core/Cargo.toml b/native/core/Cargo.toml index 8dc8d73273f..79146aeeb03 100644 --- a/native/core/Cargo.toml +++ b/native/core/Cargo.toml @@ -37,7 +37,7 @@ publish = false [dependencies] arrow = { workspace = true } bytes = { workspace = true } -parquet = { workspace = true, default-features = false, features = ["experimental", "arrow", "snap", "lz4", "zstd", "flate2-zlib-rs"] } +parquet = { workspace = true, default-features = false, features = ["experimental", "arrow", "arrow_canonical_extension_types", "snap", "lz4", "zstd", "flate2-zlib-rs"] } futures = { workspace = true } mimalloc = { version = "*", default-features = false, optional = true } tikv-jemallocator = { version = "0.6.1", optional = true, features = ["disable_initial_exec_tls"] } diff --git a/native/core/src/parquet/cast_column/variant.rs b/native/core/src/parquet/cast_column/variant.rs index d2e098664a8..0047d3a07b1 100644 --- a/native/core/src/parquet/cast_column/variant.rs +++ b/native/core/src/parquet/cast_column/variant.rs @@ -131,6 +131,7 @@ fn normalize_variant_type(data_type: &DataType) -> Option { DataType::Timestamp(TimeUnit::Millisecond, timezone) => { Some(DataType::Timestamp(TimeUnit::Microsecond, timezone.clone())) } + DataType::FixedSizeBinary(_) => Some(DataType::Binary), DataType::FixedSizeList(field, _) => Some(DataType::List( normalize_field(field).unwrap_or_else(|| Arc::clone(field)), )), @@ -156,11 +157,33 @@ fn normalize_variant_type(data_type: &DataType) -> Option { } } +fn contains_uuid_extension(data_type: &DataType) -> bool { + fn field_contains_uuid(field: &FieldRef) -> bool { + (field.data_type() == &DataType::FixedSizeBinary(16) + && field.extension_type_name() == Some("arrow.uuid")) + || contains_uuid_extension(field.data_type()) + } + + match data_type { + DataType::Struct(fields) => fields.iter().any(field_contains_uuid), + DataType::List(field) + | DataType::LargeList(field) + | DataType::ListView(field) + | DataType::LargeListView(field) + | DataType::FixedSizeList(field, _) + | DataType::Map(field, _) => field_contains_uuid(field), + DataType::Dictionary(_, value_type) => contains_uuid_extension(value_type), + _ => false, + } +} + /// Normalize Arrow types that Spark's Parquet reader accepts but `VariantArray` 58.4 rejects. /// Parquet restores unsigned integers to Arrow unsigned arrays and millisecond timestamps at their -/// annotated unit; embedded Arrow schemas may also restore dictionary arrays and fixed-size lists. -/// Spark decodes dictionaries, widens integers and timestamps, and treats fixed-size lists as -/// ordinary Variant arrays. +/// annotated unit; it also exposes unannotated FIXED_LEN_BYTE_ARRAY as FixedSizeBinary. The +/// `arrow_canonical_extension_types` feature retains UUID annotations so they are not mistaken for +/// Spark-compatible Binary. Embedded Arrow schemas may restore dictionary arrays and fixed-size +/// lists. Spark decodes dictionaries, widens integers and timestamps, treats unannotated +/// fixed-length binary as Binary, and treats fixed-size lists as ordinary Variant arrays. /// arrow-rs #10416/#10417 would move the UInt8/16/32 widening into /// `VariantArray`/`unshred_variant`; remove those three arms after that ships and Comet upgrades: /// https://github.com/apache/arrow-rs/issues/10416 @@ -175,6 +198,11 @@ fn normalize_variant_type(data_type: &DataType) -> Option { /// dictionary `value` and `typed_value` children, so keep decoding those for Spark compatibility: /// https://github.com/apache/arrow-rs/pull/10810 fn normalize_variant_storage(array: &ArrayRef) -> DataFusionResult { + if contains_uuid_extension(array.data_type()) { + return Err(DataFusionError::Execution( + "Parquet UUID is not supported as a shredded Variant child".to_string(), + )); + } let Some(data_type) = normalize_variant_type(array.data_type()) else { return Ok(Arc::clone(array)); }; diff --git a/native/core/src/parquet/cast_column/variant/tests.rs b/native/core/src/parquet/cast_column/variant/tests.rs index 2756177384e..e011a1708e1 100644 --- a/native/core/src/parquet/cast_column/variant/tests.rs +++ b/native/core/src/parquet/cast_column/variant/tests.rs @@ -16,11 +16,15 @@ // under the License. use super::*; use arrow::array::{ - Decimal128Array, DictionaryArray, FixedSizeListArray, Int32Array, Int64Array, - TimestampMillisecondArray, UInt16Array, UInt32Array, UInt64Array, UInt8Array, + Decimal128Array, DictionaryArray, FixedSizeBinaryArray, FixedSizeListArray, Int32Array, + Int64Array, TimestampMillisecondArray, UInt16Array, UInt32Array, UInt64Array, UInt8Array, }; use arrow::datatypes::{Field, Fields, Int32Type}; -use parquet::variant::{VariantDecimal16, VariantType}; +use parquet::{ + arrow::parquet_to_arrow_schema, + schema::{parser::parse_message_type, types::SchemaDescriptor}, + variant::{VariantDecimal16, VariantType}, +}; fn unicode_object_keys() -> Vec { let mut keys = (0..30).map(|i| format!("k{i:02}")).collect::>(); @@ -311,6 +315,44 @@ fn test_normalize_shredded_variant_converts_fixed_size_list() { ); } +#[test] +fn test_normalize_shredded_variant_converts_fixed_size_binary() { + let bytes = (0_u8..16).collect::>(); + let typed_value: ArrayRef = Arc::new(FixedSizeBinaryArray::from(vec![Some(bytes.as_slice())])); + + let output = normalize_typed_value(typed_value, &[]); + assert_eq!(output.value(0), Variant::Binary(&bytes)); +} + +#[test] +fn test_normalize_shredded_variant_rejects_uuid_annotation() { + let bytes = (0_u8..16).collect::>(); + let typed_value: ArrayRef = Arc::new(FixedSizeBinaryArray::from(vec![Some(bytes.as_slice())])); + let parquet_schema = parse_message_type( + "message root { required fixed_len_byte_array(16) typed_value (UUID); }", + ) + .unwrap(); + let arrow_schema = + parquet_to_arrow_schema(&SchemaDescriptor::new(Arc::new(parquet_schema)), None).unwrap(); + let typed_value_field = Arc::clone(&arrow_schema.fields()[0]); + assert_eq!(typed_value_field.extension_type_name(), Some("arrow.uuid")); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&[][..])])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Arc::new(Field::new("metadata", DataType::Binary, false)), + typed_value_field, + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + + let error = normalize_variant_storage(&physical).unwrap_err(); + assert!(error.to_string().contains("Parquet UUID is not supported")); +} + #[test] fn test_normalize_shredded_variant_compacts_spark_integer_widths() { let (metadata_bytes, _) = VariantBuilder::new().finish(); diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala index d3225d9c333..16fe659ee39 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala @@ -143,6 +143,16 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS "Full native scan disabled because spark.sql.legacy.parquet.nanosAsLong is enabled") } + // Spark interprets unadjusted Parquet timestamps as TimestampType when NTZ inference is + // disabled, while Arrow exposes them as timezone-free TimestampNTZ values. The Variant schema + // does not expose shredded leaf types, so preserve Spark's vectorized-reader semantics. + if (hasVariant && !SQLConf.get.parquetInferTimestampNTZEnabled) { + withFallbackReason( + scanExec, + s"Full native scan disabled because " + + s"${SQLConf.PARQUET_INFER_TIMESTAMP_NTZ_ENABLED.key} is disabled") + } + // the scan is supported if no fallback reasons were added to the node !hasFallbackReason(scanExec) } diff --git a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala index 742ec883792..036ba0817ba 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala @@ -319,6 +319,60 @@ abstract class ParquetReadSuite extends CometTestBase { } } + test("inferTimestampNTZ disabled Variant children fall back to Spark") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempDir { dir => + val path = new Path(dir.toURI.toString, "timestamp-ntz-variant.parquet") + val parquetSchema = MessageTypeParser.parseMessageType("""message root { + | optional group v { + | required binary metadata; + | required int64 typed_value (TIMESTAMP(MICROS,false)); + | } + |} + |""".stripMargin) + val variant = sql("SELECT parse_json('0')").head().get(0) + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + val writer = createParquetWriter(parquetSchema, path, dictionaryEnabled = false) + try { + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(0, Binary.fromConstantByteArray(metadata)) + group.add(1, 1704067200123456L) + writer.write(row) + } finally { + writer.close() + } + + withTable("variant_timestamp_ntz") { + sql(s"""CREATE TABLE variant_timestamp_ntz(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + val readerConfs = Seq( + SQLConf.PARQUET_VECTORIZED_READER_ENABLED.key -> "true", + SQLConf.PARQUET_INFER_TIMESTAMP_NTZ_ENABLED.key -> "false", + SQLConf.SESSION_LOCAL_TIMEZONE.key -> "America/Los_Angeles", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false", + "spark.sql.variant.inferShreddingSchema" -> "false") + var expected = Seq.empty[Seq[Any]] + withSQLConf((readerConfs :+ (CometConf.COMET_ENABLED.key -> "false")): _*) { + expected = normalizedVariantRows(sql("SELECT v FROM variant_timestamp_ntz"), 0) + } + withSQLConf(readerConfs: _*) { + val result = sql("SELECT v FROM variant_timestamp_ntz") + assert(normalizedVariantRows(result, 0) == expected) + val plan = result.queryExecution.executedPlan + assert(collect(plan) { case _: CometNativeScanExec => true }.isEmpty) + val fallbackReasons = new ExtendedExplainInfo().getFallbackReasons(plan) + assert( + fallbackReasons.exists(_.contains(SQLConf.PARQUET_INFER_TIMESTAMP_NTZ_ENABLED.key)), + s"Expected inferTimestampNTZ fallback, found: ${fallbackReasons.mkString(", ")}") + } + } + } + } + test("native scan preserves Variant existence default pairing") { assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") @@ -329,7 +383,9 @@ abstract class ParquetReadSuite extends CometTestBase { sql("ALTER TABLE variant_defaults ADD COLUMNS(n INT DEFAULT 7)") } - withSQLConf("spark.sql.variant.pushVariantIntoScan" -> "false") { + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { val df = sql("SELECT v, n FROM variant_defaults") assert(normalizedVariantRows(df, 0) == Seq(Seq("42", 7))) val cometPlan = df.queryExecution.executedPlan @@ -370,7 +426,9 @@ abstract class ParquetReadSuite extends CometTestBase { Seq(3, null, 9), Seq(4, "null", 10))) - withSQLConf("spark.sql.variant.pushVariantIntoScan" -> "false") { + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { val df = sql(query) assert(normalizedVariantRows(df, 1) == expected) val cometPlan = df.queryExecution.executedPlan @@ -393,6 +451,52 @@ abstract class ParquetReadSuite extends CometTestBase { } } + test("native scan preserves unannotated fixed-length Variant binary") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + Seq(8, 16, 20).foreach { length => + withTempDir { dir => + val path = new Path(dir.toURI.toString, "fixed-binary-variant.parquet") + val parquetSchema = MessageTypeParser.parseMessageType(s"""message root { + | optional group v { + | required binary metadata; + | required fixed_len_byte_array($length) typed_value; + | } + |} + |""".stripMargin) + val variant = sql("SELECT parse_json('0')").head().get(0) + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + val bytes = Array.tabulate(length)(_.toByte) + val writer = createParquetWriter(parquetSchema, path, dictionaryEnabled = false) + try { + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(0, Binary.fromConstantByteArray(metadata)) + group.add(1, Binary.fromConstantByteArray(bytes)) + writer.write(row) + } finally { + writer.close() + } + + withTable("fixed_binary_variant") { + sql(s"""CREATE TABLE fixed_binary_variant(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql("SELECT v FROM fixed_binary_variant") + val expected = "\"" + Base64.getEncoder.encodeToString(bytes) + "\"" + assert(normalizedVariantRows(df, 0) == Seq(Seq(expected))) + assert(collect(df.queryExecution.executedPlan) { case _: CometNativeScanExec => + true + }.size == 1) + } + } + } + } + } + test("native scan decodes dictionary-encoded Variant storage") { assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") From 43e7e179a9805700fca46dcec4b0b215332773ba Mon Sep 17 00:00:00 2001 From: peterxcli Date: Thu, 27 Aug 2026 15:18:42 +0800 Subject: [PATCH 14/14] style: remove redundant interpolation --- .../scala/org/apache/comet/serde/operator/CometNativeScan.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala index 16fe659ee39..3f2f4b0fef6 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala @@ -149,7 +149,7 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS if (hasVariant && !SQLConf.get.parquetInferTimestampNTZEnabled) { withFallbackReason( scanExec, - s"Full native scan disabled because " + + "Full native scan disabled because " + s"${SQLConf.PARQUET_INFER_TIMESTAMP_NTZ_ENABLED.key} is disabled") }