viirya commented on code in PR #6682:
URL: https://github.com/apache/datafusion-comet/pull/6682#discussion_r4226229762
##########
spark/src/main/scala/org/apache/comet/serde/maps.scala:
##########
@@ -132,77 +133,124 @@ 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` and the
+ * equality it gives a `FLOAT` or `DOUBLE` key. The map_funcs expression audit
covers how, and
+ * where they still differ from Spark.
+ */
+private object MapBuilderSupport {
+
+ private val disableMapKeyNormalizationKey =
"spark.sql.legacy.disableMapKeyNormalization"
+
+ /**
+ * Whether `ArrayBasedMapBuilder` normalizes a `FLOAT` or `DOUBLE` key
before it compares and
+ * stores it, as Spark 4.0 and later do unless
`spark.sql.legacy.disableMapKeyNormalization` is
+ * set: `-0.0` becomes `0.0` and every `NaN` the canonical `NaN`. Otherwise
it compares boxed
+ * keys with `Double.equals`, so `NaN`s are one key but `-0.0` and `0.0` are
two. The native
+ * builders follow either rule, under different function names.
+ */
+ private def normalizesFloatKeys: Boolean =
Review Comment:
This reads `spark.sql.legacy.disableMapKeyNormalization` at planning time
and bakes the answer into the function name. Spark reads it lazily:
`keyNormalizer` is a lazy val that runs on the first `put`, on the executor,
even under whole-stage codegen. This file deliberately reads
`spark.sql.mapKeyDedupPolicy` per task in `CometExecIterator` for the same
reason. So a Dataset that is re-executed after the flag is toggled keeps the
old rule in Comet, where Spark picks up the new one. Could we pass the flag the
same way the dedup policy goes? If keeping the choice in Scala is the intended
design per #6385, could we record this as a known limitation in the map_funcs
audit, next to the dedup-policy note?
##########
native/spark-expr/src/map_funcs/map_builders.rs:
##########
@@ -0,0 +1,1526 @@
+// 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.
+//!
+//! The kernels find duplicate keys by comparing `ScalarValue`s, which compare
a `FLOAT` or
+//! `DOUBLE` by its bits. Spark compares those keys as boxed values, where
every NaN is one key,
+//! and from Spark 4.0 it normalizes them first, so `-0.0` and `0.0` are one
key too. For such a
+//! key the wrappers hand the kernel the keys as Spark compares them, put back
the keys that Spark
+//! stores in the map it returns, and name a duplicate key as Spark does.
[`MapFloatKeys`] says
+//! which rule a call follows.
+
+use crate::conversion_funcs::java_float_string;
+use crate::float_semantics::{canonicalize_nans, normalize_floats};
+use crate::SparkError;
+use arrow::array::{
+ Array, ArrayRef, AsArray, BooleanArray, BooleanBufferBuilder, ListArray,
MapArray, StructArray,
+};
+use arrow::buffer::{BooleanBuffer, NullBuffer, OffsetBuffer};
+use arrow::compute::filter;
+use arrow::compute::kernels::zip::zip;
+use arrow::datatypes::{DataType, FieldRef, Float32Type, Float64Type};
+use datafusion::common::config::MapKeyDedupPolicy;
+use datafusion::common::{exec_err, internal_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;
+
+/// How a map builder compares and stores a `FLOAT` or `DOUBLE` key. Spark's
`ArrayBasedMapBuilder`
+/// finds duplicates in a `HashMap` of boxed keys and, from Spark 4.0,
normalizes a float key
+/// before it looks it up, unless
`spark.sql.legacy.disableMapKeyNormalization` is set. The serde
+/// picks the rule for the session, and each rule has its own function name.
+#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash)]
+pub enum MapFloatKeys {
+ /// The equality of a boxed `java.lang.Double` or `java.lang.Float`, which
compares
+ /// `doubleToLongBits` or `floatToIntBits`: every NaN is one key, while
`-0.0` and `0.0` are
+ /// two. A map keeps each key as it first occurred. Spark 3.4 and 3.5, and
Spark 4.0+ with
+ /// `spark.sql.legacy.disableMapKeyNormalization`.
+ #[default]
+ Boxed,
+ /// Spark 4.0+ compares and stores the key `NormalizeFloatingNumbers`
gives: `-0.0` becomes
+ /// `0.0`, and every NaN the canonical NaN. `map_from_arrays` stores the
keys it was given
+ /// when none of them repeats, because `ArrayBasedMapBuilder.from` then
returns its input.
+ Normalized,
+}
+
+/// Spark-compatible `map_from_arrays(keys, values)`.
+#[derive(Debug, Default, PartialEq, Eq, Hash)]
+pub struct SparkMapFromArrays {
+ inner: DataFusionMapFromArrays,
+ float_keys: MapFloatKeys,
+}
+
+impl SparkMapFromArrays {
+ pub fn new(float_keys: MapFloatKeys) -> Self {
+ Self {
+ inner: DataFusionMapFromArrays::default(),
+ float_keys,
+ }
+ }
+}
+
+impl ScalarUDFImpl for SparkMapFromArrays {
+ fn name(&self) -> &str {
+ match self.float_keys {
+ MapFloatKeys::Boxed => self.inner.name(),
+ MapFloatKeys::Normalized => "map_from_arrays_normalized_keys",
+ }
+ }
+
+ fn signature(&self) -> &Signature {
+ self.inner.signature()
+ }
+
+ fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
+ self.inner.return_type(arg_types)
+ }
+
+ fn return_field_from_args(&self, args: ReturnFieldArgs) ->
Result<FieldRef> {
+ self.inner.return_field_from_args(args)
+ }
+
+ fn invoke_with_args(&self, args: ScalarFunctionArgs) ->
Result<ColumnarValue> {
+ invoke_list_builder(
+ &self.inner,
+ args,
+ KeyLayout::Arrays,
+ self.float_keys,
+ |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,
+ float_keys: MapFloatKeys,
+}
+
+impl SparkMapFromEntries {
+ pub fn new(float_keys: MapFloatKeys) -> Self {
+ Self {
+ inner: DataFusionMapFromEntries::default(),
+ float_keys,
+ }
+ }
+}
+
+impl ScalarUDFImpl for SparkMapFromEntries {
+ fn name(&self) -> &str {
+ match self.float_keys {
+ MapFloatKeys::Boxed => self.inner.name(),
+ MapFloatKeys::Normalized => "map_from_entries_normalized_keys",
+ }
+ }
+
+ fn signature(&self) -> &Signature {
+ self.inner.signature()
+ }
+
+ fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
+ self.inner.return_type(arg_types)
+ }
+
+ fn return_field_from_args(&self, args: ReturnFieldArgs) ->
Result<FieldRef> {
+ self.inner.return_field_from_args(args)
+ }
+
+ fn invoke_with_args(&self, args: ScalarFunctionArgs) ->
Result<ColumnarValue> {
+ invoke_list_builder(
+ &self.inner,
+ args,
+ KeyLayout::Entries,
+ self.float_keys,
+ |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<DataType> {
+ self.inner.return_type(arg_types)
+ }
+
+ fn return_field_from_args(&self, args: ReturnFieldArgs) ->
Result<FieldRef> {
+ self.inner.return_field_from_args(args)
+ }
+
+ fn invoke_with_args(&self, args: ScalarFunctionArgs) ->
Result<ColumnarValue> {
+ // 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,
+ layout: KeyLayout,
+ float_keys: MapFloatKeys,
+ validate: impl FnOnce(&[ColumnarValue], bool) -> Result<()>,
+) -> Result<ColumnarValue> {
+ // 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;
+ let float_keys = FloatKeyRewrite::prepare(&mut args.args, layout,
float_keys)?;
+ let result = validate(&args.args, last_value_wins).and_then(|()| {
+ inner
+ .invoke_with_args(args)
+ .map_err(|error| as_spark_error(error, DuplicateKeyFormat::Bare))
+ });
+ let result = match &float_keys {
+ None => result?,
+ Some(keys) => keys.restore(result.map_err(|error|
keys.restate_duplicate(error))?)?,
+ };
+ match (all_scalar, result) {
+ (true, ColumnarValue::Array(array)) => Ok(ColumnarValue::Scalar(
+ ScalarValue::try_from_array(&array, 0)?,
+ )),
+ (_, result) => Ok(result),
+ }
+}
+
+/// Where a builder finds its keys among its arguments.
+#[derive(Clone, Copy, PartialEq, Eq)]
+enum KeyLayout {
+ /// `map_from_arrays(keys, values)`: the elements of the first list.
+ Arrays,
+ /// `map_from_entries(entries)`: the first field of the list's structs.
+ Entries,
+}
+
+/// The `FLOAT` or `DOUBLE` keys of one call. The kernel compares keys by
their bits, so it is
+/// handed them as Spark compares them, with every NaN made canonical and,
under
+/// [`MapFloatKeys::Normalized`], `-0.0` made `0.0`. Its result then gets back
the keys Spark
+/// stores.
+struct FloatKeyRewrite {
+ layout: KeyLayout,
+ rule: MapFloatKeys,
+ /// The flat keys as the call passed them.
+ original: ArrayRef,
+ /// The flat keys as Spark compares them, when that changes the bits of
any of them.
+ compared: Option<ArrayRef>,
+ /// Each row's range of flat keys.
+ offsets: OffsetBuffer<i32>,
+ /// The rows Spark builds a map for. Every other row is a NULL map, whose
keys it never sees.
+ built: BooleanBuffer,
+}
+
+impl FloatKeyRewrite {
+ /// Hands `args` the keys as Spark compares them, when the keys are
`FLOAT` or `DOUBLE`. Returns
+ /// `None` for any other key type.
+ fn prepare(
+ args: &mut [ColumnarValue],
+ layout: KeyLayout,
+ rule: MapFloatKeys,
+ ) -> Result<Option<Self>> {
+ let Some((original, offsets, built)) = float_keys_of(args, layout)
else {
+ return Ok(None);
+ };
+ let compared = match rule {
+ MapFloatKeys::Boxed => canonicalize_nans(&original),
+ MapFloatKeys::Normalized => normalize_floats(&original),
+ };
+ // Keys that are not NaN, or not `-0.0` under normalization, compare
the same either way,
+ // and then the kernel's own result is already Spark's. `ArrayData`
compares primitive
+ // values byte for byte, so this sees their bits.
+ let compared = if original.to_data() == compared.to_data() {
+ None
+ } else {
+ replace_keys(args, layout, Arc::clone(&compared))?;
+ Some(compared)
+ };
+ Ok(Some(Self {
+ layout,
+ rule,
+ original,
+ compared,
+ offsets,
+ built,
+ }))
+ }
+
+ /// Replaces the keys of the kernel's map with the ones Spark stores. The
kernel keeps the first
+ /// occurrence of each key, as Spark does, but it stored them as they were
compared.
+ fn restore(&self, result: ColumnarValue) -> Result<ColumnarValue> {
+ let Some(compared) = &self.compared else {
+ return Ok(result);
+ };
+ // `map_from_entries` on Spark 4.0+ always stores the normalized key.
+ if self.rule == MapFloatKeys::Normalized && self.layout ==
KeyLayout::Entries {
+ return Ok(result);
+ }
+ let ColumnarValue::Array(array) = &result else {
+ return Ok(result);
+ };
+ let (Some(map), DataType::Map(field, ordered)) = (array.as_map_opt(),
array.data_type())
Review Comment:
If the result isn't an array (line 327) or isn't a map, `restore` returns
the kernel's result as is. That result holds the canonicalized or normalized
keys, so it would be a wrong answer with no error. Both cases should be
unreachable, but the length check a few lines below already raises
`internal_err!` for a mismatch. Could these two do the same, so that a future
kernel change fails loudly rather than returning the wrong keys?
##########
native/spark-expr/src/map_funcs/map_builders.rs:
##########
@@ -0,0 +1,1526 @@
+// 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.
+//!
+//! The kernels find duplicate keys by comparing `ScalarValue`s, which compare
a `FLOAT` or
+//! `DOUBLE` by its bits. Spark compares those keys as boxed values, where
every NaN is one key,
+//! and from Spark 4.0 it normalizes them first, so `-0.0` and `0.0` are one
key too. For such a
+//! key the wrappers hand the kernel the keys as Spark compares them, put back
the keys that Spark
+//! stores in the map it returns, and name a duplicate key as Spark does.
[`MapFloatKeys`] says
+//! which rule a call follows.
+
+use crate::conversion_funcs::java_float_string;
+use crate::float_semantics::{canonicalize_nans, normalize_floats};
+use crate::SparkError;
+use arrow::array::{
+ Array, ArrayRef, AsArray, BooleanArray, BooleanBufferBuilder, ListArray,
MapArray, StructArray,
+};
+use arrow::buffer::{BooleanBuffer, NullBuffer, OffsetBuffer};
+use arrow::compute::filter;
+use arrow::compute::kernels::zip::zip;
+use arrow::datatypes::{DataType, FieldRef, Float32Type, Float64Type};
+use datafusion::common::config::MapKeyDedupPolicy;
+use datafusion::common::{exec_err, internal_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;
+
+/// How a map builder compares and stores a `FLOAT` or `DOUBLE` key. Spark's
`ArrayBasedMapBuilder`
+/// finds duplicates in a `HashMap` of boxed keys and, from Spark 4.0,
normalizes a float key
+/// before it looks it up, unless
`spark.sql.legacy.disableMapKeyNormalization` is set. The serde
+/// picks the rule for the session, and each rule has its own function name.
+#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash)]
+pub enum MapFloatKeys {
+ /// The equality of a boxed `java.lang.Double` or `java.lang.Float`, which
compares
+ /// `doubleToLongBits` or `floatToIntBits`: every NaN is one key, while
`-0.0` and `0.0` are
+ /// two. A map keeps each key as it first occurred. Spark 3.4 and 3.5, and
Spark 4.0+ with
+ /// `spark.sql.legacy.disableMapKeyNormalization`.
+ #[default]
+ Boxed,
+ /// Spark 4.0+ compares and stores the key `NormalizeFloatingNumbers`
gives: `-0.0` becomes
+ /// `0.0`, and every NaN the canonical NaN. `map_from_arrays` stores the
keys it was given
+ /// when none of them repeats, because `ArrayBasedMapBuilder.from` then
returns its input.
+ Normalized,
+}
+
+/// Spark-compatible `map_from_arrays(keys, values)`.
+#[derive(Debug, Default, PartialEq, Eq, Hash)]
+pub struct SparkMapFromArrays {
+ inner: DataFusionMapFromArrays,
+ float_keys: MapFloatKeys,
+}
+
+impl SparkMapFromArrays {
+ pub fn new(float_keys: MapFloatKeys) -> Self {
+ Self {
+ inner: DataFusionMapFromArrays::default(),
+ float_keys,
+ }
+ }
+}
+
+impl ScalarUDFImpl for SparkMapFromArrays {
+ fn name(&self) -> &str {
+ match self.float_keys {
+ MapFloatKeys::Boxed => self.inner.name(),
+ MapFloatKeys::Normalized => "map_from_arrays_normalized_keys",
+ }
+ }
+
+ fn signature(&self) -> &Signature {
+ self.inner.signature()
+ }
+
+ fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
+ self.inner.return_type(arg_types)
+ }
+
+ fn return_field_from_args(&self, args: ReturnFieldArgs) ->
Result<FieldRef> {
+ self.inner.return_field_from_args(args)
+ }
+
+ fn invoke_with_args(&self, args: ScalarFunctionArgs) ->
Result<ColumnarValue> {
+ invoke_list_builder(
+ &self.inner,
+ args,
+ KeyLayout::Arrays,
+ self.float_keys,
+ |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,
+ float_keys: MapFloatKeys,
+}
+
+impl SparkMapFromEntries {
+ pub fn new(float_keys: MapFloatKeys) -> Self {
+ Self {
+ inner: DataFusionMapFromEntries::default(),
+ float_keys,
+ }
+ }
+}
+
+impl ScalarUDFImpl for SparkMapFromEntries {
+ fn name(&self) -> &str {
+ match self.float_keys {
+ MapFloatKeys::Boxed => self.inner.name(),
+ MapFloatKeys::Normalized => "map_from_entries_normalized_keys",
+ }
+ }
+
+ fn signature(&self) -> &Signature {
+ self.inner.signature()
+ }
+
+ fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
+ self.inner.return_type(arg_types)
+ }
+
+ fn return_field_from_args(&self, args: ReturnFieldArgs) ->
Result<FieldRef> {
+ self.inner.return_field_from_args(args)
+ }
+
+ fn invoke_with_args(&self, args: ScalarFunctionArgs) ->
Result<ColumnarValue> {
+ invoke_list_builder(
+ &self.inner,
+ args,
+ KeyLayout::Entries,
+ self.float_keys,
+ |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<DataType> {
+ self.inner.return_type(arg_types)
+ }
+
+ fn return_field_from_args(&self, args: ReturnFieldArgs) ->
Result<FieldRef> {
+ self.inner.return_field_from_args(args)
+ }
+
+ fn invoke_with_args(&self, args: ScalarFunctionArgs) ->
Result<ColumnarValue> {
+ // 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,
+ layout: KeyLayout,
+ float_keys: MapFloatKeys,
+ validate: impl FnOnce(&[ColumnarValue], bool) -> Result<()>,
+) -> Result<ColumnarValue> {
+ // 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;
+ let float_keys = FloatKeyRewrite::prepare(&mut args.args, layout,
float_keys)?;
+ let result = validate(&args.args, last_value_wins).and_then(|()| {
+ inner
+ .invoke_with_args(args)
+ .map_err(|error| as_spark_error(error, DuplicateKeyFormat::Bare))
+ });
+ let result = match &float_keys {
+ None => result?,
+ Some(keys) => keys.restore(result.map_err(|error|
keys.restate_duplicate(error))?)?,
+ };
+ match (all_scalar, result) {
+ (true, ColumnarValue::Array(array)) => Ok(ColumnarValue::Scalar(
+ ScalarValue::try_from_array(&array, 0)?,
+ )),
+ (_, result) => Ok(result),
+ }
+}
+
+/// Where a builder finds its keys among its arguments.
+#[derive(Clone, Copy, PartialEq, Eq)]
+enum KeyLayout {
+ /// `map_from_arrays(keys, values)`: the elements of the first list.
+ Arrays,
+ /// `map_from_entries(entries)`: the first field of the list's structs.
+ Entries,
+}
+
+/// The `FLOAT` or `DOUBLE` keys of one call. The kernel compares keys by
their bits, so it is
+/// handed them as Spark compares them, with every NaN made canonical and,
under
+/// [`MapFloatKeys::Normalized`], `-0.0` made `0.0`. Its result then gets back
the keys Spark
+/// stores.
+struct FloatKeyRewrite {
+ layout: KeyLayout,
+ rule: MapFloatKeys,
+ /// The flat keys as the call passed them.
+ original: ArrayRef,
+ /// The flat keys as Spark compares them, when that changes the bits of
any of them.
+ compared: Option<ArrayRef>,
+ /// Each row's range of flat keys.
+ offsets: OffsetBuffer<i32>,
+ /// The rows Spark builds a map for. Every other row is a NULL map, whose
keys it never sees.
+ built: BooleanBuffer,
+}
+
+impl FloatKeyRewrite {
+ /// Hands `args` the keys as Spark compares them, when the keys are
`FLOAT` or `DOUBLE`. Returns
+ /// `None` for any other key type.
+ fn prepare(
+ args: &mut [ColumnarValue],
+ layout: KeyLayout,
+ rule: MapFloatKeys,
+ ) -> Result<Option<Self>> {
+ let Some((original, offsets, built)) = float_keys_of(args, layout)
else {
+ return Ok(None);
+ };
+ let compared = match rule {
+ MapFloatKeys::Boxed => canonicalize_nans(&original),
+ MapFloatKeys::Normalized => normalize_floats(&original),
+ };
+ // Keys that are not NaN, or not `-0.0` under normalization, compare
the same either way,
+ // and then the kernel's own result is already Spark's. `ArrayData`
compares primitive
+ // values byte for byte, so this sees their bits.
+ let compared = if original.to_data() == compared.to_data() {
Review Comment:
For every float key type, this allocates a full canonicalized copy of the
keys on every batch and then compares it, even when no key changes. The PR
description puts the common-case cost at one pass and a byte comparison, but
the allocation comes on top of that. Could we first scan for a non-canonical
NaN, or a `-0.0` under `Normalized`, and only build `compared` when one turns
up? That keeps the common case to a single read-only pass. It would also help
to add a `Float64` key case to `benches/map_from_arrays.rs`, so we can see the
overhead next to the `Int32` and `Utf8` cases.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]