diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 58c04f406a7..ef1b04720b7 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -373,6 +373,8 @@ jobs: org.apache.spark.sql.comet.ParquetDatetimeRebaseV2Suite org.apache.spark.sql.comet.ParquetEncryptionITCase org.apache.comet.exec.CometNativeReaderSuite + org.apache.comet.CometVariantProjectionSuite + org.apache.spark.sql.CometVariantShreddingSuite org.apache.comet.CometIcebergNativeSuite org.apache.comet.CometIcebergEncryptionSuite org.apache.comet.CometIcebergRewriteActionSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 5731b2f5be7..999b4d04bda 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -127,6 +127,8 @@ jobs: org.apache.spark.sql.comet.ParquetDatetimeRebaseV2Suite org.apache.spark.sql.comet.ParquetEncryptionITCase org.apache.comet.exec.CometNativeReaderSuite + org.apache.comet.CometVariantProjectionSuite + org.apache.spark.sql.CometVariantShreddingSuite org.apache.comet.CometIcebergNativeSuite org.apache.comet.CometIcebergEncryptionSuite org.apache.comet.CometIcebergRewriteActionSuite diff --git a/dev/diffs/4.1.3.diff b/dev/diffs/4.1.3.diff index d0289d0de06..5763f9be3c7 100644 --- a/dev/diffs/4.1.3.diff +++ b/dev/diffs/4.1.3.diff @@ -1208,6 +1208,20 @@ index e4b5e10f7c3..c6efde09c8a 100644 protected val baseResourcePath = { // use the same way as `SQLQueryTestSuite` to get the resource path +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/ResolveDefaultColumnsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/ResolveDefaultColumnsSuite.scala +index cb9d0909554..084d6515e8b 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/ResolveDefaultColumnsSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/ResolveDefaultColumnsSuite.scala +@@ -284,7 +284,8 @@ class ResolveDefaultColumnsSuite extends QueryTest with SharedSparkSession { + withTable("t") { + sql("CREATE TABLE t(v VARIANT DEFAULT parse_json('1')) USING PARQUET") + sql("INSERT INTO t VALUES(DEFAULT)") +- checkAnswer(sql("select v from t"), sql("select parse_json('1')").collect()) ++ // Native unshredding may use a different integer width for the same Variant value. ++ assert(sql("select v from t").collect().map(_.get(0).toString).toSeq == Seq("1")) + } + } + diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala index 74cdee49e55..f7452c9abb7 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala diff --git a/docs/source/user-guide/latest/datatypes.md b/docs/source/user-guide/latest/datatypes.md index c2b942961c2..839d8b11ee1 100644 --- a/docs/source/user-guide/latest/datatypes.md +++ b/docs/source/user-guide/latest/datatypes.md @@ -103,9 +103,23 @@ functions, and hashing a `CalendarInterval`. Remaining work is tracked by ## Variant -| Type | Status | Notes | -| ------------- | ------ | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `VariantType` | 🔜 | Spark 4.0+. Native scan support is tracked by [#4295](https://github.com/apache/datafusion-comet/issues/4295); shredded Parquet read/write by [#3983](https://github.com/apache/datafusion-comet/issues/3983). | +| Type | Status | Notes | +| ------------- | ------ | ------------------------------------------------------------------------------------------------------------------------------- | +| `VariantType` | ⚠️ | Spark 4.0+. Native Parquet scans support direct projection of top-level Variant columns. Non-null existence defaults fall back. | + +Direct projection requires explicit configuration on every supported Spark version: +`spark.sql.variant.allowReadingShredded=true` (defaults to false in Spark 4.0) and +`spark.sql.variant.pushVariantIntoScan=false` (defaults to true in Spark 4.1+), with the default +Parquet timestamp inference settings. Support for Spark's whole-value pushdown rewrite is tracked +by [#5519](https://github.com/apache/datafusion-comet/issues/5519). Nested Variant columns, pushed-down +Variant field extraction, expressions, writes, shuffle and spill, Python operators, encrypted +files, and Iceberg scans fall back to Spark. Spark also handles columnar-to-row conversion of +the native scan output and strict reads with `allowReadingShredded=false`. Broader +support is tracked by [#4295](https://github.com/apache/datafusion-comet/issues/4295) and +[#3983](https://github.com/apache/datafusion-comet/issues/3983). + +Shredded reconstruction can be slower than Spark's reader; see the +[focused scan and allocation measurements in PR #5868](https://github.com/apache/datafusion-comet/pull/5868). ## Other diff --git a/native/common/src/error.rs b/native/common/src/error.rs index 41773237cba..cd912a41994 100644 --- a/native/common/src/error.rs +++ b/native/common/src/error.rs @@ -21,6 +21,11 @@ use std::sync::Arc; #[derive(thiserror::Error, Debug, Clone)] pub enum SparkError { + #[error( + "[MALFORMED_VARIANT] Variant binary is malformed. Please check the data source is valid." + )] + MalformedVariant, + // This list was generated from the Spark code. Many of the exceptions are not yet used by Comet #[error("[CAST_INVALID_INPUT] The value '{value}' of the type \"{from_type}\" cannot be cast to \"{to_type}\" \ because it is malformed. Correct the value as per the syntax, or change its target type. \ @@ -301,6 +306,7 @@ impl SparkError { /// Get the error type name for JSON serialization pub(crate) fn error_type_name(&self) -> &'static str { match self { + SparkError::MalformedVariant => "MalformedVariant", SparkError::CastInvalidValue { .. } => "CastInvalidValue", SparkError::InvalidInputInCastToDatetime { .. } => "InvalidInputInCastToDatetime", SparkError::NumericValueOutOfRange { .. } => "NumericValueOutOfRange", @@ -662,7 +668,8 @@ impl SparkError { | SparkError::InvalidIndexOfZero => "org/apache/spark/SparkArrayIndexOutOfBoundsException", // RuntimeException - SparkError::CannotParseDecimal + SparkError::MalformedVariant + | SparkError::CannotParseDecimal | SparkError::DuplicatedMapKey { .. } | SparkError::NullMapKey | SparkError::MapKeyValueDiffSizes @@ -726,6 +733,7 @@ impl SparkError { /// Returns the Spark error class code for this error pub(crate) fn error_class(&self) -> Option<&'static str> { match self { + SparkError::MalformedVariant => Some("MALFORMED_VARIANT"), // Cast errors SparkError::CastInvalidValue { .. } => Some("CAST_INVALID_INPUT"), SparkError::InvalidInputInCastToDatetime { .. } => Some("CAST_INVALID_INPUT"), diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index d66638ba32d..c8f15cf8424 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -1000,6 +1000,19 @@ impl PhysicalPlanner { } } + /// Only constant literals are supported as scan defaults. + 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()); + } + Err(GeneralError("Expected a literal scan default".to_string())) + } + /// Create a DataFusion physical sort expression from Spark physical expression fn create_sort_expr<'a>( &'a self, @@ -1658,43 +1671,34 @@ 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()) - })?; - 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(); + if common.default_values.len() != common.default_values_indexes.len() { + return Err(GeneralError( + "Scan default values and indexes have different lengths".to_string(), + )); + } + let default_values = if common.default_values.is_empty() { + None + } else { 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) + common + .default_values + .iter() + .zip(&common.default_values_indexes) + .map(|(expr, offset)| { + let idx = usize::try_from(*offset).map_err(|_| { + GeneralError(format!("Invalid scan default index {offset}")) + })?; + let field = required_schema.fields().get(idx).ok_or_else(|| { + GeneralError(format!( + "Scan default index {idx} is outside schema" + )) + })?; + let value = + self.create_default_value(expr, Arc::clone(&required_schema))?; + Ok((Column::new(field.name(), idx), value)) }) - .collect(), + .collect::, ExecutionError>>()?, ) - } else { - None }; // Get one file from this partition (we know it's not empty due to early return above) @@ -5146,6 +5150,39 @@ mod tests { max_frame_size: usize, } + #[test] + fn scan_default_rejects_struct_expressions() { + let planner = PhysicalPlanner::new(Arc::new(SessionContext::new()), 0); + let storage = DataType::Struct(Fields::from(vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ])); + let field = Field::new("v", storage.clone(), true).with_extension_type(VariantType); + let schema = Arc::new(Schema::new(vec![field.clone()])); + let bytes = |value| Expr { + expr_struct: Some(ExprStruct::Literal(spark_expression::Literal { + value: Some(literal::Value::BytesVal(value)), + datatype: Some(spark_expression::DataType { + type_id: spark_expression::data_type::DataTypeId::Bytes as i32, + type_info: None, + }), + is_null: false, + })), + ..Default::default() + }; + let value = spark_expression::CreateNamedStruct { + names: vec!["value".to_string(), "metadata".to_string()], + values: vec![bytes(vec![0]), bytes(vec![1, 0, 0])], + }; + let default_expr = |value| Expr { + expr_struct: Some(ExprStruct::CreateNamedStruct(value)), + ..Default::default() + }; + assert!(planner + .create_default_value(&default_expr(value), schema) + .is_err()); + } + #[test] fn spark_variant_schema_preserves_field_metadata() { let schema = convert_spark_types_to_arrow_schema(&[spark_operator::SparkStructField { diff --git a/native/core/src/parquet/cast_column/variant.rs b/native/core/src/parquet/cast_column/variant.rs index 4e727ecab31..35287733c16 100644 --- a/native/core/src/parquet/cast_column/variant.rs +++ b/native/core/src/parquet/cast_column/variant.rs @@ -20,15 +20,16 @@ use arrow::{ 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::{DataFusionError, Result as DataFusionResult}; +use datafusion_comet_common::SparkError; use parquet::variant::{ - unshred_variant, ListBuilder, MetadataBuilder, ObjectBuilder, ParentState, - ReadOnlyMetadataBuilder, ValueBuilder, Variant, VariantArray, VariantMetadata, + unshred_variant, ListBuilder, MetadataBuilder, ObjectBuilder, ObjectFieldBuilder, ParentState, + ReadOnlyMetadataBuilder, ValueBuilder, Variant, VariantArray, VariantBuilderExt, + VariantDecimal4, VariantDecimal8, VariantMetadata, WritableMetadataBuilder, }; use std::{ panic::{catch_unwind, AssertUnwindSafe}, @@ -57,22 +58,49 @@ pub(super) fn normalize_variant_array( } // VariantArray resolves metadata/value/typed_value by name, so the reader's child order is - // irrelevant. Legacy Spark residuals must be put in Arrow order before the single upstream - // unshred call; the whole output is then put back in the order expected by released Spark 4. + // irrelevant. Legacy Spark residuals must be put in Arrow order before unshredding; + // shredded output is then rebuilt with Spark's byte encoding. let array = normalize_variant_storage(array)?; let variant = VariantArray::try_new(array.as_ref())?; - let prepared = prepare_variant_for_unshredding(&variant)?; - let unshredded = unshred_variant(&prepared)?; - let value = unshredded.value_column(); - let value = cast(value.as_ref(), &DataType::Binary)?; - let metadata = cast(unshredded.metadata_column().as_ref(), &DataType::Binary)?; - let value = reorder_variant_values(&value, &metadata, unshredded.inner().nulls())?; - - Ok(Arc::new(StructArray::try_new( - fields.clone(), - vec![value, metadata], - unshredded.inner().nulls().cloned(), - )?)) + let normalize = |metadata: Option<&ArrayRef>| -> DataFusionResult { + let extended = extend_shredded_metadata(&variant, metadata)?; + let prepared = prepare_variant_for_unshredding(&variant, extended.as_ref().or(metadata))?; + let unshredded = unshred_variant(&prepared).map_err(|error| match error { + ArrowError::InvalidArgumentError(_) => { + DataFusionError::from(SparkError::MalformedVariant) + } + error => error.into(), + })?; + let (value, metadata) = if variant.typed_value_column().is_some() { + let value = cast(unshredded.value_column().as_ref(), &DataType::Binary)?; + let metadata = cast(unshredded.metadata_column().as_ref(), &DataType::Binary)?; + rebuild_spark_variant(&variant, &value, &metadata)? + } else { + // Spark passes unshredded bytes through, including dictionary order, unused keys, + // and wide scalar encodings. Preparation above still validates legacy input. + let mut value = cast(variant.value_column().as_ref(), &DataType::Binary)?; + let metadata = cast(variant.metadata_column().as_ref(), &DataType::Binary)?; + if variant.inner().null_count() != 0 { + value = arrow::compute::nullif( + value.as_ref(), + &arrow::compute::is_null(variant.inner())?, + )?; + } + (value, metadata) + }; + Ok(Arc::new(StructArray::try_new( + fields.clone(), + vec![value, metadata], + unshredded.inner().nulls().cloned(), + )?)) + }; + match normalize(None) { + Ok(array) => Ok(array), + Err(error) => match canonicalize_spark_empty_key_metadata(&variant)? { + Some(metadata) => normalize(Some(&metadata)), + None => Err(error), + }, + } } /// Arrow Variant compute rejects some storage types that Spark's Parquet reader accepts. @@ -178,7 +206,9 @@ fn normalize_variant_storage(array: &ArrayRef) -> DataFusionResult { fn rewrite_shredding_state( state: &StructArray, metadata: &BinaryArray, + target_metadata: Option<&BinaryArray>, metadata_rows: &[Option], + allow_missing: bool, ) -> DataFusionResult<(ArrayRef, bool)> { if state.len() != metadata_rows.len() { return Err(DataFusionError::Execution( @@ -195,9 +225,40 @@ fn rewrite_shredding_state( let mut columns = state.columns().to_vec(); let mut changed = false; + let value_index = state + .fields() + .iter() + .position(|field| field.name() == "value"); + let typed = state.column_by_name("typed_value"); + for (index, row) in metadata_rows.iter().enumerate() { + if row.is_some() + && (state.is_null(index) + || (!allow_missing + && value_index.is_none_or(|column| state.column(column).is_null(index)) + && typed.is_none_or(|column| column.is_null(index)))) + { + return Err(SparkError::MalformedVariant.into()); + } + } + + // Spark gives scalar/array typed_value precedence over a redundant residual. Arrow + // rejects both being present, so remove the ignored residual before validating it. + if let (Some(index), Some(typed)) = (value_index, typed) { + if !matches!(typed.data_type(), DataType::Struct(_)) + && (0..state.len()).any(|row| { + active_rows[row].is_some() && typed.is_valid(row) && columns[index].is_valid(row) + }) + { + let present = arrow::compute::is_not_null(typed)?; + columns[index] = arrow::compute::nullif(columns[index].as_ref(), &present)?; + fields[index] = Arc::new(fields[index].as_ref().clone().with_nullable(true)); + changed = true; + } + } + if let Some(index) = fields.iter().position(|field| field.name() == "value") { let (value, value_changed) = - rewrite_residual_values(&columns[index], metadata, &active_rows)?; + rewrite_residual_values(&columns[index], metadata, target_metadata, &active_rows)?; if value_changed { fields[index] = Arc::new( fields[index] @@ -220,7 +281,7 @@ fn rewrite_shredding_state( .map(|(row, metadata)| columns[index].is_valid(row).then_some(*metadata).flatten()) .collect::>(); let (typed_value, typed_changed) = - rewrite_typed_value(&columns[index], metadata, &typed_rows)?; + rewrite_typed_value(&columns[index], metadata, target_metadata, &typed_rows)?; if typed_changed { fields[index] = Arc::new( fields[index] @@ -249,6 +310,7 @@ fn rewrite_shredding_state( fn rewrite_residual_values( value: &ArrayRef, metadata: &BinaryArray, + target_metadata: Option<&BinaryArray>, metadata_rows: &[Option], ) -> DataFusionResult<(ArrayRef, bool)> { let binary = cast(value.as_ref(), &DataType::Binary)?; @@ -276,18 +338,33 @@ fn rewrite_residual_values( let rebuilt = catch_unwind(AssertUnwindSafe( || -> Result>, ArrowError> { - let metadata = VariantMetadata::try_new(metadata.value(*metadata_row))?; + let source = metadata.value(*metadata_row); + let target = target_metadata + .map(|metadata| metadata.value(*metadata_row)) + .filter(|target| *target != source) + .map(VariantMetadata::try_new) + .transpose()?; + let metadata = if target.is_some() { + // The empty-key workaround validated every original dictionary entry. + VariantMetadata::new(source) + } else { + VariantMetadata::try_new(source)? + }; let variant = Variant::new_with_metadata(metadata.clone(), binary.value(index)); - if is_compatible_variant(&variant, VariantObjectKeyOrder::ArrowUtf8) { + let arrow_ordered = + is_compatible_variant(&variant, VariantObjectKeyOrder::ArrowUtf8); + if arrow_ordered && target.is_none() { return Ok(None); } - if !is_compatible_variant(&variant, VariantObjectKeyOrder::SparkUtf16) { + 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(variant_bytes( - &metadata, + target.as_ref().unwrap_or(&metadata), variant, VariantObjectKeyOrder::ArrowUtf8, )?)) @@ -338,6 +415,7 @@ fn rewrite_list_typed_value( array: &ArrayRef, list: &L, metadata: &BinaryArray, + target_metadata: Option<&BinaryArray>, metadata_rows: &[Option], ) -> DataFusionResult<(ArrayRef, bool)> { let child_rows = list_metadata_rows(list, metadata_rows)?; @@ -347,7 +425,8 @@ fn rewrite_list_typed_value( list.values().data_type() )) })?; - let (values, changed) = rewrite_shredding_state(values, metadata, &child_rows)?; + let (values, changed) = + rewrite_shredding_state(values, metadata, target_metadata, &child_rows, false)?; if !changed { return Ok((Arc::clone(array), false)); } @@ -395,6 +474,7 @@ fn rewrite_list_typed_value( fn rewrite_typed_value( typed_value: &ArrayRef, metadata: &BinaryArray, + target_metadata: Option<&BinaryArray>, metadata_rows: &[Option], ) -> DataFusionResult<(ArrayRef, bool)> { match typed_value.data_type() { @@ -412,7 +492,7 @@ fn rewrite_typed_value( )) })?; let (child, child_changed) = - rewrite_shredding_state(child, metadata, metadata_rows)?; + rewrite_shredding_state(child, metadata, target_metadata, metadata_rows, true)?; if child_changed { fields[index] = Arc::new( fields[index] @@ -440,41 +520,71 @@ fn rewrite_typed_value( typed_value, typed_value.as_list::(), metadata, + target_metadata, metadata_rows, ), DataType::LargeList(_) => rewrite_list_typed_value( typed_value, typed_value.as_list::(), metadata, + target_metadata, metadata_rows, ), DataType::ListView(_) => rewrite_list_typed_value( typed_value, typed_value.as_list_view::(), metadata, + target_metadata, metadata_rows, ), DataType::LargeListView(_) => rewrite_list_typed_value( typed_value, typed_value.as_list_view::(), metadata, + target_metadata, metadata_rows, ), _ => Ok((Arc::clone(typed_value), false)), } } -fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult { - if variant.typed_value_column().is_none() { - return Ok(variant.clone()); - } - +fn prepare_variant_for_unshredding( + variant: &VariantArray, + target_metadata: Option<&ArrayRef>, +) -> DataFusionResult { let metadata = cast(variant.metadata_column().as_ref(), &DataType::Binary)?; let metadata = metadata.as_binary::(); let metadata_rows = (0..variant.len()) .map(|index| variant.inner().is_valid(index).then_some(index)) .collect::>(); - let (array, changed) = rewrite_shredding_state(variant.inner(), metadata, &metadata_rows)?; + let (array, changed) = rewrite_shredding_state( + variant.inner(), + metadata, + target_metadata.map(|metadata| metadata.as_binary::()), + &metadata_rows, + false, + )?; + if let Some(metadata) = target_metadata { + let array = array.as_struct(); + let mut fields = array.fields().to_vec(); + let mut columns = array.columns().to_vec(); + let index = fields + .iter() + .position(|field| field.name() == "metadata") + .unwrap(); + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(DataType::Binary), + ); + columns[index] = Arc::clone(metadata); + return Ok(VariantArray::try_new(&StructArray::try_new( + fields.into(), + columns, + array.nulls().cloned(), + )?)?); + } if changed { Ok(VariantArray::try_new(array.as_ref())?) } else { @@ -482,6 +592,433 @@ fn prepare_variant_for_unshredding(variant: &VariantArray) -> DataFusionResult, +) -> DataFusionResult> { + fn collect_keys<'a>(typed: &'a DataType, keys: &mut Vec<&'a str>) { + let children = match typed { + DataType::Struct(fields) => { + keys.extend(fields.iter().map(|field| field.name().as_str())); + fields.iter().collect::>() + } + DataType::List(field) + | DataType::LargeList(field) + | DataType::ListView(field) + | DataType::LargeListView(field) => vec![field], + _ => return, + }; + for field in children { + if let DataType::Struct(state) = field.data_type() { + if let Some(typed) = state.iter().find(|field| field.name() == "typed_value") { + collect_keys(typed.data_type(), keys); + } + } + } + } + + let mut keys = Vec::new(); + if let Some(typed) = variant.typed_value_column() { + collect_keys(typed.data_type(), &mut keys); + } + if keys.is_empty() { + return Ok(None); + } + keys.sort_unstable(); + keys.dedup(); + let metadata = cast( + metadata.unwrap_or(variant.metadata_column()).as_ref(), + &DataType::Binary, + )?; + let metadata = metadata.as_binary::(); + let mut output: Option = None; + for index in 0..variant.len() { + if variant.inner().is_null(index) { + if let Some(output) = &mut output { + output.append_option(metadata.is_valid(index).then(|| metadata.value(index))); + } + continue; + } + if metadata.is_null(index) { + return Err(SparkError::MalformedVariant.into()); + } + let dictionary = VariantMetadata::try_new(metadata.value(index))?; + if keys.iter().any(|key| dictionary.get_entry(key).is_none()) { + let mut names = dictionary + .iter() + .chain(keys.iter().copied()) + .collect::>(); + names.sort_unstable(); + names.dedup(); + let mut builder = WritableMetadataBuilder::from_iter(names); + builder.finish(); + output + .get_or_insert_with(|| binary_prefix_builder(metadata, index)) + .append_value(builder.into_inner()); + } else if let Some(output) = &mut output { + output.append_value(metadata.value(index)); + } + } + Ok(output.map(|mut output| Arc::new(output.finish()) as ArrayRef)) +} + +/// Spark writes unsorted dictionaries with equal offsets for empty object keys. Arrow's +/// validator rejects these, so retry with sorted metadata and remap every residual field ID. +/// TODO: Remove this workaround once an arrow-rs release includes +/// https://github.com/apache/arrow-rs/pull/10352; tracked by +/// https://github.com/apache/datafusion-comet/issues/5477. +fn canonicalize_spark_empty_key_metadata( + variant: &VariantArray, +) -> DataFusionResult> { + let metadata = cast(variant.metadata_column().as_ref(), &DataType::Binary)?; + let metadata = metadata.as_binary::(); + let mut output: Option = None; + for index in 0..variant.len() { + let replacement = if variant.inner().is_null(index) + || metadata.is_null(index) + || VariantMetadata::try_new(metadata.value(index)).is_ok() + { + None + } else { + let replacement = catch_unwind(AssertUnwindSafe( + || -> Result>, ArrowError> { + let bytes = metadata.value(index); + let original = VariantMetadata::new(bytes); + let mut names = original.iter_try().collect::, _>>()?; + if !names.contains(&"") { + return Ok(None); + } + // Accept only Spark's encoding of otherwise valid, unique field names. + let mut source = WritableMetadataBuilder::from_iter(names.iter().copied()); + source.finish(); + let mut source = source.into_inner(); + source[0] &= !0x10; + if source != bytes { + return Ok(None); + } + names.sort_unstable(); + if names.windows(2).any(|names| names[0] == names[1]) { + return Ok(None); + } + let mut metadata = WritableMetadataBuilder::from_iter(names); + metadata.finish(); + let metadata = metadata.into_inner(); + VariantMetadata::try_new(&metadata)?; + Ok(Some(metadata)) + }, + )); + let Ok(Ok(Some(replacement))) = replacement else { + return Ok(None); + }; + Some(replacement) + }; + if replacement.is_some() && output.is_none() { + output = Some(binary_prefix_builder(metadata, index)); + } + if let Some(output) = &mut output { + if metadata.is_null(index) { + output.append_null(); + } else { + output.append_value( + replacement + .as_deref() + .unwrap_or_else(|| metadata.value(index)), + ); + } + } + } + Ok(output.map(|mut output| Arc::new(output.finish()) as ArrayRef)) +} + +/// Spark's ShreddingUtils.rebuild uses a fresh dictionary in traversal order, with its sorted +/// flag unset. Keep Arrow's builders for the wire format and adapt only object-key ordering. +/// UTF-16 ordering removal is tracked by https://github.com/apache/datafusion-comet/issues/5474. +#[derive(Debug, Default)] +struct SparkOutputMetadata { + dictionary: WritableMetadataBuilder, + sort_keys: Vec, +} + +impl MetadataBuilder for SparkOutputMetadata { + fn try_upsert_field_name(&mut self, name: &str) -> Result { + let id = self.dictionary.upsert_field_name(name); + if id as usize == self.sort_keys.len() { + self.sort_keys.push(spark_sort_key(name)); + } + Ok(id) + } + + fn field_name(&self, id: usize) -> &str { + &self.sort_keys[id] + } + + fn num_field_names(&self) -> usize { + self.sort_keys.len() + } + + fn truncate_field_names(&mut self, size: usize) { + self.sort_keys.truncate(size); + self.dictionary.truncate_field_names(size); + } + + fn finish(&mut self) -> usize { + self.dictionary.finish() + } +} + +struct SparkValueBuilder<'a> { + value: &'a mut ValueBuilder, + metadata: &'a mut SparkOutputMetadata, +} + +impl VariantBuilderExt for SparkValueBuilder<'_> { + type State<'a> + = () + where + Self: 'a; + + fn append_null(&mut self) { + self.append_value(Variant::Null); + } + + fn append_value<'m, 'v>(&mut self, value: impl Into>) { + ValueBuilder::append_variant( + ParentState::variant(self.value, self.metadata), + value.into(), + ); + } + + fn try_new_list(&mut self) -> Result, ArrowError> { + Ok(ListBuilder::new( + ParentState::variant(self.value, self.metadata), + true, + )) + } + + fn try_new_object(&mut self) -> Result, ArrowError> { + Ok(ObjectBuilder::new( + ParentState::variant(self.value, self.metadata), + true, + )) + } +} + +fn spark_sort_key(name: &str) -> String { + name.encode_utf16() + .map(|unit| char::from_u32(0x10000 + u32::from(unit)).unwrap()) + .collect() +} + +fn binary_value(array: &ArrayRef, row: usize) -> Result<&[u8], ArrowError> { + match array.data_type() { + DataType::Binary => Ok(array.as_binary::().value(row)), + DataType::LargeBinary => Ok(array.as_binary::().value(row)), + DataType::BinaryView => Ok(array.as_binary_view().value(row)), + data_type => Err(ArrowError::InvalidArgumentError(format!( + "Expected Variant binary storage, got {data_type}" + ))), + } +} + +fn spark_typed_scalar<'m, 'v>(value: Variant<'m, 'v>) -> Variant<'m, 'v> { + match value { + Variant::Int8(_) | Variant::Int16(_) | Variant::Int32(_) | Variant::Int64(_) => { + let value = value.as_int64().unwrap(); + 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) + } + } + Variant::Decimal16(decimal) => { + if let Ok(decimal) = VariantDecimal4::try_from(decimal) { + Variant::Decimal4(decimal) + } else if let Ok(decimal) = VariantDecimal8::try_from(decimal) { + Variant::Decimal8(decimal) + } else { + value + } + } + Variant::Decimal8(decimal) => VariantDecimal4::try_from(decimal) + .map(Variant::Decimal4) + .unwrap_or(value), + Variant::String(s) => Variant::from(s), + Variant::Float(v) if v.is_nan() => Variant::Float(f32::NAN), + Variant::Double(v) if v.is_nan() => Variant::Double(f64::NAN), + _ => value, + } +} + +/// Arrow provides the decoded typed values and validates the shredding states. The original +/// state retains Spark's traversal order and distinguishes typed scalars (which Spark narrows) +/// from residual scalars (whose existing encoding Spark preserves). +fn append_spark_variant( + builder: &mut impl VariantBuilderExt, + value: Variant<'_, '_>, + state: Option<(&StructArray, usize)>, + source_metadata: &VariantMetadata<'_>, +) -> Result<(), ArrowError> { + let typed = state.and_then(|(state, row)| { + state + .column_by_name("typed_value") + .filter(|typed| typed.is_valid(row)) + .map(|typed| (typed, row)) + }); + let residual = state.and_then(|(state, row)| { + state + .column_by_name("value") + .filter(|value| value.is_valid(row)) + .map(|value| (value, row)) + }); + let value = if typed.is_none() { + match residual { + Some((value, row)) => { + Variant::new_with_metadata(source_metadata.clone(), binary_value(value, row)?) + } + None => value, + } + } else { + value + }; + match value { + Variant::Object(object) => { + let mut builder = builder.try_new_object()?; + if let Some((typed, row)) = typed { + let fields = typed.as_struct(); + for (field, child) in fields.fields().iter().zip(fields.columns()) { + let child = child.as_struct(); + if ["typed_value", "value"].iter().any(|name| { + child + .column_by_name(name) + .is_some_and(|value| value.is_valid(row)) + }) { + let value = object.get(field.name()).ok_or_else(|| { + ArrowError::InvalidArgumentError("Missing unshredded field".into()) + })?; + append_spark_variant( + &mut ObjectFieldBuilder::new(field.name(), &mut builder), + value, + Some((child, row)), + source_metadata, + )?; + } + } + if let Some((value, row)) = residual { + let Variant::Object(residual) = Variant::new_with_metadata( + source_metadata.clone(), + binary_value(value, row)?, + ) else { + return Err(ArrowError::InvalidArgumentError( + "Expected residual object".into(), + )); + }; + for entry in residual.iter_try() { + let (name, value) = entry?; + append_spark_variant( + &mut ObjectFieldBuilder::new(name, &mut builder), + value, + None, + source_metadata, + )?; + } + } + } else { + for entry in object.iter_try() { + let (name, value) = entry?; + append_spark_variant( + &mut ObjectFieldBuilder::new(name, &mut builder), + value, + None, + source_metadata, + )?; + } + } + builder.finish(); + } + Variant::List(list) => { + let mut builder = builder.try_new_list()?; + let elements = typed.map(|(typed, row)| { + macro_rules! elements { + ($list:expr) => {{ + let list = $list; + (list.values().as_struct(), list.element_range(row)) + }}; + } + match typed.data_type() { + DataType::List(_) => elements!(typed.as_list::()), + DataType::LargeList(_) => elements!(typed.as_list::()), + DataType::ListView(_) => elements!(typed.as_list_view::()), + DataType::LargeListView(_) => elements!(typed.as_list_view::()), + _ => unreachable!("validated shredded list"), + } + }); + for (index, value) in list.iter().enumerate() { + let state = elements + .as_ref() + .map(|(states, range)| (*states, range.start + index)); + append_spark_variant(&mut builder, value, state, source_metadata)?; + } + builder.finish(); + } + value => builder.append_value(if typed.is_some() { + spark_typed_scalar(value) + } else { + value + }), + } + Ok(()) +} + +fn rebuild_spark_variant( + source: &VariantArray, + value: &ArrayRef, + metadata: &ArrayRef, +) -> DataFusionResult<(ArrayRef, ArrayRef)> { + let mut values = BinaryBuilder::new(); + let mut dictionaries = BinaryBuilder::new(); + for row in 0..source.len() { + if source.is_null(row) { + values.append_null(); + dictionaries.append_null(); + continue; + } + let rebuilt = catch_unwind(AssertUnwindSafe(|| -> Result<_, ArrowError> { + let value = Variant::try_new(binary_value(metadata, row)?, binary_value(value, row)?)?; + // The preparation pass already validated/canonicalized legacy input metadata. + let original = VariantMetadata::new(binary_value(source.metadata_column(), row)?); + let mut output = ValueBuilder::new(); + let mut dictionary = SparkOutputMetadata::default(); + append_spark_variant( + &mut SparkValueBuilder { + value: &mut output, + metadata: &mut dictionary, + }, + value, + Some((source.inner(), row)), + &original, + )?; + dictionary.finish(); + let mut metadata = dictionary.dictionary.into_inner(); + metadata[0] &= !0x10; + Ok((output.into_inner(), metadata)) + })) + .map_err(|_| SparkError::MalformedVariant)? + .map_err(|_| SparkError::MalformedVariant)?; + values.append_value(rebuilt.0); + dictionaries.append_value(rebuilt.1); + } + Ok((Arc::new(values.finish()), Arc::new(dictionaries.finish()))) +} + /// Supplies sort-only field names whose Rust ordering matches Java `String.compareTo` ordering. /// Field IDs still come from the original metadata dictionary. #[derive(Debug)] @@ -492,15 +1029,7 @@ struct SparkMetadataBuilder<'a, 'm> { 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(); + let sort_keys = metadata.iter().map(spark_sort_key).collect(); Self { metadata, sort_keys, @@ -616,14 +1145,16 @@ fn variant_bytes( Ok(value_builder.into_inner()) } -/// Released Spark 4 profiles search object fields in Java UTF-16 order. Convert whole-value output -/// to that order until #5474 can remove this rewrite after every supported profile includes -/// SPARK-58949. Values already in the requested order remain byte-for-byte unchanged. +/// Released Spark 4 profiles search object fields in Java UTF-16 order. Values already in that +/// order remain byte-for-byte unchanged. +/// TODO: Remove this output rewrite once every supported Spark profile includes SPARK-58949. +/// Retain input conversion for historical Spark files with UTF-16 object-key ordering. /// https://github.com/apache/datafusion-comet/issues/5474 +#[cfg(test)] fn reorder_variant_values( value: &ArrayRef, metadata: &ArrayRef, - parent_nulls: Option<&NullBuffer>, + parent_nulls: Option<&arrow::buffer::NullBuffer>, ) -> DataFusionResult { let original = value; let value = value.as_binary::(); @@ -641,9 +1172,7 @@ fn reorder_variant_values( continue; } if value.is_null(index) { - return Err(DataFusionError::Execution(format!( - "Variant value is null at row {index}" - ))); + return Err(SparkError::MalformedVariant.into()); } if metadata.is_null(index) { return Err(DataFusionError::Execution(format!( diff --git a/native/core/src/parquet/cast_column/variant/tests.rs b/native/core/src/parquet/cast_column/variant/tests.rs index 46acc2b56ab..52f33ecc3eb 100644 --- a/native/core/src/parquet/cast_column/variant/tests.rs +++ b/native/core/src/parquet/cast_column/variant/tests.rs @@ -18,7 +18,7 @@ use super::*; use arrow::{ array::{Int64Array, ListArray}, - buffer::OffsetBuffer, + buffer::{NullBuffer, OffsetBuffer}, datatypes::{Field, Fields}, }; use parquet::variant::{ @@ -231,6 +231,46 @@ fn assert_spark_unicode_output(output: &StructArray) { assert_spark_unicode_variant(Variant::new(metadata.value(0), value.value(0))); } +#[test] +fn shredded_scalar_bytes_match_spark_and_residual_bytes_remain_wide() { + // Golden bytes from Spark VariantBuilder.appendLong, including each width boundary. + for (number, width) in [(1_i64, 1), (-129, 2), (32768, 4), (2147483648, 8)] { + let metadata: ArrayRef = Arc::new(BinaryArray::from(vec![&[1, 1, 0, 1, b'z'][..]])); + let mut wide = vec![0x18]; + wide.extend_from_slice(&number.to_le_bytes()); + for typed in [true, false] { + let physical: ArrayRef = Arc::new(StructArray::new( + vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("value", DataType::Binary, true), + Field::new("typed_value", DataType::Int64, true), + ] + .into(), + vec![ + Arc::clone(&metadata), + Arc::new(BinaryArray::from(vec![(!typed).then_some(wide.as_slice())])), + Arc::new(Int64Array::from(vec![typed.then_some(number)])), + ], + None, + )); + let output = normalize_variant_array(&physical, &target_field(false)).unwrap(); + let output = output.as_struct(); + assert_eq!(output.column(1).as_binary::().value(0), [1, 0, 0]); + let mut expected = vec![match width { + 1 => 0x0c, + 2 => 0x10, + 4 => 0x14, + _ => 0x18, + }]; + expected.extend_from_slice(&number.to_le_bytes()[..width]); + assert_eq!( + output.column(0).as_binary::().value(0), + if typed { &expected } else { &wide } + ); + } + } +} + #[test] fn normalize_full_shredding_reorders_children_and_preserves_parent_nulls() { let mut builder = VariantArrayBuilder::new(3); @@ -269,8 +309,8 @@ fn normalize_full_shredding_reorders_children_and_preserves_parent_nulls() { assert!(output.is_null(1)); let variant = VariantArray::try_new(output).unwrap(); - assert_eq!(variant.value(0), Variant::from(10_i64)); - assert_eq!(variant.value(2), Variant::from(30_i64)); + assert_eq!(variant.value(0), Variant::Int8(10)); + assert_eq!(variant.value(2), Variant::Int8(30)); } #[test] @@ -316,13 +356,128 @@ fn normalize_fully_shredded_object_orders_for_spark() { assert_spark_unicode_output(output.as_struct()); } +#[test] +fn normalize_shredded_objects_extend_metadata_and_preserve_missing_fields() { + let mut builder = VariantBuilder::new(); + builder.new_object().with_field("z", 9_i64).finish(); + let (metadata, residual) = builder.finish(); + let empty_metadata = [1, 0, 0]; + let field_a: ArrayRef = Arc::new(StructArray::new( + vec![ + Field::new("value", DataType::Binary, true), + Field::new("typed_value", DataType::Int64, true), + ] + .into(), + vec![ + Arc::new(BinaryArray::from(vec![None, None, None, Some(&[0_u8][..])])), + Arc::new(Int64Array::from(vec![None, None, Some(1), None])), + ], + None, + )); + let field_b: ArrayRef = Arc::new(StructArray::new( + vec![Field::new("typed_value", DataType::Int64, true)].into(), + vec![Arc::new(Int64Array::from(vec![None, None, None, Some(2)]))], + None, + )); + let typed: ArrayRef = Arc::new(StructArray::new( + vec![ + Field::new("a", field_a.data_type().clone(), false), + Field::new("b", field_b.data_type().clone(), false), + ] + .into(), + vec![field_a, field_b], + None, + )); + let input: ArrayRef = Arc::new(StructArray::new( + vec![ + Field::new("metadata", DataType::Binary, true), + Field::new("value", DataType::Binary, true), + Field::new("typed_value", typed.data_type().clone(), false), + ] + .into(), + vec![ + Arc::new(BinaryArray::from(vec![ + None, + Some(empty_metadata.as_slice()), + Some(metadata.as_slice()), + Some(empty_metadata.as_slice()), + ])), + Arc::new(BinaryArray::from(vec![ + None, + None, + Some(residual.as_slice()), + None, + ])), + typed, + ], + Some(NullBuffer::from(vec![false, true, true, true])), + )); + let output = normalize_variant_array(&input, &target_field(true)).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert!(output.is_null(0)); + for (row, expected) in [ + (1, vec![]), + ( + 2, + vec![("a", Variant::Int8(1)), ("z", Variant::from(9_i64))], + ), + (3, vec![("a", Variant::Null), ("b", Variant::Int8(2))]), + ] { + let Variant::Object(object) = output.value(row) else { + panic!("expected object") + }; + assert_eq!(object.iter().collect::>(), expected); + } +} + +#[test] +fn normalize_rejects_missing_required_shredding_states() { + let wrap = |typed: ArrayRef| -> ArrayRef { + Arc::new(StructArray::new( + vec![ + Field::new("metadata", DataType::Binary, false), + Field::new("typed_value", typed.data_type().clone(), true), + ] + .into(), + vec![Arc::new(BinaryArray::from(vec![&[1_u8, 0, 0][..]])), typed], + None, + )) + }; + let missing: ArrayRef = Arc::new(Int64Array::from(vec![None])); + let mut inputs = vec![wrap(Arc::clone(&missing))]; + for nulls in [None, Some(NullBuffer::new_null(1))] { + let state: ArrayRef = Arc::new(StructArray::new( + vec![Field::new("typed_value", DataType::Int64, true)].into(), + vec![Arc::clone(&missing)], + nulls.clone(), + )); + inputs.push(wrap(Arc::new(ListArray::new( + Arc::new(Field::new("item", state.data_type().clone(), true)), + OffsetBuffer::from_lengths([1]), + Arc::clone(&state), + None, + )))); + if nulls.is_some() { + inputs.push(wrap(Arc::new(StructArray::new( + vec![Field::new("a", state.data_type().clone(), true)].into(), + vec![state], + None, + )))); + } + } + for input in inputs { + let error = normalize_variant_array(&input, &target_field(false)).unwrap_err(); + assert!(error.to_string().contains("MALFORMED_VARIANT"), "{error}"); + } +} + #[test] fn canonical_and_shredded_values_normalize_equally() { let mut builder = VariantArrayBuilder::new(6); - builder.new_object().with_field("known", 1_i64).finish(); + builder.new_object().with_field("known", 1_i8).finish(); builder .new_object() - .with_field("known", 2_i64) + .with_field("known", 2_i8) .with_field("extra", 3_i64) .finish(); builder @@ -344,7 +499,7 @@ fn canonical_and_shredded_values_normalize_equally() { ) .unwrap(); - let prepared = prepare_variant_for_unshredding(&shredded).unwrap(); + let prepared = prepare_variant_for_unshredding(&shredded, None).unwrap(); assert!(Arc::ptr_eq( shredded.value_column(), prepared.value_column() @@ -371,7 +526,7 @@ fn canonical_and_shredded_values_normalize_equally() { } #[test] -fn normalize_unshredded_variant_orders_for_spark_and_is_idempotent() { +fn normalize_unshredded_variant_preserves_bytes_and_is_idempotent() { let keys = unicode_object_keys(); let mut builder = VariantBuilder::new().with_field_names(keys.iter().map(String::as_str)); let mut object = builder.new_object(); @@ -397,11 +552,14 @@ fn normalize_unshredded_variant_orders_for_spark_and_is_idempotent() { ); let first = normalize_variant_array(&physical, &target_field(false)).unwrap(); - assert_spark_unicode_output(first.as_struct()); let first_value = first.as_struct().column(0).as_binary::().value(0); + assert_eq!(first_value, value); + assert_eq!( + first.as_struct().column(1).as_binary::().value(0), + metadata + ); let second = normalize_variant_array(&first, &target_field(false)).unwrap(); - assert_spark_unicode_output(second.as_struct()); assert_eq!( second.as_struct().column(0).as_binary::().value(0), first_value @@ -459,7 +617,8 @@ fn unchanged_values_reuse_buffers_and_still_validate() { ); } let (output, changed) = - rewrite_residual_values(&values, metadata.as_binary::(), &[Some(0), None]).unwrap(); + rewrite_residual_values(&values, metadata.as_binary::(), None, &[Some(0), None]) + .unwrap(); assert!(!changed); assert!(Arc::ptr_eq(&values, &output)); @@ -478,6 +637,7 @@ fn unchanged_values_reuse_buffers_and_still_validate() { assert!(rewrite_residual_values( &values, missing_metadata.as_binary::(), + None, &[Some(0), Some(1)], ) .is_err()); @@ -511,6 +671,7 @@ fn lazy_rewrites_preserve_prefix_nulls_and_suffix() { let (output, changed) = rewrite_residual_values( &mixed, metadata.as_binary::(), + None, &[Some(0), None, Some(2), None], ) .unwrap(); @@ -525,25 +686,30 @@ fn lazy_rewrites_preserve_prefix_nulls_and_suffix() { ); } -// Run explicitly with --ignored --nocapture; fixture construction is outside the timed loop. +// Run explicitly with --release --features jemalloc -- --ignored --nocapture. +// Fixture construction is outside the timed loop; every row repeats its dictionary. #[test] #[ignore] fn benchmark_variant_buffer_reuse() { use std::{hint::black_box, time::Instant}; let rows = 4096; let payload = "x".repeat(4096); - let mut strings = VariantArrayBuilder::new(rows); let mut objects = VariantArrayBuilder::new(rows); + let mut empty_keys = VariantArrayBuilder::new(rows); for _ in 0..rows { - strings.append_variant(Variant::from(payload.as_str())); objects .new_object() - .with_field("known", 1_i64) + .with_field("known", 1_i8) .with_field("payload", payload.as_str()) .finish(); + empty_keys + .new_object() + .with_field("known", 1_i8) + .with_field("", payload.as_str()) + .finish(); } - let strings = strings.build(); let objects = objects.build(); + let empty_keys = empty_keys.build(); let shredded = shred_variant( &objects, &DataType::Struct(Fields::from(vec![Field::new( @@ -553,30 +719,67 @@ fn benchmark_variant_buffer_reuse() { )])), ) .unwrap(); - let target = target_field(false); - let DataType::Struct(fields) = target.data_type() else { - unreachable!() - }; - let strings: ArrayRef = Arc::new(StructArray::new( - fields.clone(), + let full = shred_variant( + &objects, + &DataType::Struct(Fields::from(vec![ + Field::new("known", DataType::Int64, true), + Field::new("payload", DataType::Utf8, true), + ])), + ) + .unwrap(); + let empty_metadata: ArrayRef = Arc::new(BinaryArray::from_iter_values((0..rows).map(|row| { + let mut metadata = binary_value(empty_keys.metadata_column(), row) + .unwrap() + .to_vec(); + metadata[0] &= !0x10; + metadata + }))); + let empty: ArrayRef = Arc::new(StructArray::new( vec![ - cast(strings.value_column().as_ref(), &DataType::Binary).unwrap(), - cast(strings.metadata_column().as_ref(), &DataType::Binary).unwrap(), + Field::new("metadata", DataType::Binary, false), + Field::new( + "value", + empty_keys.value_column().data_type().clone(), + false, + ), + Field::new("typed_value", DataType::Int64, true), + ] + .into(), + vec![ + empty_metadata, + Arc::clone(empty_keys.value_column()), + Arc::new(Int64Array::from(vec![None; rows])), ], None, )); - let shredded: ArrayRef = Arc::new(shredded.into_inner()); - for (name, input) in [("canonical", strings), ("partially_shredded", shredded)] { + let target = target_field(false); + let cases: [(&str, ArrayRef); 4] = [ + ("canonical", Arc::new(objects.into_inner())), + ("partially_shredded", Arc::new(shredded.into_inner())), + ("fully_shredded", Arc::new(full.into_inner())), + ("empty_key", empty), + ]; + for (name, input) in cases { for _ in 0..3 { black_box(normalize_variant_array(&input, &target).unwrap()); } + #[cfg(all(feature = "jemalloc", not(feature = "mimalloc")))] + let allocated = tikv_jemalloc_ctl::thread::allocatedp::read().unwrap(); + #[cfg(all(feature = "jemalloc", not(feature = "mimalloc")))] + let before = allocated.get(); let start = Instant::now(); for _ in 0..30 { black_box(normalize_variant_array(black_box(&input), &target).unwrap()); } + let elapsed = start.elapsed(); + #[cfg(all(feature = "jemalloc", not(feature = "mimalloc")))] + eprintln!( + "{name}: {} allocator bytes/row", + (allocated.get() - before) / (30 * rows as u64) + ); eprintln!( "{name}: {:.3} ms/batch, {rows} rows, 4096-byte payload", - start.elapsed().as_secs_f64() * 1000.0 / 30.0 + elapsed.as_secs_f64() * 1000.0 / 30.0 ); } } @@ -638,14 +841,21 @@ fn normalize_nested_list_residuals_use_their_root_metadata() { } object.finish(); let (metadata, value) = builder.finish(); - let metadata_array: ArrayRef = Arc::new(BinaryArray::from(vec![Some(metadata.as_slice())])); - let value: ArrayRef = Arc::new(BinaryArray::from(vec![Some(value.as_slice())])); - let value = reorder_variant_values(&value, &metadata_array, None).unwrap(); - (metadata, value.as_binary::().value(0).to_vec()) + let mut spark_metadata = WritableMetadataBuilder::from_iter(keys.iter().copied()); + spark_metadata.finish(); + let mut spark_metadata = spark_metadata.into_inner(); + spark_metadata[0] &= !0x10; + let value = variant_bytes( + &VariantMetadata::new(&spark_metadata), + Variant::new(&metadata, &value), + VariantObjectKeyOrder::SparkUtf16, + ) + .unwrap(); + (spark_metadata, value) } - let (metadata0, value0) = legacy_row(&["a", "\u{e000}", "😀"]); - let (metadata1, value1) = legacy_row(&["b", "zz", "\u{ffff}", "𐀀"]); + let (metadata0, value0) = legacy_row(&["a", "\u{e000}", "😀", ""]); + let (metadata1, value1) = legacy_row(&["b", "zz", "\u{ffff}", "𐀀", ""]); let states: ArrayRef = Arc::new( StructArray::try_new( Fields::from(vec![Field::new("value", DataType::Binary, true)]), @@ -691,7 +901,61 @@ fn normalize_nested_list_residuals_use_their_root_metadata() { 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)); + // Output slots follow Spark UTF-16 ordering, so Arrow's UTF-8 binary search cannot + // be used to look up supplementary characters in the normalized object. + let fields = object.iter().collect::>(); + assert_eq!(fields[key].as_int64(), Some(index as i64 + 2)); + assert_eq!(fields[""].as_int64(), Some(index as i64 + 3)); + } +} + +#[test] +fn normalize_spark_empty_key_metadata_rejects_other_malformed_encodings() { + // Spark dictionary ["z", "", "a"], deliberately requiring field ID remapping. + let metadata = [1, 3, 0, 1, 1, 2, b'z', b'a']; + let mut builder = VariantBuilder::new(); + let mut object = builder.new_object(); + object.insert("z", 1_i64); + object.insert("", 2_i64); + object.insert("a", 3_i64); + object.finish(); + let (canonical_metadata, canonical_value) = builder.finish(); + let value = variant_bytes( + &VariantMetadata::new(&metadata), + Variant::new(&canonical_metadata, &canonical_value), + VariantObjectKeyOrder::ArrowUtf8, + ) + .unwrap(); + let normalize = |metadata: &[u8]| { + let physical: ArrayRef = Arc::new(StructArray::new( + vec![ + Field::new("value", DataType::Binary, false), + Field::new("metadata", DataType::Binary, false), + ] + .into(), + vec![ + Arc::new(BinaryArray::from(vec![value.as_slice()])), + Arc::new(BinaryArray::from(vec![metadata])), + ], + None, + )); + normalize_variant_array(&physical, &target_field(false)) + }; + let output = normalize(&metadata).unwrap(); + let output = VariantArray::try_new(output.as_ref()).unwrap(); + assert_eq!( + output.value(0), + Variant::new(&canonical_metadata, &canonical_value) + ); + + for malformed in [ + vec![1, 3, 0, 1, 1, 2, 0xff, b'a'], // Invalid UTF-8. + vec![1, 3, 0, 2, 1, 2, b'z', b'a'], // Decreasing offsets. + vec![1, 3, 0, 1, 1, 2, b'z', b'z'], // Duplicate dictionary keys. + vec![1, 3, 0, 1, 1, 3, b'z', b'a'], // Out-of-bounds offset. + vec![1, 3, 0, 1, 1, 2, b'z', b'a', 0], // Unexpected trailing bytes. + ] { + assert!(normalize(&malformed).is_err(), "accepted {malformed:?}"); } } diff --git a/native/core/src/parquet/parquet_exec/variant_tests.rs b/native/core/src/parquet/parquet_exec/variant_tests.rs index 641529c40dd..ec857019656 100644 --- a/native/core/src/parquet/parquet_exec/variant_tests.rs +++ b/native/core/src/parquet/parquet_exec/variant_tests.rs @@ -34,7 +34,7 @@ use parquet::{ }, file::{properties::WriterProperties, writer::SerializedFileWriter}, schema::types::{Type as ParquetType, TypePtr}, - variant::{Variant, VariantArray, VariantBuilder, VariantDecimal16}, + variant::{Variant, VariantArray, VariantBuilder, VariantDecimal4}, }; use std::{fs::File, path::PathBuf}; fn required_variant_schema() -> SchemaRef { @@ -275,7 +275,7 @@ async fn variant_scan_uses_parquet_physical_types_instead_of_arrow_schema_hints( let output = write_and_scan_shredded_variant(decimal, false).await; assert_eq!( output.value(0), - Variant::Decimal16(VariantDecimal16::try_new(123, 2).unwrap()) + Variant::Decimal4(VariantDecimal4::try_new(123, 2).unwrap()) ); let date64: ArrayRef = Arc::new(Date64Array::from(vec![86_400_000])); @@ -344,7 +344,7 @@ async fn variant_scan_reads_wide_physical_decimal_as_decimal128() { for (index, value) in [123, -123].into_iter().enumerate() { assert_eq!( output.value(index), - Variant::Decimal16(VariantDecimal16::try_new(value, 2).unwrap()) + Variant::Decimal4(VariantDecimal4::try_new(value, 2).unwrap()) ); } } 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 58eab77e3b5..a4a3663aa00 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -840,7 +840,8 @@ case class CometExecRule(session: SparkSession) case writeFiles: WriteFilesExec => Seq(writeFiles.child) case other => Seq(other) } - if ((op.output ++ dataProducingChildren.flatMap(_.output)).exists(attr => + if (!op.isInstanceOf[CometScanExec] && + (op.output ++ dataProducingChildren.flatMap(_.output)).exists(attr => containsVariantType(attr.dataType))) { 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 b07052a1877..1c3d54be9de 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometScanRule.scala @@ -332,6 +332,14 @@ case class CometScanRule(session: SparkSession) withFallbackReason(scanExec, "Native Parquet scan does not support encryption") return None } + // TODO: Remove this fallback once DataFusion can ignore embedded Arrow schema hints and + // preserve Spark's ENUM inference without losing Parquet decryption state. + // https://github.com/apache/datafusion-comet/issues/5477 + if (encryptionEnabled(hadoopConf) && + scanExec.requiredSchema.exists(field => isVariantType(field.dataType))) { + withFallbackReason(scanExec, "Native Parquet Variant scans do not support encryption") + 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. @@ -1068,8 +1076,12 @@ 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) + // Admit Variant only at a required root in ordinary Parquet. Recursive and Iceberg type + // checks continue to use CometScanTypeChecker's stricter support rules. + 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 6bb3f0dcd59..d076c14f746 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EliminateRedundantTransitions.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EliminateRedundantTransitions.scala @@ -25,12 +25,13 @@ import org.apache.spark.sql.catalyst.util.sideBySide import org.apache.spark.sql.comet.{CometCollectLimitExec, CometColumnarToRowExec, CometIcebergWriteExec, CometMapInBatchExec, CometNativeColumnarToRowExec, CometNativeWriteExec, CometPlan, CometSparkToColumnarExec} import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.comet.shims.{MapInBatchInfo, ShimCometMapInBatch} +import org.apache.spark.sql.comet.util.Utils.containsVariantType import org.apache.spark.sql.execution.{ColumnarToRowExec, RowToColumnarExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.QueryStageExec import org.apache.spark.sql.execution.exchange.ReusedExchangeExec import org.apache.comet.CometConf -import org.apache.comet.CometSparkSessionExtensions.withInfo +import org.apache.comet.CometSparkSessionExtensions.{withFallbackReason, withInfo} import org.apache.comet.serde.NativeOptIn import org.apache.comet.shims.ShimSQLConf @@ -257,7 +258,18 @@ case class EliminateRedundantTransitions(session: SparkSession) } else { matchMapInArrow(plan) .orElse(matchMapInPandas(plan)) - .flatMap(info => extractColumnarChild(info.child).map(child => (info, child))) + .flatMap { info => + // TODO: Remove this guard once Comet Python operators preserve Variant identity + // and Spark's Arrow layout for both input and output. + // https://github.com/apache/datafusion-comet/issues/5437 + if ((info.output ++ info.child.output).exists(attr => + containsVariantType(attr.dataType))) { + withFallbackReason(plan, "Comet Python operators do not support type VariantType") + None + } else { + extractColumnarChild(info.child).map(child => (info, child)) + } + } } } } @@ -266,10 +278,19 @@ case class EliminateRedundantTransitions(session: SparkSession) * Creates an appropriate columnar to row transition operator. * * If native columnar to row conversion is enabled and the schema is supported, uses - * CometNativeColumnarToRowExec. Otherwise falls back to CometColumnarToRowExec. + * CometNativeColumnarToRowExec. Variant uses Spark's conversion; other unsupported schemas use + * CometColumnarToRowExec. */ private def createColumnarToRowExec(child: SparkPlan): SparkPlan = { val schema = child.schema + // TODO: Remove this fallback once Comet columnar-to-row conversion supports Variant getters + // and Spark's Variant UnsafeRow encoding. + // https://github.com/apache/datafusion-comet/issues/5436 + if (containsVariantType(schema)) { + return withFallbackReason( + ColumnarToRowExec(child), + "Native columnar-to-row conversion does not support type VariantType") + } val useNative = CometConf.COMET_NATIVE_COLUMNAR_TO_ROW_ENABLED.get() && CometNativeColumnarToRowExec.supportsSchema(schema) 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 52c6959cdc6..de8652ad90c 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometNativeScan.scala @@ -23,7 +23,7 @@ import scala.collection.mutable.ListBuffer import scala.jdk.CollectionConverters._ import org.apache.spark.internal.Logging -import org.apache.spark.sql.catalyst.expressions.{AttributeReference, Expression, Literal} +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, 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} @@ -51,6 +51,30 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS // like "file_size" could collide with a real column of the same name. Prefix to avoid it. private[comet] val constantMetadataFieldPrefix = "_comet_metadata_" + private val unsupportedDefaultReason = + "Full native scan disabled because one or more column default values are not supported" + + private[comet] def serializeExistenceDefaultValues( + schema: StructType, + output: Seq[Attribute]): Option[(Seq[Expr], Seq[java.lang.Long])] = { + val defaults = getExistenceDefaultValues(schema).iterator + .zip(schema.fields.iterator) + .zipWithIndex + .collect { + case ((value, field), index) if value != null => + val expression = if (isVariantType(field.dataType)) { + // Spark's vectorized reader cannot materialize a non-null Variant default. + None + } else { + Some(Literal.create(value, field.dataType)) + } + expression.flatMap(exprToProto(_, output)).map(_ -> java.lang.Long.valueOf(index)) + } + .toSeq + // Never drop a value independently of its index: that would shift every later default. + if (defaults.forall(_.isDefined)) Some(defaults.flatten.unzip) else None + } + /** * Build synthetic constant-metadata field names, uniquified against `reservedNames` (physical * data and partition schema names): DataFusion substitutes partition constants BY NAME, so a @@ -115,6 +139,28 @@ 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) + } + + if (scanExec.requiredSchema.exists(field => isVariantType(field.dataType))) { + // Spark's strict legacy reader owns malformed-layout errors (SPARK-47546). + // TODO: Remove this guard once the native reader implements Spark's strict Variant layout + // validation and malformed-input errors when allowReadingShredded=false. + if (!SQLConf.get.getConfString("spark.sql.variant.allowReadingShredded").toBoolean) { + withFallbackReason(scanExec, "Native Variant scans require allowReadingShredded=true") + } + // These settings change the interpretation of shredded timestamp children, whose types + // are not visible in the logical Variant schema at planning time. + // TODO: Remove this guard once the native reader receives these settings and applies + // Spark's timestamp inference to shredded Variant children. + if (SQLConf.get.legacyParquetNanosAsLong || !SQLConf.get.parquetInferTimestampNTZEnabled) { + withFallbackReason( + scanExec, + "Native Variant scans require default Parquet timestamp inference") + } + } + // the scan is supported if no fallback reasons were added to the node !hasFallbackReason(scanExec) } @@ -168,23 +214,13 @@ 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)) => + 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). @@ -211,11 +247,8 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS val partitionSchema = schema2Proto(partitionSchemaFields) val requiredSchema = schema2Proto(scan.requiredSchema) - // Spark's required schema can prune a Variant column, including one nested under an - // unrequested struct, while the complete relation schema still contains that unsupported - // type. Exclude unread roots and replace requested roots with their already-validated, - // pruned required fields so Variant never enters the native reader data schema. A requested - // Variant is rejected by CometScanRule and CometExecRule before reaching this point. + // Retain the pruned required field for a requested Variant root, including a struct whose + // Variant child was pruned. Entirely unread Variant roots never enter the native schema. val nativeDataSchema = StructType(scan.relation.dataSchema.fields.flatMap { field => if (containsVariantType(field.dataType)) { scan.requiredSchema.fields.find(requiredField => 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 8392aa76af2..01fb233d8db 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 @@ -56,13 +56,12 @@ 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 whole-value Variant reader does not support this pushed + // 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 + // Outside direct Parquet projection, Comet has no native execution path for Spark 4's + // `VariantType`. 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. def isVariantType(dt: DataType): Boolean = dt.isInstanceOf[VariantType] diff --git a/spark/src/main/spark-4.x/org/apache/spark/sql/comet/shims/ShimSparkErrorConverter.scala b/spark/src/main/spark-4.x/org/apache/spark/sql/comet/shims/ShimSparkErrorConverter.scala index 7397745885c..30e25423357 100644 --- a/spark/src/main/spark-4.x/org/apache/spark/sql/comet/shims/ShimSparkErrorConverter.scala +++ b/spark/src/main/spark-4.x/org/apache/spark/sql/comet/shims/ShimSparkErrorConverter.scala @@ -288,6 +288,9 @@ trait ShimSparkErrorConverter { case "CannotParseDecimal" => Some(QueryExecutionErrors.cannotParseDecimalError()) + case "MalformedVariant" => + Some(QueryExecutionErrors.malformedVariant()) + case "InvalidUtf8String" => val hexStr = UTF8String.fromString(params("hexString").toString) Some(QueryExecutionErrors.invalidUTF8StringError(hexStr)) 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 c1c018398bc..ec50d49bd38 100644 --- a/spark/src/test/resources/sql-tests/expressions/misc/variant.sql +++ b/spark/src/test/resources/sql-tests/expressions/misc/variant.sql @@ -15,11 +15,12 @@ -- 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. +-- Checks Variant pruning and fallback with Spark's strict unshredded reader. -- MinSparkVersion: 4.0 +-- Config: spark.sql.variant.allowReadingShredded=false +-- Config: spark.sql.variant.pushVariantIntoScan=false +-- Config: spark.sql.variant.writeShredding.enabled=false statement CREATE TABLE test_variant(id INT, v VARIANT, tail STRING) USING parquet @@ -47,16 +48,16 @@ SELECT id, tail FROM test_variant WHERE tail IS NOT NULL ORDER BY id query expect_fallback(Native operators do not support schemas containing type VariantType) SELECT CAST(id AS VARIANT) FROM test_variant -query expect_fallback(type VariantType) +query expect_fallback(Native Variant scans require allowReadingShredded=true) SELECT id, v FROM test_variant ORDER BY id -query expect_fallback(type VariantType) +query expect_fallback(Native Variant scans require allowReadingShredded=true) SELECT variant_get(v, '$.a', 'int') AS a FROM test_variant ORDER BY id -query expect_fallback(type VariantType) +query expect_fallback(Native Variant scans require allowReadingShredded=true) SELECT id FROM test_variant WHERE variant_get(v, '$.a', 'int') = 1 -query expect_fallback(type VariantType) +query expect_fallback(Native Variant scans require allowReadingShredded=true) SELECT COUNT(*) FROM test_variant WHERE v IS NOT NULL statement diff --git a/spark/src/test/scala/org/apache/comet/CometVariantProjectionSuite.scala b/spark/src/test/scala/org/apache/comet/CometVariantProjectionSuite.scala new file mode 100644 index 00000000000..8bb9f62a72c --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/CometVariantProjectionSuite.scala @@ -0,0 +1,390 @@ +/* + * 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. + */ + +package org.apache.comet + +import org.apache.hadoop.fs.Path +import org.apache.parquet.example.data.simple.SimpleGroup +import org.apache.parquet.io.api.Binary +import org.apache.parquet.schema.MessageTypeParser +import org.apache.spark.SparkConf +import org.apache.spark.sql.{CometTestBase, DataFrame, Row} +import org.apache.spark.sql.comet.CometNativeColumnarToRowExec +import org.apache.spark.sql.comet.CometNativeScanExec +import org.apache.spark.sql.comet.util.Utils +import org.apache.spark.sql.execution.{ColumnarToRowExec, CommandResultExec, ProjectExec, SparkPlan} +import org.apache.spark.sql.execution.command.DataWritingCommandExec +import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.{IntegerType, StructField, StructType} + +import org.apache.comet.serde.operator.CometNativeScan + +class CometVariantProjectionSuite extends CometTestBase { + override protected def sparkConf: SparkConf = super.sparkConf + .set(SQLConf.USE_V1_SOURCE_LIST.key, "parquet") + .set(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key, "false") + .set("spark.sql.variant.allowReadingShredded", "true") + .set("spark.sql.variant.pushVariantIntoScan", "false") + + private def withVariantFile(query: String)(check: String => Unit): Unit = { + assume(Utils.variantType.isDefined, "VariantType requires Spark 4.0+") + withTempPath { dir => + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + sql(query).coalesce(1).write.parquet(dir.getCanonicalPath) + } + check(dir.getCanonicalPath) + } + } + + private def checkVariantAnswer(df: DataFrame, expected: Seq[Row]): SparkPlan = { + // VariantVal equality compares both value and metadata bytes, including integer widths. + checkAnswer(df, expected) + df.queryExecution.executedPlan + } + + private def sparkRows(df: => DataFrame): Seq[Row] = { + var rows = Seq.empty[Row] + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + rows = df.collect().toSeq + } + rows + } + + private def checkNative(df: => DataFrame, expected: Option[Seq[Row]] = None): Unit = { + val plan = checkVariantAnswer(df, expected.getOrElse(sparkRows(df))) + checkCometOperators(plan, classOf[ColumnarToRowExec]) + assert(collect(plan) { case scan: CometNativeScanExec => scan }.nonEmpty, plan.toString) + assert(collect(plan) { case c: CometNativeColumnarToRowExec => c }.isEmpty, plan.toString) + } + + private def checkScanFallbackPlan(df: DataFrame, reason: String): Unit = { + val plan = df.queryExecution.executedPlan + assert(new ExtendedExplainInfo().getFallbackReasons(plan).exists(_.contains(reason))) + assert(collect(plan) { case scan: CometNativeScanExec => scan }.isEmpty, plan.toString) + } + + private def checkScanFallback(df: => DataFrame, reason: String): Unit = { + val (_, plan) = checkSparkAnswerAndFallbackReason(df, reason) + assert(collect(plan) { case scan: CometNativeScanExec => scan }.isEmpty, plan.toString) + } + + test("direct Variant projection preserves values and siblings") { + withVariantFile(""" + SELECT id, parse_json(json) AS v, id + 10 AS tail FROM VALUES + (1, '{"a":1,"nested":{"b":[true,null,2.5]}}'), + (2, '[1,"text",false,{"x":2}]'), + (3, '42'), (4, '"text"'), (5, 'null'), (6, NULL), + (7, '{}'), (8, '[]') AS input(id, json) + """) { path => + checkNative(spark.read.parquet(path).select("v")) + checkNative(spark.read.parquet(path).select("id", "v", "tail")) + } + withVariantFile("SELECT 1 AS id, CAST(NULL AS VARIANT) AS v") { path => + checkNative(spark.read.parquet(path)) + } + } + + test("Variant objects with empty keys match Spark") { + for (shredding <- Seq("false", "true")) { + withSQLConf("spark.sql.variant.writeShredding.enabled" -> shredding) { + withVariantFile(""" + SELECT id, parse_json(json) AS v FROM VALUES + (1, '{"":1}'), (2, '{"z":1,"":2,"a":{"":3}}'), + (3, '[{"z":4,"":5},{"":6}]'), (4, NULL) AS input(id, json) + """) { path => + checkNative(spark.read.parquet(path)) + } + } + } + } + + test("shredded Variant missing fields and redundant residuals match Spark") { + // The metadata dictionary may omit keys that only occur in typed_value. + withVariantFile(""" + SELECT named_struct('metadata', X'010000', 'typed_value', + named_struct('a', named_struct('value', residual, 'typed_value', a), + 'b', named_struct('typed_value', b))) AS v + FROM VALUES (CAST(NULL AS BINARY), CAST(NULL AS INT), CAST(NULL AS INT)), + (X'00', NULL, NULL), (NULL, 1, NULL), (NULL, NULL, 2), + (X'00', NULL, 2), (NULL, 3, 4) AS input(residual, a, b) + """) { path => + checkNative(spark.read.schema("v VARIANT").parquet(path)) + } + for (typed <- Seq( + "'typed scalar'", + "array(named_struct('typed_value', named_struct('inner', " + + "named_struct('typed_value', 7))))")) { + // Spark ignores the residual for a present scalar or array typed_value. + withVariantFile(s""" + SELECT named_struct('metadata', X'010000', 'value', X'FF', + 'typed_value', $typed) AS v + """) { path => + checkNative(spark.read.schema("v VARIANT").parquet(path)) + } + } + } + + test("shredded Variant scalar encodings and dictionary traversal match Spark bytes") { + for (typed <- Seq( + "CAST(n AS BIGINT)", + "CAST(n AS DECIMAL(38, 2))", + "CAST(n AS STRING)", + "array(named_struct('typed_value', CAST(n AS BIGINT)))")) { + withVariantFile(s""" + SELECT named_struct('metadata', X'01010006756E75736564', 'typed_value', $typed) AS v + FROM VALUES (-2147483649L), (-32769L), (-129L), (-128L), (0L), (127L), + (128L), (32767L), (32768L), (2147483647L), (2147483648L) AS input(n) + """) { path => + checkNative(spark.read.schema("v VARIANT").parquet(path)) + } + } + // Source IDs are z=0, unused=1. Spark visits b, its child, a, then residual z; + // the residual's wide integer remains wide while typed integers are narrowed. + withVariantFile(""" + SELECT named_struct('metadata', X'01020001077A756E75736564', + 'value', X'0201000009180900000000000000', 'typed_value', + named_struct('b', named_struct('typed_value', named_struct('inner', + named_struct('typed_value', 1))), 'a', named_struct('typed_value', 2))) AS v + """) { path => + checkNative(spark.read.schema("v VARIANT").parquet(path)) + } + // Merely having typed_value in the schema triggers reconstruction, even when it is null. + for (typed <- Seq("", ", 'typed_value', CAST(NULL AS INT)")) { + withVariantFile(s""" + SELECT named_struct('metadata', X'01010006756E75736564', + 'value', X'180100000000000000' $typed) AS v + """) { path => + checkNative(spark.read.schema("v VARIANT").parquet(path)) + } + } + } + + test("malformed shredded Variant values report Spark's error class") { + for (typed <- Seq( + "CAST(NULL AS INT)", + "array(named_struct('typed_value', CAST(NULL AS INT)))", + "named_struct('a', CAST(NULL AS STRUCT))")) { + withVariantFile(s""" + SELECT named_struct('metadata', X'010000', 'typed_value', $typed) AS v + """) { path => + val df = spark.read.schema("v VARIANT").parquet(path) + assert(collect(df.queryExecution.executedPlan) { case scan: CometNativeScanExec => + scan + }.nonEmpty) + val (sparkError, cometError) = checkSparkAnswerMaybeThrows(df) + for (error <- Seq(sparkError, cometError)) { + val causes = Iterator.iterate(error.get)(_.getCause).takeWhile(_ != null).toSeq + assert(!causes.exists(_.isInstanceOf[CometNativeException])) + assert( + causes + .collect { case e: org.apache.spark.SparkThrowable => + e.getErrorClass + } + .contains("MALFORMED_VARIANT")) + } + } + } + } + + test("non-null Variant existence defaults fall back to Spark") { + assume(Utils.variantType.isDefined, "VariantType requires Spark 4.0+") + val schema = StructType( + Seq( + StructField("id", IntegerType), + StructField("before", IntegerType).withExistenceDefaultValue("11"), + StructField("v", Utils.variantType.get) + .withExistenceDefaultValue("parse_json('{\"default\":42}')"), + StructField("tail", IntegerType).withExistenceDefaultValue("99"))) + assert(CometNativeScan.serializeExistenceDefaultValues(schema, Seq.empty).isEmpty) + for (query <- Seq( + "SELECT 1 AS id", + "SELECT 1 AS id, CAST(NULL AS VARIANT) AS v, 7 AS tail", + "SELECT 1 AS id, parse_json('{\"present\":true}') AS v, 7 AS tail")) { + withVariantFile(query) { path => + val df = spark.read.schema(schema).parquet(path) + checkScanFallbackPlan(df, "one or more column default values are not supported") + if (query == "SELECT 1 AS id") { + val (sparkError, cometError) = checkSparkAnswerMaybeThrows(df) + assert(sparkError.nonEmpty && cometError.nonEmpty) + } else { + checkAnswer(df, sparkRows(spark.read.schema(schema).parquet(path))) + } + } + } + } + + test("null Variant existence defaults preserve later default indexes") { + assume(Utils.variantType.isDefined, "VariantType requires Spark 4.0+") + val schema = StructType( + Seq( + StructField("id", IntegerType), + StructField("v", Utils.variantType.get).withExistenceDefaultValue("NULL"), + StructField("tail", IntegerType).withExistenceDefaultValue("99"))) + withVariantFile("SELECT 1 AS id") { path => + checkNative(spark.read.schema(schema).parquet(path), Some(Seq(Row(1, null, 99)))) + } + } + + test("Variant projection uses shared Unicode field matching") { + withSQLConf(SQLConf.CASE_SENSITIVE.key -> "false") { + for ((physical, logical) <- Seq("MÜNCHEN" -> "münchen", "K" -> "k", "ſ" -> "s")) { + withVariantFile(s"""SELECT parse_json('{"a":1}') AS `$physical`, 7 AS `Ü`""") { path => + val schema = StructType( + Seq(StructField(logical, Utils.variantType.get), StructField("ü", IntegerType))) + checkNative(spark.read.schema(schema).parquet(path)) + } + } + } + } + + test("unread Variant roots and nested fields are pruned from native scans") { + withVariantFile(""" + SELECT 1 AS id, parse_json('{"a":1}') AS v, + named_struct('n', 7, 'v', parse_json('[1,2]')) AS s + """) { path => + checkNative(spark.read.parquet(path).select("id")) + checkNative(spark.read.parquet(path).select("s.n")) + checkScanFallback(spark.read.parquet(path).select("s"), "VariantType") + } + for (nested <- Seq("array(parse_json('1'))", "map('key', parse_json('1'))")) { + withVariantFile(s"SELECT $nested AS nested") { path => + checkScanFallback(spark.read.parquet(path), "VariantType") + } + } + } + + test("Variant scans preserve strict reader and timestamp inference fallbacks") { + withSQLConf("spark.sql.variant.writeShredding.enabled" -> "false") { + withVariantFile("SELECT parse_json('{\"a\":1}') AS v") { path => + withSQLConf("spark.sql.variant.allowReadingShredded" -> "false") { + checkScanFallback(spark.read.parquet(path), "allowReadingShredded=true") + } + for (setting <- Seq( + "spark.sql.legacy.parquet.nanosAsLong" -> "true", + "spark.sql.parquet.inferTimestampNTZ.enabled" -> "false")) { + withSQLConf(setting) { + checkScanFallback(spark.read.parquet(path), "default Parquet timestamp inference") + } + } + withSQLConf("spark.sql.variant.pushVariantIntoScan" -> "true") { + checkScanFallback( + spark.read.parquet(path).selectExpr("variant_get(v, '$.a', 'int')"), + "VariantType") + } + } + } + } + + test("Variant consumers fall back above a native scan") { + withVariantFile("SELECT 1 AS id, parse_json('{\"a\":1}') AS v") { path => + withSQLConf("spark.sql.variant.pushVariantIntoScan" -> "false") { + val (_, plan) = checkSparkAnswerAndFallbackReason( + spark.read.parquet(path).selectExpr("variant_get(v, '$.a', 'int')"), + "Native operators do not support schemas containing type VariantType") + assert(collect(plan) { case p: ProjectExec => p }.nonEmpty) + assert(collect(plan) { case s: CometNativeScanExec => s }.nonEmpty) + } + val expected = sparkRows(spark.read.parquet(path)) + val plan = checkVariantAnswer(spark.read.parquet(path).repartition(2), expected) + assert(collect(plan) { case s: ShuffleExchangeExec => s }.nonEmpty) + assert(collect(plan) { case s: CometNativeScanExec => s }.nonEmpty) + + withTempView("variant_source") { + spark.read.parquet(path).createOrReplaceTempView("variant_source") + withTempPath { output => + withTable("variant_copy") { + sql( + "CREATE TABLE variant_copy (id INT, v VARIANT) USING parquet " + + s"LOCATION '${output.getCanonicalPath}'") + withSQLConf( + CometConf.COMET_NATIVE_PARQUET_WRITE_ENABLED.key -> "true", + CometConf.getOperatorAllowIncompatConfigKey( + classOf[DataWritingCommandExec]) -> "true") { + val command = sql("INSERT INTO variant_copy SELECT * FROM variant_source") + val plan = command.queryExecution.executedPlan + .asInstanceOf[CommandResultExec] + .commandPhysicalPlan + assert( + collect(plan) { case write: DataWritingCommandExec => write }.nonEmpty, + plan.toString) + assert( + new ExtendedExplainInfo() + .getFallbackReasons(plan) + .exists(_.contains( + "Native operators do not support schemas containing type VariantType"))) + checkNative(spark.read.parquet(output.getCanonicalPath)) + } + } + } + } + } + } + + test("encrypted Variant scans fall back to Spark") { + withSQLConf( + "parquet.crypto.factory.class" -> + "org.apache.parquet.crypto.keytools.PropertiesDrivenCryptoFactory", + "parquet.encryption.kms.client.class" -> + "org.apache.parquet.crypto.keytools.mocks.InMemoryKMS", + "parquet.encryption.key.list" -> "variantKey: MDEyMzQ1Njc4OTAxMjM0NQ==", + "parquet.encryption.uniform.key" -> "variantKey") { + withVariantFile("SELECT parse_json('{\"a\":1}') AS v") { path => + checkScanFallback(spark.read.parquet(path), "Variant scans do not support encryption") + } + } + } + + test("strict Variant reader preserves malformed layout errors") { + assume(Utils.variantType.isDefined, "VariantType requires Spark 4.0+") + withTempPath { file => + val physical = MessageTypeParser.parseMessageType("""message root { + optional group v { + required binary value; + optional binary metadata; + } + }""") + val writer = createParquetWriter(physical, new Path(file.toURI)) + try { + val row = new SimpleGroup(physical) + row + .addGroup("v") + .append("value", Binary.fromConstantByteArray(Array[Byte](0))) + .append("metadata", Binary.fromConstantByteArray(Array[Byte](1, 0, 0))) + writer.write(row) + } finally { + writer.close() + } + withSQLConf("spark.sql.variant.allowReadingShredded" -> "false") { + val df = spark.read + .schema(StructType(Seq(StructField("v", Utils.variantType.get)))) + .parquet(file.getCanonicalPath) + checkScanFallbackPlan(df, "allowReadingShredded=true") + val error = intercept[Exception](df.collect()) + assert( + Iterator + .iterate[Throwable](error)(_.getCause) + .takeWhile(_ != null) + .exists(cause => + Option(cause.getMessage).exists( + _.contains("INVALID_VARIANT_FROM_PARQUET.NULLABLE_OR_NOT_BINARY_FIELD")))) + } + } + } +} diff --git a/spark/src/test/scala/org/apache/comet/vector/NativeUtilSuite.scala b/spark/src/test/scala/org/apache/comet/vector/NativeUtilSuite.scala index ec9dde945cf..e24a396c5f6 100644 --- a/spark/src/test/scala/org/apache/comet/vector/NativeUtilSuite.scala +++ b/spark/src/test/scala/org/apache/comet/vector/NativeUtilSuite.scala @@ -35,7 +35,7 @@ import org.apache.spark.sql.CometTestBase import org.apache.spark.sql.comet.CometExec import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.vectorized.ConstantColumnVector -import org.apache.spark.sql.types.{IntegerType, StringType, StructField, StructType} +import org.apache.spark.sql.types.{BinaryType, IntegerType, StringType, StructField, StructType} import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} import org.apache.comet.CometConf @@ -432,7 +432,20 @@ class NativeUtilSuite extends CometTestBase { val iterator = CometExec.getCometIterator(Array.empty[Object], 1, plan, 1, 0) try { assert(iterator.hasNext) - Iterator.single(iterator.next().column(0).dataType()) + val column = iterator.next().column(0) + assert(column.getChild(0).dataType() == BinaryType) + assert(column.getChild(1).dataType() == BinaryType) + // Spark 3.x has no getVariant method; the suite still compiles for that profile. + val value = classOf[ColumnVector] + .getMethod("getVariant", classOf[Int]) + .invoke(column, Int.box(0)) + assert( + value.getClass + .getMethod("getValue") + .invoke(value) + .asInstanceOf[Array[Byte]] + .sameElements(Array[Byte](0))) + Iterator.single(column.dataType()) } finally { iterator.close() } diff --git a/spark/src/test/spark-4.x/org/apache/spark/sql/CometVariantShreddingSuite.scala b/spark/src/test/spark-4.x/org/apache/spark/sql/CometVariantShreddingSuite.scala new file mode 100644 index 00000000000..c5359c3512c --- /dev/null +++ b/spark/src/test/spark-4.x/org/apache/spark/sql/CometVariantShreddingSuite.scala @@ -0,0 +1,53 @@ +/* + * 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. + */ + +package org.apache.spark.sql + +import java.io.File + +import org.apache.spark.SparkConf +import org.apache.spark.sql.comet.CometNativeScanExec +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.test.TestSparkSession + +import org.apache.comet.{CometConf, CometSparkSessionExtensions} + +/** Run Spark's unchanged reconstruction assertions in Comet's regular Spark 4 CI jobs. */ +class CometVariantShreddingSuite extends VariantShreddingSuite { + override protected def sparkConf: SparkConf = super.sparkConf + .set(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key, "false") + .set(CometConf.COMET_ENABLED.key, "true") + .set(CometConf.COMET_EXEC_ENABLED.key, "true") + .set(CometConf.COMET_ONHEAP_ENABLED.key, "true") + .set(CometConf.COMET_SHUFFLE_ENABLED.key, "false") + + override protected def createSparkSession: TestSparkSession = { + val session = super.createSparkSession + new CometSparkSessionExtensions().apply(session.extensions) + session + } + + override def checkExpr(path: File, expr: String, expected: Any*): Unit = { + super.checkExpr(path, expr, expected: _*) + if (expr == "v" && !isPushEnabled) { + val plan = read(path).queryExecution.executedPlan + assert(plan.collect { case scan: CometNativeScanExec => scan }.nonEmpty, plan.toString) + } + } +} diff --git a/spark/src/test/spark-4.x/org/apache/spark/sql/benchmark/CometVariantReadBenchmark.scala b/spark/src/test/spark-4.x/org/apache/spark/sql/benchmark/CometVariantReadBenchmark.scala new file mode 100644 index 00000000000..11a43a8d9a8 --- /dev/null +++ b/spark/src/test/spark-4.x/org/apache/spark/sql/benchmark/CometVariantReadBenchmark.scala @@ -0,0 +1,114 @@ +/* + * 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. + */ + +package org.apache.spark.sql.benchmark + +import org.apache.spark.benchmark.Benchmark +import org.apache.spark.sql.Encoders +import org.apache.spark.sql.comet.CometNativeScanExec +import org.apache.spark.types.variant.VariantBuilder + +import org.apache.comet.CometConf + +/** + * Matched warm local scans with repeated Parquet dictionaries. Both readers hash every returned + * Variant's value and metadata bytes. Includes row conversion and consumption; excludes writes. + * Run with -Pspark-4.0 or -Pspark-4.1 using benchmark-org.apache.spark.sql.benchmark. + * CometVariantReadBenchmark [rows] [payloadBytes] [--reverse-cases]. + */ +object CometVariantReadBenchmark extends CometBenchmarkBase { + override def runCometBenchmark(args: Array[String]): Unit = { + val sizes = args.filterNot(_ == "--reverse-cases") + val rows = sizes.headOption.map(_.toInt).getOrElse(100000) + val payloadBytes = sizes.lift(1).map(_.toInt).getOrElse(1024) + val payload = "x" * payloadBytes + val readers = if (args.contains("--reverse-cases")) Seq(true, false) else Seq(false, true) + runBenchmark("Variant scans with repeated dictionaries") { + withSQLConf( + "spark.sql.sources.useV1SourceList" -> "parquet", + "spark.sql.adaptive.enabled" -> "false", + CometConf.COMET_EXEC_ENABLED.key -> "true", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true", + CometConf.COMET_ONHEAP_ENABLED.key -> "true", + "spark.sql.variant.allowReadingShredded" -> "true", + "spark.sql.variant.pushVariantIntoScan" -> "false") { + for (shape <- Seq("canonical", "partially shredded", "fully shredded", "empty key")) { + val key = if (shape == "empty key") "" else "payload" + val json = + if (shape == "canonical") s"""{"known":1,"$key":"$payload"}""" + else s"""{"$key":"$payload"}""" + // Parquet metadata includes shredded keys too; residual IDs use that dictionary. + val builder = new VariantBuilder(false) + Seq("known", key).foreach(builder.addKey) + builder.appendVariant(VariantBuilder.parseJson(json, false)) + val value = builder.result() + def binary(bytes: Array[Byte]): String = + "X'" + bytes.map(b => f"${b & 0xff}%02X").mkString + "'" + val metadata = binary(value.getMetadata) + val residual = binary(value.getValue) + val fields = shape match { + case "canonical" => s"'metadata', $metadata, 'value', $residual" + case "fully shredded" => + s"""'metadata', $metadata, 'typed_value', named_struct( + |'known', named_struct('typed_value', 1), + |'payload', named_struct('typed_value', '$payload'))""".stripMargin + case _ => + s"""'metadata', $metadata, 'value', $residual, 'typed_value', + |named_struct('known', named_struct('typed_value', 1))""".stripMargin + } + withTempPath { dir => + spark + .sql(s"SELECT named_struct($fields) AS v FROM range($rows)") + .coalesce(1) + .write + .option("parquet.enable.dictionary", "true") + .parquet(dir.getCanonicalPath) + def read() = spark.read.schema("v VARIANT").parquet(dir.getCanonicalPath) + val expected = read().head() + val benchmark = new Benchmark( + s"Variant $shape: $payloadBytes payload bytes", + rows, + minNumIters = 5, + output = output) + for (enabled <- readers) { + withSQLConf(CometConf.COMET_ENABLED.key -> enabled.toString) { + val df = read() + assert(df.head() == expected) + assert(collect(df.queryExecution.executedPlan) { case scan: CometNativeScanExec => + scan + }.nonEmpty == enabled) + } + benchmark.addCase(if (enabled) "Comet" else "Spark") { _ => + withSQLConf(CometConf.COMET_ENABLED.key -> enabled.toString) { + read() + .mapPartitions { rows => + Iterator.single( + rows.foldLeft(0L)((sum, row) => sum + row.get(0).hashCode())) + }(Encoders.scalaLong) + .collect() + } + } + } + benchmark.run() + } + } + } + } + } +} 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 8d8617fae33..2a3a0f9920e 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.{LongType, StructField, StructType, VariantType} import org.apache.spark.sql.vectorized.ColumnarBatch import org.apache.comet.{CometConf, ExtendedExplainInfo} @@ -105,6 +105,31 @@ class CometMapInBatchSuite extends CometTestBase { } } + test("Variant inputs and outputs keep Python operators on Spark") { + val plain = Seq(AttributeReference("id", LongType)()) + val variant = Seq(AttributeReference("v", VariantType)()) + withSQLConf(CometConf.COMET_PYARROW_UDF_ENABLED.key -> "true") { + for ((input, output) <- Seq(variant -> plain, plain -> variant)) { + val udf = stubPythonUDF.copy( + children = input, + dataType = StructType(output.map(attr => StructField(attr.name, attr.dataType)))) + val plan = MapInArrowExec( + udf, + output, + ColumnarToRowExec(StubCometLeaf(input)), + isBarrier = false, + profile = None) + val rewritten = EliminateRedundantTransitions(spark).apply(plan) + assert(rewritten.isInstanceOf[MapInArrowExec]) + assert(!rewritten.exists(_.isInstanceOf[CometMapInBatchExec])) + assert( + new ExtendedExplainInfo() + .getFallbackReasons(rewritten) + .exists(_.contains("Comet Python operators do not support type VariantType"))) + } + } + } + test("rule annotates operator with opt-in hint when feature is disabled") { withSQLConf(CometConf.COMET_PYARROW_UDF_ENABLED.key -> "false") { val rewritten = EliminateRedundantTransitions(spark).apply(buildPlan())