diff --git a/docs/source/contributor-guide/expression-audits/map_funcs.md b/docs/source/contributor-guide/expression-audits/map_funcs.md index fc9c5584c9e..533338545d5 100644 --- a/docs/source/contributor-guide/expression-audits/map_funcs.md +++ b/docs/source/contributor-guide/expression-audits/map_funcs.md @@ -45,9 +45,14 @@ ## map_from_arrays - Spark 3.4.3 (audited 2026-05-27): identical to 3.5.8. -- Spark 3.5.8 (audited 2026-05-27): baseline. `MapFromArrays(left, right) extends BinaryExpression with NullIntolerant`; Spark uses `ArrayBasedMapBuilder` to detect duplicate keys (subject to `spark.sql.mapKeyDedupPolicy`) and rejects null keys with `RuntimeException("Cannot use null as map key")`. Comet `CometMapFromArrays` wraps the inputs in `CaseWhen(IsNotNull(left) AND IsNotNull(right), map(left, right), null)` so NULL-array inputs return NULL rather than triggering the previously reported native crash ([#3327](https://github.com/apache/datafusion-comet/issues/3327)). +- Spark 3.5.8 (audited 2026-05-27): baseline. `MapFromArrays(left, right) extends BinaryExpression with NullIntolerant`; Spark uses `ArrayBasedMapBuilder` to detect duplicate keys (subject to `spark.sql.mapKeyDedupPolicy`) and rejects null keys with `RuntimeException("Cannot use null as map key")`. Comet `CometMapFromArrays` wires the native `map_from_arrays` from `datafusion-spark`, which is null intolerant the same way, so NULL-array inputs return NULL rather than triggering the previously reported native crash ([#3327](https://github.com/apache/datafusion-comet/issues/3327)). The serde still wraps the call in `CASE WHEN left IS NOT NULL THEN map_from_arrays(left, right) END`, for evaluation order rather than the result: `BinaryExpression.eval` never evaluates `right` for a row whose `left` is NULL, and DataFusion evaluates a THEN branch only on the rows its WHEN selected, so under ANSI a failing cast in the values array does not run for such a row. A single `left IS NOT NULL AND right IS NOT NULL` guard would not give that, since DataFusion's `AND` evaluates its right side on the whole batch unless the left side is false on all or most rows. `right` needs no guard: Spark evaluates it whenever `left` is not NULL, and the native function returns a NULL map for a NULL `right`. - Spark 4.0.1 (audited 2026-05-27): semantics unchanged; `NullIntolerant` trait replaced by `nullIntolerant: Boolean`. - Spark 4.1.1 (audited 2026-05-27): identical to 4.0.1. +- `ArrayBasedMapBuilder` semantics, reproduced natively rather than falling back ([#4680](https://github.com/apache/datafusion-comet/issues/4680)): a `NULL` key raises `NULL_MAP_KEY`, and a duplicate key follows `spark.sql.mapKeyDedupPolicy` (`EXCEPTION` raises `DUPLICATED_MAP_KEY` naming the key, `LAST_WIN` keeps the last value in the slot of the key's first occurrence). Spark checks each row's key and value lengths before inserting any of its entries, then inserts them one at a time, so the builders report whichever of a length mismatch, a `NULL` key and a duplicate key comes first, in the first row that has one. A row whose key and value arrays differ in length raises `SparkError::MapKeyValueDiffSizes`, which `ShimSparkErrorConverter` turns into the `_LEGACY_ERROR_TEMP_2128` Spark raises (`The key array and value array of MapData must have the same length`). +- The duplicate-key policy reaches the native session as `datafusion.spark.map_key_dedup_policy`, which `datafusion-spark`'s map kernels read. `CometExecIterator.serializeCometSQLConfs` reads `spark.sql.mapKeyDedupPolicy` when it builds the native plan for a task, so materializing or explaining a plan does not fix it and a Dataset re-executed after a change to the setting uses the new value. The same applies to `map_from_entries` and `str_to_map`. +- Known limitation: on Spark 4.0+, `ArrayBasedMapBuilder` normalizes a floating-point key before comparing it (`keyNormalizer`, added in 4.0 with `spark.sql.legacy.disableMapKeyNormalization`), so `-0.0` and `+0.0` are one key and all `NaN`s are one key; the native builder compares the raw bits and keeps `-0.0` and `+0.0`, and `NaN`s with different bit patterns, apart. `from` returns the input arrays untouched when no key repeated, so the stored keys match Spark either way and only duplicate detection diverges. Spark 3.4 and 3.5 do not normalize, so they already match. `spark.comet.exec.strictFloatingPoint=true` marks the expression `Incompatible` for a floating-point key type, and the projection falls back to Spark (`map_builders_strict_fp.sql`). +- Known limitation: Spark reads `spark.sql.mapKeyDedupPolicy` into `ArrayBasedMapBuilder`, a lazy field of the map expression, so _when_ it reads it depends on how the projection runs. Outside whole-stage codegen (the flag off, or a projection wider than `spark.sql.codegen.maxFields`) the projection is rebuilt in every task and the setting is read again on each action, which is what Comet does. Inside whole-stage codegen Spark creates the builder once on the driver, in the first action, and keeps it, so a Dataset re-executed after a change to the setting still builds its maps under the policy it started with, where Comet uses the new one. Comet cannot tell the two apart: it replaces the operator before `CollapseCodegenStages` runs, so the plan it sees carries no record of which path Spark would have taken. Matching the whole-stage case instead would mean returning a map where Spark raises `DUPLICATED_MAP_KEY` in the other three configurations, so the loud divergence is preferred over the silent one. Only a Dataset that is executed more than once across a change to the setting is affected. +- Known limitation: the NULL guard serializes the keys a second time inside the `map_from_arrays` call, so a nondeterministic keys expression such as `IF(monotonically_increasing_id() % 2 = 0, array(1), NULL)` would advance independently in each copy and the result would drift from Spark ([#5781](https://github.com/apache/datafusion-comet/issues/5781)). `CometMapFromArrays` declines a nondeterministic keys expression as `Unsupported` and the projection falls back to Spark. The values are serialized once and evaluated on the rows Spark evaluates them on, so a nondeterministic values expression stays native. ## map_from_entries @@ -55,6 +60,8 @@ - Spark 3.5.8 (audited 2026-05-27): baseline. `MapFromEntries(child) extends UnaryExpression with NullIntolerant`; expects an array of structs and produces a map. Wired as `CometScalarFunction("map_from_entries")`. - Spark 4.0.1 (audited 2026-05-27): semantics unchanged; trait refactor. - Spark 4.1.1 (audited 2026-05-27): identical to 4.0.1. +- The same `ArrayBasedMapBuilder` semantics and policy handling as `map_from_arrays` above, without the length check. A NULL entries array, or one holding a NULL entry, gives a NULL map without inserting any entry, so its keys are not checked. +- Known limitation: on Spark 4.0+, `ArrayBasedMapBuilder` normalizes a floating-point key before comparing it (`keyNormalizer`, added in 4.0 with `spark.sql.legacy.disableMapKeyNormalization`), so `-0.0` and `+0.0` are one key and all `NaN`s are one key; the native builder compares the raw bits and keeps `-0.0` and `+0.0`, and `NaN`s with different bit patterns, apart. Unlike `map_from_arrays`, this expression always calls `build()`, so Spark stores the normalized key and returns `+0.0` for a `-0.0` key where Comet returns `-0.0`. Spark 3.4 and 3.5 do not normalize, so they already match. `spark.comet.exec.strictFloatingPoint=true` marks the expression `Incompatible` for a floating-point key type, and `CodegenDispatchFallback` then runs Spark's own code for it through the JVM codegen dispatcher (`map_builders_strict_fp.sql`). - Known limitation: input arrays where the struct's key or value type contains `BinaryType` are marked `Incompatible` and fall back unless `spark.comet.expression.MapFromEntries.allowIncompatible=true`. ## map_keys @@ -79,7 +86,7 @@ ## str_to_map - Spark 3.4.3 (audited 2026-05-27): identical to 3.5.8. -- Spark 3.5.8 (audited 2026-05-27): baseline. `StringToMap(text, pairDelim, keyValueDelim) extends TernaryExpression`; splits `text` on `pairDelim`, then each pair on `keyValueDelim` (default `","` and `":"`). Uses `ArrayBasedMapBuilder` for duplicate-key handling. Wired as `CometScalarFunction("str_to_map")`. +- Spark 3.5.8 (audited 2026-05-27): baseline. `StringToMap(text, pairDelim, keyValueDelim) extends TernaryExpression`; splits `text` on `pairDelim`, then each pair on `keyValueDelim` (default `","` and `":"`). Uses `ArrayBasedMapBuilder` for duplicate-key handling. Wired as `CometScalarFunction("str_to_map")`, which follows `spark.sql.mapKeyDedupPolicy` the way `map_from_arrays` does (see there). - Spark 4.0.1 (audited 2026-05-27): `inputTypes` widened to `StringTypeNonCSAICollation`; uses `CollationAwareUTF8String.splitSQL` with a `collationId`. Runtime unchanged for `UTF8_BINARY`. - Spark 4.1.1 (audited 2026-05-27): adds the `legacySplitTruncate` flag (driven by `spark.sql.legacy.truncateForEmptyRegexSplit`) to both `splitSQL` calls. The Comet native impl always behaves as if the flag were false, so `CometStrToMap` reads the config by string key and reports `Incompatible` when it is enabled; the `CodegenDispatchFallback` trait then routes the expression through the JVM codegen dispatcher rather than falling the whole projection back to Spark. Non-UTF8_BINARY collations on the input or the delimiters are handled the same way. diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index faf4d5dfda0..14459b9dbd1 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -59,8 +59,6 @@ use datafusion_spark::function::datetime::to_utc_timestamp::SparkToUtcTimestamp; use datafusion_spark::function::hash::crc32::SparkCrc32; use datafusion_spark::function::hash::sha1::SparkSha1; use datafusion_spark::function::hash::sha2::SparkSha2; -use datafusion_spark::function::map::map_from_entries::MapFromEntries; -use datafusion_spark::function::map::str_to_map::SparkStrToMap; use datafusion_spark::function::math::expm1::SparkExpm1; use datafusion_spark::function::math::factorial::SparkFactorial; use datafusion_spark::function::math::hex::SparkHex; @@ -115,7 +113,7 @@ use crate::execution::memory_pools::logging_pool::LoggingMemoryPool; use crate::execution::spark_config::{ SparkConfig, COMET_DEBUG_ENABLED, COMET_DEBUG_MEMORY, COMET_EXPLAIN_NATIVE_ENABLED, COMET_MAX_TEMP_DIRECTORY_SIZE, COMET_PARQUET_ROW_FILTER_PUSHDOWN_ENABLED, - COMET_TRACING_ENABLED, SPARK_EXECUTOR_CORES, + COMET_TRACING_ENABLED, SPARK_EXECUTOR_CORES, SPARK_MAP_KEY_DEDUP_POLICY, }; use crate::parquet::encryption_support::{CometEncryptionFactory, ENCRYPTION_FACTORY_ID}; use crate::parquet::parquet_support::CometObjectStoreRegistry; @@ -761,6 +759,14 @@ fn prepare_datafusion_session_context( session_config.set_str("datafusion.execution.parquet.reorder_filters", "true"); } + // The duplicate-key policy of the native map constructors, which DataFusion spells + // `datafusion.spark.map_key_dedup_policy` with the same `EXCEPTION` / `LAST_WIN` values (see + // the map_funcs expression audit). Set before the `spark.comet.datafusion.*` testing escape + // hatch pass-through below, so an explicit override of the DataFusion key still wins. + if let Some(policy) = spark_config.get(SPARK_MAP_KEY_DEDUP_POLICY) { + session_config.options_mut().spark.map_key_dedup_policy = policy.parse()?; + } + // Pass through DataFusion configs from Spark. // e.g: spark-shell --conf spark.comet.datafusion.sql_parser.parse_float_as_decimal=true // becomes datafusion.sql_parser.parse_float_as_decimal=true @@ -802,7 +808,6 @@ fn register_datafusion_spark_function(session_ctx: &SessionContext) { session_ctx.register_udf(ScalarUDF::new_from_impl(SparkBitwiseNot::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkHex::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkWidthBucket::default())); - session_ctx.register_udf(ScalarUDF::new_from_impl(MapFromEntries::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkCrc32::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkLuhnCheck::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkSpace::default())); @@ -810,7 +815,6 @@ fn register_datafusion_spark_function(session_ctx: &SessionContext) { session_ctx.register_udf(ScalarUDF::new_from_impl(SparkArrayContains::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkArrayRepeat::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkBin::default())); - session_ctx.register_udf(ScalarUDF::new_from_impl(SparkStrToMap::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkUrlDecode::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkUrlEncode::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkTryUrlDecode::default())); diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 8f030da455b..08f6a25b015 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -3792,7 +3792,10 @@ impl PhysicalPlanner { fun_expr, args.to_vec(), Arc::new(Field::new(fun_name, data_type.clone(), true)), - Arc::new(ConfigOptions::default()), + // The session's options rather than DataFusion's defaults, so a kernel that reads one + // (the map constructors read `datafusion.spark.map_key_dedup_policy`) sees what + // `prepare_datafusion_session_context` set. + Arc::clone(self.session_ctx.copied_config().options()), )); // DF53 changed some UDFs (e.g. md5) to return StringViewArray at execution diff --git a/native/core/src/execution/spark_config.rs b/native/core/src/execution/spark_config.rs index 4c2811cb5de..573e1e9544f 100644 --- a/native/core/src/execution/spark_config.rs +++ b/native/core/src/execution/spark_config.rs @@ -25,6 +25,8 @@ pub(crate) const COMET_DEBUG_MEMORY: &str = "spark.comet.debug.memory"; pub(crate) const COMET_PARQUET_ROW_FILTER_PUSHDOWN_ENABLED: &str = "spark.comet.parquet.rowFilterPushdown.enabled"; pub(crate) const SPARK_EXECUTOR_CORES: &str = "spark.executor.cores"; +/// Spark's duplicate map key policy, forwarded to `datafusion.spark.map_key_dedup_policy`. +pub(crate) const SPARK_MAP_KEY_DEDUP_POLICY: &str = "spark.sql.mapKeyDedupPolicy"; pub(crate) trait SparkConfig { fn get_bool(&self, name: &str) -> bool; diff --git a/native/spark-expr/src/comet_scalar_funcs.rs b/native/spark-expr/src/comet_scalar_funcs.rs index 8fe19f0aad5..bd281a7e283 100644 --- a/native/spark-expr/src/comet_scalar_funcs.rs +++ b/native/spark-expr/src/comet_scalar_funcs.rs @@ -32,8 +32,8 @@ use crate::{ EvalMode, SparkArrayPositionFunc, SparkArraySlice, SparkArraysOverlap, SparkContains, SparkDateDiff, SparkDateFromUnixDate, SparkDateTrunc, SparkDayOfWeek, SparkFlatten, SparkIcebergBucket, SparkIcebergTemporalTransform, SparkIcebergTruncate, SparkMakeDate, - SparkMakeInterval, SparkMakeTime, SparkMapExtract, SparkNextDay, SparkSecondsToTimestamp, - SparkSizeFunc, SparkWeekDay, + SparkMakeInterval, SparkMakeTime, SparkMapExtract, SparkMapFromArrays, SparkMapFromEntries, + SparkNextDay, SparkSecondsToTimestamp, SparkSizeFunc, SparkStrToMap, SparkWeekDay, }; use arrow::datatypes::DataType; use datafusion::common::{DataFusionError, Result as DataFusionResult}; @@ -341,9 +341,12 @@ fn all_scalar_functions() -> Vec> { // returns the value itself rather than a one-element list (#5795). It carries the same // `element_at` alias so both registry entries the override replaces point here. Arc::new(ScalarUDF::new_from_impl(SparkMapExtract::default())), + Arc::new(ScalarUDF::new_from_impl(SparkMapFromArrays::default())), + Arc::new(ScalarUDF::new_from_impl(SparkMapFromEntries::default())), Arc::new(ScalarUDF::new_from_impl(SparkNextDay::default())), Arc::new(ScalarUDF::new_from_impl(SparkSecondsToTimestamp::default())), Arc::new(ScalarUDF::new_from_impl(SparkSizeFunc::default())), + Arc::new(ScalarUDF::new_from_impl(SparkStrToMap::default())), Arc::new(ScalarUDF::new_from_impl(JsonArrayLength::default())), ] } diff --git a/native/spark-expr/src/lib.rs b/native/spark-expr/src/lib.rs index 0d7675a64e0..dd1e05f32a8 100644 --- a/native/spark-expr/src/lib.rs +++ b/native/spark-expr/src/lib.rs @@ -61,7 +61,9 @@ pub mod jvm_udf; mod conditional_funcs; mod conversion_funcs; mod map_funcs; -pub use map_funcs::{spark_map_sort, SparkMapExtract}; +pub use map_funcs::{ + spark_map_sort, SparkMapExtract, SparkMapFromArrays, SparkMapFromEntries, SparkStrToMap, +}; mod math_funcs; mod nondetermenistic_funcs; pub mod url_funcs; diff --git a/native/spark-expr/src/map_funcs/map_builders.rs b/native/spark-expr/src/map_funcs/map_builders.rs new file mode 100644 index 00000000000..3e1b464f1e6 --- /dev/null +++ b/native/spark-expr/src/map_funcs/map_builders.rs @@ -0,0 +1,873 @@ +// 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. + +//! Spark-compatible `map_from_arrays`, `map_from_entries` and `str_to_map`. +//! +//! The `datafusion-spark` kernels build the `MapArray` the way Spark's `ArrayBasedMapBuilder` +//! does: row by row, checking a row's key and value lengths and then its keys in order, under the +//! duplicate-key policy in `datafusion.spark.map_key_dedup_policy`. These wrappers add the one +//! check the kernels skip, a `NULL` key, and restate the kernels' errors as the `SparkError`s the +//! JVM side turns back into Spark's own: +//! +//! - a row whose key and value arrays differ in length raises `SparkError::MapKeyValueDiffSizes`, +//! which reaches the user as Spark's `_LEGACY_ERROR_TEMP_2128`; +//! - a `NULL` key raises `NULL_MAP_KEY` and, under `EXCEPTION`, a duplicate key raises +//! `DUPLICATED_MAP_KEY` naming the key. Spark inserts entries one at a time, so whichever comes +//! first in the row decides which of the two it reports. +//! +//! `str_to_map` builds its keys by splitting a string, so it needs only the duplicate-key +//! restatement. + +use crate::SparkError; +use arrow::array::{Array, ArrayRef, AsArray, ListArray}; +use arrow::buffer::{NullBuffer, OffsetBuffer}; +use arrow::datatypes::{DataType, FieldRef}; +use datafusion::common::config::MapKeyDedupPolicy; +use datafusion::common::{exec_err, DataFusionError, HashSet, Result, ScalarValue}; +use datafusion::logical_expr::{ + ColumnarValue, ReturnFieldArgs, ScalarFunctionArgs, ScalarUDFImpl, Signature, +}; +use datafusion_spark::function::map::map_from_arrays::MapFromArrays as DataFusionMapFromArrays; +use datafusion_spark::function::map::map_from_entries::MapFromEntries as DataFusionMapFromEntries; +use datafusion_spark::function::map::str_to_map::SparkStrToMap as DataFusionStrToMap; +use std::sync::Arc; + +/// Spark-compatible `map_from_arrays(keys, values)`. +#[derive(Debug, Default, PartialEq, Eq, Hash)] +pub struct SparkMapFromArrays { + inner: DataFusionMapFromArrays, +} + +impl ScalarUDFImpl for SparkMapFromArrays { + fn name(&self) -> &str { + self.inner.name() + } + + fn signature(&self) -> &Signature { + self.inner.signature() + } + + fn return_type(&self, arg_types: &[DataType]) -> Result { + self.inner.return_type(arg_types) + } + + fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result { + self.inner.return_field_from_args(args) + } + + fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result { + invoke_list_builder(&self.inner, args, |args, last_value_wins| match args { + [ColumnarValue::Array(keys), ColumnarValue::Array(values)] => { + validate_map_from_arrays(keys, values, last_value_wins) + } + other => exec_err!("map_from_arrays expects 2 arguments, got {}", other.len()), + }) + } +} + +/// Spark-compatible `map_from_entries(entries)`. +#[derive(Debug, Default, PartialEq, Eq, Hash)] +pub struct SparkMapFromEntries { + inner: DataFusionMapFromEntries, +} + +impl ScalarUDFImpl for SparkMapFromEntries { + fn name(&self) -> &str { + self.inner.name() + } + + fn signature(&self) -> &Signature { + self.inner.signature() + } + + fn return_type(&self, arg_types: &[DataType]) -> Result { + self.inner.return_type(arg_types) + } + + fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result { + self.inner.return_field_from_args(args) + } + + fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result { + invoke_list_builder(&self.inner, args, |args, last_value_wins| match args { + [ColumnarValue::Array(entries)] => validate_map_from_entries(entries, last_value_wins), + other => exec_err!("map_from_entries expects 1 argument, got {}", other.len()), + }) + } +} + +/// Spark-compatible `str_to_map(text[, pair_delim[, key_value_delim]])`. +#[derive(Debug, Default, PartialEq, Eq, Hash)] +pub struct SparkStrToMap { + inner: DataFusionStrToMap, +} + +impl ScalarUDFImpl for SparkStrToMap { + fn name(&self) -> &str { + self.inner.name() + } + + fn signature(&self) -> &Signature { + self.inner.signature() + } + + fn return_type(&self, arg_types: &[DataType]) -> Result { + self.inner.return_type(arg_types) + } + + fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result { + self.inner.return_field_from_args(args) + } + + fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result { + // Splitting a string cannot produce a NULL key, so only the duplicate-key error needs + // restating here. + self.inner + .invoke_with_args(args) + .map_err(|error| as_spark_error(error, DuplicateKeyFormat::Quoted)) + } +} + +/// Runs `map_from_arrays` or `map_from_entries`: `validate` sees the arguments as arrays whose +/// entries start at offset zero and rejects a `NULL` key, which the kernel would store without a +/// word, then the kernel builds the maps and its own errors are restated. +fn invoke_list_builder( + inner: &dyn ScalarUDFImpl, + mut args: ScalarFunctionArgs, + validate: impl FnOnce(&[ColumnarValue], bool) -> Result<()>, +) -> Result { + // The kernel evaluates an all-scalar call once and returns a scalar, which DataFusion then + // broadcasts to the batch; build that one row here too rather than `number_rows` copies. + let all_scalar = args + .args + .iter() + .all(|arg| matches!(arg, ColumnarValue::Scalar(_))); + if all_scalar { + args.number_rows = 1; + } + expand_scalars(&mut args)?; + rebase_sliced_lists(&mut args)?; + let last_value_wins = + args.config_options.spark.map_key_dedup_policy == MapKeyDedupPolicy::LastWin; + validate(&args.args, last_value_wins)?; + let result = inner + .invoke_with_args(args) + .map_err(|error| as_spark_error(error, DuplicateKeyFormat::Bare))?; + match (all_scalar, result) { + (true, ColumnarValue::Array(array)) => Ok(ColumnarValue::Scalar( + ScalarValue::try_from_array(&array, 0)?, + )), + (_, result) => Ok(result), + } +} + +/// Materializes scalar arguments at `number_rows`, so the validation indexes rows the same way +/// the kernel does. +fn expand_scalars(args: &mut ScalarFunctionArgs) -> Result<()> { + let number_rows = args.number_rows; + for arg in args.args.iter_mut() { + if let ColumnarValue::Scalar(scalar) = arg { + *arg = ColumnarValue::Array(scalar.to_array_of_size(number_rows)?); + } + } + Ok(()) +} + +/// Rebases a list argument whose entries do not start at offset zero, which the kernels mishandle +/// (apache/datafusion#25419): they read each row's entries at its own offset but apply the mask +/// that selects the surviving keys from the start of the list's values, so a list sliced past its +/// first row pairs values with keys from earlier rows. A `LIMIT ... OFFSET` above a projection is +/// enough to produce one. A slice that only drops trailing rows needs nothing, since Arrow's +/// `filter` accepts a predicate shorter than the array it filters. Shifting the offsets and slicing +/// the values leaves the data where it is. +/// +/// Delete once Comet moves to a DataFusion release carrying apache/datafusion#25431. +fn rebase_sliced_lists(args: &mut ScalarFunctionArgs) -> Result<()> { + for arg in args.args.iter_mut() { + let ColumnarValue::Array(array) = arg else { + continue; + }; + let (DataType::List(field), Some(list)) = (array.data_type(), array.as_list_opt::()) + else { + continue; + }; + let offsets = list.value_offsets(); + let (first, last) = (offsets[0], offsets[offsets.len() - 1]); + if first == 0 { + continue; + } + let rebased = ListArray::try_new( + Arc::clone(field), + OffsetBuffer::new(offsets.iter().map(|offset| offset - first).collect()), + list.values().slice(first as usize, (last - first) as usize), + list.nulls().cloned(), + )?; + *arg = ColumnarValue::Array(Arc::new(rebased)); + } + Ok(()) +} + +/// Rejects a `NULL` key in a row `map_from_arrays` builds a map from, raising what Spark raises +/// for that row. The kernel already checks each row's lengths and duplicate keys in Spark's order, +/// so a batch where no such row holds a `NULL` key is left to it. A row whose keys or values array +/// is NULL gives a NULL map before Spark builds anything, so its keys do not count. +fn validate_map_from_arrays( + keys: &ArrayRef, + values: &ArrayRef, + last_value_wins: bool, +) -> Result<()> { + // A `NULL`-typed argument makes every row a NULL map, which never reaches the builder. + if matches!(keys.data_type(), DataType::Null) || matches!(values.data_type(), DataType::Null) { + return Ok(()); + } + let (keys, values) = (as_list(keys)?, as_list(values)?); + let Some(key_nulls) = key_nulls(keys.values()) else { + return Ok(()); + }; + let (key_offsets, value_offsets) = (keys.value_offsets(), values.value_offsets()); + let row_is_built = |row: usize| keys.is_valid(row) && values.is_valid(row); + let Some(failing_row) = first_row_with_null_key(&key_nulls, key_offsets, row_is_built) else { + return Ok(()); + }; + // `failing_row` fails one way or another, so walk the rows up to it in Spark's order: a row's + // lengths are checked before any of its keys. + let mut seen = HashSet::new(); + for row in (0..=failing_row).filter(|&row| row_is_built(row)) { + let (start, end) = (key_offsets[row] as usize, key_offsets[row + 1] as usize); + if end - start != (value_offsets[row + 1] - value_offsets[row]) as usize { + return Err(SparkError::MapKeyValueDiffSizes.into()); + } + check_keys_in_order( + keys.values(), + start, + end, + &key_nulls, + last_value_wins, + &mut seen, + )?; + } + Ok(()) +} + +/// Rejects a `NULL` key in a row `map_from_entries` builds a map from, raising what Spark raises +/// for that row. A row is not built when its entries array is NULL or holds a NULL `struct` +/// element: Spark returns a NULL map for both without inserting any entry. +fn validate_map_from_entries(entries: &ArrayRef, last_value_wins: bool) -> Result<()> { + if matches!(entries.data_type(), DataType::Null) { + return Ok(()); + } + let entries = as_list(entries)?; + let Some(structs) = entries.values().as_struct_opt() else { + return exec_err!( + "map_from_entries: expected array>, got {:?}", + entries.values().data_type() + ); + }; + let keys = structs.column(0); + let Some(key_nulls) = key_nulls(keys) else { + return Ok(()); + }; + // Only a `NULL` key inside a non-NULL entry can fail, and one AND and a popcount rule that + // out, which covers the common batch whose only `NULL` keys sit under `NULL` entries. + let entry_nulls = structs.nulls(); + let null_entries = entry_nulls.map_or(0, NullBuffer::null_count); + let null_entries_or_keys = + NullBuffer::union(Some(&key_nulls), entry_nulls).map_or(0, |nulls| nulls.null_count()); + if null_entries_or_keys == null_entries { + return Ok(()); + } + let offsets = entries.value_offsets(); + let row_is_built = |row: usize| { + let (start, end) = (offsets[row] as usize, offsets[row + 1] as usize); + entries.is_valid(row) + && entry_nulls.is_none_or(|nulls| nulls.slice(start, end - start).null_count() == 0) + }; + let Some(failing_row) = first_row_with_null_key(&key_nulls, offsets, row_is_built) else { + return Ok(()); + }; + let mut seen = HashSet::new(); + for row in (0..=failing_row).filter(|&row| row_is_built(row)) { + let (start, end) = (offsets[row] as usize, offsets[row + 1] as usize); + check_keys_in_order(keys, start, end, &key_nulls, last_value_wins, &mut seen)?; + } + Ok(()) +} + +/// The first row that `row_is_built` accepts and whose keys include a `NULL`. +fn first_row_with_null_key( + key_nulls: &NullBuffer, + offsets: &[i32], + row_is_built: impl Fn(usize) -> bool, +) -> Option { + (0..offsets.len() - 1).find(|&row| { + let (start, end) = (offsets[row] as usize, offsets[row + 1] as usize); + key_nulls.slice(start, end - start).null_count() > 0 && row_is_built(row) + }) +} + +/// Walks one row's keys in the order Spark's `ArrayBasedMapBuilder` inserts them, so whichever of +/// a `NULL` key and a duplicate key comes first is the one reported, as Spark reports it. +fn check_keys_in_order( + keys: &ArrayRef, + start: usize, + end: usize, + key_nulls: &NullBuffer, + last_value_wins: bool, + seen: &mut HashSet, +) -> Result<()> { + seen.clear(); + for index in start..end { + if key_nulls.is_null(index) { + return Err(SparkError::NullMapKey.into()); + } + // `LAST_WIN` overwrites a duplicate rather than raising, so only the `NULL` check is + // left to do in that mode. + if last_value_wins { + continue; + } + let key = ScalarValue::try_from_array(keys, index)?.compacted(); + if let Some(duplicate) = seen.replace(key) { + return Err(SparkError::DuplicatedMapKey { + key: duplicate.to_string(), + } + .into()); + } + } + Ok(()) +} + +/// Comet hands every Spark `ArrayType` to native code as a `List`. +fn as_list(array: &ArrayRef) -> Result<&ListArray> { + match array.as_list_opt::() { + Some(list) => Ok(list), + None => exec_err!("expected a list argument, got {:?}", array.data_type()), + } +} + +/// The nulls of a map key array, or `None` when no key is `NULL`. Logical nulls, so a `NullArray` +/// (which has no null buffer) and a dictionary whose values hold a `NULL` count too. +fn key_nulls(keys: &ArrayRef) -> Option { + keys.logical_nulls().filter(|nulls| nulls.null_count() > 0) +} + +/// How the upstream kernel renders the offending key in its duplicate-key message. +#[derive(Clone, Copy)] +enum DuplicateKeyFormat { + /// The map builders write the key unquoted, as Spark does. + Bare, + /// `str_to_map` single-quotes it. + Quoted, +} + +/// The message the kernels' shared map builder raises for a row whose key and value arrays differ +/// in length. +const LENGTH_MISMATCH_MESSAGE: &str = + "keys and values lists in the same row must have equal lengths"; + +/// Restates the upstream length-mismatch and duplicate-key errors as `SparkError`s, so the JVM +/// side raises Spark's own errors (naming the same key, for a duplicate). Any other error is +/// passed through. The `*_reports_*` tests pin the wordings parsed here against the kernels +/// themselves, so an upstream rewording fails there rather than silently downgrading the error to +/// a generic execution failure. +fn as_spark_error(error: DataFusionError, key_format: DuplicateKeyFormat) -> DataFusionError { + let message = error.to_string(); + if message.contains(LENGTH_MISMATCH_MESSAGE) { + return SparkError::MapKeyValueDiffSizes.into(); + } + match duplicate_map_key(&message, key_format) { + Some(key) => SparkError::DuplicatedMapKey { key }.into(), + None => error, + } +} + +/// The key named by `datafusion-spark`'s duplicate-key message. +fn duplicate_map_key(message: &str, key_format: DuplicateKeyFormat) -> Option { + let (open, close) = match key_format { + DuplicateKeyFormat::Bare => ("[DUPLICATED_MAP_KEY] Duplicate map key ", " was found"), + DuplicateKeyFormat::Quoted => ("[DUPLICATED_MAP_KEY] Duplicate map key '", "' was found"), + }; + let (_, tail) = message.split_once(open)?; + let (key, _) = tail.rsplit_once(close)?; + Some(key.to_string()) +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::{Int32Array, MapArray, StringArray, StructArray}; + use arrow::datatypes::{Field, Fields, Int32Type}; + use datafusion::common::config::ConfigOptions; + + /// `[[1, 2], [3]]`-shaped keys, with `nulls` marking whole rows NULL. + fn int_list(values: Int32Array, offsets: &[i32], nulls: Option) -> ArrayRef { + let field = Arc::new(Field::new("item", DataType::Int32, true)); + Arc::new(ListArray::new( + field, + OffsetBuffer::new(offsets.to_vec().into()), + Arc::new(values), + nulls, + )) + } + + fn string_list(values: StringArray, offsets: &[i32], nulls: Option) -> ArrayRef { + let field = Arc::new(Field::new("item", DataType::Utf8, true)); + Arc::new(ListArray::new( + field, + OffsetBuffer::new(offsets.to_vec().into()), + Arc::new(values), + nulls, + )) + } + + /// `array>`, with `element_nulls` marking NULL entries. + fn entry_list( + keys: Int32Array, + values: StringArray, + offsets: &[i32], + element_nulls: Option, + ) -> ArrayRef { + let fields = Fields::from(vec![ + Field::new("key", DataType::Int32, true), + Field::new("value", DataType::Utf8, true), + ]); + let structs = StructArray::new( + fields.clone(), + vec![Arc::new(keys), Arc::new(values)], + element_nulls, + ); + let field = Arc::new(Field::new("item", DataType::Struct(fields), true)); + Arc::new(ListArray::new( + field, + OffsetBuffer::new(offsets.to_vec().into()), + Arc::new(structs), + None, + )) + } + + fn invoke_values( + udf: &dyn ScalarUDFImpl, + args: Vec, + number_rows: usize, + policy: MapKeyDedupPolicy, + ) -> Result { + let arg_fields: Vec = args + .iter() + .enumerate() + .map(|(i, arg)| Arc::new(Field::new(format!("arg{i}"), arg.data_type(), true))) + .collect(); + let scalar_arguments: Vec> = vec![None; args.len()]; + let return_field = udf.return_field_from_args(ReturnFieldArgs { + arg_fields: &arg_fields, + scalar_arguments: &scalar_arguments, + })?; + let mut config = ConfigOptions::default(); + config.spark.map_key_dedup_policy = policy; + udf.invoke_with_args(ScalarFunctionArgs { + args, + arg_fields, + number_rows, + return_field, + config_options: Arc::new(config), + }) + } + + fn invoke( + udf: &dyn ScalarUDFImpl, + args: Vec, + policy: MapKeyDedupPolicy, + ) -> Result { + let number_rows = args.first().map(|arg| arg.len()).unwrap_or(0); + let args = args.into_iter().map(ColumnarValue::Array).collect(); + invoke_values(udf, args, number_rows, policy) + } + + fn invoke_err( + udf: &dyn ScalarUDFImpl, + args: Vec, + policy: MapKeyDedupPolicy, + ) -> String { + invoke(udf, args, policy).unwrap_err().to_string() + } + + fn map_result(value: ColumnarValue) -> MapArray { + match value { + ColumnarValue::Array(array) => array.as_map().clone(), + ColumnarValue::Scalar(scalar) => { + scalar.to_array().expect("scalar to array").as_map().clone() + } + } + } + + #[test] + fn map_from_arrays_ignores_null_key_in_a_null_row() { + // Row 0's keys array is NULL, so Spark returns a NULL map without inspecting its keys. + let keys = int_list( + Int32Array::from(vec![None, Some(1)]), + &[0, 1, 2], + Some(NullBuffer::from(vec![false, true])), + ); + let values = string_list( + StringArray::from(vec![Some("a"), Some("b")]), + &[0, 1, 2], + None, + ); + let result = map_result( + invoke( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ) + .unwrap(), + ); + assert!(result.is_null(0)); + assert_eq!(result.value_offsets(), &[0, 0, 1]); + } + + /// The serde guards only the keys array, so a row whose values array is NULL reaches the + /// wrapper. Spark returns a NULL map for it without looking at its keys. + #[test] + fn map_from_arrays_ignores_null_key_in_a_row_with_null_values() { + let keys = int_list(Int32Array::from(vec![None, Some(1)]), &[0, 1, 2], None); + let values = string_list( + StringArray::from(vec![Some("a"), Some("b")]), + &[0, 1, 2], + Some(NullBuffer::from(vec![false, true])), + ); + let result = map_result( + invoke( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ) + .unwrap(), + ); + assert!(result.is_null(0)); + assert_eq!(result.value_offsets(), &[0, 0, 1]); + } + + /// Pins the upstream duplicate-key message `duplicate_map_key` parses. The list builders write + /// a string key unquoted, as Spark does; `str_to_map` quotes its key, which is why it goes + /// through a different `DuplicateKeyFormat`. + #[test] + fn map_from_arrays_reports_the_duplicate_key_unquoted() { + let keys = string_list(StringArray::from(vec![Some("a"), Some("a")]), &[0, 2], None); + let values = int_list(Int32Array::from(vec![1, 2]), &[0, 2], None); + let err = invoke_err( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ); + assert!( + err.contains("[DUPLICATED_MAP_KEY] Cannot create map with duplicate keys: a."), + "{err}" + ); + } + + /// Pins the upstream length-mismatch message `as_spark_error` recognizes. + #[test] + fn map_from_arrays_reports_a_length_mismatch() { + let keys = int_list(Int32Array::from(vec![1, 2]), &[0, 2], None); + let values = string_list(StringArray::from(vec![Some("a")]), &[0, 1], None); + let err = invoke_err( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ); + assert!(err.contains("[MAP_KEY_VALUE_DIFF_SIZES]"), "{err}"); + } + + /// Spark checks a row's lengths only when it reaches that row, so a duplicate key in an + /// earlier row is reported ahead of a later row's length mismatch. + #[test] + fn map_from_arrays_reports_a_duplicate_before_a_later_length_mismatch() { + let keys = int_list(Int32Array::from(vec![1, 1, 2]), &[0, 2, 3], None); + let values = string_list( + StringArray::from(vec![Some("a"), Some("b"), Some("c"), Some("d")]), + &[0, 2, 4], + None, + ); + let err = invoke_err( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ); + assert!(err.contains("[DUPLICATED_MAP_KEY]"), "{err}"); + } + + /// Spark checks a row's lengths before inserting any of its keys, so a length mismatch is + /// reported ahead of a `NULL` key in the same row. + #[test] + fn map_from_arrays_reports_a_length_mismatch_before_a_null_key() { + let keys = int_list(Int32Array::from(vec![None, Some(1)]), &[0, 2], None); + let values = string_list(StringArray::from(vec![Some("a")]), &[0, 1], None); + let err = invoke_err( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ); + assert!(err.contains("[MAP_KEY_VALUE_DIFF_SIZES]"), "{err}"); + } + + /// Spark inserts entries one at a time, so a duplicate at an earlier index is reported even + /// though a `NULL` key follows it. + #[test] + fn map_from_arrays_reports_a_duplicate_before_a_later_null_key() { + let keys = int_list( + Int32Array::from(vec![Some(1), Some(1), None]), + &[0, 3], + None, + ); + let values = string_list( + StringArray::from(vec![Some("a"), Some("b"), Some("c")]), + &[0, 3], + None, + ); + let err = invoke_err( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ); + assert!(err.contains("[DUPLICATED_MAP_KEY]"), "{err}"); + } + + /// The mirror case: the `NULL` comes first, so it is the one reported. + #[test] + fn map_from_arrays_reports_a_null_key_before_a_later_duplicate() { + let keys = int_list( + Int32Array::from(vec![None, Some(1), Some(1)]), + &[0, 3], + None, + ); + let values = string_list( + StringArray::from(vec![Some("a"), Some("b"), Some("c")]), + &[0, 3], + None, + ); + let err = invoke_err( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ); + assert!(err.contains("[NULL_MAP_KEY]"), "{err}"); + } + + /// A duplicate in an earlier row wins over a `NULL` key in a later one. + #[test] + fn map_from_arrays_reports_the_first_offending_row() { + let keys = int_list( + Int32Array::from(vec![Some(1), Some(1), None, Some(2)]), + &[0, 2, 4], + None, + ); + let values = string_list( + StringArray::from(vec![Some("a"), Some("b"), Some("c"), Some("d")]), + &[0, 2, 4], + None, + ); + let err = invoke_err( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::Exception, + ); + assert!(err.contains("[DUPLICATED_MAP_KEY]"), "{err}"); + } + + /// Under `LAST_WIN` a duplicate is not an error, so a `NULL` key is still reported. + #[test] + fn last_win_still_rejects_a_null_key_after_a_duplicate() { + let keys = int_list( + Int32Array::from(vec![Some(1), Some(1), None]), + &[0, 3], + None, + ); + let values = string_list( + StringArray::from(vec![Some("a"), Some("b"), Some("c")]), + &[0, 3], + None, + ); + let err = invoke_err( + &SparkMapFromArrays::default(), + vec![keys, values], + MapKeyDedupPolicy::LastWin, + ); + assert!(err.contains("[NULL_MAP_KEY]"), "{err}"); + } + + /// An all-scalar call builds its one row once and returns a scalar, as the kernel alone does, + /// rather than `number_rows` copies of it. + #[test] + fn map_from_arrays_evaluates_an_all_scalar_call_once() { + let keys = int_list(Int32Array::from(vec![1, 2]), &[0, 2], None); + let values = string_list(StringArray::from(vec![Some("a"), Some("b")]), &[0, 2], None); + let scalar = |array: ArrayRef| { + ColumnarValue::Scalar(ScalarValue::try_from_array(&array, 0).expect("scalar")) + }; + let result = invoke_values( + &SparkMapFromArrays::default(), + vec![scalar(keys), scalar(values)], + 3, + MapKeyDedupPolicy::Exception, + ) + .unwrap(); + let ColumnarValue::Scalar(ScalarValue::Map(map)) = result else { + panic!("expected a scalar map, got {result:?}"); + }; + assert_eq!(map.value_offsets(), &[0, 2]); + } + + /// A `LIMIT ... OFFSET` above a projection hands the kernel a sliced list. A slice past the + /// first row would read a preceding row's key without `rebase_sliced_lists`; a head slice + /// needs no rebase. + #[test] + fn map_from_arrays_reads_the_right_row_of_a_sliced_list() { + let keys = int_list(Int32Array::from(vec![10, 20]), &[0, 1, 2], None); + let values = string_list( + StringArray::from(vec![Some("100"), Some("200")]), + &[0, 1, 2], + None, + ); + for (row, key, value) in [(0, 10, "100"), (1, 20, "200")] { + let result = map_result( + invoke( + &SparkMapFromArrays::default(), + vec![keys.slice(row, 1), values.slice(row, 1)], + MapKeyDedupPolicy::Exception, + ) + .unwrap(), + ); + assert_eq!(result.len(), 1); + assert_eq!( + result.keys().as_primitive::().values().as_ref(), + &[key] + ); + assert_eq!(result.values().as_string::().value(0), value); + } + } + + #[test] + fn map_from_entries_reads_the_right_row_of_a_sliced_list() { + let entries = entry_list( + Int32Array::from(vec![10, 20]), + StringArray::from(vec![Some("100"), Some("200")]), + &[0, 1, 2], + None, + ); + let result = map_result( + invoke( + &SparkMapFromEntries::default(), + vec![entries.slice(1, 1)], + MapKeyDedupPolicy::Exception, + ) + .unwrap(), + ); + assert_eq!(result.len(), 1); + assert_eq!( + result.keys().as_primitive::().values().as_ref(), + &[20] + ); + assert_eq!(result.values().as_string::().value(0), "200"); + } + + #[test] + fn map_from_entries_ignores_a_null_entry() { + // A NULL struct element makes the whole row a NULL map, so its NULL key is never a key. + let entries = entry_list( + Int32Array::from(vec![None, Some(2)]), + StringArray::from(vec![None, Some("b")]), + &[0, 1, 2], + Some(NullBuffer::from(vec![false, true])), + ); + let result = map_result( + invoke( + &SparkMapFromEntries::default(), + vec![entries], + MapKeyDedupPolicy::Exception, + ) + .unwrap(), + ); + assert!(result.is_null(0)); + assert_eq!(result.value_offsets(), &[0, 0, 1]); + } + + /// A NULL entry makes the whole row a NULL map, so a `NULL` key in another entry of the same + /// row is never inserted either. + #[test] + fn map_from_entries_ignores_a_null_key_beside_a_null_entry() { + let entries = entry_list( + Int32Array::from(vec![None, None, Some(3)]), + StringArray::from(vec![None, Some("b"), Some("c")]), + &[0, 2, 3], + Some(NullBuffer::from(vec![false, true, true])), + ); + let result = map_result( + invoke( + &SparkMapFromEntries::default(), + vec![entries], + MapKeyDedupPolicy::Exception, + ) + .unwrap(), + ); + assert!(result.is_null(0)); + assert_eq!(result.value_offsets(), &[0, 0, 1]); + } + + #[test] + fn map_from_entries_reports_a_duplicate_before_a_later_null_key() { + let entries = entry_list( + Int32Array::from(vec![Some(1), Some(1), None]), + StringArray::from(vec![Some("a"), Some("b"), Some("c")]), + &[0, 3], + None, + ); + let err = invoke_err( + &SparkMapFromEntries::default(), + vec![entries], + MapKeyDedupPolicy::Exception, + ); + assert!(err.contains("[DUPLICATED_MAP_KEY]"), "{err}"); + } + + /// Pins the quoted duplicate-key message `str_to_map` raises. + #[test] + fn str_to_map_reports_the_duplicate_key() { + let text: ArrayRef = Arc::new(StringArray::from(vec![Some("a:1,b:2,a:3")])); + let err = invoke_err( + &SparkStrToMap::default(), + vec![text], + MapKeyDedupPolicy::Exception, + ); + assert!( + err.contains("[DUPLICATED_MAP_KEY] Cannot create map with duplicate keys: a."), + "{err}" + ); + } + + #[test] + fn duplicate_map_key_ignores_unrelated_errors() { + assert_eq!( + duplicate_map_key("Execution error: something else", DuplicateKeyFormat::Bare), + None + ); + assert_eq!( + duplicate_map_key( + "Execution error: something else", + DuplicateKeyFormat::Quoted + ), + None + ); + } +} diff --git a/native/spark-expr/src/map_funcs/mod.rs b/native/spark-expr/src/map_funcs/mod.rs index 644466d0320..1db58917350 100644 --- a/native/spark-expr/src/map_funcs/mod.rs +++ b/native/spark-expr/src/map_funcs/mod.rs @@ -15,7 +15,9 @@ // specific language governing permissions and limitations // under the License. +mod map_builders; mod map_extract; mod map_sort; +pub use map_builders::{SparkMapFromArrays, SparkMapFromEntries, SparkStrToMap}; pub use map_extract::SparkMapExtract; pub use map_sort::spark_map_sort; diff --git a/spark/src/main/scala/org/apache/comet/CometExecIterator.scala b/spark/src/main/scala/org/apache/comet/CometExecIterator.scala index 4da95af18d0..8d888dc09bc 100644 --- a/spark/src/main/scala/org/apache/comet/CometExecIterator.scala +++ b/spark/src/main/scala/org/apache/comet/CometExecIterator.scala @@ -388,6 +388,15 @@ object CometExecIterator extends Logging { CometConf.COMET_PARQUET_ROW_FILTER_PUSHDOWN_ENABLED.key, CometConf.COMET_PARQUET_ROW_FILTER_PUSHDOWN_ENABLED.get(SQLConf.get).toString) + // The native map constructors (map_from_arrays, map_from_entries, str_to_map) resolve + // duplicate keys with this policy, which the native side reads as + // `datafusion.spark.map_key_dedup_policy`. Read here, when the native plan for a task is + // built, which is where Spark's `ArrayBasedMapBuilder` reads it for a projection outside + // whole-stage codegen. See the note on `map_from_arrays` in the map_funcs expression audit. + builder.putEntries( + SQLConf.MAP_KEY_DEDUP_POLICY.key, + SQLConf.get.getConf(SQLConf.MAP_KEY_DEDUP_POLICY).toString) + builder.build().toByteArray } diff --git a/spark/src/main/scala/org/apache/comet/serde/maps.scala b/spark/src/main/scala/org/apache/comet/serde/maps.scala index 51fa428b543..70b3ac17946 100644 --- a/spark/src/main/scala/org/apache/comet/serde/maps.scala +++ b/spark/src/main/scala/org/apache/comet/serde/maps.scala @@ -23,8 +23,9 @@ import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ +import org.apache.comet.CometConf.COMET_EXEC_STRICT_FLOATING_POINT import org.apache.comet.DataTypeSupport.isComplexType -import org.apache.comet.serde.QueryPlanSerde.{createBinaryExpr, exprToProtoInternal, hasNonDefaultStringCollation, scalarFunctionExprToProto} +import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, hasNonDefaultStringCollation, scalarFunctionExprToProto} import org.apache.comet.shims.CometTypeShim /** @@ -132,80 +133,101 @@ object CometMapExtract extends CometExpressionSerde[GetMapValue] { } } -private object MapKeyDedupPolicySupport { - val incompatibleReason: String = - s"`${SQLConf.MAP_KEY_DEDUP_POLICY.key}` is set to " + - s"`${SQLConf.MapKeyDedupPolicy.LAST_WIN}`; Comet's native map construction " + - "does not implement LAST_WIN dedup semantics." - - val nullKeyReason: String = - "Spark rejects a `NULL` element inside the keys array with a `RuntimeException`" + - " (`Cannot use null as map key`); Comet's native `map_from_arrays` / `map_from_entries`" + - " does not detect a per-element `NULL` key and produces a map with a `NULL` key instead" + - " ([#4680](https://github.com/apache/datafusion-comet/issues/4680))." - - def isLastWin: Boolean = - SQLConf.get - .getConf(SQLConf.MAP_KEY_DEDUP_POLICY) - .toString - .equalsIgnoreCase(SQLConf.MapKeyDedupPolicy.LAST_WIN.toString) +/** + * Shared gate for the native map constructors (`map_from_arrays`, `map_from_entries`), which + * reproduce Spark's `ArrayBasedMapBuilder`, including `spark.sql.mapKeyDedupPolicy`. The + * map_funcs expression audit covers how, and where they still differ from Spark. + */ +private object MapBuilderSupport { + + /** Floating-point keys on Spark 4.0 and later: see the map_funcs expression audit. */ + val floatingPointKeyNote: String = + "On Spark 4.0 and later, `ArrayBasedMapBuilder` normalizes a floating-point map key before " + + "comparing it, so `-0.0` counts as the same key as `+0.0` and all `NaN`s count as one " + + "key. Comet's native map construction compares the raw Arrow values, so a map built from " + + "both `-0.0` and `+0.0` keeps two entries where Spark reports a duplicate key. " + + "`map_from_entries` also stores the normalized key, so Spark returns `+0.0` for a `-0.0` " + + "key where Comet returns `-0.0`; `map_from_arrays` keeps the original keys in both " + + "engines when nothing repeated. Spark 3.4 and 3.5 do not normalize at all, so they match " + + s"Comet already. Set `${COMET_EXEC_STRICT_FLOATING_POINT.key}=true` to keep a " + + "floating-point map key off the native path." + + val strictFloatingPointKeyReason: String = + s"When `${COMET_EXEC_STRICT_FLOATING_POINT.key}=true`, map construction on a floating-point " + + "key is not 100% compatible with Spark" + + /** + * `ArrayBasedMapBuilder` keys its dedup map on `TypeUtils.getInterpretedOrdering` once the key + * type contains a string, so under `UTF8_LCASE` the keys `'a'` and `'A'` are one key. The + * native builders compare the raw Arrow bytes and would keep both, missing the duplicate that + * Spark reports (or, under `LAST_WIN`, the overwrite Spark performs). `MapKeySupport` declines + * a collated key for `map_extract` for the same reason. + */ + val collationKeyReason: String = + "Comet's native map construction compares string keys as `UTF8_BINARY`, so it cannot honour " + + "a non-default collation when it looks for a duplicate key." + + /** The support level for a map constructor whose result has key type `keyType`. */ + def keySupport(keyType: DataType): SupportLevel = + if (hasNonDefaultStringCollation(keyType)) { + Incompatible(Some(collationKeyReason)) + } else { + SupportLevel + .strictFloatingPointReason(keyType, "Map construction on a floating-point key") + .map(reason => Incompatible(Some(reason))) + .getOrElse(Compatible(None)) + } } object CometMapFromArrays extends CometExpressionSerde[MapFromArrays] { + private val nondeterministicKeysReason: String = + "a nondeterministic operand as the keys array: the native NULL guard serializes the keys " + + "twice, and the two copies of a stateful expression drift apart" + override def getIncompatibleReasons(): Seq[String] = - Seq(MapKeyDedupPolicySupport.incompatibleReason) + Seq(MapBuilderSupport.collationKeyReason, MapBuilderSupport.strictFloatingPointKeyReason) + + override def getUnsupportedReasons(): Seq[String] = Seq(nondeterministicKeysReason) override def getCompatibleNotes(): Seq[String] = - Seq(MapKeyDedupPolicySupport.nullKeyReason) + Seq(MapBuilderSupport.floatingPointKeyNote) - override def getSupportLevel(expr: MapFromArrays): SupportLevel = { - if (MapKeyDedupPolicySupport.isLastWin) { - Incompatible(Some(MapKeyDedupPolicySupport.incompatibleReason)) + override def getSupportLevel(expr: MapFromArrays): SupportLevel = + if (!expr.left.deterministic) { + Unsupported(Some(nondeterministicKeysReason)) } else { - Compatible(None) + MapBuilderSupport.keySupport(expr.dataType.keyType) } - } + /** + * `CASE WHEN keys IS NOT NULL THEN map_from_arrays(keys, values) END`. The native function + * already returns a NULL map for a NULL input array; the guard is about evaluation order, since + * Spark never evaluates `values` for a row whose `keys` is NULL (see the map_funcs expression + * audit). It serializes `keys` a second time, which is why `getSupportLevel` declines a + * nondeterministic `keys`: https://github.com/apache/datafusion-comet/issues/5781. + */ override def convert( expr: MapFromArrays, inputs: Seq[Attribute], binding: Boolean): Option[ExprOuterClass.Expr] = { val keysExpr = exprToProtoInternal(expr.left, inputs, binding) val valuesExpr = exprToProtoInternal(expr.right, inputs, binding) - val keyType = expr.left.dataType.asInstanceOf[ArrayType].elementType - val valueType = expr.right.dataType.asInstanceOf[ArrayType].elementType - val returnType = MapType(keyType = keyType, valueType = valueType) for { - andBinaryExprProto <- createAndBinaryExpr(expr, inputs, binding) - mapFromArraysExprProto <- scalarFunctionExprToProto("map", keysExpr, valuesExpr) - nullLiteralExprProto <- exprToProtoInternal(Literal(null, returnType), inputs, binding) + keysNotNullExprProto <- exprToProtoInternal(IsNotNull(expr.left), inputs, binding) + mapFromArraysExprProto <- scalarFunctionExprToProto("map_from_arrays", keysExpr, valuesExpr) } yield { - val caseWhenExprProto = ExprOuterClass.CaseWhen + val keysGuardProto = ExprOuterClass.CaseWhen .newBuilder() - .addWhen(andBinaryExprProto) + .addWhen(keysNotNullExprProto) .addThen(mapFromArraysExprProto) - .setElseExpr(nullLiteralExprProto) .build() ExprOuterClass.Expr .newBuilder() - .setCaseWhen(caseWhenExprProto) + .setCaseWhen(keysGuardProto) .build() } } - - private def createAndBinaryExpr( - expr: MapFromArrays, - inputs: Seq[Attribute], - binding: Boolean): Option[ExprOuterClass.Expr] = { - createBinaryExpr( - expr, - IsNotNull(expr.left), - IsNotNull(expr.right), - inputs, - binding, - (builder, binaryExpr) => builder.setAnd(binaryExpr)) - } } object CometMapFromEntries @@ -217,20 +239,22 @@ object CometMapFromEntries "`BinaryType` is not supported as a map value in `map_from_entries`" override def getIncompatibleReasons(): Seq[String] = - Seq(keyUnsupportedReason, valueUnsupportedReason, MapKeyDedupPolicySupport.incompatibleReason) + Seq( + keyUnsupportedReason, + valueUnsupportedReason, + MapBuilderSupport.collationKeyReason, + MapBuilderSupport.strictFloatingPointKeyReason) override def getCompatibleNotes(): Seq[String] = - Seq(MapKeyDedupPolicySupport.nullKeyReason) + Seq(MapBuilderSupport.floatingPointKeyNote) override def getSupportLevel(expr: MapFromEntries): SupportLevel = { if (SupportLevel.containsType(expr.dataType.keyType, classOf[BinaryType])) { Incompatible(Some(keyUnsupportedReason)) } else if (SupportLevel.containsType(expr.dataType.valueType, classOf[BinaryType])) { Incompatible(Some(valueUnsupportedReason)) - } else if (MapKeyDedupPolicySupport.isLastWin) { - Incompatible(Some(MapKeyDedupPolicySupport.incompatibleReason)) } else { - Compatible(None) + MapBuilderSupport.keySupport(expr.dataType.keyType) } } } diff --git a/spark/src/test/resources/sql-tests/expressions/map/map_builders_collation.sql b/spark/src/test/resources/sql-tests/expressions/map/map_builders_collation.sql new file mode 100644 index 00000000000..3673b854b98 --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/map/map_builders_collation.sql @@ -0,0 +1,52 @@ +-- 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. + +-- MinSparkVersion: 4.0 + +-- Spark 4.0+ supports string collations. `ArrayBasedMapBuilder` keys its dedup map on +-- `TypeUtils.getInterpretedOrdering` once the key type contains a string, so under `UTF8_LCASE` +-- the keys 'a' and 'A' are one key and Spark raises `DUPLICATED_MAP_KEY`. Comet's native +-- builders compare the raw Arrow bytes and would keep both, so both constructors decline a +-- collated key type outright, whether or not a given row actually collides. +-- +-- The keys below are distinct under `UTF8_LCASE` so both engines return a map and the queries +-- can check where the expression ran. `CometMapFromArrays` has no codegen dispatcher, so it +-- falls back to Spark; `CometMapFromEntries` mixes in `CodegenDispatchFallback`, so it stays in +-- the Comet pipeline running Spark's own generated code. +-- +-- `size` wraps each call so the projection's output type is an `int`. A map with a collated key +-- is not a supported Comet output type, and that check runs first: returning the map itself +-- takes the whole plan off Comet with no expression-level reason, testing nothing here. + +statement +CREATE TABLE test_map_builders_collation(k string) USING parquet + +statement +INSERT INTO test_map_builders_collation VALUES ('a'), ('b') + +query expect_fallback(cannot honour a non-default collation) +SELECT size(map_from_arrays( + array(CAST(k AS STRING COLLATE UTF8_LCASE), + CAST(concat(k, 'z') AS STRING COLLATE UTF8_LCASE)), + array(1, 2))) +FROM test_map_builders_collation + +query expect_dispatch(map_from_entries) +SELECT size(map_from_entries(array( + struct(CAST(k AS STRING COLLATE UTF8_LCASE) AS key, 1 AS value), + struct(CAST(concat(k, 'z') AS STRING COLLATE UTF8_LCASE) AS key, 2 AS value)))) +FROM test_map_builders_collation diff --git a/spark/src/test/resources/sql-tests/expressions/map/map_builders_strict_fp.sql b/spark/src/test/resources/sql-tests/expressions/map/map_builders_strict_fp.sql new file mode 100644 index 00000000000..c50f04de726 --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/map/map_builders_strict_fp.sql @@ -0,0 +1,43 @@ +-- 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. + +-- With `spark.comet.exec.strictFloatingPoint` on, the native map constructors decline a +-- floating-point key type, which they compare by its raw bits where Spark 4.0+ normalizes it (see +-- the map_funcs expression audit). Each constructor goes a different way: `CometMapFromArrays` has +-- no codegen dispatcher, so the projection falls back to Spark, while `CometMapFromEntries` mixes +-- in `CodegenDispatchFallback`, so it stays in the Comet pipeline running Spark's own generated +-- code. A floating-point value type is not affected. + +-- Config: spark.comet.exec.strictFloatingPoint=true + +statement +CREATE TABLE test_map_builders_strict_fp(k double, v int) USING parquet + +statement +INSERT INTO test_map_builders_strict_fp VALUES (1.0D, 1), (double('-0.0'), 2), (2.5D, 3) + +query expect_fallback(Map construction on a floating-point key) +SELECT map_from_arrays(array(k), array(v)) FROM test_map_builders_strict_fp + +query expect_dispatch(map_from_entries) +SELECT map_from_entries(array(struct(k, v))) FROM test_map_builders_strict_fp + +query expect_native(map_from_arrays) +SELECT map_from_arrays(array(v), array(k)) FROM test_map_builders_strict_fp + +query expect_native(map_from_entries) +SELECT map_from_entries(array(struct(v, k))) FROM test_map_builders_strict_fp diff --git a/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays.sql b/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays.sql index 178c07f432a..fea572a0e11 100644 --- a/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays.sql +++ b/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays.sql @@ -58,4 +58,33 @@ query SELECT map_from_arrays(array('a'), NULL) query -SELECT map_from_arrays(NULL, NULL) \ No newline at end of file +SELECT map_from_arrays(NULL, NULL) + +-- Spark's ArrayBasedMapBuilder rejects a NULL key and, under the default +-- `spark.sql.mapKeyDedupPolicy` = `EXCEPTION`, a duplicate key. It inserts a row's entries one at +-- a time, so whichever of the two comes first in the row is the one reported. +-- `map_from_arrays_dedup_policy.sql` covers `LAST_WIN`. + +query expect_error(NULL_MAP_KEY) +SELECT map_from_arrays(array('a', NULL), array(1, 2)) + +-- the NULL comes first, so it is reported although the key after it repeats it +query expect_error(NULL_MAP_KEY) +SELECT map_from_arrays(array(CAST(NULL AS STRING), NULL), array(1, 2)) + +query expect_error(DUPLICATED_MAP_KEY) +SELECT map_from_arrays(array('a', 'a'), array(1, 2)) + +-- the duplicate comes first, so it is reported although a NULL key follows it +query expect_error(DUPLICATED_MAP_KEY) +SELECT map_from_arrays(array('a', 'a', NULL), array(1, 2, 3)) + +-- key and value arrays of different lengths. Spark reports this through a legacy condition, +-- `_LEGACY_ERROR_TEMP_2128` in every version Comet supports; matching on the message keeps the +-- fixture readable. +query expect_error(must have the same length) +SELECT map_from_arrays(array('a', 'b'), array(1)) + +-- and in the other direction +query expect_error(must have the same length) +SELECT map_from_arrays(array('a'), array(1, 2)) diff --git a/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays_dedup_policy.sql b/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays_dedup_policy.sql index fffaf5f9a92..25406740c99 100644 --- a/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays_dedup_policy.sql +++ b/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays_dedup_policy.sql @@ -15,10 +15,9 @@ -- specific language governing permissions and limitations -- under the License. --- Verifies that `map_from_arrays` falls back to Spark when `spark.sql.mapKeyDedupPolicy` is set --- to `LAST_WIN`. Spark's ArrayBasedMapBuilder keeps the last occurrence of each duplicate key; --- Comet's native `map` scalar has no LAST_WIN path, so it must fall back. The default `EXCEPTION` --- mode agrees with Comet and is covered by `map_from_arrays.sql`. +-- Verifies that `map_from_arrays` runs natively under `spark.sql.mapKeyDedupPolicy` = `LAST_WIN` +-- and keeps the last value for each duplicate key. The default `EXCEPTION` mode is covered by +-- `map_from_arrays.sql`. -- Config: spark.sql.mapKeyDedupPolicy=LAST_WIN @@ -29,13 +28,23 @@ statement INSERT INTO test_map_from_arrays_dedup VALUES (array('a', 'b', 'c'), array(1, 2, 3)), (array('a', 'a', 'b'), array(1, 2, 3)), - (array('x', 'x'), array(10, 20)) + (array('x', 'x'), array(10, 20)), + (array('a', 'a', 'a'), array(1, 2, 3)), + (array('a', 'b', 'a'), array(1, 2, 3)), + (array('a', 'a', 'b'), array(1, NULL, 3)), + (array(), array()), + (NULL, array(99)) --- literal duplicate keys under LAST_WIN: Spark keeps the last value; Comet must fall back. -query expect_fallback(mapKeyDedupPolicy) +-- literal arguments, for the all-scalar path +query SELECT map_from_arrays(array('a', 'a', 'b'), array(1, 2, 3)) --- column input falls back the same way; the incompat branch is triggered by the SQLConf value, --- not per-row content. -query expect_fallback(mapKeyDedupPolicy) -SELECT map_from_arrays(k, v) FROM test_map_from_arrays_dedup +-- A repeated key keeps the position of its first occurrence and takes its last value, as +-- `ArrayBasedMapBuilder` does, so ('a', 'b', 'a') gives {a -> 3, b -> 2}; a NULL can be the value +-- that wins. Maps compare equal in any entry order, so `map_keys` and `map_values` pin the order. +query expect_native(map_from_arrays) +SELECT map_keys(map_from_arrays(k, v)), map_values(map_from_arrays(k, v)) FROM test_map_from_arrays_dedup + +-- LAST_WIN does not weaken the NULL key check +query expect_error(NULL_MAP_KEY) +SELECT map_from_arrays(array('a', NULL), array(1, 2)) diff --git a/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays_nondeterministic_child.sql b/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays_nondeterministic_child.sql new file mode 100644 index 00000000000..e2d7433e4b4 --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/map/map_from_arrays_nondeterministic_child.sql @@ -0,0 +1,56 @@ +-- 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. + +-- `CometMapFromArrays` reproduces Spark's evaluation order with a +-- `CASE WHEN keys IS NOT NULL THEN map_from_arrays(keys, values) END` guard, which serializes the +-- keys a second time. A stateful keys expression advances each copy independently: the guard's +-- copy sees every row while the constructor's copy sees only the rows the guard selected, so the +-- result would silently drift from Spark (#5781). The serde declines a nondeterministic keys +-- expression and the projection falls back to Spark, which evaluates it once. + +statement +CREATE TABLE test_map_from_arrays_nondet(_1 int) USING parquet + +statement +INSERT INTO test_map_from_arrays_nondet SELECT id FROM range(0, 16) + +-- Spark returns {1 -> 2} on every row whose keys array is non-NULL and NULL on the rest. +query expect_fallback(nondeterministic operand) +SELECT _1, map_from_arrays(IF(monotonically_increasing_id() % 2 = 0, array(1), CAST(NULL AS ARRAY)), array(2)) AS m +FROM test_map_from_arrays_nondet + +-- A non-nullable stateful keys expression is declined too, rather than relying on the guard +-- matching every row. +query expect_fallback(nondeterministic operand) +SELECT _1, map_from_arrays(array(monotonically_increasing_id()), array(2)) AS m +FROM test_map_from_arrays_nondet + +-- The values are serialized once, inside the call, and evaluated only on the rows whose keys +-- array is non-NULL, which are the rows Spark evaluates them on. So a stateful values expression +-- stays native and numbers the same rows Spark does. +query expect_native(map_from_arrays) +SELECT _1, map_from_arrays(IF(_1 % 2 = 0, array(1), CAST(NULL AS ARRAY)), array(monotonically_increasing_id())) AS m +FROM test_map_from_arrays_nondet + +query expect_native(map_from_arrays) +SELECT _1, map_from_arrays(array(1), IF(monotonically_increasing_id() % 2 = 0, array(2), CAST(NULL AS ARRAY))) AS m +FROM test_map_from_arrays_nondet + +-- A deterministic nullable keys expression stays on the native guarded path. +query expect_native(map_from_arrays) +SELECT _1, map_from_arrays(IF(_1 % 2 = 0, array(1), CAST(NULL AS ARRAY)), array(2)) AS m +FROM test_map_from_arrays_nondet diff --git a/spark/src/test/resources/sql-tests/expressions/map/map_from_entries.sql b/spark/src/test/resources/sql-tests/expressions/map/map_from_entries.sql index 74723509334..3c320917723 100644 --- a/spark/src/test/resources/sql-tests/expressions/map/map_from_entries.sql +++ b/spark/src/test/resources/sql-tests/expressions/map/map_from_entries.sql @@ -21,6 +21,12 @@ CREATE TABLE test_map_from_entries(entries array>) statement INSERT INTO test_map_from_entries VALUES (array(struct('a', 1), struct('b', 2), struct('c', 3))), (array()), (NULL) +-- A NULL entry makes the whole map NULL, so no key of the row is inserted, not even a NULL one. +statement +INSERT INTO test_map_from_entries VALUES + (array(CAST(NULL AS struct), struct('b' AS key, 2 AS value))), + (array(CAST(NULL AS struct), struct(CAST(NULL AS STRING) AS key, 2 AS value))) + query SELECT map_from_entries(entries) FROM test_map_from_entries @@ -35,3 +41,18 @@ SELECT map_from_entries(array(struct(10, cast('x' as binary)))) -- literal arguments query spark_answer_only SELECT map_from_entries(array(struct('x', 10), struct('y', 20), struct('z', 30))) + +-- Spark's ArrayBasedMapBuilder rejects a NULL key and, under the default +-- `spark.sql.mapKeyDedupPolicy` = `EXCEPTION`, a duplicate key. It inserts a row's entries one at +-- a time, so whichever of the two comes first in the row is the one reported. +-- `map_from_entries_dedup_policy.sql` covers `LAST_WIN`. + +query expect_error(NULL_MAP_KEY) +SELECT map_from_entries(array(struct(CAST(NULL AS STRING), 1), struct('b', 2))) + +query expect_error(DUPLICATED_MAP_KEY) +SELECT map_from_entries(array(struct('a', 1), struct('a', 2))) + +-- the duplicate comes first, so it is reported although a NULL key follows it +query expect_error(DUPLICATED_MAP_KEY) +SELECT map_from_entries(array(struct('a', 1), struct('a', 2), struct(CAST(NULL AS STRING), 3))) diff --git a/spark/src/test/resources/sql-tests/expressions/map/map_from_entries_dedup_policy.sql b/spark/src/test/resources/sql-tests/expressions/map/map_from_entries_dedup_policy.sql index feba7951933..906252bd0a7 100644 --- a/spark/src/test/resources/sql-tests/expressions/map/map_from_entries_dedup_policy.sql +++ b/spark/src/test/resources/sql-tests/expressions/map/map_from_entries_dedup_policy.sql @@ -15,15 +15,11 @@ -- specific language governing permissions and limitations -- under the License. --- Verifies that `map_from_entries` falls back to Spark when `spark.sql.mapKeyDedupPolicy` is set --- to `LAST_WIN`. `CometMapFromEntries` mixes in `CodegenDispatchFallback`, so its native --- `Incompatible` normally routes through the JVM codegen dispatcher; we disable the dispatcher --- here so the incompat branch surfaces as a genuine Spark fallback rather than in-pipeline --- codegen. The default `EXCEPTION` mode agrees with Comet and is covered by +-- Verifies that `map_from_entries` runs natively under `spark.sql.mapKeyDedupPolicy` = `LAST_WIN` +-- and keeps the last value for each duplicate key. The default `EXCEPTION` mode is covered by -- `map_from_entries.sql`. -- Config: spark.sql.mapKeyDedupPolicy=LAST_WIN --- Config: spark.comet.exec.scalaUDF.codegen.enabled=false statement CREATE TABLE test_map_from_entries_dedup(entries array>) USING parquet @@ -32,13 +28,24 @@ statement INSERT INTO test_map_from_entries_dedup VALUES (array(struct('a', 1), struct('b', 2), struct('c', 3))), (array(struct('a', 1), struct('a', 2), struct('b', 3))), - (array(struct('x', 10), struct('x', 20))) + (array(struct('x', 10), struct('x', 20))), + (array(struct('a', 1), struct('a', 2), struct('a', 3))), + (array(struct('a', 1), struct('b', 2), struct('a', 3))), + (array(struct('a', 1), struct('a', CAST(NULL AS INT)), struct('b', 3))), + (array()), + (NULL) --- literal duplicate keys under LAST_WIN: Spark keeps the last value; Comet must fall back. -query expect_fallback(mapKeyDedupPolicy) +-- literal arguments, for the all-scalar path +query SELECT map_from_entries(array(struct('a', 1), struct('a', 2), struct('b', 3))) --- column input falls back the same way; the incompat branch is triggered by the SQLConf value, --- not per-row content. -query expect_fallback(mapKeyDedupPolicy) -SELECT map_from_entries(entries) FROM test_map_from_entries_dedup +-- A repeated key keeps the position of its first occurrence and takes its last value, as +-- `ArrayBasedMapBuilder` does, so ('a', 'b', 'a') gives {a -> 3, b -> 2}; a NULL can be the value +-- that wins. Maps compare equal in any entry order, so `map_keys` and `map_values` pin the order. +-- `expect_native` also rules out the JVM codegen dispatcher, which a plain `query` would accept. +query expect_native(map_from_entries) +SELECT map_keys(map_from_entries(entries)), map_values(map_from_entries(entries)) FROM test_map_from_entries_dedup + +-- LAST_WIN does not weaken the NULL key check +query expect_error(NULL_MAP_KEY) +SELECT map_from_entries(array(struct(CAST(NULL AS STRING), 1), struct('b', 2))) diff --git a/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_disabled.sql b/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_disabled.sql index a390e18e091..5b1ea65e21a 100644 --- a/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_disabled.sql +++ b/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_disabled.sql @@ -17,18 +17,13 @@ -- Config: spark.comet.exec.scalaUDF.codegen.enabled=false -- Config: spark.sql.legacy.truncateForEmptyRegexSplit=true --- Config: spark.sql.mapKeyDedupPolicy=LAST_WIN -- Config: spark.comet.expression.StringToMap.allowIncompatible=false --- Config: spark.comet.expression.MapFromEntries.allowIncompatible=false statement -CREATE TABLE routing_map_legacy(s STRING, e ARRAY>) USING parquet +CREATE TABLE routing_map_legacy(s STRING) USING parquet statement -INSERT INTO routing_map_legacy VALUES ('a:1,b:2', array(named_struct('key', 'a', 'value', 1))), (NULL, NULL) +INSERT INTO routing_map_legacy VALUES ('a:1,b:2'), (NULL) query expect_fallback(str_to_map: spark.comet.exec.scalaUDF.codegen.enabled=false) SELECT str_to_map(s) FROM routing_map_legacy - -query expect_fallback(map_from_entries: spark.comet.exec.scalaUDF.codegen.enabled=false) -SELECT map_from_entries(e) FROM routing_map_legacy diff --git a/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_enabled.sql b/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_enabled.sql index afbfc95dba4..77cec6b40cf 100644 --- a/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_enabled.sql +++ b/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_enabled.sql @@ -17,18 +17,13 @@ -- Config: spark.comet.exec.scalaUDF.codegen.enabled=true -- Config: spark.sql.legacy.truncateForEmptyRegexSplit=true --- Config: spark.sql.mapKeyDedupPolicy=LAST_WIN -- Config: spark.comet.expression.StringToMap.allowIncompatible=false --- Config: spark.comet.expression.MapFromEntries.allowIncompatible=false statement -CREATE TABLE routing_map_legacy(s STRING, e ARRAY>) USING parquet +CREATE TABLE routing_map_legacy(s STRING) USING parquet statement -INSERT INTO routing_map_legacy VALUES ('a:1,b:2', array(named_struct('key', 'a', 'value', 1))), (NULL, NULL) +INSERT INTO routing_map_legacy VALUES ('a:1,b:2'), (NULL) query expect_dispatch(str_to_map) SELECT str_to_map(s) FROM routing_map_legacy - -query expect_dispatch(map_from_entries) -SELECT map_from_entries(e) FROM routing_map_legacy diff --git a/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_opt_in.sql b/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_opt_in.sql index 91b6ecf1bbf..a0c0575b53f 100644 --- a/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_opt_in.sql +++ b/spark/src/test/resources/sql-tests/expressions/map/routing_map_legacy_opt_in.sql @@ -17,18 +17,13 @@ -- Config: spark.comet.exec.scalaUDF.codegen.enabled=false -- Config: spark.sql.legacy.truncateForEmptyRegexSplit=true --- Config: spark.sql.mapKeyDedupPolicy=LAST_WIN -- Config: spark.comet.expression.StringToMap.allowIncompatible=true --- Config: spark.comet.expression.MapFromEntries.allowIncompatible=true statement -CREATE TABLE routing_map_legacy(s STRING, e ARRAY>) USING parquet +CREATE TABLE routing_map_legacy(s STRING) USING parquet statement -INSERT INTO routing_map_legacy VALUES ('a:1,b:2', array(named_struct('key', 'a', 'value', 1))), (NULL, NULL) +INSERT INTO routing_map_legacy VALUES ('a:1,b:2'), (NULL) query expect_native(str_to_map) SELECT str_to_map(s) FROM routing_map_legacy - -query expect_native(map_from_entries) -SELECT map_from_entries(e) FROM routing_map_legacy diff --git a/spark/src/test/resources/sql-tests/expressions/map/str_to_map.sql b/spark/src/test/resources/sql-tests/expressions/map/str_to_map.sql index 7db1242fd4e..481ab398913 100644 --- a/spark/src/test/resources/sql-tests/expressions/map/str_to_map.sql +++ b/spark/src/test/resources/sql-tests/expressions/map/str_to_map.sql @@ -70,10 +70,11 @@ SELECT str_to_map('a') query SELECT str_to_map('a=1&b=2&c=3', '&', '=') --- Duplicate keys: EXCEPTION policy (Spark 3.0+ default) --- TODO: Add LAST_WIN policy tests when spark.sql.mapKeyDedupPolicy config is supported --- query --- SELECT str_to_map('a:1,b:2,a:3') +-- Duplicate keys under the default EXCEPTION policy; `str_to_map_dedup_policy.sql` covers +-- LAST_WIN. Spark's message names the key unquoted, while the upstream kernel's quotes it, so this +-- only matches once Comet has converted the error to Spark's. +query expect_error(Duplicate map key a was found) +SELECT str_to_map('a:1,b:2,a:3') -- NULL input returns NULL query diff --git a/spark/src/test/resources/sql-tests/expressions/map/str_to_map_dedup_policy.sql b/spark/src/test/resources/sql-tests/expressions/map/str_to_map_dedup_policy.sql new file mode 100644 index 00000000000..8af746886f1 --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/map/str_to_map_dedup_policy.sql @@ -0,0 +1,43 @@ +-- 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. + +-- Verifies that `str_to_map` runs natively under `spark.sql.mapKeyDedupPolicy` = `LAST_WIN` and +-- keeps the last value for each duplicate key. The default `EXCEPTION` mode is covered by +-- `str_to_map.sql`. + +-- Config: spark.sql.mapKeyDedupPolicy=LAST_WIN + +statement +CREATE TABLE test_str_to_map_dedup(s string) USING parquet + +statement +INSERT INTO test_str_to_map_dedup VALUES + ('a:1,b:2,a:3'), + ('a:1,b:2,c:3'), + ('x:1,x:2,x:3'), + (NULL) + +-- literal arguments, for the all-scalar path +query +SELECT str_to_map('a:1,b:2,a:3') + +-- `a` keeps the position of its first occurrence and takes its last value, as +-- `ArrayBasedMapBuilder` does, so 'a:1,b:2,a:3' gives {a -> 3, b -> 2}. Maps compare equal in any +-- entry order, so `map_keys` and `map_values` pin the order. `expect_native` also rules out the +-- JVM codegen dispatcher, which a plain `query` would accept. +query expect_native(str_to_map) +SELECT map_keys(str_to_map(s)), map_values(str_to_map(s)) FROM test_str_to_map_dedup diff --git a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala index ebbdce406a3..1bc71f18806 100644 --- a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala @@ -19,10 +19,10 @@ package org.apache.comet -import scala.util.Random +import scala.util.{Random, Try} import org.apache.hadoop.fs.Path -import org.apache.spark.sql.CometTestBase +import org.apache.spark.sql.{CometTestBase, DataFrame, Row} import org.apache.spark.sql.catalyst.expressions.ArrayContains import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf @@ -126,6 +126,196 @@ class CometMapExpressionSuite extends CometTestBase { } } + // Spark builds both `map_from_arrays` and `map_from_entries` through `ArrayBasedMapBuilder`, + // which rejects a NULL key and resolves a duplicate key by `spark.sql.mapKeyDedupPolicy`. The + // native builders reproduce both (see the map_funcs expression audit), so the engines must agree + // on the answer and on the error. Each query reads a column so constant folding cannot evaluate + // it on the driver. https://github.com/apache/datafusion-comet/issues/4680 + private def withMapBuilderTable(f: String => Unit): Unit = { + val table = "map_builder_input" + withTable(table) { + sql(s"CREATE TABLE $table(k INT, v STRING) USING parquet") + sql(s"INSERT INTO $table VALUES (1, 'a'), (2, 'b'), (3, 'c')") + f(table) + } + } + + test("map_from_arrays - null key is rejected") { + withMapBuilderTable { table => + val exception = checkSparkError( + sql(s"SELECT map_from_arrays(array(k, CAST(NULL AS INT)), array(v, v)) FROM $table"), + "NULL_MAP_KEY") + assert(exception.getMessage.contains("Cannot use null as map key")) + } + } + + // Spark never evaluates the values argument for a row whose keys array is NULL, so a failing + // cast there must not run natively either; the serde's NULL guard is for this (see the map_funcs + // expression audit). The rows with keys outnumber the row without, all in one batch, which is + // the shape where a single `keys IS NOT NULL AND values IS NOT NULL` guard would still evaluate + // the cast on every row. + test("map_from_arrays - a null keys array skips the values expression under ANSI") { + withSQLConf(SQLConf.ANSI_ENABLED.key -> "true") { + withTable("map_short_circuit") { + // One partition, so every row lands in the same file and the same batch. + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(0, 5, 1, 1) + .selectExpr( + "IF(id = 0, CAST(NULL AS ARRAY), array(CAST(id AS INT))) AS k", + "IF(id = 0, 'bad', CAST(id AS STRING)) AS v") + .write + .format("parquet") + .saveAsTable("map_short_circuit") + } + checkSparkAnswerAndOperator( + sql("SELECT map_from_arrays(k, array(CAST(v AS INT))) FROM map_short_circuit")) + } + } + } + + test("map_from_arrays - a null input array gives a null map") { + withMapBuilderTable { table => + checkSparkAnswerAndOperator( + sql(s"""SELECT map_from_arrays(CASE WHEN k > 1 THEN array(k) END, array(v)), + | map_from_arrays(array(k), CASE WHEN k > 2 THEN array(v) END) + |FROM $table""".stripMargin)) + } + } + + test("map_from_arrays - key and value arrays of different lengths are rejected") { + withMapBuilderTable { table => + // Spark reports this through a legacy condition rather than a named one, but the number is + // the same in every version Comet supports (checked in 3.4.3, 3.5.8 and 4.1.3). + checkSparkError( + sql(s"SELECT map_from_arrays(array(k, k + 1), array(v)) FROM $table"), + "_LEGACY_ERROR_TEMP_2128") + } + } + + test("map_from_arrays - a duplicate key is rejected under EXCEPTION") { + withMapBuilderTable { table => + withSQLConf(SQLConf.MAP_KEY_DEDUP_POLICY.key -> "EXCEPTION") { + // One row, so both engines name the same offending key. + val exception = checkSparkError( + sql( + s"SELECT map_from_arrays(array(k, k), array(v, concat(v, 'x'))) FROM $table " + + "WHERE k = 2"), + "DUPLICATED_MAP_KEY") + assert(exception.getMessage.contains("Duplicate map key 2 was found")) + } + } + } + + // Outside whole-stage codegen Spark reads `spark.sql.mapKeyDedupPolicy` again on every action, + // and so does Comet; the map_funcs expression audit covers this and the one divergence, inside + // whole-stage codegen. Each sequence plans `df` under one policy, runs it under the other, then + // runs it again under the first, always collecting `df` itself so the plan it materialized is + // the one executed. The first run fails if the policy were fixed when the plan is converted, + // the second if it were fixed by the first run. + // https://github.com/apache/datafusion-comet/issues/4680 + test("map constructors read the dedup policy on every action") { + // AQE converts each query stage as it runs, which hides when the setting is read. + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + withTempPath { dir => + val path = dir.getCanonicalPath + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark.range(0, 1, 1, 1).write.parquet(path) + } + // The two ways a projection runs outside whole-stage codegen: the flag is off, or the + // projection is wider than `spark.sql.codegen.maxFields`. + val outsideWholeStageCodegen = Seq( + (Seq(SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false"), Seq.empty[String]), + (Seq.empty[(String, String)], (1 to 100).map(i => s"id + $i AS c$i"))) + for ((codegenConf, padding) <- outsideWholeStageCodegen; + engine <- Seq("Spark", "Comet"); + (first, second) <- Seq(("LAST_WIN", "EXCEPTION"), ("EXCEPTION", "LAST_WIN"))) { + val cometEnabled = CometConf.COMET_ENABLED.key -> (engine == "Comet").toString + withSQLConf((codegenConf :+ cometEnabled): _*) { + val df = mapPolicyQuery(path, padding) + withSQLConf(SQLConf.MAP_KEY_DEDUP_POLICY.key -> first) { + val plan = df.queryExecution.executedPlan + if (engine == "Comet") { + checkCometOperators(stripAQEPlan(plan)) + } + } + withSQLConf(SQLConf.MAP_KEY_DEDUP_POLICY.key -> second) { + assertMapPolicy(df, second, engine) + } + withSQLConf(SQLConf.MAP_KEY_DEDUP_POLICY.key -> first) { + assertMapPolicy(df, first, engine) + } + } + } + } + } + } + + /** + * One row whose three map constructors each see the duplicate key `0`, with `padding` extra + * columns for callers that need the projection to exceed `spark.sql.codegen.maxFields`. + */ + private def mapPolicyQuery(path: String, padding: Seq[String]): DataFrame = + spark.read + .parquet(path) + .selectExpr(Seq( + "map_from_arrays(array(id, id), array(1, 2)) AS a", + "map_from_entries(array(struct(id, 1), struct(id, 2))) AS e", + "str_to_map(concat(CAST(id AS STRING), ':1,', CAST(id AS STRING), ':2')) AS s") ++ + padding: _*) + + /** Collects `df` itself, so the run reuses the plan `df` has already materialized. */ + private def assertMapPolicy(df: DataFrame, policy: String, engine: String): Unit = + policy match { + case "LAST_WIN" => + // The three maps only, leaving out any padding. + val maps = df.collect().toSeq.map(row => Row(row.get(0), row.get(1), row.get(2))) + assert(maps == Seq(Row(Map(0L -> 2), Map(0L -> 2), Map("0" -> "2"))), engine) + case "EXCEPTION" => + val error = + structuredError(Try(df.collect()).failed.toOption, engine, "DUPLICATED_MAP_KEY") + assert(error.getErrorClass == "DUPLICATED_MAP_KEY", s"$engine: $error") + } + + test("map_from_entries - null key is rejected") { + withMapBuilderTable { table => + val exception = checkSparkError( + sql(s"SELECT map_from_entries(array(struct(CAST(NULL AS INT), v))) FROM $table"), + "NULL_MAP_KEY") + assert(exception.getMessage.contains("Cannot use null as map key")) + } + } + + test("map_from_entries - a duplicate key is rejected under EXCEPTION") { + withMapBuilderTable { table => + withSQLConf(SQLConf.MAP_KEY_DEDUP_POLICY.key -> "EXCEPTION") { + // `struct` names a column argument after the column, so both entries need explicit field + // names for `array` to see one struct type. + val exception = checkSparkError( + sql( + "SELECT map_from_entries(array(struct(k AS key, v AS value), " + + s"struct(k AS key, concat(v, 'x') AS value))) FROM $table WHERE k = 2"), + "DUPLICATED_MAP_KEY") + assert(exception.getMessage.contains("Duplicate map key 2 was found")) + } + } + } + + // A native `LIMIT ... OFFSET` slices the batch, so the constructors see lists whose entries do + // not start at offset zero, which the upstream kernels mishandle (apache/datafusion#25419). No + // row has a NULL keys array, so `map_from_arrays`'s NULL guard passes the sliced batch through + // instead of compacting it. + test("map constructors on a sliced list read the visible entries") { + val rows = (0 until 20).map { i => + (i, Seq((s"a$i", i), (s"b$i", i * 10)), Seq(s"a$i", s"b$i"), Seq(i, i * 10)) + } + withParquetTable(rows, "t") { + checkSparkAnswerAndOperator( + "SELECT _1, map_from_entries(_2), map_from_arrays(_3, _4) " + + "FROM (SELECT * FROM t ORDER BY _1 LIMIT 15 OFFSET 5)") + } + } + test("size with map input") { withTempDir { dir => withTempView("t1") { diff --git a/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala b/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala index 9740a4a5468..19dbf01fd8c 100644 --- a/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala +++ b/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala @@ -459,20 +459,8 @@ abstract class CometTestBase errorClass: String): SparkThrowable with Throwable = { checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan)) val (sparkError, cometError) = checkSparkAnswerMaybeThrows(df) - - def structuredError( - error: Option[Throwable], - engine: String): SparkThrowable with Throwable = { - val failure = error.getOrElse(fail(s"$engine did not fail with $errorClass")) - val chain = causeChain(failure) - assert(!chain.exists(_.isInstanceOf[CometNativeException]), s"$engine: $failure") - chain.collect { case e: SparkThrowable with Throwable => e }.lastOption.getOrElse { - fail(s"$engine did not throw a SparkThrowable: $failure") - } - } - - val expected = structuredError(sparkError, "Spark") - val actual = structuredError(cometError, "Comet") + val expected = structuredError(sparkError, "Spark", errorClass) + val actual = structuredError(cometError, "Comet", errorClass) assert(expected.getErrorClass == errorClass) assert(actual.getClass == expected.getClass) assert(actual.getErrorClass == errorClass) @@ -480,6 +468,24 @@ abstract class CometTestBase actual } + /** + * The last `SparkThrowable` in the cause chain of `error`, which `engine` raised where + * `errorClass` was expected. Fails when there is no error, when no `SparkThrowable` is in the + * chain, or when a `CometNativeException` is anywhere in it, which means a native error reached + * the user without being converted to Spark's. + */ + protected def structuredError( + error: Option[Throwable], + engine: String, + errorClass: String): SparkThrowable with Throwable = { + val failure = error.getOrElse(fail(s"$engine did not fail with $errorClass")) + val chain = causeChain(failure) + assert(!chain.exists(_.isInstanceOf[CometNativeException]), s"$engine: $failure") + chain.collect { case e: SparkThrowable with Throwable => e }.lastOption.getOrElse { + fail(s"$engine did not throw a SparkThrowable: $failure") + } + } + /** * Compares the Comet DataFrame result against the expected Spark answer, using labels that * correctly identify which side is Comet and which is Spark. This avoids the misleading "Spark