diff --git a/native/core/Cargo.toml b/native/core/Cargo.toml index 8dc8d73273f..79146aeeb03 100644 --- a/native/core/Cargo.toml +++ b/native/core/Cargo.toml @@ -37,7 +37,7 @@ publish = false [dependencies] arrow = { workspace = true } bytes = { workspace = true } -parquet = { workspace = true, default-features = false, features = ["experimental", "arrow", "snap", "lz4", "zstd", "flate2-zlib-rs"] } +parquet = { workspace = true, default-features = false, features = ["experimental", "arrow", "arrow_canonical_extension_types", "snap", "lz4", "zstd", "flate2-zlib-rs"] } futures = { workspace = true } mimalloc = { version = "*", default-features = false, optional = true } tikv-jemallocator = { version = "0.6.1", optional = true, features = ["disable_initial_exec_tls"] } diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index bb218bf108c..e19a9907265 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -758,6 +758,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(); @@ -784,6 +785,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 @@ -800,11 +802,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 9d2a03f29a3..5d2f7d95100 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -44,7 +44,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, ShuffleWriterDestination, ShuffleWriterExec}, }; use crate::jvm_bridge::{jni_call, JVMClasses, JavaShufflePartitionPusher, ShufflePartitionPusher}; @@ -116,7 +116,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}; @@ -941,6 +941,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, @@ -1598,44 +1617,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 @@ -3878,15 +3891,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(); @@ -4805,7 +4820,8 @@ mod tests { }; 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; @@ -4818,7 +4834,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; @@ -4831,6 +4850,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; @@ -5250,6 +5270,85 @@ mod tests { assert_eq!(recorder.pushes.load(Ordering::Relaxed), 1); } + #[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/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..435da3fe030 100644 --- a/native/core/src/execution/utils.rs +++ b/native/core/src/execution/utils.rs @@ -19,28 +19,48 @@ 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, 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; 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(self.data_type())?); + 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. @@ -55,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(self.data_type())?); + std::ptr::write(array_ptr, ffi_array); + std::ptr::write(schema_ptr, ffi_schema); } } @@ -65,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 1cc928d1d59..502622c880f 100644 --- a/native/core/src/parquet/cast_column.rs +++ b/native/core/src/parquet/cast_column.rs @@ -14,6 +14,9 @@ // KIND, either express or implied. See the License for the // specific language governing permissions and limitations // under the License. +mod variant; + +use self::variant::normalize_variant_array; use arrow::{ array::{ make_array, Array, ArrayRef, LargeListArray, ListArray, MapArray, StructArray, @@ -23,11 +26,9 @@ use arrow::{ datatypes::{DataType, FieldRef, Schema, TimeUnit}, record_batch::RecordBatch, }; - -use crate::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::{ + format::DEFAULT_CAST_OPTIONS, DataFusionError, Result as DataFusionResult, ScalarValue, +}; use datafusion::logical_expr::ColumnarValue; use datafusion::physical_expr::PhysicalExpr; use std::{ @@ -36,6 +37,11 @@ use std::{ sync::Arc, }; +use crate::{ + execution::serde::is_variant_field, + parquet::parquet_support::{spark_parquet_convert, SparkParquetOptions}, +}; + /// Returns true if two DataTypes are structurally equivalent (same data layout) /// but may differ in field names within nested types. fn types_differ_only_in_field_names(physical: &DataType, logical: &DataType) -> bool { @@ -260,6 +266,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 +367,86 @@ impl PhysicalExpr for CometCastColumnExpr { #[cfg(test)] mod tests { use super::*; - use arrow::array::{Array, Int32Array, StringArray}; - use arrow::datatypes::{Field, Fields}; + use arrow::{ + array::{ + Array, AsArray, BinaryArray, DictionaryArray, Int32Array, Int64Array, StringArray, + TimestampMillisecondArray, + }, + compute::cast, + datatypes::{Field, Fields, Int32Type}, + }; use datafusion::physical_expr::expressions::Column; + use parquet::variant::{Variant, VariantArray, VariantArrayBuilder, VariantType}; + + #[test] + 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 = 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), + 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_i8)); + assert_eq!(variant.value(2), Variant::from(30_i8)); + } #[test] fn test_cast_timestamp_micros_to_millis_array() { diff --git a/native/core/src/parquet/cast_column/variant.rs b/native/core/src/parquet/cast_column/variant.rs new file mode 100644 index 00000000000..43d8e412dba --- /dev/null +++ b/native/core/src/parquet/cast_column/variant.rs @@ -0,0 +1,1419 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. +use arrow::{ + array::{ + make_array, Array, ArrayRef, AsArray, BinaryArray, BinaryBuilder, ListLikeArray, + StructArray, + }, + buffer::NullBuffer, + compute::{cast, cast_with_options}, + datatypes::{DataType, FieldRef, TimeUnit, DECIMAL128_MAX_PRECISION}, + error::ArrowError, +}; +use datafusion::common::{ + format::DEFAULT_CAST_OPTIONS, DataFusionError, Result as DataFusionResult, +}; +use parquet::variant::{ + unshred_variant, BorrowedShreddingState, ListBuilder, MetadataBuilder, ObjectBuilder, + ParentState, ReadOnlyMetadataBuilder, ValueBuilder, Variant, VariantArray, VariantBuilder, + VariantDecimal4, VariantDecimal8, VariantMetadata, WritableMetadataBuilder, +}; +use std::{ + collections::HashSet, + panic::{catch_unwind, AssertUnwindSafe}, + sync::Arc, +}; + +pub(super) fn normalize_variant_array( + array: &ArrayRef, + target_field: &FieldRef, +) -> DataFusionResult { + let DataType::Struct(fields) = target_field.data_type() else { + return Err(DataFusionError::Execution( + "Variant extension field must use Struct storage".to_string(), + )); + }; + if fields.len() != 2 + || fields[0].name() != "value" + || fields[1].name() != "metadata" + || fields + .iter() + .any(|field| field.data_type() != &DataType::Binary) + { + return Err(DataFusionError::Execution( + "Variant output must contain Binary children [value, metadata]".to_string(), + )); + } + + let array = normalize_variant_storage(array)?; + let variant = VariantArray::try_new(array.as_ref())?; + let was_shredded = variant.typed_value_field().is_some(); + let unshredded = unshred_variant_for_spark(&variant)?; + let value = unshredded.value_field().ok_or_else(|| { + DataFusionError::Execution("Unshredded Variant is missing its value field".to_string()) + })?; + let value = cast(value.as_ref(), &DataType::Binary)?; + let metadata = cast(unshredded.metadata_field().as_ref(), &DataType::Binary)?; + let (value, metadata) = if was_shredded { + rebuild_shredded_variant_for_spark(&variant, &value, &metadata, unshredded.inner().nulls())? + } else { + let value = reorder_variant_values( + &value, + &metadata, + unshredded.inner().nulls(), + VariantObjectKeyOrder::SparkUtf16, + false, + )?; + (value, metadata) + }; + let output = StructArray::try_new( + fields.clone(), + vec![value, metadata], + unshredded.inner().nulls().cloned(), + )?; + Ok(Arc::new(output)) +} + +fn unshred_variant_for_spark(variant: &VariantArray) -> DataFusionResult { + let first_error = match unshred_variant(variant) { + Ok(array) => return Ok(array), + Err(error) => DataFusionError::from(error), + }; + + if let Ok(prepared) = prepare_variant_for_unshredding(variant, None) { + if let Ok(array) = unshred_variant(&prepared) { + return Ok(array); + } + } + let Some(metadata) = canonicalize_spark_empty_key_metadata(variant)? else { + return Err(first_error); + }; + let Ok(prepared) = prepare_variant_for_unshredding(variant, Some(metadata.as_binary::())) + else { + return Err(first_error); + }; + match unshred_variant(&prepared) { + Ok(array) => Ok(array), + Err(_) => Err(first_error), + } +} + +fn normalize_variant_type(data_type: &DataType) -> Option { + fn normalize_field(field: &FieldRef) -> Option { + normalize_variant_type(field.data_type()) + .map(|data_type| Arc::new(field.as_ref().clone().with_data_type(data_type))) + } + + match data_type { + DataType::Dictionary(_, value_type) => { + Some(normalize_variant_type(value_type).unwrap_or_else(|| value_type.as_ref().clone())) + } + DataType::UInt8 => Some(DataType::Int16), + DataType::UInt16 => Some(DataType::Int32), + DataType::UInt32 => Some(DataType::Int64), + // Spark reads Parquet UINT_64 as Decimal(20, 0). This is lossless for the full range and + // lets the existing Spark-compatible rebuild choose the Variant decimal width per value. + DataType::UInt64 => Some(DataType::Decimal128(20, 0)), + // Arrow chooses Decimal256 from the physical byte width, but Spark's DecimalType is + // precision-based and stores every supported precision (<= 38) in 128 bits. + DataType::Decimal256(precision, scale) if *precision <= DECIMAL128_MAX_PRECISION => { + Some(DataType::Decimal128(*precision, *scale)) + } + DataType::Timestamp(TimeUnit::Millisecond, timezone) => { + Some(DataType::Timestamp(TimeUnit::Microsecond, timezone.clone())) + } + DataType::FixedSizeBinary(_) => Some(DataType::Binary), + DataType::FixedSizeList(field, _) => Some(DataType::List( + normalize_field(field).unwrap_or_else(|| Arc::clone(field)), + )), + DataType::List(field) => normalize_field(field).map(DataType::List), + DataType::LargeList(field) => normalize_field(field).map(DataType::LargeList), + DataType::ListView(field) => normalize_field(field).map(DataType::ListView), + DataType::LargeListView(field) => normalize_field(field).map(DataType::LargeListView), + DataType::Struct(fields) => { + let mut changed = false; + let fields = fields + .iter() + .map(|field| match normalize_field(field) { + Some(field) => { + changed = true; + field + } + None => Arc::clone(field), + }) + .collect::>(); + changed.then(|| DataType::Struct(fields.into())) + } + _ => None, + } +} + +fn contains_uuid_extension(data_type: &DataType) -> bool { + fn field_contains_uuid(field: &FieldRef) -> bool { + (field.data_type() == &DataType::FixedSizeBinary(16) + && field.extension_type_name() == Some("arrow.uuid")) + || contains_uuid_extension(field.data_type()) + } + + match data_type { + DataType::Struct(fields) => fields.iter().any(field_contains_uuid), + DataType::List(field) + | DataType::LargeList(field) + | DataType::ListView(field) + | DataType::LargeListView(field) + | DataType::FixedSizeList(field, _) + | DataType::Map(field, _) => field_contains_uuid(field), + DataType::Dictionary(_, value_type) => contains_uuid_extension(value_type), + _ => false, + } +} + +/// Normalize Arrow types that Spark's Parquet reader accepts but `VariantArray` 58.4 rejects. +/// Parquet restores unsigned integers to Arrow unsigned arrays and millisecond timestamps at their +/// annotated unit; it also exposes unannotated FIXED_LEN_BYTE_ARRAY as FixedSizeBinary. The +/// `arrow_canonical_extension_types` feature retains UUID annotations so they are not mistaken for +/// Spark-compatible Binary. Embedded Arrow schemas may restore dictionary arrays and fixed-size +/// lists. Spark decodes dictionaries, widens integers and timestamps, treats unannotated +/// fixed-length binary as Binary, and treats fixed-size lists as ordinary Variant arrays. +/// arrow-rs #10416/#10417 would move the UInt8/16/32 widening into +/// `VariantArray`/`unshred_variant`; remove those three arms after that ships and Comet upgrades: +/// https://github.com/apache/arrow-rs/issues/10416 +/// https://github.com/apache/arrow-rs/pull/10417 +/// Arrow #50622/#50810 instead proposes removing unsigned Parquet `typed_value` mappings; keep +/// compatibility until upstream settles that schema boundary: +/// https://github.com/apache/arrow/issues/50622 +/// https://github.com/apache/arrow/pull/50810 +/// Spark maps Parquet UINT_64 to Decimal(20,0), so its UInt64 arm remains Spark-specific: +/// https://github.com/apache/spark/blob/v4.0.4/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/parquet/ParquetSchemaConverter.scala#L281-L296 +/// arrow-rs #10810 covers encoded canonical metadata only. Embedded Arrow schemas can also restore +/// dictionary `value` and `typed_value` children, so keep decoding those for Spark compatibility: +/// https://github.com/apache/arrow-rs/pull/10810 +fn normalize_variant_storage(array: &ArrayRef) -> DataFusionResult { + if contains_uuid_extension(array.data_type()) { + return Err(DataFusionError::Execution( + "Parquet UUID is not supported as a shredded Variant child".to_string(), + )); + } + let Some(data_type) = normalize_variant_type(array.data_type()) else { + return Ok(Arc::clone(array)); + }; + Ok(cast_with_options( + array.as_ref(), + &data_type, + &DEFAULT_CAST_OPTIONS, + )?) +} + +/// Arrow fully validates residual `value` fields while unshredding. Spark versions before +/// SPARK-58949 wrote their object keys in Java UTF-16 order, so rewrite every reachable legacy +/// residual to Arrow's UTF-8 order before calling the upstream unshredder. Keep this input-side +/// compatibility for historical files after #5474 removes the output-side UTF-16 rewrite: +/// https://github.com/apache/datafusion-comet/issues/5474 +/// `metadata_rows` carries the root metadata row through nested lists. +fn rewrite_shredding_state( + state: &StructArray, + source_metadata: &BinaryArray, + target_metadata: &BinaryArray, + metadata_rows: &[Option], + remap_metadata: bool, +) -> DataFusionResult<(ArrayRef, bool)> { + if state.len() != metadata_rows.len() { + return Err(DataFusionError::Execution( + "Variant shredding state and metadata row mapping have different lengths".to_string(), + )); + } + let active_rows = metadata_rows + .iter() + .enumerate() + .map(|(index, row)| state.is_valid(index).then_some(*row).flatten()) + .collect::>(); + let mut fields = state.fields().iter().cloned().collect::>(); + let mut columns = state.columns().to_vec(); + let mut changed = false; + + if let Some(index) = fields.iter().position(|field| field.name() == "value") { + let (value, value_changed) = rewrite_residual_values( + &columns[index], + source_metadata, + target_metadata, + &active_rows, + remap_metadata, + )?; + if value_changed { + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(value.data_type().clone()), + ); + columns[index] = value; + changed = true; + } + } + + if let Some(index) = fields + .iter() + .position(|field| field.name() == "typed_value") + { + let typed_rows = active_rows + .iter() + .enumerate() + .map(|(row, metadata)| columns[index].is_valid(row).then_some(*metadata).flatten()) + .collect::>(); + let (typed_value, typed_changed) = rewrite_typed_value( + &columns[index], + source_metadata, + target_metadata, + &typed_rows, + remap_metadata, + )?; + if typed_changed { + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(typed_value.data_type().clone()), + ); + columns[index] = typed_value; + changed = true; + } + } + + if !changed { + return Ok((Arc::new(state.clone()), false)); + } + Ok(( + Arc::new(StructArray::try_new( + fields.into(), + columns, + state.nulls().cloned(), + )?), + true, + )) +} + +fn rewrite_residual_values( + value: &ArrayRef, + source_metadata: &BinaryArray, + target_metadata: &BinaryArray, + metadata_rows: &[Option], + remap_metadata: bool, +) -> DataFusionResult<(ArrayRef, bool)> { + let binary = cast(value.as_ref(), &DataType::Binary)?; + let binary = binary.as_binary::(); + let mut output = BinaryBuilder::new(); + let mut changed = false; + + for (index, metadata_row) in metadata_rows.iter().enumerate() { + if binary.is_null(index) { + output.append_null(); + continue; + } + let Some(metadata_row) = metadata_row else { + output.append_value(binary.value(index)); + continue; + }; + if source_metadata.is_null(*metadata_row) || target_metadata.is_null(*metadata_row) { + return Err(DataFusionError::Execution(format!( + "Variant metadata is null at row {metadata_row}" + ))); + } + + let rebuilt = catch_unwind(AssertUnwindSafe( + || -> Result>, ArrowError> { + let source = VariantMetadata::new(source_metadata.value(*metadata_row)); + let target = VariantMetadata::new(target_metadata.value(*metadata_row)); + let variant = Variant::new_with_metadata(source, binary.value(index)); + let arrow_ordered = + is_compatible_variant(&variant, VariantObjectKeyOrder::ArrowUtf8); + if !remap_metadata && arrow_ordered { + return Ok(None); + } + if !arrow_ordered + && !is_compatible_variant(&variant, VariantObjectKeyOrder::SparkUtf16) + { + return Err(ArrowError::InvalidArgumentError( + "Variant residual is neither UTF-8 nor Spark UTF-16 ordered".to_string(), + )); + } + Ok(Some(arrow_variant_bytes(&target, variant)?)) + }, + )) + .map_err(|_| { + DataFusionError::Execution(format!("Invalid Variant residual at row {metadata_row}")) + })??; + changed |= rebuilt.is_some(); + output.append_value(rebuilt.as_deref().unwrap_or_else(|| binary.value(index))); + } + + if changed { + Ok((Arc::new(output.finish()), true)) + } else { + Ok((Arc::clone(value), false)) + } +} + +fn list_metadata_rows( + list: &L, + parent_rows: &[Option], +) -> DataFusionResult>> { + let mut child_rows = vec![None; list.values().len()]; + for (index, metadata_row) in parent_rows.iter().enumerate() { + let Some(metadata_row) = metadata_row else { + continue; + }; + for child_index in list.element_range(index) { + match child_rows[child_index] { + Some(existing) if existing != *metadata_row => { + return Err(DataFusionError::Execution( + "A shared Variant list child refers to different metadata rows".to_string(), + )); + } + _ => child_rows[child_index] = Some(*metadata_row), + } + } + } + Ok(child_rows) +} + +fn rewrite_list_typed_value( + array: &ArrayRef, + list: &L, + source_metadata: &BinaryArray, + target_metadata: &BinaryArray, + metadata_rows: &[Option], + remap_metadata: bool, +) -> DataFusionResult<(ArrayRef, bool)> { + let child_rows = list_metadata_rows(list, metadata_rows)?; + let values = list.values().as_struct_opt().ok_or_else(|| { + DataFusionError::Execution(format!( + "Invalid shredded Variant list values: expected Struct, got {}", + list.values().data_type() + )) + })?; + let (values, changed) = rewrite_shredding_state( + values, + source_metadata, + target_metadata, + &child_rows, + remap_metadata, + )?; + if !changed { + return Ok((Arc::clone(array), false)); + } + + let data_type = match array.data_type() { + DataType::List(field) => DataType::List(Arc::new( + field + .as_ref() + .clone() + .with_data_type(values.data_type().clone()), + )), + DataType::LargeList(field) => DataType::LargeList(Arc::new( + field + .as_ref() + .clone() + .with_data_type(values.data_type().clone()), + )), + DataType::ListView(field) => DataType::ListView(Arc::new( + field + .as_ref() + .clone() + .with_data_type(values.data_type().clone()), + )), + DataType::LargeListView(field) => DataType::LargeListView(Arc::new( + field + .as_ref() + .clone() + .with_data_type(values.data_type().clone()), + )), + data_type => { + return Err(DataFusionError::Execution(format!( + "Expected a Variant list, got {data_type}" + ))); + } + }; + let data = array + .to_data() + .into_builder() + .data_type(data_type) + .child_data(vec![values.to_data()]) + .build()?; + Ok((make_array(data), true)) +} + +fn rewrite_typed_value( + typed_value: &ArrayRef, + source_metadata: &BinaryArray, + target_metadata: &BinaryArray, + metadata_rows: &[Option], + remap_metadata: bool, +) -> DataFusionResult<(ArrayRef, bool)> { + match typed_value.data_type() { + DataType::Struct(_) => { + let object = typed_value.as_struct(); + let mut fields = object.fields().iter().cloned().collect::>(); + let mut columns = object.columns().to_vec(); + let mut changed = false; + for (index, column) in object.columns().iter().enumerate() { + let child = column.as_struct_opt().ok_or_else(|| { + DataFusionError::Execution(format!( + "Invalid shredded Variant object field '{}': expected Struct, got {}", + fields[index].name(), + column.data_type() + )) + })?; + let (child, child_changed) = rewrite_shredding_state( + child, + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + )?; + if child_changed { + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(child.data_type().clone()), + ); + columns[index] = child; + changed = true; + } + } + if !changed { + return Ok((Arc::clone(typed_value), false)); + } + Ok(( + Arc::new(StructArray::try_new( + fields.into(), + columns, + object.nulls().cloned(), + )?), + true, + )) + } + DataType::List(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list::(), + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + ), + DataType::LargeList(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list::(), + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + ), + DataType::ListView(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list_view::(), + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + ), + DataType::LargeListView(_) => rewrite_list_typed_value( + typed_value, + typed_value.as_list_view::(), + source_metadata, + target_metadata, + metadata_rows, + remap_metadata, + ), + _ => Ok((Arc::clone(typed_value), false)), + } +} + +fn prepare_variant_for_unshredding( + variant: &VariantArray, + target_metadata: Option<&BinaryArray>, +) -> DataFusionResult { + if variant.typed_value_field().is_none() { + return Ok(variant.clone()); + } + let source_metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let source_metadata = source_metadata.as_binary::(); + let remap_metadata = target_metadata.is_some(); + let target_metadata = target_metadata.unwrap_or(source_metadata); + let metadata_rows = (0..variant.len()) + .map(|index| variant.inner().is_valid(index).then_some(index)) + .collect::>(); + let (array, _) = rewrite_shredding_state( + variant.inner(), + source_metadata, + target_metadata, + &metadata_rows, + remap_metadata, + )?; + let array = array.as_struct(); + let mut fields = array.fields().iter().cloned().collect::>(); + let mut columns = array.columns().to_vec(); + let metadata_index = fields + .iter() + .position(|field| field.name() == "metadata") + .unwrap(); + fields[metadata_index] = Arc::new( + fields[metadata_index] + .as_ref() + .clone() + .with_data_type(DataType::Binary), + ); + columns[metadata_index] = Arc::new(target_metadata.clone()); + let array = StructArray::try_new(fields.into(), columns, array.nulls().cloned())?; + Ok(VariantArray::try_new(&array)?) +} + +/// Arrow 58.4 rejects Spark metadata dictionaries containing an empty object key when Spark's +/// insertion-order encoding leaves equal offsets. On an upstream validation failure, canonicalize +/// only that exact case; the recursive rewrite above remaps every residual to the new field IDs. +/// https://github.com/apache/arrow-rs/pull/10352 +fn canonicalize_spark_empty_key_metadata( + variant: &VariantArray, +) -> DataFusionResult> { + let metadata = cast(variant.metadata_field().as_ref(), &DataType::Binary)?; + let metadata = metadata.as_binary::(); + let mut output = BinaryBuilder::new(); + let mut changed = false; + + for index in 0..variant.len() { + if variant.inner().is_null(index) || metadata.is_null(index) { + if metadata.is_null(index) { + output.append_null(); + } else { + output.append_value(metadata.value(index)); + } + continue; + } + let metadata_bytes = metadata.value(index); + if VariantMetadata::try_new(metadata_bytes).is_ok() { + output.append_value(metadata_bytes); + continue; + } + + let replacement = catch_unwind(AssertUnwindSafe( + || -> Result>, ArrowError> { + let old_metadata = VariantMetadata::new(metadata_bytes); + let mut names = old_metadata + .iter_try() + .map(|name| name.map(str::to_string)) + .collect::, _>>()?; + if !names.iter().any(String::is_empty) { + return Ok(None); + } + let mut source = + WritableMetadataBuilder::from_iter(names.iter().map(String::as_str)); + source.finish(); + let mut source = source.into_inner(); + source[0] &= !0x10; + if source != metadata_bytes { + return Ok(None); + } + names.sort_unstable(); + if names.windows(2).any(|names| names[0] == names[1]) { + return Ok(None); + } + let metadata = VariantBuilder::new() + .with_field_names(names.iter().map(String::as_str)) + .finish() + .0; + VariantMetadata::try_new(&metadata)?; + Ok(Some(metadata)) + }, + )); + let Ok(Ok(Some(replacement))) = replacement else { + return Ok(None); + }; + output.append_value(replacement); + changed = true; + } + + if changed { + Ok(Some(Arc::new(output.finish()))) + } else { + Ok(None) + } +} + +/// Supplies sort-only field names whose Rust ordering matches Java `String.compareTo` ordering. +/// The original metadata dictionary still supplies the field IDs written to the Variant value. +#[derive(Debug)] +struct SparkMetadataBuilder<'a, 'm> { + metadata: &'a VariantMetadata<'m>, + sort_keys: Vec, +} + +impl<'a, 'm> SparkMetadataBuilder<'a, 'm> { + fn new(metadata: &'a VariantMetadata<'m>) -> Self { + let sort_keys = metadata + .iter() + .map(|field_name| { + field_name + .encode_utf16() + .map(|unit| char::from_u32(0x10000 + u32::from(unit)).unwrap()) + .collect() + }) + .collect(); + Self { + metadata, + sort_keys, + } + } +} + +impl MetadataBuilder for SparkMetadataBuilder<'_, '_> { + fn try_upsert_field_name(&mut self, field_name: &str) -> Result { + self.metadata + .get_entry(field_name) + .map(|(field_id, _)| field_id) + .ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Field name '{field_name}' not found in metadata dictionary" + )) + }) + } + + fn field_name(&self, field_id: usize) -> &str { + &self.sort_keys[field_id] + } + + fn num_field_names(&self) -> usize { + self.metadata.len() + } + + fn truncate_field_names(&mut self, new_size: usize) { + debug_assert_eq!(self.metadata.len(), new_size); + } + + fn finish(&mut self) -> usize { + self.metadata.size() + } +} + +#[derive(Clone, Copy)] +enum VariantObjectKeyOrder { + ArrowUtf8, + SparkUtf16, +} + +fn is_compatible_variant(variant: &Variant<'_, '_>, order: VariantObjectKeyOrder) -> bool { + match variant { + Variant::Object(object) => { + let mut previous = None; + object.iter().all(|(name, value)| { + let ordered = previous + .map(|previous: &str| match order { + VariantObjectKeyOrder::ArrowUtf8 => previous <= name, + VariantObjectKeyOrder::SparkUtf16 => { + previous.encode_utf16().cmp(name.encode_utf16()) + != std::cmp::Ordering::Greater + } + }) + .unwrap_or(true); + previous = Some(name); + ordered && is_compatible_variant(&value, order) + }) + } + Variant::List(list) => list + .iter() + .all(|value| is_compatible_variant(&value, order)), + _ => true, + } +} + +fn compact_spark_integer(value: i64) -> Variant<'static, 'static> { + if let Ok(value) = i8::try_from(value) { + Variant::Int8(value) + } else if let Ok(value) = i16::try_from(value) { + Variant::Int16(value) + } else if let Ok(value) = i32::try_from(value) { + Variant::Int32(value) + } else { + Variant::Int64(value) + } +} + +fn compact_spark_typed_variant<'m, 'v>(variant: Variant<'m, 'v>) -> Variant<'m, 'v> { + match variant { + Variant::Int16(value) => compact_spark_integer(value.into()), + Variant::Int32(value) => compact_spark_integer(value.into()), + Variant::Int64(value) => compact_spark_integer(value), + Variant::Decimal8(value) => i32::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal4::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal4) + .unwrap_or(Variant::Decimal8(value)), + Variant::Decimal16(value) => i32::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal4::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal4) + .or_else(|| { + i64::try_from(value.integer()) + .ok() + .and_then(|integer| VariantDecimal8::try_new(integer, value.scale()).ok()) + .map(Variant::Decimal8) + }) + .unwrap_or(Variant::Decimal16(value)), + Variant::Float(value) if value.is_nan() => Variant::Float(f32::from_bits(0x7fc0_0000)), + Variant::Double(value) if value.is_nan() => { + Variant::Double(f64::from_bits(0x7ff8_0000_0000_0000)) + } + variant => variant, + } +} + +fn arrow_variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + match variant { + Variant::Object(object) => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + true, + ); + for (name, value) in object.iter() { + let value = arrow_variant_bytes(metadata, value)?; + builder + .try_insert_bytes(name, Variant::new_with_metadata(metadata.clone(), &value))?; + } + builder.finish(); + } + Variant::List(list) => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + let mut builder = ListBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + true, + ); + for value in list.iter() { + let value = arrow_variant_bytes(metadata, value)?; + builder.append_value_bytes(Variant::new_with_metadata(metadata.clone(), &value)); + } + builder.finish(); + } + variant => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + ValueBuilder::try_append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + variant, + )?; + } + } + Ok(value_builder.into_inner()) +} + +/// Re-encode a residual Variant against `metadata`. Scalar widths are intentionally preserved, +/// matching Spark's `VariantBuilder.appendVariant` behavior. +fn spark_variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + match variant { + Variant::Object(object) => { + let mut metadata_builder = SparkMetadataBuilder::new(metadata); + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for (name, value) in object.iter() { + let value = spark_variant_bytes(metadata, value)?; + builder + .try_insert_bytes(name, Variant::new_with_metadata(metadata.clone(), &value))?; + } + builder.finish(); + } + Variant::List(list) => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + let mut builder = ListBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for value in list.iter() { + let value = spark_variant_bytes(metadata, value)?; + builder.append_value_bytes(Variant::new_with_metadata(metadata.clone(), &value)); + } + builder.finish(); + } + variant => { + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + ValueBuilder::try_append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + variant, + )?; + } + } + Ok(value_builder.into_inner()) +} + +fn variant_binary_value(array: &ArrayRef, index: usize) -> Result, ArrowError> { + if array.is_null(index) { + return Ok(None); + } + let value = match array.data_type() { + DataType::Binary => array.as_binary::().value(index), + DataType::LargeBinary => array.as_binary::().value(index), + DataType::BinaryView => array.as_binary_view().value(index), + data_type => { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant value must be binary-like, got {data_type}" + ))) + } + }; + Ok(Some(value)) +} + +fn shredding_state_has_value(state: &BorrowedShreddingState<'_>, index: usize) -> bool { + state + .typed_value_field() + .is_some_and(|array| array.is_valid(index)) + || state + .value_field() + .is_some_and(|array| array.is_valid(index)) +} + +fn collect_spark_field_name(name: &str, field_names: &mut Vec, seen: &mut HashSet) { + if seen.insert(name.to_string()) { + field_names.push(name.to_string()); + } +} + +fn collect_residual_field_names( + variant: Variant<'_, '_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + match variant { + Variant::Object(object) => { + for (name, value) in object.iter() { + collect_spark_field_name(name, field_names, seen); + collect_residual_field_names(value, field_names, seen)?; + } + } + Variant::List(list) => { + for value in list.iter() { + collect_residual_field_names(value, field_names, seen)?; + } + } + _ => {} + } + Ok(()) +} + +fn collect_list_field_names( + list: &L, + index: usize, + source_metadata: &VariantMetadata<'_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + let values = list.values().as_struct(); + let state = BorrowedShreddingState::try_from(values)?; + for element_index in list.element_range(index) { + collect_shredded_field_names( + state.clone(), + element_index, + source_metadata, + field_names, + seen, + )?; + } + Ok(()) +} + +fn collect_shredded_field_names( + state: BorrowedShreddingState<'_>, + index: usize, + source_metadata: &VariantMetadata<'_>, + field_names: &mut Vec, + seen: &mut HashSet, +) -> Result<(), ArrowError> { + let Some(typed_value) = state + .typed_value_field() + .filter(|array| array.is_valid(index)) + else { + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + collect_residual_field_names( + Variant::new_with_metadata(source_metadata.clone(), value), + field_names, + seen, + )?; + } + } + return Ok(()); + }; + + match typed_value.data_type() { + DataType::Struct(_) => { + let object = typed_value.as_struct(); + for (field, column) in object.fields().iter().zip(object.columns()) { + let child = column.as_struct_opt().ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Invalid shredded Variant object field '{}': expected Struct, got {}", + field.name(), + column.data_type() + )) + })?; + if child.is_null(index) { + return Err(ArrowError::InvalidArgumentError(format!( + "Shredded Variant object field '{}' is null", + field.name() + ))); + } + let child_state = BorrowedShreddingState::try_from(child)?; + if shredding_state_has_value(&child_state, index) { + collect_spark_field_name(field.name(), field_names, seen); + collect_shredded_field_names( + child_state, + index, + source_metadata, + field_names, + seen, + )?; + } + } + + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + let Variant::Object(residual) = + Variant::new_with_metadata(source_metadata.clone(), value) + else { + return Err(ArrowError::InvalidArgumentError( + "Partially shredded Variant object has a non-object value".to_string(), + )); + }; + for (name, value) in residual.iter() { + if object.fields().iter().any(|field| field.name() == name) { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant field '{name}' appears in both value and typed_value" + ))); + } + collect_spark_field_name(name, field_names, seen); + collect_residual_field_names(value, field_names, seen)?; + } + } + } + } + DataType::List(_) => collect_list_field_names( + typed_value.as_list::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::LargeList(_) => collect_list_field_names( + typed_value.as_list::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::ListView(_) => collect_list_field_names( + typed_value.as_list_view::(), + index, + source_metadata, + field_names, + seen, + )?, + DataType::LargeListView(_) => collect_list_field_names( + typed_value.as_list_view::(), + index, + source_metadata, + field_names, + seen, + )?, + _ => {} + } + Ok(()) +} + +fn spark_typed_variant_bytes( + metadata: &VariantMetadata<'_>, + variant: Variant<'_, '_>, +) -> Result, ArrowError> { + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = ReadOnlyMetadataBuilder::new(metadata); + ValueBuilder::try_append_variant( + ParentState::variant(&mut value_builder, &mut metadata_builder), + compact_spark_typed_variant(variant), + )?; + Ok(value_builder.into_inner()) +} + +fn spark_list_bytes( + list: &L, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Variant::List(semantic) = semantic else { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant list did not unshred to a list".to_string(), + )); + }; + let semantic = semantic.iter_try().collect::, _>>()?; + let element_range = list.element_range(index); + if element_range.len() != semantic.len() { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant list length changed while unshredding".to_string(), + )); + } + + let values = list.values().as_struct(); + let state = BorrowedShreddingState::try_from(values)?; + let mut elements = Vec::with_capacity(semantic.len()); + for (element_index, semantic) in element_range.zip(semantic) { + elements.push(spark_shredded_variant_bytes( + state.clone(), + element_index, + semantic, + source_metadata, + target_metadata, + )?); + } + + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = ReadOnlyMetadataBuilder::new(target_metadata); + let mut builder = ListBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for element in elements { + builder.append_value_bytes(Variant::new_with_metadata( + target_metadata.clone(), + &element, + )); + } + builder.finish(); + Ok(value_builder.into_inner()) +} + +fn spark_object_bytes( + state: BorrowedShreddingState<'_>, + object: &StructArray, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Variant::Object(semantic) = semantic else { + return Err(ArrowError::InvalidArgumentError( + "Shredded Variant object did not unshred to an object".to_string(), + )); + }; + let mut entries = Vec::new(); + for (field, column) in object.fields().iter().zip(object.columns()) { + let child = column.as_struct_opt().ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Invalid shredded Variant object field '{}': expected Struct, got {}", + field.name(), + column.data_type() + )) + })?; + if child.is_null(index) { + return Err(ArrowError::InvalidArgumentError(format!( + "Shredded Variant object field '{}' is null", + field.name() + ))); + } + let child_state = BorrowedShreddingState::try_from(child)?; + if shredding_state_has_value(&child_state, index) { + let value = semantic.get(field.name()).ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "Unshredded Variant is missing field '{}'", + field.name() + )) + })?; + entries.push(( + field.name().to_string(), + spark_shredded_variant_bytes( + child_state, + index, + value, + source_metadata, + target_metadata, + )?, + )); + } + } + + if let Some(value) = state.value_field() { + if let Some(value) = variant_binary_value(value, index)? { + let Variant::Object(residual) = + Variant::new_with_metadata(source_metadata.clone(), value) + else { + return Err(ArrowError::InvalidArgumentError( + "Partially shredded Variant object has a non-object value".to_string(), + )); + }; + for (name, value) in residual.iter() { + if object.fields().iter().any(|field| field.name() == name) { + return Err(ArrowError::InvalidArgumentError(format!( + "Variant field '{name}' appears in both value and typed_value" + ))); + } + entries.push(( + name.to_string(), + spark_variant_bytes(target_metadata, value)?, + )); + } + } + } + + let mut value_builder = ValueBuilder::new(); + let mut metadata_builder = SparkMetadataBuilder::new(target_metadata); + let mut builder = ObjectBuilder::new( + ParentState::variant(&mut value_builder, &mut metadata_builder), + false, + ); + for (name, value) in entries { + builder.try_insert_bytes( + &name, + Variant::new_with_metadata(target_metadata.clone(), &value), + )?; + } + builder.finish(); + Ok(value_builder.into_inner()) +} + +// ponytail: recursive child buffers can be O(depth²); use a streaming encoder only if deeply +// nested Variant profiles show this compatibility path is a bottleneck. +fn spark_shredded_variant_bytes( + state: BorrowedShreddingState<'_>, + index: usize, + semantic: Variant<'_, '_>, + source_metadata: &VariantMetadata<'_>, + target_metadata: &VariantMetadata<'_>, +) -> Result, ArrowError> { + let Some(typed_value) = state + .typed_value_field() + .filter(|array| array.is_valid(index)) + else { + return match state.value_field() { + Some(value) => match variant_binary_value(value, index)? { + Some(value) => spark_variant_bytes( + target_metadata, + Variant::new_with_metadata(source_metadata.clone(), value), + ), + None => Err(ArrowError::InvalidArgumentError( + "Shredded Variant has neither value nor typed_value".to_string(), + )), + }, + None => Err(ArrowError::InvalidArgumentError( + "Shredded Variant has neither value nor typed_value".to_string(), + )), + }; + }; + + match typed_value.data_type() { + DataType::Struct(_) => spark_object_bytes( + state, + typed_value.as_struct(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::List(_) => spark_list_bytes( + typed_value.as_list::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::LargeList(_) => spark_list_bytes( + typed_value.as_list::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::ListView(_) => spark_list_bytes( + typed_value.as_list_view::(), + index, + semantic, + source_metadata, + target_metadata, + ), + DataType::LargeListView(_) => spark_list_bytes( + typed_value.as_list_view::(), + index, + semantic, + source_metadata, + target_metadata, + ), + _ => spark_typed_variant_bytes(target_metadata, semantic), + } +} + +/// Arrow's unshredder preserves the source metadata but rebuilds object slots in UTF-8 order. +/// Rebuild from the physical shredding state so Spark's metadata insertion order, UTF-16 object +/// headers, typed scalar widths, and residual scalar bytes all remain compatible. +fn rebuild_shredded_variant_for_spark( + source: &VariantArray, + value: &ArrayRef, + metadata: &ArrayRef, + parent_nulls: Option<&NullBuffer>, +) -> DataFusionResult<(ArrayRef, ArrayRef)> { + let source_metadata = cast(source.metadata_field().as_ref(), &DataType::Binary)?; + let source_metadata = source_metadata.as_binary::(); + let source_state = source.shredding_state().borrow(); + let value = value.as_binary::(); + let metadata = metadata.as_binary::(); + let mut value_output = BinaryBuilder::new(); + let mut metadata_output = BinaryBuilder::new(); + + for index in 0..value.len() { + if parent_nulls.is_some_and(|nulls| nulls.is_null(index)) { + value_output.append_null(); + metadata_output.append_null(); + continue; + } + if value.is_null(index) || metadata.is_null(index) { + return Err(DataFusionError::Execution(format!( + "Variant value or metadata is null at row {index}" + ))); + } + + let (rebuilt_value, rebuilt_metadata) = catch_unwind(AssertUnwindSafe( + || -> Result<(Vec, Vec), ArrowError> { + let source_metadata = VariantMetadata::new(source_metadata.value(index)); + let semantic_metadata = VariantMetadata::try_new(metadata.value(index))?; + let semantic = Variant::new_with_metadata(semantic_metadata, value.value(index)); + let mut field_names = Vec::new(); + collect_shredded_field_names( + source_state.clone(), + index, + &source_metadata, + &mut field_names, + &mut HashSet::new(), + ) + .map_err(|error| { + ArrowError::InvalidArgumentError(format!( + "Failed to collect Spark Variant metadata: {error}" + )) + })?; + + let mut metadata_builder = + WritableMetadataBuilder::from_iter(field_names.iter().map(String::as_str)); + metadata_builder.finish(); + let mut rebuilt_metadata = metadata_builder.into_inner(); + // Spark's VariantBuilder never marks its insertion-ordered dictionary as sorted. + rebuilt_metadata[0] &= !0x10; + let target = VariantMetadata::new(&rebuilt_metadata); + let rebuilt_value = spark_shredded_variant_bytes( + source_state.clone(), + index, + semantic, + &source_metadata, + &target, + ) + .map_err(|error| { + ArrowError::InvalidArgumentError(format!( + "Failed to rebuild Spark Variant value: {error}" + )) + })?; + Ok((rebuilt_value, rebuilt_metadata)) + }, + )) + .map_err(|_| { + DataFusionError::Execution(format!("Invalid shredded Variant at row {index}")) + })??; + value_output.append_value(rebuilt_value); + metadata_output.append_value(rebuilt_metadata); + } + + Ok(( + Arc::new(value_output.finish()), + Arc::new(metadata_output.finish()), + )) +} + +/// Reorder object keys for either Arrow's UTF-8 order or Spark's Java UTF-16 order. Preserve +/// already-compatible values byte-for-byte and retain the original metadata dictionary. +/// SPARK-58949 fixes this mismatch and retains legacy lookup compatibility. Once every Spark +/// profile supported by Comet contains that fix, #5474 can remove this output-side UTF-16 rewrite; +/// the input-side residual rewrite above remains necessary for historical Spark-written files. +/// https://issues.apache.org/jira/browse/SPARK-58949 +/// https://github.com/apache/datafusion-comet/issues/5474 +/// https://github.com/apache/parquet-java/issues/3735 +fn reorder_variant_values( + value: &ArrayRef, + metadata: &ArrayRef, + parent_nulls: Option<&NullBuffer>, + order: VariantObjectKeyOrder, + allow_null_value: bool, +) -> DataFusionResult { + let value = value.as_any().downcast_ref::().unwrap(); + let metadata = metadata.as_any().downcast_ref::().unwrap(); + let mut output = BinaryBuilder::new(); + + for index in 0..value.len() { + if parent_nulls.is_some_and(|nulls| nulls.is_null(index)) { + output.append_null(); + continue; + } + if value.is_null(index) { + if allow_null_value { + output.append_null(); + continue; + } + return Err(DataFusionError::Execution(format!( + "Variant value is null at row {index}" + ))); + } + if metadata.is_null(index) { + return Err(DataFusionError::Execution(format!( + "Variant metadata is null at row {index}" + ))); + } + + let rebuilt = catch_unwind(AssertUnwindSafe(|| -> DataFusionResult>> { + // Spark encodes empty object keys with equal metadata offsets, which Arrow 58.4's + // full validator rejects. Keep shallow parsing and all accesses inside this boundary. + // https://github.com/apache/arrow-rs/blob/58.4.0/parquet-variant/src/variant/metadata.rs#L307-L317 + // Upstream fix: https://github.com/apache/arrow-rs/pull/10352 + let metadata = VariantMetadata::new(metadata.value(index)); + let variant = Variant::new_with_metadata(metadata.clone(), value.value(index)); + if is_compatible_variant(&variant, order) { + return Ok(None); + } + let value = match order { + VariantObjectKeyOrder::ArrowUtf8 => arrow_variant_bytes(&metadata, variant)?, + VariantObjectKeyOrder::SparkUtf16 => spark_variant_bytes(&metadata, variant)?, + }; + Ok(Some(value)) + })) + .map_err(|_| { + DataFusionError::Execution(format!("Invalid Variant value at row {index}")) + })??; + output.append_value(rebuilt.as_deref().unwrap_or_else(|| value.value(index))); + } + + Ok(Arc::new(output.finish())) +} + +#[cfg(test)] +mod tests; diff --git a/native/core/src/parquet/cast_column/variant/tests.rs b/native/core/src/parquet/cast_column/variant/tests.rs new file mode 100644 index 00000000000..e011a1708e1 --- /dev/null +++ b/native/core/src/parquet/cast_column/variant/tests.rs @@ -0,0 +1,1226 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. +use super::*; +use arrow::array::{ + Decimal128Array, DictionaryArray, FixedSizeBinaryArray, FixedSizeListArray, Int32Array, + Int64Array, TimestampMillisecondArray, UInt16Array, UInt32Array, UInt64Array, UInt8Array, +}; +use arrow::datatypes::{Field, Fields, Int32Type}; +use parquet::{ + arrow::parquet_to_arrow_schema, + schema::{parser::parse_message_type, types::SchemaDescriptor}, + variant::{VariantDecimal16, VariantType}, +}; + +fn unicode_object_keys() -> Vec { + let mut keys = (0..30).map(|i| format!("k{i:02}")).collect::>(); + keys.push("\u{e000}".to_string()); + keys.push("😀".to_string()); + keys +} + +fn assert_spark_unicode_variant(variant: Variant<'_, '_>) { + let Variant::Object(object) = variant else { + panic!("expected object") + }; + let fields = object.iter().collect::>(); + + assert_eq!(fields.len(), 32); + assert_eq!(fields[30].0, "😀"); + assert_eq!(fields[31].0, "\u{e000}"); + let emoji = fields + .binary_search_by(|(name, _)| name.encode_utf16().cmp("😀".encode_utf16())) + .unwrap(); + assert_eq!(fields[emoji].1.as_int64(), Some(531)); + let private_use = fields + .binary_search_by(|(name, _)| name.encode_utf16().cmp("\u{e000}".encode_utf16())) + .unwrap(); + assert_eq!(fields[private_use].1.as_int64(), Some(30)); +} + +fn assert_spark_unicode_object(output: &StructArray) { + let value = output.column(0).as_binary::(); + let metadata = output.column(1).as_binary::(); + assert_spark_unicode_variant(Variant::new(metadata.value(0), value.value(0))); +} + +fn normalize_typed_value(typed_value: ArrayRef, field_names: &[&str]) -> VariantArray { + let (metadata_bytes, _) = VariantBuilder::new() + .with_field_names(field_names.iter().copied()) + .finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + VariantArray::try_new(output.as_ref()).unwrap() +} + +#[test] +fn test_normalize_shredded_variant_widens_unsigned_values() { + let metadata_builder = VariantBuilder::new().with_field_names(["u8", "u16", "u32", "u64"]); + 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, + ), + ( + "u64", + Arc::new(UInt64Array::from(vec![u64::MAX])) as ArrayRef, + ), + ]; + let mut object_fields = Vec::with_capacity(fields.len()); + let mut object_columns = Vec::with_capacity(fields.len()); + 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))); + assert_eq!( + object.get("u64"), + Some(Variant::Decimal16( + VariantDecimal16::try_new(18_446_744_073_709_551_615_i128, 0).unwrap() + )) + ); +} + +#[test] +fn test_normalize_variant_decodes_dictionary_storage() { + let mut builder = VariantBuilder::new(); + builder.append_value(42_i64); + let (metadata_bytes, value_bytes) = builder.finish(); + let values: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let value: ArrayRef = + Arc::new(DictionaryArray::::try_new(Int32Array::from(vec![0]), values).unwrap()); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", value.data_type().clone(), false), + Field::new("metadata", DataType::Binary, false), + ]), + vec![value, metadata], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!(output.value(0), Variant::Int64(42)); + + let values: ArrayRef = Arc::new(Int64Array::from(vec![42])); + let dictionary: ArrayRef = + Arc::new(DictionaryArray::::try_new(Int32Array::from(vec![0]), values).unwrap()); + let child: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + dictionary.data_type().clone(), + false, + )]), + vec![dictionary], + None, + ) + .unwrap(), + ); + let object: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("d", child.data_type().clone(), false)]), + vec![child], + None, + ) + .unwrap(), + ); + let output = normalize_typed_value(object, &["d"]); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + assert_eq!(object.get("d"), Some(Variant::Int8(42))); +} + +#[test] +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_converts_fixed_size_binary() { + let bytes = (0_u8..16).collect::>(); + let typed_value: ArrayRef = Arc::new(FixedSizeBinaryArray::from(vec![Some(bytes.as_slice())])); + + let output = normalize_typed_value(typed_value, &[]); + assert_eq!(output.value(0), Variant::Binary(&bytes)); +} + +#[test] +fn test_normalize_shredded_variant_rejects_uuid_annotation() { + let bytes = (0_u8..16).collect::>(); + let typed_value: ArrayRef = Arc::new(FixedSizeBinaryArray::from(vec![Some(bytes.as_slice())])); + let parquet_schema = parse_message_type( + "message root { required fixed_len_byte_array(16) typed_value (UUID); }", + ) + .unwrap(); + let arrow_schema = + parquet_to_arrow_schema(&SchemaDescriptor::new(Arc::new(parquet_schema)), None).unwrap(); + let typed_value_field = Arc::clone(&arrow_schema.fields()[0]); + assert_eq!(typed_value_field.extension_type_name(), Some("arrow.uuid")); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(&[][..])])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Arc::new(Field::new("metadata", DataType::Binary, false)), + typed_value_field, + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + + let error = normalize_variant_storage(&physical).unwrap_err(); + assert!(error.to_string().contains("Parquet UUID is not supported")); +} + +#[test] +fn test_normalize_shredded_variant_compacts_spark_integer_widths() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![ + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + Some(metadata_bytes.as_slice()), + ])); + let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![ + 1, + i64::from(i8::MAX) + 1, + i64::from(i16::MAX) + 1, + i64::from(i32::MAX) + 1, + ])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", DataType::Int64, false), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!(output.value(0), Variant::Int8(1)); + assert_eq!(output.value(1), Variant::Int16(128)); + assert_eq!(output.value(2), Variant::Int32(32_768)); + assert_eq!(output.value(3), Variant::Int64(2_147_483_648)); +} + +#[test] +fn test_compact_spark_typed_variant_canonicalizes_nan() { + let Variant::Float(float) = + compact_spark_typed_variant(Variant::Float(f32::from_bits(0x7fc0_0001))) + else { + panic!("expected float") + }; + assert_eq!(float.to_bits(), 0x7fc0_0000); + + let Variant::Double(double) = + compact_spark_typed_variant(Variant::Double(f64::from_bits(0x7ff8_0000_0000_0001))) + else { + panic!("expected double") + }; + assert_eq!(double.to_bits(), 0x7ff8_0000_0000_0000); + + let (metadata, _) = VariantBuilder::new().finish(); + let metadata = VariantMetadata::new(&metadata); + let residual = f32::from_bits(0x7fc0_0001); + let bytes = spark_variant_bytes(&metadata, Variant::Float(residual)).unwrap(); + let Variant::Float(output) = Variant::new_with_metadata(metadata, &bytes) else { + panic!("expected residual float") + }; + assert_eq!(output.to_bits(), residual.to_bits()); +} + +#[test] +fn test_normalize_shredded_variant_rejects_missing_required_value() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let typed_value: ArrayRef = Arc::new(Int64Array::from(vec![None])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", DataType::Int64, true), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + assert!(normalize_variant_array(&physical, &target_field).is_err()); +} + +#[test] +fn test_normalize_shredded_variant_uses_physical_metadata_order() { + let (metadata_bytes, _) = VariantBuilder::new().with_field_names(["a", "b"]).finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let mut fields = Vec::new(); + let mut columns = Vec::new(); + for (name, value) in [("b", 2_i64), ("a", 1_i64)] { + let child = StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![value]))], + None, + ) + .unwrap(); + fields.push(Field::new(name, child.data_type().clone(), false)); + columns.push(Arc::new(child) as ArrayRef); + } + let typed_value: ArrayRef = + Arc::new(StructArray::try_new(fields.into(), columns, None).unwrap()); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_eq!( + output.column(1).as_binary::().value(0), + &[0x01, 2, 0, 1, 2, b'b', b'a'] + ); + + let mut expected = VariantBuilder::new().with_field_names(["b", "a"]); + let mut object = expected.new_object(); + object.insert("b", 2_i8); + object.insert("a", 1_i8); + object.finish(); + let (_, expected_value) = expected.finish(); + assert_eq!(output.column(0).as_binary::().value(0), expected_value); +} + +#[test] +fn test_normalize_shredded_variant_preserves_residual_scalar_width() { + let mut builder = VariantBuilder::new().with_field_names(["known", "residual"]); + let mut object = builder.new_object(); + object.insert("residual", 1_i64); + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let known: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![2]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), + vec![known], + None, + ) + .unwrap(), + ); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("value", DataType::Binary, true), + Field::new("typed_value", typed_value.data_type().clone(), true), + ]), + vec![metadata, value, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let Variant::Object(object) = output.value(0) else { + panic!("expected object") + }; + assert_eq!(object.get("known"), Some(Variant::Int8(2))); + assert_eq!(object.get("residual"), Some(Variant::Int64(1))); +} + +#[test] +fn test_normalize_shredded_variant_compacts_spark_decimal_width() { + let (metadata_bytes, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let typed_value: ArrayRef = Arc::new( + Decimal128Array::from(vec![123_i128]) + .with_precision_and_scale(38, 2) + .unwrap(), + ); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!( + output.value(0), + Variant::Decimal4(VariantDecimal4::try_new(123, 2).unwrap()) + ); +} + +#[test] +fn test_normalize_shredded_variant_uses_spark_object_key_order() { + let keys = unicode_object_keys(); + + let metadata_builder = VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); + let (metadata_bytes, _) = metadata_builder.finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + + let mut object_fields = Vec::with_capacity(keys.len()); + let mut object_columns = Vec::with_capacity(keys.len()); + for (index, key) in keys.iter().enumerate() { + let value = if key == "😀" { 531 } else { index as i64 }; + let field = Field::new("typed_value", DataType::Int64, false); + let column: ArrayRef = Arc::new(Int64Array::from(vec![value])); + let shredded_field = + StructArray::try_new(Fields::from(vec![field]), vec![column], None).unwrap(); + object_fields.push(Field::new(key, shredded_field.data_type().clone(), false)); + object_columns.push(Arc::new(shredded_field) as ArrayRef); + } + let typed_value: ArrayRef = + Arc::new(StructArray::try_new(object_fields.into(), object_columns, None).unwrap()); + let physical_fields = Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]); + let physical: ArrayRef = + Arc::new(StructArray::try_new(physical_fields, vec![metadata, typed_value], None).unwrap()); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false).with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_spark_unicode_object(output); +} + +#[test] +fn test_normalize_nested_shredded_variant_uses_spark_object_key_order() { + let keys = unicode_object_keys(); + let field_names = std::iter::once("nested") + .chain(keys.iter().map(String::as_str)) + .collect::>(); + let (metadata_bytes, _) = VariantBuilder::new().with_field_names(field_names).finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + + let mut nested_fields = Vec::with_capacity(keys.len()); + let mut nested_columns = Vec::with_capacity(keys.len()); + for (index, key) in keys.iter().enumerate() { + let value = if key == "😀" { 531 } else { index as i64 }; + let state = StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![value]))], + None, + ) + .unwrap(); + nested_fields.push(Field::new(key, state.data_type().clone(), false)); + nested_columns.push(Arc::new(state) as ArrayRef); + } + let nested_value: ArrayRef = + Arc::new(StructArray::try_new(nested_fields.into(), nested_columns, None).unwrap()); + let nested_state: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "typed_value", + nested_value.data_type().clone(), + false, + )]), + vec![nested_value], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "nested", + nested_state.data_type().clone(), + false, + )]), + vec![nested_state], + None, + ) + .unwrap(), + ); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), false), + ]), + vec![metadata, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + let Variant::Object(object) = output.value(0) else { + panic!("expected outer object") + }; + assert_spark_unicode_variant(object.get("nested").expect("nested field")); +} + +#[test] +fn test_normalize_partially_shredded_spark_object_key_order() { + let keys = unicode_object_keys(); + let mut builder = VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); + let mut object = builder.new_object(); + for (index, key) in keys.iter().enumerate().skip(1) { + object.insert(key, if key == "😀" { 531 } else { index as i64 }); + } + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let spark_value = reorder_variant_values( + &value, + &metadata, + None, + VariantObjectKeyOrder::SparkUtf16, + false, + ) + .unwrap(); + + let shredded_k00: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![0]))], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "k00", + shredded_k00.data_type().clone(), + false, + )]), + vec![shredded_k00], + None, + ) + .unwrap(), + ); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("value", DataType::Binary, true), + Field::new("typed_value", typed_value.data_type().clone(), true), + ]), + vec![metadata, spark_value, typed_value], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + assert_spark_unicode_object(output.as_struct()); +} + +#[test] +fn test_normalize_unshredded_variant_uses_spark_object_key_order() { + let keys = unicode_object_keys(); + let mut builder = VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); + let mut object = builder.new_object(); + for (index, key) in keys.iter().enumerate() { + object.insert(key, if key == "😀" { 531 } else { index as i64 }); + } + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + + let Variant::Object(canonical) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + let canonical_fields = canonical.iter().collect::>(); + assert_eq!(canonical_fields[30].0, "\u{e000}"); + assert_eq!(canonical_fields[31].0, "😀"); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let physical: ArrayRef = + Arc::new(StructArray::try_new(physical_fields, vec![value, metadata], None).unwrap()); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false).with_extension_type(VariantType), + ); + + let first = normalize_variant_array(&physical, &target_field).unwrap(); + let first_value = first + .as_struct() + .column(0) + .as_binary::() + .value(0) + .to_vec(); + assert_spark_unicode_object(first.as_struct()); + + // Spark-produced already-unshredded input is UTF-16 ordered. Normalizing it again must + // remain valid without Arrow's UTF-8-order full validation. + let second = normalize_variant_array(&first, &target_field).unwrap(); + assert_spark_unicode_object(second.as_struct()); + assert_eq!( + second.as_struct().column(0).as_binary::().value(0), + first_value + ); +} + +#[test] +fn test_normalize_spark_ordered_variant_preserves_value_bytes() { + let mut builder = VariantBuilder::new(); + let mut object = builder.new_object(); + object.insert("b", 1_i64); + object.insert("a", 2_i64); + object.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + + let Variant::Object(object) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + assert_eq!( + object.iter().map(|(name, _)| name).collect::>(), + vec!["a", "b"] + ); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let physical: ArrayRef = + Arc::new(StructArray::try_new(physical_fields, vec![value, metadata], None).unwrap()); + let target_fields = Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]); + let target_field = Arc::new( + Field::new("v", DataType::Struct(target_fields), false).with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + assert_eq!( + output.as_struct().column(0).as_binary::().value(0), + value_bytes + ); +} + +#[test] +fn test_normalize_variant_preserves_empty_object_keys() { + let mut builder = VariantBuilder::new(); + let mut object = builder.new_object(); + object.insert("", 1_i64); + let mut nested = object.new_object("nested"); + nested.insert("", 2_i64); + nested.finish(); + object.finish(); + let (mut metadata_bytes, value_bytes) = builder.finish(); + + // Spark leaves the metadata dictionary unsorted. Equal offsets encode the empty key. + metadata_bytes[0] &= !0x10; + assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]), + vec![value, metadata], + None, + ) + .unwrap(), + ); + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = output.as_struct(); + assert_eq!(output.column(0).as_binary::().value(0), value_bytes); + assert_eq!(output.column(1).as_binary::().value(0), metadata_bytes); + + let Variant::Object(object) = Variant::new(&metadata_bytes, &value_bytes) else { + panic!("expected object") + }; + assert_eq!(object.get(""), Some(Variant::from(1_i64))); + let Variant::Object(nested) = object.get("nested").unwrap() else { + panic!("expected nested object") + }; + assert_eq!(nested.get(""), Some(Variant::from(2_i64))); +} + +#[test] +fn test_compatibility_retry_rejects_non_spark_residual_order() { + let mut value_builder = VariantBuilder::new().with_field_names(["a", "b", "c"]); + let mut object = value_builder.new_object(); + object.insert("a", 1_i64); + object.insert("b", 2_i64); + object.insert("c", 3_i64); + object.finish(); + let (_, value) = value_builder.finish(); + + // Reassign the field IDs so the encoded slots read as b, a, c. That is neither canonical + // UTF-8 order nor Spark's UTF-16 order and must not be repaired by the compatibility retry. + let mut metadata_builder = WritableMetadataBuilder::from_iter(["b", "a", "c"]); + metadata_builder.finish(); + let metadata = metadata_builder.into_inner(); + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])); + + let error = rewrite_residual_values( + &value, + metadata.as_binary::(), + metadata.as_binary::(), + &[Some(0)], + false, + ) + .unwrap_err(); + assert!(error.to_string().contains("neither UTF-8 nor Spark UTF-16")); +} + +#[test] +fn test_empty_key_retry_rejects_non_spark_metadata() { + let mut metadata_builder = WritableMetadataBuilder::from_iter(["", "b", "a"]); + metadata_builder.finish(); + let mut metadata = metadata_builder.into_inner(); + metadata[0] |= 0x10; + + let mut value_builder = VariantBuilder::new(); + value_builder.append_value(1_i64); + let (_, value) = value_builder.finish(); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ]), + vec![ + Arc::new(BinaryArray::from(vec![Some(value.as_slice())])), + Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])), + ], + None, + ) + .unwrap(), + ); + let variant = VariantArray::try_new(physical.as_ref()).unwrap(); + assert!(canonicalize_spark_empty_key_metadata(&variant) + .unwrap() + .is_none()); +} + +#[test] +fn test_normalize_partially_shredded_nested_unicode_and_empty_keys() { + let keys = unicode_object_keys(); + let mut builder = VariantBuilder::new().with_field_names(["nested", "known"]); + let mut residual = builder.new_object(); + residual.insert("", -1_i64); + for (index, key) in keys.iter().enumerate() { + residual.insert(key, if key == "😀" { 531_i64 } else { index as i64 }); + } + residual.finish(); + let (metadata_bytes, value_bytes) = builder.finish(); + assert!(VariantMetadata::try_new(&metadata_bytes).is_err()); + + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value_bytes.as_slice())])); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata_bytes.as_slice())])); + let value = reorder_variant_values( + &value, + &metadata, + None, + VariantObjectKeyOrder::SparkUtf16, + false, + ) + .unwrap(); + let known: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("typed_value", DataType::Int64, false)]), + vec![Arc::new(Int64Array::from(vec![99]))], + None, + ) + .unwrap(), + ); + let nested_typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("known", known.data_type().clone(), false)]), + vec![known], + None, + ) + .unwrap(), + ); + let nested: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("value", DataType::Binary, true), + Field::new("typed_value", nested_typed_value.data_type().clone(), true), + ]), + vec![value, nested_typed_value], + None, + ) + .unwrap(), + ); + let typed_value: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new( + "nested", + nested.data_type().clone(), + false, + )]), + vec![nested], + None, + ) + .unwrap(), + ); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed_value.data_type().clone(), 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), + ); + + 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") + }; + let Variant::Object(nested) = object.get("nested").unwrap() else { + panic!("expected nested object") + }; + let fields = nested.iter().collect::>(); + assert_eq!(fields.len(), 34); + assert_eq!(fields[0].0, ""); + assert_eq!(fields[32].0, "😀"); + assert_eq!(fields[33].0, "\u{e000}"); + assert_eq!(nested.get("known").unwrap().as_int64(), Some(99)); + let emoji = fields + .binary_search_by(|(name, _)| name.encode_utf16().cmp("😀".encode_utf16())) + .unwrap(); + assert_eq!(fields[emoji].1.as_int64(), Some(531)); +} + +#[test] +fn test_normalize_nested_list_residuals_tracks_root_metadata() { + fn row(keys: &[&str]) -> DataFusionResult<(Vec, Vec)> { + let mut builder = VariantBuilder::new().with_field_names(keys.iter().copied()); + let mut object = builder.new_object(); + for (index, key) in keys.iter().enumerate() { + object.insert(key, index as i64); + } + object.finish(); + let (metadata, value) = builder.finish(); + let metadata_array: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])); + let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value.as_slice())])); + let value = reorder_variant_values( + &value, + &metadata_array, + None, + VariantObjectKeyOrder::SparkUtf16, + false, + )?; + Ok((metadata, value.as_binary::().value(0).to_vec())) + } + + let (metadata0, value0) = row(&["a", "\u{e000}", "😀"]).unwrap(); + let (metadata1, value1) = row(&["b", "zz", "\u{ffff}", "𐀀"]).unwrap(); + let states: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![Field::new("value", DataType::Binary, true)]), + vec![Arc::new(BinaryArray::from(vec![ + Some(value0.as_slice()), + Some(value1.as_slice()), + ]))], + None, + ) + .unwrap(), + ); + let list: ArrayRef = Arc::new( + arrow::array::ListArray::try_new( + Arc::new(Field::new("element", states.data_type().clone(), false)), + arrow::buffer::OffsetBuffer::new(vec![0, 1, 2].into()), + states, + None, + ) + .unwrap(), + ); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![ + Some(metadata0.as_slice()), + Some(metadata1.as_slice()), + ])); + let physical: ArrayRef = Arc::new( + StructArray::try_new( + Fields::from(vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", list.data_type().clone(), false), + ]), + vec![metadata, list], + None, + ) + .unwrap(), + ); + + let target_field = Arc::new( + Field::new( + "v", + DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])), + false, + ) + .with_extension_type(VariantType), + ); + let output = normalize_variant_array(&physical, &target_field).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + for (index, key) in ["😀", "𐀀"].into_iter().enumerate() { + let Variant::List(list) = output.value(index) else { + panic!("expected list") + }; + let Variant::Object(object) = list.get(0).unwrap() else { + panic!("expected object") + }; + assert_eq!(object.get(key).unwrap().as_int64(), Some(index as i64 + 2)); + } +} + +#[test] +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)); +} diff --git a/native/core/src/parquet/eager_page_index_reader_factory.rs b/native/core/src/parquet/eager_page_index_reader_factory.rs index 278814c4bf9..34bef7d56c6 100644 --- a/native/core/src/parquet/eager_page_index_reader_factory.rs +++ b/native/core/src/parquet/eager_page_index_reader_factory.rs @@ -45,7 +45,13 @@ //! //! Filed upstream as apache/datafusion#23978. Revert this once the opener merges its deferred //! page-index load back into `FileMetadataCache` instead of bypassing it. +//! +//! For unencrypted scans that project Variant, this factory also replaces the advisory +//! `ARROW:schema` footer entry with physical Parquet schema inference. The only added hint maps +//! Parquet ENUM leaves to Arrow Utf8, matching Spark's physical-schema interpretation: +//! 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 +use arrow::datatypes::{DataType, FieldRef, Schema}; use bytes::Bytes; use datafusion::common::Result as DFResult; use datafusion::datasource::physical_plan::parquet::metadata::DFParquetMetadata; @@ -58,9 +64,17 @@ use datafusion_datasource::PartitionedFile; use futures::future::BoxFuture; use futures::FutureExt; use object_store::ObjectStore; -use parquet::arrow::arrow_reader::ArrowReaderOptions; use parquet::arrow::async_reader::{AsyncFileReader, ParquetObjectReader}; -use parquet::file::metadata::{PageIndexPolicy, ParquetMetaData}; +use parquet::arrow::{ + arrow_reader::ArrowReaderOptions, encode_arrow_schema, parquet_to_arrow_schema, + ARROW_SCHEMA_META_KEY, +}; +use parquet::basic::{ConvertedType, LogicalType}; +use parquet::errors::{ParquetError, Result as ParquetResult}; +use parquet::file::metadata::{ + FileMetaData, KeyValue, PageIndexPolicy, ParquetMetaData, ParquetMetaDataBuilder, +}; +use parquet::schema::types::{ColumnDescPtr, SchemaDescriptor}; use std::fmt::Debug; use std::ops::Range; use std::sync::Arc; @@ -69,13 +83,19 @@ use std::sync::Arc; pub struct EagerPageIndexReaderFactory { store: Arc, metadata_cache: Arc, + skip_arrow_schema: bool, } impl EagerPageIndexReaderFactory { - pub fn new(store: Arc, metadata_cache: Arc) -> Self { + pub fn new( + store: Arc, + metadata_cache: Arc, + skip_arrow_schema: bool, + ) -> Self { Self { store, metadata_cache, + skip_arrow_schema, } } } @@ -109,6 +129,7 @@ impl ParquetFileReaderFactory for EagerPageIndexReaderFactory { partitioned_file, metadata_cache: Arc::clone(&self.metadata_cache), metadata_size_hint, + skip_arrow_schema: self.skip_arrow_schema, })) } } @@ -120,6 +141,134 @@ struct EagerPageIndexReader { partitioned_file: PartitionedFile, metadata_cache: Arc, metadata_size_hint: Option, + skip_arrow_schema: bool, +} + +fn is_enum_column(column: &ColumnDescPtr) -> bool { + matches!(column.logical_type_ref(), Some(LogicalType::Enum)) + || column.converted_type() == ConvertedType::ENUM +} + +fn spark_enum_field( + field: &FieldRef, + columns: &[ColumnDescPtr], + column_index: &mut usize, +) -> ParquetResult { + let rewrite = + |field: &FieldRef, data_type| Arc::new(field.as_ref().clone().with_data_type(data_type)); + let data_type = match field.data_type() { + DataType::Struct(fields) => DataType::Struct( + fields + .iter() + .map(|field| spark_enum_field(field, columns, column_index)) + .collect::>>()? + .into(), + ), + DataType::List(child) => DataType::List(spark_enum_field(child, columns, column_index)?), + DataType::LargeList(child) => { + DataType::LargeList(spark_enum_field(child, columns, column_index)?) + } + DataType::FixedSizeList(child, size) => { + DataType::FixedSizeList(spark_enum_field(child, columns, column_index)?, *size) + } + DataType::ListView(child) => { + DataType::ListView(spark_enum_field(child, columns, column_index)?) + } + DataType::LargeListView(child) => { + DataType::LargeListView(spark_enum_field(child, columns, column_index)?) + } + DataType::Map(child, sorted) => { + DataType::Map(spark_enum_field(child, columns, column_index)?, *sorted) + } + _ => { + let column = columns.get(*column_index).ok_or_else(|| { + ParquetError::General( + "Arrow schema contains more leaves than the Parquet schema".to_string(), + ) + })?; + *column_index += 1; + if is_enum_column(column) { + DataType::Utf8 + } else { + return Ok(Arc::clone(field)); + } + } + }; + Ok(rewrite(field, data_type)) +} + +/// Arrow maps Parquet ENUM to Binary, while Spark maps it to String. Build the smallest physical +/// schema hint needed to preserve Spark's interpretation without restoring the file's advisory +/// `ARROW:schema` types such as Date64 or Decimal256. +/// https://github.com/apache/arrow-rs/blob/58.4.0/parquet/src/arrow/schema/primitive.rs#L285-L293 +/// https://github.com/apache/spark/blob/v4.2.0/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/parquet/ParquetSchemaConverter.scala#L357-L363 +fn spark_enum_schema(schema: &SchemaDescriptor) -> ParquetResult> { + let columns = schema.columns(); + if !columns.iter().any(is_enum_column) { + return Ok(None); + } + + let arrow_schema = parquet_to_arrow_schema(schema, None)?; + let mut column_index = 0; + let fields = arrow_schema + .fields() + .iter() + .map(|field| spark_enum_field(field, columns, &mut column_index)) + .collect::>>()?; + if column_index != columns.len() { + return Err(ParquetError::General( + "Parquet schema contains more leaves than the Arrow schema".to_string(), + )); + } + Ok(Some(Schema::new_with_metadata( + fields, + arrow_schema.metadata().clone(), + ))) +} + +/// Ignore the file's Arrow schema hint so ambiguous leaves such as Date64 retain Spark semantics. +/// Add back only a physical-schema-derived hint for Parquet ENUM, which Spark reads as String. +fn with_spark_arrow_schema(metadata: Arc) -> ParquetResult> { + let file = metadata.file_metadata(); + let has_arrow_schema = file.key_value_metadata().is_some_and(|key_values| { + key_values + .iter() + .any(|key_value| key_value.key == ARROW_SCHEMA_META_KEY) + }); + let enum_schema = spark_enum_schema(file.schema_descr())?; + if !has_arrow_schema && enum_schema.is_none() { + return Ok(metadata); + } + + let mut key_values = file + .key_value_metadata() + .into_iter() + .flatten() + .filter(|key_value| key_value.key != ARROW_SCHEMA_META_KEY) + .cloned() + .collect::>(); + if let Some(schema) = enum_schema { + key_values.push(KeyValue { + key: ARROW_SCHEMA_META_KEY.to_string(), + value: Some(encode_arrow_schema(&schema)), + }); + } + + let file = FileMetaData::new( + file.version(), + file.num_rows(), + file.created_by().map(str::to_owned), + Some(key_values), + file.schema_descr_ptr(), + file.column_orders().cloned(), + ); + Ok(Arc::new( + ParquetMetaDataBuilder::new(file) + .set_row_groups(metadata.row_groups().to_vec()) + .set_column_index(metadata.column_index().cloned()) + .set_offset_index(metadata.offset_index().cloned()) + .build(), + )) } impl AsyncFileReader for EagerPageIndexReader { @@ -151,17 +300,19 @@ impl AsyncFileReader for EagerPageIndexReader { let metadata_cache = Arc::clone(&self.metadata_cache); let store = Arc::clone(&self.store); let metadata_size_hint = self.metadata_size_hint; + let skip_arrow_schema = self.skip_arrow_schema; async move { let file_decryption_properties = options .and_then(|o| o.file_decryption_properties()) .map(Arc::clone); + let encrypted = file_decryption_properties.is_some(); let page_index_policy = if file_decryption_properties.is_none() { Some(PageIndexPolicy::Optional) } else { options.map(|o| o.column_index_policy()) }; - DFParquetMetadata::new(store.as_ref(), &object_meta) + let metadata = DFParquetMetadata::new(store.as_ref(), &object_meta) .with_decryption_properties(file_decryption_properties) .with_file_metadata_cache(Some(metadata_cache)) .with_metadata_size_hint(metadata_size_hint) @@ -173,7 +324,12 @@ impl AsyncFileReader for EagerPageIndexReader { "Failed to fetch metadata for file {}: {e}", object_meta.location, )) - }) + })?; + Ok(if skip_arrow_schema && !encrypted { + with_spark_arrow_schema(metadata)? + } else { + metadata + }) } .boxed() } 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/parquet_exec.rs b/native/core/src/parquet/parquet_exec.rs index 1308ce97fca..4659f8a478b 100644 --- a/native/core/src/parquet/parquet_exec.rs +++ b/native/core/src/parquet/parquet_exec.rs @@ -16,6 +16,7 @@ // under the License. use crate::execution::operators::ExecutionError; +use crate::execution::serde::is_variant_field; use crate::parquet::eager_page_index_reader_factory::EagerPageIndexReaderFactory; use crate::parquet::encryption_support::{CometEncryptionConfig, ENCRYPTION_FACTORY_ID}; use crate::parquet::parquet_support::SparkParquetOptions; @@ -164,8 +165,12 @@ pub(crate) fn init_datasource_exec( let runtime_env = session_ctx.runtime_env(); let store = runtime_env.object_store(&object_store_url)?; let metadata_cache = runtime_env.cache_manager.get_file_metadata_cache(); + let skip_arrow_schema = required_schema + .fields() + .iter() + .any(|field| is_variant_field(field)); parquet_source = parquet_source.with_parquet_file_reader_factory(Arc::new( - EagerPageIndexReaderFactory::new(store, metadata_cache), + EagerPageIndexReaderFactory::new(store, metadata_cache, skip_arrow_schema), )); // Route data filters through `try_pushdown_filters` rather than calling @@ -291,16 +296,161 @@ fn get_options( #[cfg(test)] mod tests { use super::*; - use arrow::array::Int32Array; - use arrow::datatypes::{DataType, Field, Schema}; + use arrow::array::{ + ArrayRef, BinaryArray, Date64Array, Decimal256Array, Int32Array, StructArray, + }; + use arrow::datatypes::{i256, DataType, Field, Fields, Schema}; use arrow::record_batch::RecordBatch; use datafusion::datasource::physical_plan::parquet::metadata::CachedParquetMetaData; use datafusion::physical_plan::ExecutionPlan; use datafusion_comet_spark_expr::test_common::file_util::get_temp_filename; use futures::StreamExt; - use parquet::arrow::ArrowWriter; + use parquet::arrow::{arrow_writer::ArrowWriterOptions, ArrowWriter}; + use parquet::basic::{LogicalType, Repetition, Type as PhysicalType}; + use parquet::data_type::{ + ByteArray, ByteArrayType, DataType as ParquetDataType, FixedLenByteArray, + FixedLenByteArrayType, + }; use parquet::file::properties::{EnabledStatistics, WriterProperties}; + use parquet::file::writer::SerializedFileWriter; + use parquet::schema::types::{Type as ParquetType, TypePtr}; + use parquet::variant::{Variant, VariantArray, VariantBuilder, VariantDecimal4, VariantType}; use std::fs::File; + use std::path::PathBuf; + + fn required_variant_schema() -> SchemaRef { + Arc::new(Schema::new(vec![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)])) + } + + async fn write_and_scan_shredded_variant( + typed_value: ArrayRef, + coerce_types: bool, + ) -> VariantArray { + let (metadata, _) = VariantBuilder::new().finish(); + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata.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 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![physical]).unwrap(); + + let filename = get_temp_filename(); + let file = File::create(&filename).unwrap(); + let properties = WriterProperties::builder() + .set_coerce_types(coerce_types) + .build(); + let mut writer = ArrowWriter::try_new_with_options( + file, + file_schema, + ArrowWriterOptions::new().with_properties(properties), + ) + .unwrap(); + writer.write(&batch).unwrap(); + writer.close().unwrap(); + + scan_variant_file(filename).await + } + + fn write_variant_typed_value( + typed_value: TypePtr, + values: &[T::T], + ) -> PathBuf { + let filename = get_temp_filename(); + let file = File::create(&filename).unwrap(); + let metadata = Arc::new( + ParquetType::primitive_type_builder("metadata", PhysicalType::BYTE_ARRAY) + .with_repetition(Repetition::REQUIRED) + .build() + .unwrap(), + ); + let variant = Arc::new( + ParquetType::group_type_builder("v") + .with_repetition(Repetition::REQUIRED) + .with_logical_type(Some(LogicalType::Variant { + specification_version: None, + })) + .with_fields(vec![metadata, typed_value]) + .build() + .unwrap(), + ); + let schema = Arc::new( + ParquetType::group_type_builder("schema") + .with_fields(vec![variant]) + .build() + .unwrap(), + ); + let mut writer = SerializedFileWriter::new(file, schema, Default::default()).unwrap(); + let mut row_group = writer.next_row_group().unwrap(); + + let (metadata, _) = VariantBuilder::new().finish(); + let metadata = (0..values.len()) + .map(|_| ByteArray::from(metadata.clone())) + .collect::>(); + let mut column = row_group.next_column().unwrap().unwrap(); + column + .typed::() + .write_batch(&metadata, None, None) + .unwrap(); + column.close().unwrap(); + + let mut column = row_group.next_column().unwrap().unwrap(); + column.typed::().write_batch(values, None, None).unwrap(); + column.close().unwrap(); + row_group.close().unwrap(); + writer.close().unwrap(); + filename + } + + async fn scan_variant_file(filename: PathBuf) -> VariantArray { + let partitioned_file = + PartitionedFile::from_path(filename.to_string_lossy().into_owned()).unwrap(); + let session_ctx = Arc::new(SessionContext::new()); + let scan = init_datasource_exec( + required_variant_schema(), + None, + None, + ObjectStoreUrl::local_filesystem(), + vec![vec![partitioned_file]], + None, + None, + None, + "UTC", + true, + false, + false, + false, + &session_ctx, + false, + false, + false, + ) + .unwrap(); + let mut stream = scan.execute(0, session_ctx.task_ctx()).unwrap(); + let batch = stream.next().await.unwrap().unwrap(); + assert!(stream.next().await.is_none()); + VariantArray::try_new(batch.column(0).as_ref()).unwrap() + } // Regression test for #4990: a fresh `TableParquetOptions::new()` ignored session-level // `datafusion.execution.parquet.*` settings entirely, so `spark.comet.datafusion. @@ -432,4 +582,95 @@ mod tests { "cached metadata must include the page index" ); } + + #[tokio::test] + async fn variant_scan_uses_parquet_physical_types_instead_of_arrow_schema_hints() { + let decimal: ArrayRef = Arc::new( + Decimal256Array::from(vec![i256::from_i128(123)]) + .with_precision_and_scale(38, 2) + .unwrap(), + ); + let output = write_and_scan_shredded_variant(decimal, false).await; + assert_eq!( + output.value(0), + Variant::Decimal4(VariantDecimal4::try_new(123, 2).unwrap()) + ); + + let date64: ArrayRef = Arc::new(Date64Array::from(vec![86_400_000])); + let output = write_and_scan_shredded_variant(Arc::clone(&date64), false).await; + assert_eq!(output.value(0).as_int64(), Some(86_400_000)); + + let output = write_and_scan_shredded_variant(date64, true).await; + let Variant::Date(date) = output.value(0) else { + panic!("expected DATE-annotated physical value") + }; + assert_eq!(date.to_string(), "1970-01-02"); + } + + #[tokio::test] + async fn variant_scan_preserves_parquet_enum_string_and_binary_semantics() { + for (logical_type, expected_string) in [ + (Some(LogicalType::Enum), true), + (Some(LogicalType::String), true), + (None, false), + ] { + let typed_value = Arc::new( + ParquetType::primitive_type_builder("typed_value", PhysicalType::BYTE_ARRAY) + .with_repetition(Repetition::REQUIRED) + .with_logical_type(logical_type) + .build() + .unwrap(), + ); + let filename = write_variant_typed_value::( + typed_value, + &[ByteArray::from(b"red".to_vec())], + ); + + let output = scan_variant_file(filename).await; + if expected_string { + assert_eq!(output.value(0).as_string(), Some("red")); + } else { + assert_eq!(output.value(0), Variant::Binary(b"red")); + } + } + } + + #[tokio::test] + async fn variant_scan_reads_wide_physical_decimal_as_decimal128() { + for width in [17, 32] { + let values = [123_i128, -123_i128] + .into_iter() + .map(|value| { + let mut bytes = vec![if value.is_negative() { 0xff } else { 0 }; width]; + bytes[width - 16..].copy_from_slice(&value.to_be_bytes()); + FixedLenByteArray::from(bytes) + }) + .collect::>(); + let typed_value = Arc::new( + ParquetType::primitive_type_builder( + "typed_value", + PhysicalType::FIXED_LEN_BYTE_ARRAY, + ) + .with_repetition(Repetition::REQUIRED) + .with_logical_type(Some(LogicalType::Decimal { + scale: 2, + precision: 38, + })) + .with_length(width as i32) + .with_precision(38) + .with_scale(2) + .build() + .unwrap(), + ); + let filename = write_variant_typed_value::(typed_value, &values); + + let output = scan_variant_file(filename).await; + for (index, value) in [123, -123].into_iter().enumerate() { + assert_eq!( + output.value(index), + Variant::Decimal4(VariantDecimal4::try_new(value, 2).unwrap()) + ); + } + } + } } diff --git a/native/core/src/parquet/schema_adapter.rs b/native/core/src/parquet/schema_adapter.rs index c6586b4681e..d395ae95007 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 @@ -1103,13 +1119,13 @@ 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, ArrayRef, BinaryArray, Date32Array, Decimal128Array, FixedSizeListArray, + Float32Array, Float64Array, Int32Array, Int64Array, StringArray, StructArray, + TimestampMicrosecondArray, TimestampMillisecondArray, UInt16Array, UInt32Array, UInt8Array, }; - use arrow::datatypes::SchemaRef; - use arrow::datatypes::{DataType, Field, Schema}; + use arrow::buffer::NullBuffer; + use arrow::datatypes::{DataType, Field, Fields, Schema, SchemaRef}; use arrow::record_batch::RecordBatch; use datafusion::common::DataFusionError; use datafusion::datasource::listing::PartitionedFile; @@ -1122,7 +1138,8 @@ 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; @@ -1679,16 +1696,218 @@ 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(); + 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(required_variant_field(name, false)); + } + + 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(()) + } + + #[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()?; 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 c69602fc80b..322ede3c688 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, ShimCometWindowGroupLimit, ShimSubqueryBroadcast} +import org.apache.comet.shims.{CometTypeShim, ShimCometStreaming, ShimCometWindowGroupLimit, ShimSubqueryBroadcast} object CometExecRule { @@ -124,6 +124,7 @@ object CometExecRule { */ case class CometExecRule(session: SparkSession) extends Rule[SparkPlan] + with CometTypeShim with ShimSubqueryBroadcast { private lazy val showTransformations = CometConf.COMET_EXPLAIN_TRANSFORMATIONS.get() @@ -733,17 +734,27 @@ 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 ++ dataProducingChildren.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 // 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/CometScanRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala index a524da3af92..020d29652b1 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala @@ -295,10 +295,17 @@ case class CometScanRule(session: SparkSession) if (!CometNativeScan.isSupported(scanExec)) { return None } - if (encryptionEnabled(hadoopConf) && !isEncryptionConfigSupported(hadoopConf)) { + val encrypted = encryptionEnabled(hadoopConf) + if (encrypted && !isEncryptionConfigSupported(hadoopConf)) { withFallbackReason(scanExec, "Native Parquet scan does not support encryption") return None } + if (encrypted && scanExec.requiredSchema.exists(field => isVariantType(field.dataType))) { + withFallbackReason( + scanExec, + "Native Parquet scan does not support encrypted Variant columns") + return None + } // input_file_name, input_file_block_start, and input_file_block_length read from // InputFileBlockHolder, a thread-local set by Spark's FileScanRDD. The native DataFusion // scan does not use FileScanRDD, so these expressions would return empty/default values. @@ -963,8 +970,23 @@ 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) + // The physical Parquet name is unavailable here, so keep a Variant-bearing scan with a + // non-ASCII required name on Spark until the native adapter implements Spark's Unicode + // case-insensitive comparison. Physical Unicode names resolved to ASCII here remain tracked by + // https://github.com/apache/datafusion-comet/issues/5495. + if (!session.sessionState.conf.caseSensitiveAnalysis && + scanExec.requiredSchema.exists(field => isVariantType(field.dataType)) && + scanExec.requiredSchema.exists(field => field.name.exists(_ > '\u007f'))) { + withFallbackReason( + scanExec, + "Native Parquet scan does not support case-insensitive Unicode column names in " + + "Variant scans") + return false + } + 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/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/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..3f2f4b0fef6 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala @@ -23,13 +23,13 @@ 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} 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 @@ -50,14 +50,30 @@ 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 - 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 + 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 */ @@ -102,6 +118,41 @@ 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) + } + + val hasVariant = + scanExec.requiredSchema.fields.exists(field => isVariantType(field.dataType)) + + // Spark's strict mode validates the legacy two-field Variant layout and reports SPARK-47546 + // errors itself: https://issues.apache.org/jira/browse/SPARK-47546 + if (hasVariant && + !SQLConf.get.getConfString("spark.sql.variant.allowReadingShredded").toBoolean) { + withFallbackReason( + scanExec, + "Full native scan disabled because Spark's strict unshredded Variant reader is enabled") + } + + // Spark interprets TIMESTAMP(NANOS) leaves as raw longs in this legacy mode. The logical + // Variant schema does not expose whether a file contains such a shredded child, so preserve + // Spark semantics by falling back before reading any requested Variant value. + if (hasVariant && SQLConf.get.legacyParquetNanosAsLong) { + withFallbackReason( + scanExec, + "Full native scan disabled because spark.sql.legacy.parquet.nanosAsLong is enabled") + } + + // Spark interprets unadjusted Parquet timestamps as TimestampType when NTZ inference is + // disabled, while Arrow exposes them as timezone-free TimestampNTZ values. The Variant schema + // does not expose shredded leaf types, so preserve Spark's vectorized-reader semantics. + if (hasVariant && !SQLConf.get.parquetInferTimestampNTZEnabled) { + withFallbackReason( + scanExec, + "Full native scan disabled because " + + s"${SQLConf.PARQUET_INFER_TIMESTAMP_NTZ_ENABLED.key} is disabled") + } + // the scan is supported if no fallback reasons were added to the node !hasFallbackReason(scanExec) } @@ -153,23 +204,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). @@ -198,8 +241,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..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 { @@ -39,6 +40,15 @@ 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 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 f48955a7da5..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) @@ -49,17 +51,42 @@ 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) + + // 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 328254b340a..69a98502d8f 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,27 @@ -- 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 +-- Config: spark.sql.variant.pushVariantIntoScan=false +-- Config: spark.sql.variant.allowReadingShredded=true +-- Config: spark.sql.variant.forceShreddingSchemaForTest=k00 BIGINT 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('{"b": "hello", "a": 1}'), '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'), + (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. @@ -43,8 +49,31 @@ 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 + +-- 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 + +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 +84,92 @@ 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 + +-- 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 top-level and nested +-- shredded fields so the native scan reconstructs their residual values before Spark's lookup. +statement +SET spark.sql.variant.forceShreddingSchemaForTest=k00 BIGINT, nested STRUCT + +statement +SET spark.sql.variant.writeShredding.enabled=true + +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}')), + (parse_json( + '{"":-2,"k00":0,"nested":{"known":99,"":-1,"k00":0,"k01":1,"k02":2,"k03":3,"k04":4,"k05":5,"k06":6,"k07":7,"k08":8,"k09":9,"k10":10,"k11":11,"k12":12,"k13":13,"k14":14,"k15":15,"k16":16,"k17":17,"k18":18,"k19":19,"k20":20,"k21":21,"k22":22,"k23":23,"k24":24,"k25":25,"k26":26,"k27":27,"k28":28,"k29":29,"\uE000":30,"😀":532}}')) + +statement +SET spark.sql.variant.writeShredding.enabled=false + +statement +SET spark.sql.variant.allowReadingShredded=true + +query +SELECT v FROM test_variant_unicode + +query expect_fallback(type VariantType) +SELECT variant_get(v, '$.😀', 'bigint'), variant_get(v, '$.nested.😀', 'bigint') +FROM test_variant_unicode + +-- 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=b BIGINT, a BIGINT, d DECIMAL(38,2) + +statement +SET spark.sql.variant.writeShredding.enabled=true + +statement +CREATE TABLE test_variant_typed_bytes(v VARIANT) USING parquet + +statement +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_typed_bytes + statement CREATE TABLE test_variant_struct(id INT, s STRUCT, tail STRING) USING parquet diff --git a/spark/src/test/scala/org/apache/comet/exec/CometNativeReaderSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometNativeReaderSuite.scala index fd7d4ee3f46..3145ba72097 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometNativeReaderSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometNativeReaderSuite.scala @@ -24,6 +24,8 @@ import org.scalatest.Tag import org.apache.hadoop.conf.Configuration import org.apache.hadoop.fs.Path +import org.apache.parquet.crypto.DecryptionPropertiesFactory +import org.apache.parquet.crypto.keytools.PropertiesDrivenCryptoFactory import org.apache.parquet.hadoop.{ParquetFileReader, ParquetWriter} import org.apache.parquet.hadoop.api.WriteSupport import org.apache.parquet.hadoop.api.WriteSupport.WriteContext @@ -37,7 +39,7 @@ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{IntegerType, LongType, NullType, StringType, StructType} import org.apache.comet.CometConf -import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus +import org.apache.comet.CometSparkSessionExtensions.{isSpark40Plus, isSpark41Plus} class CometNativeReaderSuite extends CometTestBase with AdaptiveSparkPlanHelper { @@ -68,6 +70,87 @@ class CometNativeReaderSuite extends CometTestBase with AdaptiveSparkPlanHelper } } + test("case-insensitive Unicode Variant column names fall back") { + assume(isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempPath { path => + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + "spark.sql.variant.writeShredding.enabled" -> "false") { + sql("""SELECT parse_json('42') AS `É`, parse_json('43') AS `Σ`, + |parse_json('44') AS `K`, parse_json('45') AS `ſ`""".stripMargin).write + .parquet(path.toString) + } + + val table = s"variant_unicode_columns_${System.currentTimeMillis()}" + withTable(table) { + sql( + s"CREATE TABLE $table (`é` VARIANT, `σ` VARIANT, `k` VARIANT, `s` VARIANT) " + + s"USING parquet OPTIONS (path '$path')") + withSQLConf( + SQLConf.CASE_SENSITIVE.key -> "false", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val (_, cometPlan) = checkSparkAnswerAndFallbackReason( + sql(s"SELECT `é`, `σ`, `k`, `s` FROM $table"), + "case-insensitive Unicode column names in Variant scans") + assert(collect(cometPlan) { case scan: CometNativeScanExec => scan }.isEmpty) + } + } + } + } + + test("case-insensitive Unicode sibling of Variant column falls back") { + assume(isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempPath { path => + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + "spark.sql.variant.writeShredding.enabled" -> "false") { + sql("SELECT parse_json('42') AS v, 7 AS `É`").write.parquet(path.toString) + } + + val table = s"variant_unicode_sibling_${System.currentTimeMillis()}" + withTable(table) { + sql(s"CREATE TABLE $table (v VARIANT, `é` INT) USING parquet OPTIONS (path '$path')") + withSQLConf( + SQLConf.CASE_SENSITIVE.key -> "false", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val (_, cometPlan) = checkSparkAnswerAndFallbackReason( + sql(s"SELECT v, `é` FROM $table"), + "case-insensitive Unicode column names in Variant scans") + assert(collect(cometPlan) { case scan: CometNativeScanExec => scan }.isEmpty) + } + } + } + } + + test("Variant scan falls back when Parquet encryption is configured") { + assume(isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempPath { path => + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + "spark.sql.variant.writeShredding.enabled" -> "false") { + sql("SELECT parse_json('42') AS v").write.parquet(path.toString) + } + + withParquetTable(path.toString, "encrypted_variant") { + withSQLConf( + DecryptionPropertiesFactory.CRYPTO_FACTORY_CLASS_PROPERTY_NAME -> + classOf[PropertiesDrivenCryptoFactory].getName, + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val (_, cometPlan) = checkSparkAnswerAndFallbackReason( + sql("SELECT v FROM encrypted_variant"), + "encrypted Variant columns") + assert(collect(cometPlan) { case scan: CometNativeScanExec => scan }.isEmpty) + } + } + } + } + test("native reader duplicate fields in case-insensitive mode") { withTempPath { path => // Write parquet with columns A, B, b (B and b are duplicates case-insensitively) 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 53a57ed2abe..036ba0817ba 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 @@ -37,9 +38,10 @@ import org.apache.parquet.hadoop.example.ExampleParquetWriter import org.apache.parquet.io.api.Binary import org.apache.parquet.schema.MessageTypeParser import org.apache.spark.SparkException -import org.apache.spark.sql.{CometTestBase, DataFrame, Row} +import org.apache.spark.sql.{AnalysisException, CometTestBase, DataFrame, Row} import org.apache.spark.sql.catalyst.util.DateTimeUtils -import org.apache.spark.sql.comet.{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,11 +49,22 @@ import org.apache.spark.sql.types._ import com.google.common.primitives.UnsignedLong -import org.apache.comet.CometConf +import org.apache.comet.{CometConf, CometSparkSessionExtensions, ExtendedExplainInfo} +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( @@ -85,6 +98,539 @@ abstract class ParquetReadSuite extends CometTestBase { } } + test("native scan projects Variant through a Spark-compatible vector") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + 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('{"b": "hello", "a": 10}'), '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) => + normalizedVariantRows(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(normalizedVariantRows(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")) + } + } + } + } + + test("native scan honors Spark's shredded Variant reader configuration") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTable("variant_reader_mode") { + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.writeShredding.enabled" -> "true", + "spark.sql.variant.forceShreddingSchemaForTest" -> "a BIGINT") { + sql("CREATE TABLE variant_reader_mode(v VARIANT) USING parquet") + sql("INSERT INTO variant_reader_mode VALUES (parse_json('{\"a\":1}'))") + } + + val key = "spark.sql.variant.allowReadingShredded" + val conf = SQLConf.get + val original = conf.getAllConfs.get(key) + try { + Seq(None, Some(false), Some(true)).foreach { configured => + configured match { + case Some(value) => conf.setConfString(key, value.toString) + case None => conf.unsetConf(key) + } + val allowShredded = conf.getConfString(key).toBoolean + + withSQLConf( + "spark.sql.variant.pushVariantIntoScan" -> "false", + "spark.sql.variant.inferShreddingSchema" -> "false") { + val result = sql("SELECT v FROM variant_reader_mode") + val nativeScans = collect(result.queryExecution.executedPlan) { + case _: CometNativeScanExec => true + } + + if (allowShredded) { + assert(normalizedVariantRows(result, 0) == Seq(Seq("{\"a\":1}"))) + assert(nativeScans.size == 1) + } else { + assert(nativeScans.isEmpty) + val error = intercept[SparkException](result.collect()).getCause + assert(error.isInstanceOf[AnalysisException]) + assert( + error.asInstanceOf[AnalysisException].getErrorClass == + "INVALID_VARIANT_FROM_PARQUET.WRONG_NUM_FIELDS") + } + } + } + } finally { + original match { + case Some(value) => conf.setConfString(key, value) + case None => conf.unsetConf(key) + } + } + } + } + + test("nanosAsLong Variant children fall back to Spark") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempDir { dir => + val rawNanos = Seq(1704067200123000000L, 1704067200123456789L) + val variant = sql("SELECT parse_json('0')").head().get(0) + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + + Seq(true, false).foreach { adjustedToUtc => + val path = new Path(dir.toURI.toString, s"nanos-$adjustedToUtc.parquet") + val parquetSchema = MessageTypeParser.parseMessageType(s"""message root { + | optional group v { + | optional binary value; + | required binary metadata; + | optional int64 typed_value (TIMESTAMP(NANOS,$adjustedToUtc)); + | } + |} + |""".stripMargin) + val writer = ExampleParquetWriter + .builder(path) + .withType(parquetSchema) + .withConf(spark.sessionState.newHadoopConf()) + .build() + + try { + rawNanos.foreach { value => + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(1, Binary.fromConstantByteArray(metadata)) + group.add(2, value) + writer.write(row) + } + } finally { + writer.close() + } + } + + withTable("variant_nanos_as_long") { + sql(s"""CREATE TABLE variant_nanos_as_long(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + withSQLConf( + SQLConf.PARQUET_VECTORIZED_READER_ENABLED.key -> "true", + "spark.sql.legacy.parquet.nanosAsLong" -> "true", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false", + "spark.sql.variant.inferShreddingSchema" -> "false") { + val result = sql("SELECT v FROM variant_nanos_as_long") + assert( + normalizedVariantRows(result, 0) == + rawNanos.flatMap(value => Seq.fill(2)(Seq(value.toString)))) + + val plan = result.queryExecution.executedPlan + assert(collect(plan) { case _: CometNativeScanExec => true }.isEmpty) + val fallbackReasons = new ExtendedExplainInfo().getFallbackReasons(plan) + assert( + fallbackReasons.exists(_.contains("spark.sql.legacy.parquet.nanosAsLong")), + s"Expected nanosAsLong fallback, found: ${fallbackReasons.mkString(", ")}") + } + } + } + } + + test("inferTimestampNTZ disabled Variant children fall back to Spark") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempDir { dir => + val path = new Path(dir.toURI.toString, "timestamp-ntz-variant.parquet") + val parquetSchema = MessageTypeParser.parseMessageType("""message root { + | optional group v { + | required binary metadata; + | required int64 typed_value (TIMESTAMP(MICROS,false)); + | } + |} + |""".stripMargin) + val variant = sql("SELECT parse_json('0')").head().get(0) + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + val writer = createParquetWriter(parquetSchema, path, dictionaryEnabled = false) + try { + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(0, Binary.fromConstantByteArray(metadata)) + group.add(1, 1704067200123456L) + writer.write(row) + } finally { + writer.close() + } + + withTable("variant_timestamp_ntz") { + sql(s"""CREATE TABLE variant_timestamp_ntz(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + val readerConfs = Seq( + SQLConf.PARQUET_VECTORIZED_READER_ENABLED.key -> "true", + SQLConf.PARQUET_INFER_TIMESTAMP_NTZ_ENABLED.key -> "false", + SQLConf.SESSION_LOCAL_TIMEZONE.key -> "America/Los_Angeles", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false", + "spark.sql.variant.inferShreddingSchema" -> "false") + var expected = Seq.empty[Seq[Any]] + withSQLConf((readerConfs :+ (CometConf.COMET_ENABLED.key -> "false")): _*) { + expected = normalizedVariantRows(sql("SELECT v FROM variant_timestamp_ntz"), 0) + } + withSQLConf(readerConfs: _*) { + val result = sql("SELECT v FROM variant_timestamp_ntz") + assert(normalizedVariantRows(result, 0) == expected) + val plan = result.queryExecution.executedPlan + assert(collect(plan) { case _: CometNativeScanExec => true }.isEmpty) + val fallbackReasons = new ExtendedExplainInfo().getFallbackReasons(plan) + assert( + fallbackReasons.exists(_.contains(SQLConf.PARQUET_INFER_TIMESTAMP_NTZ_ENABLED.key)), + s"Expected inferTimestampNTZ fallback, found: ${fallbackReasons.mkString(", ")}") + } + } + } + } + + test("native scan preserves Variant existence default pairing") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + 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.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql("SELECT v, n FROM variant_defaults") + assert(normalizedVariantRows(df, 0) == Seq(Seq("42", 7))) + val cometPlan = df.queryExecution.executedPlan + 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.allowReadingShredded" -> "true", + "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 preserves unannotated fixed-length Variant binary") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + Seq(8, 16, 20).foreach { length => + withTempDir { dir => + val path = new Path(dir.toURI.toString, "fixed-binary-variant.parquet") + val parquetSchema = MessageTypeParser.parseMessageType(s"""message root { + | optional group v { + | required binary metadata; + | required fixed_len_byte_array($length) typed_value; + | } + |} + |""".stripMargin) + val variant = sql("SELECT parse_json('0')").head().get(0) + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + val bytes = Array.tabulate(length)(_.toByte) + val writer = createParquetWriter(parquetSchema, path, dictionaryEnabled = false) + try { + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(0, Binary.fromConstantByteArray(metadata)) + group.add(1, Binary.fromConstantByteArray(bytes)) + writer.write(row) + } finally { + writer.close() + } + + withTable("fixed_binary_variant") { + sql(s"""CREATE TABLE fixed_binary_variant(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql("SELECT v FROM fixed_binary_variant") + val expected = "\"" + Base64.getEncoder.encodeToString(bytes) + "\"" + assert(normalizedVariantRows(df, 0) == Seq(Seq(expected))) + assert(collect(df.queryExecution.executedPlan) { case _: CometNativeScanExec => + true + }.size == 1) + } + } + } + } + } + + test("native scan decodes dictionary-encoded Variant storage") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempDir { dir => + val path = new Path(dir.toURI.toString, "dictionary-variant.parquet") + val parquetSchema = MessageTypeParser.parseMessageType("""message root { + | optional group v { + | optional binary value; + | required binary metadata; + | optional int64 typed_value; + | } + |} + |""".stripMargin) + val valueField = new ArrowField( + "value", + new FieldType( + true, + ArrowType.Binary.INSTANCE, + new DictionaryEncoding(0L, false, new ArrowType.Int(32, true))), + Collections.emptyList[ArrowField]()) + val metadataField = new ArrowField( + "metadata", + new FieldType( + false, + ArrowType.Binary.INSTANCE, + new DictionaryEncoding(1L, false, new ArrowType.Int(32, true))), + Collections.emptyList[ArrowField]()) + val typedValueField = new ArrowField( + "typed_value", + new FieldType( + true, + new ArrowType.Int(64, true), + new DictionaryEncoding(2L, false, new ArrowType.Int(32, true))), + Collections.emptyList[ArrowField]()) + val variantField = new ArrowField( + "v", + new FieldType( + true, + ArrowType.Struct.INSTANCE, + null, + Collections.singletonMap("ARROW:extension:name", "arrow.parquet.variant")), + Seq(valueField, metadataField, typedValueField).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 4).foreach { index => + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + if (index % 2 == 0) { + group.add(0, Binary.fromConstantByteArray(value)) + } + group.add(1, Binary.fromConstantByteArray(metadata)) + if (index % 2 != 0) { + group.add(2, 42L) + } + writer.write(row) + } + } finally { + writer.close() + } + + withTable("dictionary_variant") { + sql(s"""CREATE TABLE dictionary_variant(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql("SELECT v FROM dictionary_variant") + assert(normalizedVariantRows(df, 0) == Seq.fill(4)(Seq("42"))) + assert(collect(df.queryExecution.executedPlan) { case _: CometNativeScanExec => + true + }.size == 1) + } + } + } + } + + test("native scan widens shredded Variant UINT_64") { + assume(CometSparkSessionExtensions.isSpark40Plus, "VariantType requires Spark 4.0+") + + withTempDir { dir => + val path = new Path(dir.toURI.toString, "uint64-variant.parquet") + val parquetSchema = MessageTypeParser.parseMessageType("""message root { + | optional group v { + | required binary metadata; + | required int64 typed_value (UINT_64); + | } + |} + |""".stripMargin) + val variant = sql("SELECT parse_json('0')").head().get(0) + val metadata = + variant.getClass.getMethod("getMetadata").invoke(variant).asInstanceOf[Array[Byte]] + val writer = createParquetWriter(parquetSchema, path, dictionaryEnabled = false) + try { + val row = new SimpleGroup(parquetSchema) + val group = row.addGroup(0) + group.add(0, Binary.fromConstantByteArray(metadata)) + group.add(1, -1L) + writer.write(row) + } finally { + writer.close() + } + + withTable("uint64_variant") { + sql(s"""CREATE TABLE uint64_variant(v VARIANT) + |USING parquet LOCATION '${dir.getCanonicalPath}'""".stripMargin) + withSQLConf( + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + val df = sql("SELECT v FROM uint64_variant") + assert(normalizedVariantRows(df, 0) == Seq(Seq("18446744073709551615"))) + assert(collect(df.queryExecution.executedPlan) { case _: CometNativeScanExec => + true + }.size == 1) + } + } + } + } + // 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())