This is an automated email from the ASF dual-hosted git repository.

exmy pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/gluten.git


The following commit(s) were added to refs/heads/main by this push:
     new 3fff963906 [GLUTEN-13064][CH] Fix lambda capture type mismatch in 
array higher-order functions (#13065)
3fff963906 is described below

commit 3fff96390688a2659a8e9948405d809cd1f590d4
Author: exmy <[email protected]>
AuthorDate: Mon Sep 21 14:40:57 2026 +0800

    [GLUTEN-13064][CH] Fix lambda capture type mismatch in array higher-order 
functions (#13065)
    
    * [GLUTEN-13064][CH] Fix lambda capture type mismatch in array higher-order 
functions
    
    CH declares split results as Array(Nullable(String)) while Spark infers 
non-nullable lambda arguments, so native function capture rejects the array 
element column. Align array element types with the lambda argument types for 
filter, transform with index, aggregate and zip_with, and add a regression test.
---
 .../execution/GlutenFunctionValidateSuite.scala    | 75 ++++++++++++++++++++++
 .../arrayHighOrderFunctions.cpp                    | 62 ++++++++++++------
 2 files changed, 118 insertions(+), 19 deletions(-)

diff --git 
a/backends-clickhouse/src/test/scala/org/apache/gluten/execution/GlutenFunctionValidateSuite.scala
 
b/backends-clickhouse/src/test/scala/org/apache/gluten/execution/GlutenFunctionValidateSuite.scala
index 6ee8dbd598..3ec0a70c3a 100644
--- 
a/backends-clickhouse/src/test/scala/org/apache/gluten/execution/GlutenFunctionValidateSuite.scala
+++ 
b/backends-clickhouse/src/test/scala/org/apache/gluten/execution/GlutenFunctionValidateSuite.scala
@@ -869,6 +869,81 @@ class GlutenFunctionValidateSuite extends 
GlutenClickHouseWholeStageTransformerS
     }
   }
 
+  test("array functions with lambda on nullable element array") {
+    withTable("tb_split_array", "tb_null_element_array") {
+      sql("create table tb_split_array(s string) using parquet")
+      sql("""
+            |insert into tb_split_array values
+            |('a_1,b_2'), ('b_1,c_2'), ('a_3'), ('a,,b'), (null)
+            |""".stripMargin)
+
+      sql("create table tb_null_element_array(a array<string>) using parquet")
+      sql("""
+            |insert into tb_null_element_array values
+            |(array('a', null)), (array(null)), (array()), (null)
+            |""".stripMargin)
+
+      // The CH backend declares split's result as Array(Nullable(String)) 
while Spark infers the
+      // lambda argument type as String, so the array element type must be 
aligned with the lambda
+      // argument type to avoid an incompatible type exception in native 
function capture.
+      val filter_sql =
+        """
+          |select filter(split(s, ','), x -> split(x, '_')[0] = 'a')
+          |from tb_split_array
+          |""".stripMargin
+      runQueryAndCompare(filter_sql)(checkGlutenPlan[ProjectExecTransformer])
+
+      // The filter path with an index argument is covered by the same 
alignment.
+      val filter_with_index_sql =
+        """
+          |select filter(split(s, ','), (x, i) -> i = 0 and x is not null)
+          |from tb_split_array
+          |""".stripMargin
+      
runQueryAndCompare(filter_with_index_sql)(checkGlutenPlan[ProjectExecTransformer])
+
+      val transform_sql =
+        """
+          |select transform(split(s, ','), (x, i) -> concat(x, cast(i as 
string)))
+          |from tb_split_array
+          |""".stripMargin
+      
runQueryAndCompare(transform_sql)(checkGlutenPlan[ProjectExecTransformer])
+
+      val aggregate_sql =
+        """
+          |select aggregate(split(s, ','), '', (acc, x) -> concat(acc, x))
+          |from tb_split_array
+          |""".stripMargin
+      
runQueryAndCompare(aggregate_sql)(checkGlutenPlan[ProjectExecTransformer])
+
+      val zip_with_sql =
+        """
+          |select zip_with(split(s, ','), split(s, ','), (x, y) -> concat(x, 
y))
+          |from tb_split_array
+          |""".stripMargin
+      runQueryAndCompare(zip_with_sql)(checkGlutenPlan[ProjectExecTransformer])
+
+      // Aligning the element type may narrow Array(Nullable(String)) to 
Array(String). Spark
+      // declares split's elements as non nullable, so the alignment must 
neither produce nor lose
+      // NULL elements, and empty string elements must be kept as empty 
strings.
+      val narrow_element_type_sql =
+        """
+          |select filter(split(s, ','), x -> x is null),
+          |       filter(split(s, ','), x -> x = '')
+          |from tb_split_array
+          |""".stripMargin
+      
runQueryAndCompare(narrow_element_type_sql)(checkGlutenPlan[ProjectExecTransformer])
+
+      // When the element type is nullable, NULL elements must survive the 
alignment.
+      val null_element_sql =
+        """
+          |select filter(a, x -> x is null),
+          |       transform(a, x -> x)
+          |from tb_null_element_array
+          |""".stripMargin
+      
runQueryAndCompare(null_element_sql)(checkGlutenPlan[ProjectExecTransformer])
+    }
+  }
+
   test("array aggregate with nested struct and nulls") {
     withTable("tb_array_complex") {
       sql("create table tb_array_complex(items array<struct<v:int, w:double>>) 
using parquet")
diff --git 
a/cpp-ch/local-engine/Parser/scalar_function_parser/arrayHighOrderFunctions.cpp 
b/cpp-ch/local-engine/Parser/scalar_function_parser/arrayHighOrderFunctions.cpp
index 278aeab226..9ba89765c2 100644
--- 
a/cpp-ch/local-engine/Parser/scalar_function_parser/arrayHighOrderFunctions.cpp
+++ 
b/cpp-ch/local-engine/Parser/scalar_function_parser/arrayHighOrderFunctions.cpp
@@ -37,6 +37,21 @@ namespace local_engine
 {
 using namespace DB;
 
+/// Align the element type of an array with the corresponding lambda argument 
type and keep the
+/// nullability of the array itself. ClickHouse's function capture requires an 
appended array
+/// element column to have exactly the same type as the lambda argument, 
otherwise it throws
+/// "Cannot capture column ... incompatible type". The array element types 
declared by Spark and by
+/// the CH backend may differ, e.g. `split` returns Array(Nullable(String)) in 
CH while Spark
+/// infers the lambda argument type as String.
+static const DB::ActionsDAG::Node * alignArrayElementType(
+    DB::ActionsDAG & actions_dag, const DB::ActionsDAG::Node * array_node, 
const DB::DataTypePtr & element_type, DB::ContextPtr context)
+{
+    DataTypePtr dst_array_type = std::make_shared<DataTypeArray>(element_type);
+    if (array_node->result_type->isNullable())
+        dst_array_type = std::make_shared<DataTypeNullable>(dst_array_type);
+    return ActionsDAGUtil::convertNodeTypeIfNeeded(actions_dag, array_node, 
dst_array_type, context);
+}
+
 class FunctionParserArrayFilter : public FunctionParser
 {
 public:
@@ -55,9 +70,17 @@ public:
     parse(const substrait::Expression_ScalarFunction & substrait_func, 
DB::ActionsDAG & actions_dag) const override
     {
         auto ch_func_name = getCHFunctionName(substrait_func);
+        auto lambda_args = collectLambdaArguments(parser_context, 
substrait_func.arguments()[1].value().scalar_function());
         auto parsed_args = parseFunctionArguments(substrait_func, actions_dag);
         assert(parsed_args.size() == 2);
-        if (collectLambdaArguments(parser_context, 
substrait_func.arguments()[1].value().scalar_function()).size() == 1)
+
+        /// Convert Array(T) to Array(U) if needed, Array(T) is the type of 
the first argument of filter,
+        /// U is the first argument type of the lambda function. In some cases 
Array(T) is not equal to
+        /// Array(U), e.g. CH's split returns Array(Nullable(String)) while 
the lambda argument type is
+        /// String. The difference of both types will result in runtime 
exceptions in function capture.
+        parsed_args[0] = alignArrayElementType(actions_dag, parsed_args[0], 
lambda_args.front().type, getContext());
+
+        if (lambda_args.size() == 1)
             return toFunctionNode(actions_dag, ch_func_name, {parsed_args[1], 
parsed_args[0]});
 
         /// filter with index argument.
@@ -98,11 +121,7 @@ public:
             /// U is the argument type of lambda function. In some cases 
Array(T) is not equal to Array(U).
             /// e.g. in the second query of 
https://github.com/apache/gluten/issues/6561, T is String, and U is 
Nullable(String)
             /// The difference of both types will result in runtime exceptions 
in function capture.
-            const auto & src_array_type = parsed_args[0]->result_type;
-            DataTypePtr dst_array_type = 
std::make_shared<DataTypeArray>(lambda_args.front().type);
-            if (src_array_type->isNullable())
-                dst_array_type = 
std::make_shared<DataTypeNullable>(dst_array_type);
-            const auto * dst_array_arg = 
ActionsDAGUtil::convertNodeTypeIfNeeded(actions_dag, parsed_args[0], 
dst_array_type, getContext());
+            const auto * dst_array_arg = alignArrayElementType(actions_dag, 
parsed_args[0], lambda_args.front().type, getContext());
             return toFunctionNode(actions_dag, ch_func_name, {parsed_args[1], 
dst_array_arg});
         }
 
@@ -114,6 +133,10 @@ public:
             actions_dag,
             "range",
             {addColumnToActionsDAG(actions_dag, 
std::make_shared<DataTypeInt32>(), 0), range_end_node});
+
+        /// Convert the array element type to the lambda argument type as 
well, see the comment above.
+        parsed_args[0] = alignArrayElementType(actions_dag, parsed_args[0], 
lambda_args.front().type, getContext());
+
         return toFunctionNode(actions_dag, ch_func_name, {parsed_args[1], 
parsed_args[0], index_array_node});
     }
 };
@@ -164,17 +187,13 @@ public:
         }
 
         /// Align array element type with merge lambda argument type.
-        const auto & merge_element_type = merge_arg_types.back();
-        const auto & src_array_type = parsed_args[0]->result_type;
-        DataTypePtr dst_array_type = 
std::make_shared<DataTypeArray>(merge_element_type);
-        if (src_array_type->isNullable())
-            dst_array_type = 
std::make_shared<DataTypeNullable>(dst_array_type);
-        const auto * array_col_node = 
ActionsDAGUtil::convertNodeTypeIfNeeded(actions_dag, parsed_args[0], 
dst_array_type, getContext());
+        const auto * array_col_node = alignArrayElementType(actions_dag, 
parsed_args[0], merge_arg_types.back(), getContext());
 
         /// arrayFold cannot accept nullable(array)
         if (parsed_args[0]->result_type->isNullable())
         {
-            array_col_node = toFunctionNode(actions_dag, "assumeNotNull", 
{parsed_args[0]});
+            /// Use the converted node, otherwise the element type alignment 
above will be lost.
+            array_col_node = toFunctionNode(actions_dag, "assumeNotNull", 
{array_col_node});
         }
         const auto * func_node = parsed_args.size() == 4
             ? toFunctionNode(actions_dag, ch_func_name, {parsed_args[2], 
array_col_node, parsed_args[1], parsed_args[3]})
@@ -223,11 +242,7 @@ public:
         /// In case lambda argument types are T, and the array has type 
Array(Nullable(T)) or Nullable(Array(Nullable(T))).
         /// We need to convert the array type to Array(T) or 
Nullable(Array(T)) to match the lambda argument types, otherwise it will cause 
runtime exceptions
         /// in function capture.
-        const auto & src_array_type = parsed_args[0]->result_type;
-        DataTypePtr dst_array_type = 
std::make_shared<DataTypeArray>(lambda_args.front().type);
-        if (src_array_type->isNullable())
-            dst_array_type = 
std::make_shared<DataTypeNullable>(dst_array_type);
-        parsed_args[0] = ActionsDAGUtil::convertNodeTypeIfNeeded(actions_dag, 
parsed_args[0], dst_array_type, getContext());
+        parsed_args[0] = alignArrayElementType(actions_dag, parsed_args[0], 
lambda_args.front().type, getContext());
 
         return toFunctionNode(actions_dag, ch_func_name, {parsed_args[1], 
parsed_args[0]});
     }
@@ -254,7 +269,16 @@ public:
         if (lambda_args.size() != 2)
             throw DB::Exception(DB::ErrorCodes::SIZES_OF_COLUMNS_DOESNT_MATCH, 
"The lambda function in zip_with must have two arguments");
 
-        const auto * array_zip_unaligned = toFunctionNode(actions_dag, 
"arrayZipUnaligned", {parsed_args[0], parsed_args[1]});
+        /// Convert Array(T) to Array(U) if needed for both array arguments, 
Array(T) is the element
+        /// type of an array argument of zip_with and U is the type of the 
corresponding lambda
+        /// argument. The difference of both types will result in runtime 
exceptions in function capture.
+        auto lambda_arg_it = lambda_args.begin();
+        DB::ActionsDAG::NodeRawConstPtrs arrays;
+        arrays.reserve(2);
+        for (size_t i = 0; i < 2; ++i, ++lambda_arg_it)
+            arrays.emplace_back(alignArrayElementType(actions_dag, 
parsed_args[i], lambda_arg_it->type, getContext()));
+
+        const auto * array_zip_unaligned = toFunctionNode(actions_dag, 
"arrayZipUnaligned", arrays);
         const auto * array_map = toFunctionNode(actions_dag, "arrayMap", 
{parsed_args[2], array_zip_unaligned});
         return convertNodeTypeIfNeeded(substrait_func, array_map, actions_dag);
     }


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to