From 45a0ed44ed9c58ede31410de077a0882e72fd4f8 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Fri, 21 Aug 2026 19:10:04 +0800 Subject: [PATCH 1/8] 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 2/8] 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 3/8] 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 4/8] 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 5/8] 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 6/8] 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 7/8] 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 8/8] 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()?;