diff --git a/native/spark-expr/src/comet_scalar_funcs.rs b/native/spark-expr/src/comet_scalar_funcs.rs index b5820144ea6..768616087e0 100644 --- a/native/spark-expr/src/comet_scalar_funcs.rs +++ b/native/spark-expr/src/comet_scalar_funcs.rs @@ -17,7 +17,7 @@ use crate::hash_funcs::*; use crate::json_funcs::JsonArrayLength; -use crate::map_funcs::spark_map_sort; +use crate::map_funcs::{spark_map_sort, SparkMapFromArrays}; use crate::math_funcs::abs::abs; use crate::math_funcs::checked_arithmetic::{checked_add, checked_div, checked_mul, checked_sub}; use crate::math_funcs::log::spark_log; @@ -253,6 +253,9 @@ pub fn create_comet_physical_fun_with_eval_mode( let func = Arc::new(crate::string_funcs::spark_get_json_object); make_comet_scalar_udf!("get_json_object", func, without data_type) } + "map" => Ok(Arc::new(ScalarUDF::new_from_impl( + SparkMapFromArrays::default(), + ))), "map_sort" => { let func = Arc::new(spark_map_sort); make_comet_scalar_udf!("spark_map_sort", func, without data_type) diff --git a/native/spark-expr/src/map_funcs/map_from_arrays.rs b/native/spark-expr/src/map_funcs/map_from_arrays.rs new file mode 100644 index 00000000000..284d5ad093b --- /dev/null +++ b/native/spark-expr/src/map_funcs/map_from_arrays.rs @@ -0,0 +1,324 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use arrow::array::{ArrayRef, AsArray}; +use arrow::datatypes::DataType; +use datafusion::common::{exec_err, utils::take_function_args, Result}; +use datafusion::functions_nested::map::MapFunc; +use datafusion::logical_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature}; + +/// Checks row boundaries before the DataFusion `map` used by CometMapFromArrays. +#[derive(Debug, Default, PartialEq, Eq, Hash)] +pub(crate) struct SparkMapFromArrays { + inner: MapFunc, +} + +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 invoke_with_args(&self, mut args: ScalarFunctionArgs) -> Result { + let has_array = args + .args + .iter() + .any(|arg| matches!(arg, ColumnarValue::Array(_))); + if has_array { + // MapFunc checks the flattened lengths, which can match even when individual + // rows have different lengths. Expand scalars to validate both mixed and batched + // operands, retaining the upstream scalar-only path and constructor behavior. + let number_rows = args.number_rows; + let arrays = args + .args + .into_iter() + .map(|arg| arg.into_array_of_size(number_rows)) + .collect::>>()?; + let [keys, values] = take_function_args(self.name(), arrays.as_slice())?; + validate_list_lengths(keys, values)?; + args.args = arrays.into_iter().map(ColumnarValue::Array).collect(); + } + self.inner.invoke_with_args(args) + } +} + +fn validate_list_lengths(keys: &ArrayRef, values: &ArrayRef) -> Result<()> { + // The upstream array path uses key offsets for both flattened children without checking + // each row's lengths. Do not let batched operands silently pair across map rows. + // CometMapFromArrays supplies the null-intolerant CASE guard around the constructor. + for row in 0..keys.len() { + if keys.is_valid(row) + && values.is_valid(row) + && list_length(keys, row)? != list_length(values, row)? + { + return exec_err!("map requires key and value lists to have the same length"); + } + } + Ok(()) +} + +fn list_length(array: &ArrayRef, row: usize) -> Result { + match array.data_type() { + DataType::List(_) => Ok(i64::from(array.as_list::().value_length(row))), + DataType::LargeList(_) => Ok(array.as_list::().value_length(row)), + DataType::FixedSizeList(_, length) => Ok(i64::from(*length)), + data_type => exec_err!("Expected List, LargeList, or FixedSizeList, got {data_type:?}"), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::{new_empty_array, Int32Array, ListArray, NullArray}; + use arrow::datatypes::{Field, Int32Type}; + use datafusion::common::{utils::SingleRowListArrayBuilder, ScalarValue}; + use datafusion::config::ConfigOptions; + use std::sync::Arc; + + fn scalar_list(values: ArrayRef) -> ColumnarValue { + ColumnarValue::Scalar(SingleRowListArrayBuilder::new(values).build_list_scalar()) + } + + fn lists(rows: Vec>>>) -> ArrayRef { + Arc::new(ListArray::from_iter_primitive::(rows)) + } + + fn invoke( + udf: &dyn ScalarUDFImpl, + args: Vec, + number_rows: usize, + ) -> Result { + let arg_types = args + .iter() + .map(ColumnarValue::data_type) + .collect::>(); + let return_field = Arc::new(Field::new("map", udf.return_type(&arg_types)?, true)); + let arg_fields = arg_types + .into_iter() + .map(|data_type| Arc::new(Field::new("arg", data_type, true))) + .collect(); + udf.invoke_with_args(ScalarFunctionArgs { + args, + arg_fields, + number_rows, + return_field, + config_options: Arc::new(ConfigOptions::default()), + }) + } + + fn assert_same_result(actual: ColumnarValue, expected: ColumnarValue, rows: usize) { + assert_eq!( + matches!(&actual, ColumnarValue::Scalar(_)), + matches!(&expected, ColumnarValue::Scalar(_)) + ); + assert_eq!( + actual.into_array_of_size(rows).unwrap().to_data(), + expected.into_array_of_size(rows).unwrap().to_data() + ); + } + + #[test] + fn mixed_inputs_broadcast_empty_and_nonempty_lists() { + let cases: Vec<(ArrayRef, ArrayRef)> = vec![ + ( + new_empty_array(&DataType::Null), + new_empty_array(&DataType::Null), + ), + ( + new_empty_array(&DataType::Int32), + new_empty_array(&DataType::Int32), + ), + ( + Arc::new(Int32Array::from(vec![1, 2])), + Arc::new(Int32Array::from(vec![Some(10), None])), + ), + ( + Arc::new(Int32Array::from(vec![1, 2])), + Arc::new(NullArray::new(2)), + ), + ]; + for (keys, values) in cases { + let keys = scalar_list(keys); + let values = scalar_list(values); + for rows in [0, 1, 3] { + for scalar_keys in [false, true] { + let key_array = keys.to_array_of_size(rows).unwrap(); + let value_array = values.to_array_of_size(rows).unwrap(); + let args = if scalar_keys { + vec![keys.clone(), ColumnarValue::Array(Arc::clone(&value_array))] + } else { + vec![ColumnarValue::Array(Arc::clone(&key_array)), values.clone()] + }; + let actual = invoke(&SparkMapFromArrays::default(), args, rows).unwrap(); + let expected = invoke( + &MapFunc::default(), + vec![ + ColumnarValue::Array(key_array), + ColumnarValue::Array(value_array), + ], + rows, + ) + .unwrap(); + assert_same_result(actual, expected, rows); + } + } + } + } + + #[test] + fn homogeneous_inputs_preserve_upstream_paths() { + let scalar_args = vec![ + scalar_list(Arc::new(Int32Array::from(vec![1]))), + scalar_list(Arc::new(Int32Array::from(vec![Some(10)]))), + ]; + let array_args = vec![ + ColumnarValue::Array(lists(vec![ + Some(vec![Some(1)]), + None, + Some(vec![Some(2), Some(3)]), + ])), + ColumnarValue::Array(lists(vec![ + Some(vec![Some(10)]), + Some(vec![]), + Some(vec![None, Some(30)]), + ])), + ]; + for args in [scalar_args, array_args] { + let actual = invoke(&SparkMapFromArrays::default(), args.clone(), 3).unwrap(); + let expected = invoke(&MapFunc::default(), args, 3).unwrap(); + assert_same_result(actual, expected, 3); + } + } + + #[test] + fn batched_inputs_reject_per_row_length_mismatch() { + // Both flattened children have four elements after broadcasting, so comparing only + // total lengths would allow values from one row to leak into the next map. + let batch = ColumnarValue::Array(lists(vec![ + Some(vec![Some(1)]), + Some(vec![Some(2), Some(3), Some(4)]), + ])); + let scalar = scalar_list(Arc::new(Int32Array::from(vec![10, 20]))); + let array = ColumnarValue::Array(scalar.to_array_of_size(2).unwrap()); + for args in [ + vec![batch.clone(), scalar.clone()], + vec![scalar, batch.clone()], + vec![batch.clone(), array.clone()], + vec![array, batch], + ] { + let err = invoke(&SparkMapFromArrays::default(), args, 2).unwrap_err(); + assert!(err.to_string().contains("same length"), "{err}"); + } + } + + #[test] + fn sliced_inputs_check_only_visible_rows() { + let keys = lists(vec![ + Some(vec![Some(0)]), + Some(vec![Some(1), Some(2)]), + Some(vec![Some(3), Some(4)]), + ]); + let values = lists(vec![ + Some(vec![Some(0), Some(0)]), + Some(vec![Some(10), None]), + Some(vec![Some(30), Some(40)]), + ]); + // The excluded first rows have different lengths and leave different starting offsets. + let actual = invoke( + &SparkMapFromArrays::default(), + vec![ + ColumnarValue::Array(keys.slice(1, 2)), + ColumnarValue::Array(values.slice(1, 2)), + ], + 2, + ) + .unwrap() + .into_array_of_size(2) + .unwrap(); + let map = actual.as_map(); + assert_eq!(map.value_offsets(), &[0, 2, 4]); + assert_eq!( + map.values().as_primitive::(), + &Int32Array::from(vec![Some(10), None, Some(30), Some(40)]) + ); + } + + #[test] + fn batched_inputs_reject_wrong_batch_length() { + let batch = ColumnarValue::Array(lists(vec![Some(vec![Some(1)])])); + let scalar = scalar_list(Arc::new(Int32Array::from(vec![10]))); + let array = ColumnarValue::Array(scalar.to_array_of_size(2).unwrap()); + for args in [ + vec![batch.clone(), scalar.clone()], + vec![scalar, batch.clone()], + vec![batch.clone(), array.clone()], + vec![array, batch.clone()], + vec![batch.clone(), batch], + ] { + let err = invoke(&SparkMapFromArrays::default(), args, 2).unwrap_err(); + assert!(err.to_string().contains("expected length 2"), "{err}"); + } + } + + #[test] + fn mixed_inputs_preserve_null_maps() { + let keys = lists(vec![None, Some(vec![Some(1)]), None]); + let values = scalar_list(Arc::new(Int32Array::from(vec![Some(10)]))); + let expected = invoke( + &MapFunc::default(), + vec![ + ColumnarValue::Array(Arc::clone(&keys)), + ColumnarValue::Array(values.to_array_of_size(3).unwrap()), + ], + 3, + ) + .unwrap(); + let actual = invoke( + &SparkMapFromArrays::default(), + vec![ColumnarValue::Array(keys), values], + 3, + ) + .unwrap(); + assert_same_result(actual, expected, 3); + + // Spark's null-intolerant CASE skips construction when either list is null. The + // length check must not introduce an error for offsets hidden by a null list row. + let nulls = lists(vec![None]); + let nonempty = lists(vec![Some(vec![Some(1)])]); + validate_list_lengths(&nulls, &nonempty).unwrap(); + validate_list_lengths(&nonempty, &nulls).unwrap(); + + let null_keys = + ColumnarValue::Scalar(ScalarValue::try_from_array(nulls.as_ref(), 0).unwrap()); + let actual = invoke( + &SparkMapFromArrays::default(), + vec![null_keys, ColumnarValue::Array(nonempty)], + 1, + ) + .unwrap() + .into_array_of_size(1) + .unwrap(); + assert!(actual.is_null(0)); + } +} diff --git a/native/spark-expr/src/map_funcs/mod.rs b/native/spark-expr/src/map_funcs/mod.rs index 7288b847a83..67c247da871 100644 --- a/native/spark-expr/src/map_funcs/mod.rs +++ b/native/spark-expr/src/map_funcs/mod.rs @@ -15,5 +15,7 @@ // specific language governing permissions and limitations // under the License. +mod map_from_arrays; mod map_sort; +pub(crate) use map_from_arrays::SparkMapFromArrays; pub use map_sort::spark_map_sort; 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..f9d3e2dbc00 100644 --- a/spark/src/main/scala/org/apache/comet/serde/maps.scala +++ b/spark/src/main/scala/org/apache/comet/serde/maps.scala @@ -24,7 +24,7 @@ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ 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 /** @@ -177,16 +177,25 @@ object CometMapFromArrays extends CometExpressionSerde[MapFromArrays] { val valueType = expr.right.dataType.asInstanceOf[ArrayType].elementType val returnType = MapType(keyType = keyType, valueType = valueType) for { - andBinaryExprProto <- createAndBinaryExpr(expr, inputs, binding) + keysNotNullExprProto <- exprToProtoInternal(IsNotNull(expr.left), inputs, binding) + valuesNotNullExprProto <- exprToProtoInternal(IsNotNull(expr.right), inputs, binding) mapFromArraysExprProto <- scalarFunctionExprToProto("map", keysExpr, valuesExpr) nullLiteralExprProto <- exprToProtoInternal(Literal(null, returnType), inputs, binding) } yield { - val caseWhenExprProto = ExprOuterClass.CaseWhen + // Spark skips the values expression when keys are null. Nested CASE guards preserve + // this evaluation order; a native AND may evaluate both operands for the whole batch. + val valuesCaseWhenExprProto = ExprOuterClass.CaseWhen .newBuilder() - .addWhen(andBinaryExprProto) + .addWhen(valuesNotNullExprProto) .addThen(mapFromArraysExprProto) .setElseExpr(nullLiteralExprProto) .build() + val caseWhenExprProto = ExprOuterClass.CaseWhen + .newBuilder() + .addWhen(keysNotNullExprProto) + .addThen(ExprOuterClass.Expr.newBuilder().setCaseWhen(valuesCaseWhenExprProto).build()) + .setElseExpr(nullLiteralExprProto) + .build() ExprOuterClass.Expr .newBuilder() .setCaseWhen(caseWhenExprProto) @@ -194,18 +203,6 @@ object CometMapFromArrays extends CometExpressionSerde[MapFromArrays] { } } - 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 diff --git a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala index f4a559b872b..f5d21150992 100644 --- a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala @@ -126,6 +126,94 @@ class CometMapExpressionSuite extends CometTestBase { } } + for (codegenEnabled <- Seq("false", "true")) { + test(s"map_from_arrays short-circuits null keys (codegen=$codegenEnabled)") { + withSQLConf( + SQLConf.ANSI_ENABLED.key -> "true", + SQLConf.USE_V1_SOURCE_LIST.key -> "parquet", + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> codegenEnabled) { + withTable("map_null_keys") { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + // Keep null and non-null keys in one batch. Only the null-key row divides by zero. + spark + .range(0, 3, 1, 1) + .selectExpr("CAST(id AS INT) AS k") + .write + .format("parquet") + .saveAsTable("map_null_keys") + } + val query = """SELECT map_from_arrays( + | CASE WHEN k = 0 THEN CAST(NULL AS ARRAY) ELSE array(1) END, + | array(1 / k)) + |FROM map_null_keys""".stripMargin + val plan = sql(query).queryExecution.executedPlan + assert(new ExtendedExplainInfo().getNativeExpressions(plan).contains("map_from_arrays")) + checkSparkAnswerAndOperator(sql(query)) + } + } + } + + test(s"map_from_arrays rejects unequal batched row lengths (codegen=$codegenEnabled)") { + withSQLConf( + SQLConf.USE_V1_SOURCE_LIST.key -> "parquet", + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> codegenEnabled) { + withTable("map_unequal_lengths") { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + // Each operand has four flattened elements, but row lengths are [1, 3] and [2, 2]. + spark + .range(0, 2, 1, 1) + .selectExpr("CAST(id AS INT) AS k") + .write + .format("parquet") + .saveAsTable("map_unequal_lengths") + } + val uneven = "CASE WHEN k = 0 THEN array(1) ELSE array(2, 3, 4) END" + for ((keys, values) <- Seq( + (uneven, "array(k, k + 10)"), + (uneven, "array(10, 20)"), + ("array(10, 20)", uneven))) { + val query = s"SELECT map_from_arrays($keys, $values) FROM map_unequal_lengths" + val plan = sql(query).queryExecution.executedPlan + assert( + new ExtendedExplainInfo().getNativeExpressions(plan).contains("map_from_arrays")) + val (sparkError, cometError) = checkSparkAnswerMaybeThrows(sql(query)) + assert(sparkError.exists(_.getMessage.contains("same length"))) + assert(cometError.exists(_.getMessage.contains("same length"))) + } + } + } + } + + test(s"map_from_arrays broadcasts scalar operands (codegen=$codegenEnabled)") { + withSQLConf( + SQLConf.USE_V1_SOURCE_LIST.key -> "parquet", + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> codegenEnabled) { + withTable("map_mixed_inputs") { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(0, 3, 1, 1) + .selectExpr("CAST(id AS INT) AS k") + .write + .format("parquet") + .saveAsTable("map_mixed_inputs") + } + val query = """SELECT + | map_from_arrays(array(10, 20), array(k, CAST(NULL AS INT))), + | map_from_arrays(array(k, k + 10), array(100, 200)), + | map_from_arrays( + | CASE WHEN k = 0 THEN CAST(NULL AS ARRAY) ELSE array(k) END, + | array(100)), + | map_from_arrays(array(100), + | CASE WHEN k = 0 THEN CAST(NULL AS ARRAY) ELSE array(k) END) + |FROM map_mixed_inputs""".stripMargin + val plan = sql(query).queryExecution.executedPlan + assert(new ExtendedExplainInfo().getNativeExpressions(plan).contains("map_from_arrays")) + checkSparkAnswerAndOperator(sql(query)) + } + } + } + } + test("size with map input") { withTempDir { dir => withTempView("t1") {