From 4066fe22cb12d1d559a946764bbee4ae12aa26a4 Mon Sep 17 00:00:00 2001 From: peterxcli Date: Tue, 4 Aug 2026 02:02:38 +0800 Subject: [PATCH 1/2] refactor: use Arrow casts for dictionary conversions --- native/core/src/parquet/parquet_support.rs | 59 +++++++-------- .../spark-expr/src/conversion_funcs/cast.rs | 72 ++++++++----------- 2 files changed, 59 insertions(+), 72 deletions(-) diff --git a/native/core/src/parquet/parquet_support.rs b/native/core/src/parquet/parquet_support.rs index 2ee1230ed8..0afafba59e 100644 --- a/native/core/src/parquet/parquet_support.rs +++ b/native/core/src/parquet/parquet_support.rs @@ -22,10 +22,10 @@ use arrow::compute::can_cast_types; use arrow::datatypes::{FieldRef, Fields}; use arrow::{ array::{ - cast::AsArray, new_null_array, types::Int32Type, types::TimestampMicrosecondType, Array, - ArrayRef, DictionaryArray, StructArray, + cast::AsArray, new_null_array, types::TimestampMicrosecondType, Array, ArrayRef, + StructArray, }, - compute::{cast_with_options, take, CastOptions}, + compute::{cast_with_options, CastOptions}, datatypes::{DataType, TimeUnit}, util::display::FormatOptions, }; @@ -169,31 +169,6 @@ fn parquet_convert_array( parquet_options: &SparkParquetOptions, ) -> DataFusionResult { use DataType::*; - let from_type = array.data_type().clone(); - - let array = match &from_type { - Dictionary(key_type, value_type) - if key_type.as_ref() == &Int32 - && (value_type.as_ref() == &Utf8 || value_type.as_ref() == &LargeUtf8) => - { - let dict_array = array - .as_any() - .downcast_ref::>() - .expect("Expected a dictionary array"); - - let casted_dictionary = DictionaryArray::::new( - dict_array.keys().clone(), - parquet_convert_array(Arc::clone(dict_array.values()), to_type, parquet_options)?, - ); - - let casted_result = match to_type { - Dictionary(_, _) => Arc::new(casted_dictionary.clone()), - _ => take(casted_dictionary.values().as_ref(), dict_array.keys(), None)?, - }; - return Ok(casted_result); - } - _ => array, - }; let from_type = array.data_type(); // Try Comet specific handlers first, then arrow-rs cast if supported, @@ -621,7 +596,6 @@ mod tests { use datafusion::execution::runtime_env::RuntimeEnv; #[cfg(not(feature = "hdfs-opendal"))] use object_store::path::Path; - #[cfg(not(feature = "hdfs-opendal"))] use std::sync::Arc; #[cfg(not(feature = "hdfs-opendal"))] use url::Url; @@ -631,6 +605,33 @@ mod tests { #[cfg(not(feature = "hdfs-opendal"))] use std::collections::HashMap; + #[test] + fn test_parquet_convert_dictionary_to_dictionary() { + use arrow::array::{DictionaryArray, Int32Array, LargeStringArray, StringArray}; + use arrow::datatypes::{DataType, Int32Type}; + + let array = Arc::new(DictionaryArray::::new( + Int32Array::from(vec![Some(0), None, Some(1), Some(0)]), + Arc::new(StringArray::from(vec!["a", "b"])), + )); + let data_type = + DataType::Dictionary(Box::new(DataType::Int16), Box::new(DataType::LargeUtf8)); + + let result = super::parquet_convert_array( + array, + &data_type, + &super::SparkParquetOptions::new_without_timezone(super::EvalMode::Legacy, false), + ) + .unwrap(); + + assert_eq!(result.data_type(), &data_type); + let values = arrow::compute::cast(&result, &DataType::LargeUtf8).unwrap(); + assert_eq!( + values.as_any().downcast_ref::().unwrap(), + &LargeStringArray::from(vec![Some("a"), None, Some("b"), Some("a")]) + ); + } + /// Parses the url, registers the object store, and returns a tuple of the object store url and object store path #[cfg(not(feature = "hdfs-opendal"))] pub(crate) fn prepare_object_store( diff --git a/native/spark-expr/src/conversion_funcs/cast.rs b/native/spark-expr/src/conversion_funcs/cast.rs index 37fddb8c11..e87644fc06 100644 --- a/native/spark-expr/src/conversion_funcs/cast.rs +++ b/native/spark-expr/src/conversion_funcs/cast.rs @@ -45,13 +45,13 @@ use arrow::array::{ new_null_array, BinaryBuilder, DictionaryArray, GenericByteArray, ListArray, MapArray, StringArray, StructArray, }; -use arrow::datatypes::{ArrowDictionaryKeyType, ArrowNativeType, DataType, Schema}; +use arrow::datatypes::{DataType, Schema}; use arrow::datatypes::{Field, Fields, GenericBinaryType}; use arrow::error::ArrowError; use arrow::{ array::{ cast::AsArray, types::Int32Type, Array, ArrayRef, Int16Array, Int32Array, Int64Array, - Int8Array, OffsetSizeTrait, PrimitiveArray, + Int8Array, OffsetSizeTrait, }, compute::{cast_with_options, take, CastOptions}, record_batch::RecordBatch, @@ -213,40 +213,6 @@ pub fn spark_cast( Ok(result) } -// copied from datafusion common scalar/mod.rs -fn dict_from_values( - values_array: ArrayRef, -) -> datafusion::common::Result { - // Create a key array with `size` elements of 0..array_len for all - // non-null value elements - let key_array: PrimitiveArray = (0..values_array.len()) - .map(|index| { - if values_array.is_valid(index) { - let native_index = K::Native::from_usize(index).ok_or_else(|| { - DataFusionError::Internal(format!( - "Can not create index of type {} from value {}", - K::DATA_TYPE, - index - )) - })?; - Ok(Some(native_index)) - } else { - Ok(None) - } - }) - .collect::>>()? - .into_iter() - .collect(); - - // create a new DictionaryArray - // - // Note: this path could be made faster by using the ArrayData - // APIs and skipping validation, if it every comes up in - // performance traces. - let dict_array = DictionaryArray::::try_new(key_array, values_array)?; - Ok(Arc::new(dict_array)) -} - pub(crate) fn cast_array( array: ArrayRef, to_type: &DataType, @@ -301,13 +267,11 @@ pub(crate) fn cast_array( return Ok(spark_cast_postprocess(casted_result, &from_type, to_type)); } _ => { - if let Dictionary(_, _) = to_type { - let dict_array = dict_from_values::(array)?; - let casted_result = cast_array(dict_array, to_type, cast_options)?; - return Ok(spark_cast_postprocess(casted_result, &from_type, to_type)); - } else { - array + if let Dictionary(_, value_type) = to_type { + let values = cast_array(array, value_type, cast_options)?; + return Ok(arrow::compute::cast(&values, to_type)?); } + array } }; @@ -908,11 +872,33 @@ fn cast_binary_to_string( #[cfg(test)] mod tests { use super::*; - use arrow::array::{BinaryArray, ListArray, NullArray, StringArray}; + use arrow::array::{BinaryArray, ListArray, NullArray, PrimitiveArray, StringArray}; use arrow::buffer::OffsetBuffer; use arrow::datatypes::TimestampMicrosecondType; use arrow::datatypes::{Field, Fields}; + #[test] + fn test_cast_to_dictionary_deduplicates_casted_values() { + let input: ArrayRef = Arc::new(StringArray::from(vec![Some("0.2"), Some("."), None])); + let data_type = DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Int32)); + + let result = cast_array( + input, + &data_type, + &SparkCastOptions::new(EvalMode::Legacy, "UTC", false), + ) + .unwrap(); + let dictionary = result.as_dictionary::(); + + assert_eq!( + dictionary.keys().iter().collect::>(), + vec![Some(0), Some(0), None] + ); + let values = dictionary.values().as_primitive::(); + assert_eq!(values.len(), 1); + assert_eq!(values.value(0), 0); + } + #[test] fn test_cast_binary_to_string_replaces_invalid_utf8_jvm_compatibly() { // Invalid bytes are replaced with U+FFFD instead of reinterpreted as an invalid `str`, From 9242ea3ab3173ed1677c0c92a174fa63a9dd87ab Mon Sep 17 00:00:00 2001 From: peterxcli Date: Wed, 5 Aug 2026 01:37:45 +0800 Subject: [PATCH 2/2] fix: handle dictionary sources and reject dictionary targets --- native/core/src/parquet/parquet_support.rs | 28 +----- .../spark-expr/src/conversion_funcs/cast.rs | 89 ++++++------------- .../comet/parquet/ParquetReadSuite.scala | 59 ++++++++++++ 3 files changed, 89 insertions(+), 87 deletions(-) diff --git a/native/core/src/parquet/parquet_support.rs b/native/core/src/parquet/parquet_support.rs index 0afafba59e..17e65b0194 100644 --- a/native/core/src/parquet/parquet_support.rs +++ b/native/core/src/parquet/parquet_support.rs @@ -596,6 +596,7 @@ mod tests { use datafusion::execution::runtime_env::RuntimeEnv; #[cfg(not(feature = "hdfs-opendal"))] use object_store::path::Path; + #[cfg(not(feature = "hdfs-opendal"))] use std::sync::Arc; #[cfg(not(feature = "hdfs-opendal"))] use url::Url; @@ -605,33 +606,6 @@ mod tests { #[cfg(not(feature = "hdfs-opendal"))] use std::collections::HashMap; - #[test] - fn test_parquet_convert_dictionary_to_dictionary() { - use arrow::array::{DictionaryArray, Int32Array, LargeStringArray, StringArray}; - use arrow::datatypes::{DataType, Int32Type}; - - let array = Arc::new(DictionaryArray::::new( - Int32Array::from(vec![Some(0), None, Some(1), Some(0)]), - Arc::new(StringArray::from(vec!["a", "b"])), - )); - let data_type = - DataType::Dictionary(Box::new(DataType::Int16), Box::new(DataType::LargeUtf8)); - - let result = super::parquet_convert_array( - array, - &data_type, - &super::SparkParquetOptions::new_without_timezone(super::EvalMode::Legacy, false), - ) - .unwrap(); - - assert_eq!(result.data_type(), &data_type); - let values = arrow::compute::cast(&result, &DataType::LargeUtf8).unwrap(); - assert_eq!( - values.as_any().downcast_ref::().unwrap(), - &LargeStringArray::from(vec![Some("a"), None, Some("b"), Some("a")]) - ); - } - /// Parses the url, registers the object store, and returns a tuple of the object store url and object store path #[cfg(not(feature = "hdfs-opendal"))] pub(crate) fn prepare_object_store( diff --git a/native/spark-expr/src/conversion_funcs/cast.rs b/native/spark-expr/src/conversion_funcs/cast.rs index e87644fc06..c17ae80f18 100644 --- a/native/spark-expr/src/conversion_funcs/cast.rs +++ b/native/spark-expr/src/conversion_funcs/cast.rs @@ -42,18 +42,17 @@ use crate::{cast_whole_num_to_binary, BinaryOutputStyle}; use crate::{EvalMode, SparkError}; use arrow::array::builder::{GenericStringBuilder, StringBuilder}; use arrow::array::{ - new_null_array, BinaryBuilder, DictionaryArray, GenericByteArray, ListArray, MapArray, - StringArray, StructArray, + new_null_array, BinaryBuilder, GenericByteArray, ListArray, MapArray, StringArray, StructArray, }; use arrow::datatypes::{DataType, Schema}; use arrow::datatypes::{Field, Fields, GenericBinaryType}; use arrow::error::ArrowError; use arrow::{ array::{ - cast::AsArray, types::Int32Type, Array, ArrayRef, Int16Array, Int32Array, Int64Array, - Int8Array, OffsetSizeTrait, + cast::AsArray, Array, ArrayRef, Int16Array, Int32Array, Int64Array, Int8Array, + OffsetSizeTrait, }, - compute::{cast_with_options, take, CastOptions}, + compute::{cast_with_options, CastOptions}, record_batch::RecordBatch, util::display::FormatOptions, }; @@ -221,6 +220,12 @@ pub(crate) fn cast_array( use DataType::*; let from_type = array.data_type().clone(); + // Spark's SQL data-type grammar cannot express Dictionary as a cast target: + // https://github.com/apache/spark/blob/v4.2.0/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4#L1477-L1525 + if matches!(to_type, Dictionary(_, _)) { + return internal_err!("Spark cannot specify dictionary types as cast targets"); + } + if &from_type == to_type { return Ok(Arc::new(array)); } @@ -235,45 +240,18 @@ pub(crate) fn cast_array( .with_timestamp_format(TIMESTAMP_FORMAT), }; - let array = match &from_type { - Dictionary(key_type, value_type) - if key_type.as_ref() == &Int32 - && (value_type.as_ref() == &Utf8 - || value_type.as_ref() == &LargeUtf8 - || value_type.as_ref() == &Binary - || value_type.as_ref() == &LargeBinary) => - { - let dict_array = array - .as_any() - .downcast_ref::>() - .expect("Expected a dictionary array"); - - let casted_result = match to_type { - Dictionary(_, to_value_type) => { - let casted_dictionary = DictionaryArray::::new( - dict_array.keys().clone(), - cast_array(Arc::clone(dict_array.values()), to_value_type, cast_options)?, - ); - Arc::new(casted_dictionary.clone()) - } - _ => { - let casted_dictionary = DictionaryArray::::new( - dict_array.keys().clone(), - cast_array(Arc::clone(dict_array.values()), to_type, cast_options)?, - ); - take(casted_dictionary.values().as_ref(), dict_array.keys(), None)? - } - }; + // Spark infers Parquet schemas from its own metadata or the Parquet MessageType, not + // ARROW:schema, so Arrow can expose a dictionary source while Spark requests its value type: + // 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 + if let Dictionary(_, value_type) = &from_type { + if matches!(value_type.as_ref(), Utf8 | LargeUtf8 | Binary | LargeBinary) { + let dictionary = array.as_any_dictionary(); + let values = cast_array(Arc::clone(dictionary.values()), to_type, cast_options)?; + let dictionary = dictionary.with_values(values); + let casted_result = cast_with_options(&dictionary, to_type, &native_cast_options)?; return Ok(spark_cast_postprocess(casted_result, &from_type, to_type)); } - _ => { - if let Dictionary(_, value_type) = to_type { - let values = cast_array(array, value_type, cast_options)?; - return Ok(arrow::compute::cast(&values, to_type)?); - } - array - } - }; + } let cast_result = match (&from_type, to_type) { // Null arrays carry no concrete values, so Arrow's native cast can change only the @@ -874,29 +852,20 @@ mod tests { use super::*; use arrow::array::{BinaryArray, ListArray, NullArray, PrimitiveArray, StringArray}; use arrow::buffer::OffsetBuffer; - use arrow::datatypes::TimestampMicrosecondType; - use arrow::datatypes::{Field, Fields}; + use arrow::datatypes::{Field, Fields, Int32Type, TimestampMicrosecondType}; #[test] - fn test_cast_to_dictionary_deduplicates_casted_values() { - let input: ArrayRef = Arc::new(StringArray::from(vec![Some("0.2"), Some("."), None])); - let data_type = DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Int32)); - - let result = cast_array( - input, - &data_type, + fn test_cast_to_dictionary_is_rejected() { + let error = cast_array( + Arc::new(StringArray::from(vec!["a"])), + &DataType::Dictionary(Box::new(DataType::Int32), Box::new(DataType::Utf8)), &SparkCastOptions::new(EvalMode::Legacy, "UTC", false), ) - .unwrap(); - let dictionary = result.as_dictionary::(); + .unwrap_err(); - assert_eq!( - dictionary.keys().iter().collect::>(), - vec![Some(0), Some(0), None] - ); - let values = dictionary.values().as_primitive::(); - assert_eq!(values.len(), 1); - assert_eq!(values.value(0), 0); + assert!(error + .to_string() + .contains("Spark cannot specify dictionary types as cast targets")); } #[test] 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 a5fc92a5eb..d0d5069554 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/ParquetReadSuite.scala @@ -22,6 +22,7 @@ package org.apache.comet.parquet import java.io.File import java.math.{BigDecimal, BigInteger} import java.time.{ZoneId, ZoneOffset} +import java.util.{Base64, Collections} import scala.reflect.ClassTag import scala.reflect.runtime.universe.TypeTag @@ -29,8 +30,11 @@ import scala.reflect.runtime.universe.TypeTag import org.scalactic.source.Position import org.scalatest.Tag +import org.apache.arrow.vector.types.pojo.{ArrowType, DictionaryEncoding, Field => ArrowField, FieldType, Schema => ArrowSchema} import org.apache.hadoop.fs.Path import org.apache.parquet.example.data.simple.SimpleGroup +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} @@ -81,6 +85,61 @@ abstract class ParquetReadSuite extends CometTestBase { } } + // 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: + // https://github.com/apache/spark/blob/v4.2.0/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/parquet/ParquetSchemaConverter.scala#L344-L350 + test("native scan casts Arrow dictionary binary values with Spark semantics") { + withTempDir { dir => + val path = new Path(dir.toURI.toString, "dictionary-binary.parquet") + val parquetSchema = MessageTypeParser.parseMessageType("""message root { + | optional binary value; + |} + |""".stripMargin) + val arrowField = new ArrowField( + "value", + new FieldType( + true, + ArrowType.Binary.INSTANCE, + new DictionaryEncoding(0L, false, new ArrowType.Int(32, true))), + Collections.emptyList[ArrowField]()) + val arrowSchema = new ArrowSchema(Collections.singletonList(arrowField)) + val metadata = Collections.singletonMap( + "ARROW:schema", + Base64.getEncoder.encodeToString(arrowSchema.serializeAsMessage())) + val writer = ExampleParquetWriter + .builder(path) + .withType(parquetSchema) + .withDictionaryEncoding(true) + .withExtraMetaData(metadata) + .withConf(spark.sessionState.newHadoopConf()) + .build() + + try { + Seq( + Array[Byte](0x66, 0x80.toByte, 0x6f), + Array[Byte](0x66, 0x80.toByte, 0x6f), + Array[Byte](0x76, 0x61, 0x6c, 0x69, 0x64)).foreach { bytes => + val row = new SimpleGroup(parquetSchema) + row.add(0, Binary.fromConstantByteArray(bytes)) + writer.write(row) + } + } finally { + writer.close() + } + + withSQLConf(SQLConf.PARQUET_BINARY_AS_STRING.key -> "true") { + withParquetTable(path.toString, "dictionary_binary") { + val (_, cometPlan) = + checkSparkAnswerAndOperator(sql("SELECT value FROM dictionary_binary")) + assert( + collect(cometPlan) { case scan: CometNativeScanExec => scan }.nonEmpty, + "Expected a CometNativeScanExec") + } + } + } + } + test("basic data types") { Seq(7, 1024).foreach { batchSize => withSQLConf(CometConf.COMET_BATCH_SIZE.key -> batchSize.toString) {