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]