From 662a932d27727a0c846222b0255f3b9fae0a8612 Mon Sep 17 00:00:00 2001 From: Rossi Sun Date: Fri, 5 Jun 2026 11:52:06 -0700 Subject: [PATCH 1/2] GH-50869: [C++][Compute] Tighten coalesce exact dispatch for decimal varargs Signed-off-by: Rossi Sun --- cpp/src/arrow/compute/expression_test.cc | 56 +++++++++++++++++++ .../arrow/compute/kernels/scalar_if_else.cc | 25 +++++++-- .../compute/kernels/scalar_if_else_test.cc | 27 +++++++++ 3 files changed, 104 insertions(+), 4 deletions(-) diff --git a/cpp/src/arrow/compute/expression_test.cc b/cpp/src/arrow/compute/expression_test.cc index 5e1f3c093ee2..64f56cb922bc 100644 --- a/cpp/src/arrow/compute/expression_test.cc +++ b/cpp/src/arrow/compute/expression_test.cc @@ -938,6 +938,62 @@ TEST(Expression, BindWithImplicitCastsForCaseWhenOnDecimal) { /*bound_out=*/nullptr, *exciting_schema); } +TEST(Expression, BindWithImplicitCastsForCoalesceOnDecimal) { + auto exciting_schema = schema( + {field("dec128_3_2", decimal128(3, 2)), field("dec128_4_1", decimal128(4, 1)), + field("dec128_4_2", decimal128(4, 2)), field("dec128_4_3", decimal128(4, 3)), + field("dec256_3_2", decimal256(3, 2))}); + + ExpectBindsTo(call("coalesce", {field_ref("dec128_3_2"), field_ref("dec128_4_2")}), + call("coalesce", {cast(field_ref("dec128_3_2"), decimal128(4, 2)), + field_ref("dec128_4_2")}), + /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec128_4_2"), field_ref("dec128_3_2")}), + call("coalesce", {field_ref("dec128_4_2"), + cast(field_ref("dec128_3_2"), decimal128(4, 2))}), + /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec128_4_1"), field_ref("dec128_3_2")}), + call("coalesce", {cast(field_ref("dec128_4_1"), decimal128(5, 2)), + cast(field_ref("dec128_3_2"), decimal128(5, 2))}), + /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec128_3_2"), field_ref("dec128_4_1")}), + call("coalesce", {cast(field_ref("dec128_3_2"), decimal128(5, 2)), + cast(field_ref("dec128_4_1"), decimal128(5, 2))}), + /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec128_3_2"), field_ref("dec128_4_3")}), + call("coalesce", {cast(field_ref("dec128_3_2"), decimal128(4, 3)), + field_ref("dec128_4_3")}), + /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec128_4_3"), field_ref("dec128_3_2")}), + call("coalesce", {field_ref("dec128_4_3"), + cast(field_ref("dec128_3_2"), decimal128(4, 3))}), + /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec128_3_2"), field_ref("dec256_3_2")}), + call("coalesce", {cast(field_ref("dec128_3_2"), decimal256(3, 2)), + field_ref("dec256_3_2")}), + /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec256_3_2"), field_ref("dec128_3_2")}), + call("coalesce", {field_ref("dec256_3_2"), + cast(field_ref("dec128_3_2"), decimal256(3, 2))}), + /*bound_out=*/nullptr, *exciting_schema); +} + +TEST(Expression, ExecuteCoalesceOnMixedDecimalTypes) { + ASSERT_OK_AND_ASSIGN( + auto input, StructArray::Make( + ArrayVector{ArrayFromJSON(decimal128(3, 2), R"(["1.23", null])"), + ArrayFromJSON(decimal128(4, 3), R"([null, "2.345"])")}, + std::vector{"left", "right"})); + Schema input_schema(input->type()->fields()); + auto expr = call("coalesce", {field_ref("left"), field_ref("right")}); + + ASSERT_OK_AND_ASSIGN(expr, expr.Bind(input_schema)); + ASSERT_OK_AND_ASSIGN(auto actual, + ExecuteScalarExpression(expr, input_schema, Datum(input))); + + AssertDatumsEqual(actual, ArrayFromJSON(decimal128(4, 3), R"(["1.230", "2.345"])")); +} + TEST(Expression, BindNestedCall) { auto expr = add(field_ref("a"), call("subtract", {call("multiply", {field_ref("b"), field_ref("c")}), diff --git a/cpp/src/arrow/compute/kernels/scalar_if_else.cc b/cpp/src/arrow/compute/kernels/scalar_if_else.cc index 1510dd9fc83a..0193fd4f5d80 100644 --- a/cpp/src/arrow/compute/kernels/scalar_if_else.cc +++ b/cpp/src/arrow/compute/kernels/scalar_if_else.cc @@ -2035,6 +2035,20 @@ struct CoalesceFunction : ScalarFunction { if (auto kernel = DispatchExactImpl(this, *types)) return kernel; return arrow::compute::detail::NoMatchingKernel(this, *types); } + + static std::shared_ptr DecimalMatchConstraint() { + static auto constraint = + MatchConstraint::Make([](const std::vector& types) -> bool { + DCHECK_GE(types.size(), 1); + DCHECK(std::all_of(types.begin(), types.end(), [](const TypeHolder& type) { + return is_decimal(type.id()); + })); + return std::all_of( + types.begin() + 1, types.end(), + [&types](const TypeHolder& type) { return type == types[0]; }); + }); + return constraint; + } }; // Helper: copy from a source value into all null slots of the output @@ -2793,9 +2807,10 @@ void AddNestedCaseWhenKernels(const std::shared_ptr& scalar_fu } void AddCoalesceKernel(const std::shared_ptr& scalar_function, - detail::GetTypeId get_id, ArrayKernelExec exec) { + detail::GetTypeId get_id, ArrayKernelExec exec, + std::shared_ptr constraint = nullptr) { ScalarKernel kernel(KernelSignature::Make({InputType(get_id.id)}, FirstType, - /*is_varargs=*/true), + /*is_varargs=*/true, std::move(constraint)), exec); kernel.null_handling = NullHandling::COMPUTED_PREALLOCATE; kernel.mem_allocation = MemAllocation::PREALLOCATE; @@ -2938,8 +2953,10 @@ void RegisterScalarIfElse(FunctionRegistry* registry) { AddPrimitiveCoalesceKernels(func, {boolean(), null(), float16()}); AddCoalesceKernel(func, Type::FIXED_SIZE_BINARY, CoalesceFunctor::Exec); - AddCoalesceKernel(func, Type::DECIMAL128, CoalesceFunctor::Exec); - AddCoalesceKernel(func, Type::DECIMAL256, CoalesceFunctor::Exec); + AddCoalesceKernel(func, Type::DECIMAL128, CoalesceFunctor::Exec, + CoalesceFunction::DecimalMatchConstraint()); + AddCoalesceKernel(func, Type::DECIMAL256, CoalesceFunctor::Exec, + CoalesceFunction::DecimalMatchConstraint()); for (const auto& ty : BaseBinaryTypes()) { AddCoalesceKernel(func, ty, GenerateTypeAgnosticVarBinaryBase(ty)); } diff --git a/cpp/src/arrow/compute/kernels/scalar_if_else_test.cc b/cpp/src/arrow/compute/kernels/scalar_if_else_test.cc index a1ef82383e29..c9f0e7d7dc63 100644 --- a/cpp/src/arrow/compute/kernels/scalar_if_else_test.cc +++ b/cpp/src/arrow/compute/kernels/scalar_if_else_test.cc @@ -3693,8 +3693,22 @@ TEST(TestCoalesce, DispatchBest) { CheckDispatchBest("coalesce", {int32(), decimal128(3, 2)}, {decimal128(12, 2), decimal128(12, 2)}); CheckDispatchBest("coalesce", {float32(), decimal128(3, 2)}, {float64(), float64()}); + CheckDispatchBest("coalesce", {decimal128(3, 2), decimal128(4, 2)}, + {decimal128(4, 2), decimal128(4, 2)}); + CheckDispatchBest("coalesce", {decimal128(4, 2), decimal128(3, 2)}, + {decimal128(4, 2), decimal128(4, 2)}); + CheckDispatchBest("coalesce", {decimal128(4, 1), decimal128(3, 2)}, + {decimal128(5, 2), decimal128(5, 2)}); + CheckDispatchBest("coalesce", {decimal128(3, 2), decimal128(4, 1)}, + {decimal128(5, 2), decimal128(5, 2)}); + CheckDispatchBest("coalesce", {decimal128(3, 2), decimal128(4, 3)}, + {decimal128(4, 3), decimal128(4, 3)}); + CheckDispatchBest("coalesce", {decimal128(4, 3), decimal128(3, 2)}, + {decimal128(4, 3), decimal128(4, 3)}); CheckDispatchBest("coalesce", {decimal128(3, 2), decimal256(3, 2)}, {decimal256(3, 2), decimal256(3, 2)}); + CheckDispatchBest("coalesce", {decimal256(3, 2), decimal128(3, 2)}, + {decimal256(3, 2), decimal256(3, 2)}); CheckDispatchBest("coalesce", {timestamp(TimeUnit::SECOND), date32()}, {timestamp(TimeUnit::SECOND), timestamp(TimeUnit::SECOND)}); CheckDispatchBest("coalesce", {timestamp(TimeUnit::SECOND), timestamp(TimeUnit::MILLI)}, @@ -3710,6 +3724,19 @@ TEST(TestCoalesce, DispatchBest) { {large_binary(), large_binary()}); } +TEST(TestCoalesce, DispatchExact) { + CheckDispatchExact("coalesce", {decimal128(3, 2), decimal128(3, 2)}); + CheckDispatchExact("coalesce", {decimal256(3, 2), decimal256(3, 2)}); + CheckDispatchExactFails("coalesce", {decimal128(3, 2), decimal128(4, 2)}); + CheckDispatchExactFails("coalesce", {decimal128(4, 2), decimal128(3, 2)}); + CheckDispatchExactFails("coalesce", {decimal128(4, 1), decimal128(3, 2)}); + CheckDispatchExactFails("coalesce", {decimal128(3, 2), decimal128(4, 1)}); + CheckDispatchExactFails("coalesce", {decimal128(3, 2), decimal128(4, 3)}); + CheckDispatchExactFails("coalesce", {decimal128(4, 3), decimal128(3, 2)}); + CheckDispatchExactFails("coalesce", {decimal128(3, 2), decimal256(3, 2)}); + CheckDispatchExactFails("coalesce", {decimal256(3, 2), decimal128(3, 2)}); +} + template class TestChooseNumeric : public ::testing::Test {}; template From f051dd398ed92fcd9dd74438502a00f0e6ffbadb Mon Sep 17 00:00:00 2001 From: Rossi Sun Date: Wed, 2 Sep 2026 12:49:00 +0800 Subject: [PATCH 2/2] GH-50869: [C++][Compute] Address coalesce review feedback Signed-off-by: Rossi Sun --- cpp/src/arrow/compute/expression_test.cc | 10 +++++- cpp/src/arrow/compute/kernel.cc | 16 ++++++++++ cpp/src/arrow/compute/kernel.h | 7 ++++ cpp/src/arrow/compute/kernel_test.cc | 17 ++++++++++ .../arrow/compute/kernels/scalar_if_else.cc | 32 ++----------------- .../compute/kernels/scalar_if_else_test.cc | 6 ++++ 6 files changed, 58 insertions(+), 30 deletions(-) diff --git a/cpp/src/arrow/compute/expression_test.cc b/cpp/src/arrow/compute/expression_test.cc index 64f56cb922bc..b4ae405b35a9 100644 --- a/cpp/src/arrow/compute/expression_test.cc +++ b/cpp/src/arrow/compute/expression_test.cc @@ -942,7 +942,7 @@ TEST(Expression, BindWithImplicitCastsForCoalesceOnDecimal) { auto exciting_schema = schema( {field("dec128_3_2", decimal128(3, 2)), field("dec128_4_1", decimal128(4, 1)), field("dec128_4_2", decimal128(4, 2)), field("dec128_4_3", decimal128(4, 3)), - field("dec256_3_2", decimal256(3, 2))}); + field("dec256_3_2", decimal256(3, 2)), field("dec256_4_1", decimal256(4, 1))}); ExpectBindsTo(call("coalesce", {field_ref("dec128_3_2"), field_ref("dec128_4_2")}), call("coalesce", {cast(field_ref("dec128_3_2"), decimal128(4, 2)), @@ -976,6 +976,14 @@ TEST(Expression, BindWithImplicitCastsForCoalesceOnDecimal) { call("coalesce", {field_ref("dec256_3_2"), cast(field_ref("dec128_3_2"), decimal256(3, 2))}), /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec256_4_1"), field_ref("dec128_3_2")}), + call("coalesce", {cast(field_ref("dec256_4_1"), decimal256(5, 2)), + cast(field_ref("dec128_3_2"), decimal256(5, 2))}), + /*bound_out=*/nullptr, *exciting_schema); + ExpectBindsTo(call("coalesce", {field_ref("dec128_3_2"), field_ref("dec256_4_1")}), + call("coalesce", {cast(field_ref("dec128_3_2"), decimal256(5, 2)), + cast(field_ref("dec256_4_1"), decimal256(5, 2))}), + /*bound_out=*/nullptr, *exciting_schema); } TEST(Expression, ExecuteCoalesceOnMixedDecimalTypes) { diff --git a/cpp/src/arrow/compute/kernel.cc b/cpp/src/arrow/compute/kernel.cc index addbb29edd26..dda1a5f8bda7 100644 --- a/cpp/src/arrow/compute/kernel.cc +++ b/cpp/src/arrow/compute/kernel.cc @@ -519,6 +519,22 @@ std::shared_ptr DecimalsHaveSameScale() { return instance; } +std::shared_ptr AllTypesAreIdenticalFrom(size_t first_type_index) { + return MatchConstraint::Make( + [first_type_index](const std::vector& types) -> bool { + DCHECK_LT(first_type_index, types.size()); + return std::all_of(types.begin() + first_type_index + 1, types.end(), + [&types, first_type_index](const TypeHolder& type) { + return type == types[first_type_index]; + }); + }); +} + +std::shared_ptr AllTypesAreIdentical() { + static auto instance = AllTypesAreIdenticalFrom(/*first_type_index=*/0); + return instance; +} + // ---------------------------------------------------------------------- // KernelSignature diff --git a/cpp/src/arrow/compute/kernel.h b/cpp/src/arrow/compute/kernel.h index 0d4f9d6ff436..239a03a86bb6 100644 --- a/cpp/src/arrow/compute/kernel.h +++ b/cpp/src/arrow/compute/kernel.h @@ -365,6 +365,13 @@ class ARROW_EXPORT MatchConstraint { /// \brief Constraint that all input types are decimal types and have the same scale. ARROW_EXPORT std::shared_ptr DecimalsHaveSameScale(); +/// \brief Constraint that all input types are identical. +ARROW_EXPORT std::shared_ptr AllTypesAreIdentical(); + +/// \brief Constraint that all input types starting at first_type_index are identical. +ARROW_EXPORT std::shared_ptr AllTypesAreIdenticalFrom( + size_t first_type_index); + /// \brief Holds the input types, optional match constraint and output type of the kernel. /// /// VarArgs functions with minimum N arguments should pass up to N input types to be diff --git a/cpp/src/arrow/compute/kernel_test.cc b/cpp/src/arrow/compute/kernel_test.cc index 9317ae7a42d1..5aad1effc82e 100644 --- a/cpp/src/arrow/compute/kernel_test.cc +++ b/cpp/src/arrow/compute/kernel_test.cc @@ -341,6 +341,23 @@ TEST(MatchConstraint, DecimalsHaveSameScale) { decimal128(precision, scale + 1)})); } +TEST(MatchConstraint, AllTypesAreIdentical) { + auto c = AllTypesAreIdentical(); + constexpr int32_t precision = 12, scale = 2; + ASSERT_TRUE(c->Matches({int8()})); + ASSERT_TRUE(c->Matches({decimal128(precision, scale), decimal128(precision, scale), + decimal128(precision, scale)})); + ASSERT_FALSE( + c->Matches({decimal128(precision, scale), decimal128(precision + 1, scale)})); + ASSERT_FALSE( + c->Matches({decimal128(precision, scale), decimal128(precision, scale + 1)})); + ASSERT_FALSE(c->Matches({decimal128(precision, scale), decimal256(precision, scale)})); + + auto skip_first = AllTypesAreIdenticalFrom(/*first_type_index=*/1); + ASSERT_TRUE(skip_first->Matches({boolean(), utf8(), utf8()})); + ASSERT_FALSE(skip_first->Matches({boolean(), utf8(), binary()})); +} + // ---------------------------------------------------------------------- // KernelSignature diff --git a/cpp/src/arrow/compute/kernels/scalar_if_else.cc b/cpp/src/arrow/compute/kernels/scalar_if_else.cc index 0193fd4f5d80..1d8e6c1f62ef 100644 --- a/cpp/src/arrow/compute/kernels/scalar_if_else.cc +++ b/cpp/src/arrow/compute/kernels/scalar_if_else.cc @@ -1494,18 +1494,6 @@ struct CaseWhenFunction : ScalarFunction { if (auto kernel = DispatchExactImpl(this, *types)) return kernel; return arrow::compute::detail::NoMatchingKernel(this, *types); } - - // For case_when exact dispatch, all value arguments must have identical DataType. - static std::shared_ptr AllValueTypesMatchConstraint() { - static auto constraint = - MatchConstraint::Make([](const std::vector& types) -> bool { - DCHECK_GE(types.size(), 2); - return std::all_of( - types.begin() + 2, types.end(), - [&types](const TypeHolder& type) { return type == types[1]; }); - }); - return constraint; - } }; // Implement a 'case when' (SQL)/'select' (NumPy) function for any scalar conditions @@ -2035,20 +2023,6 @@ struct CoalesceFunction : ScalarFunction { if (auto kernel = DispatchExactImpl(this, *types)) return kernel; return arrow::compute::detail::NoMatchingKernel(this, *types); } - - static std::shared_ptr DecimalMatchConstraint() { - static auto constraint = - MatchConstraint::Make([](const std::vector& types) -> bool { - DCHECK_GE(types.size(), 1); - DCHECK(std::all_of(types.begin(), types.end(), [](const TypeHolder& type) { - return is_decimal(type.id()); - })); - return std::all_of( - types.begin() + 1, types.end(), - [&types](const TypeHolder& type) { return type == types[0]; }); - }); - return constraint; - } }; // Helper: copy from a source value into all null slots of the output @@ -2926,7 +2900,7 @@ void RegisterScalarIfElse(FunctionRegistry* registry) { { auto func = std::make_shared( "case_when", Arity::VarArgs(/*min_args=*/2), case_when_doc); - auto all_value_types_match = CaseWhenFunction::AllValueTypesMatchConstraint(); + auto all_value_types_match = AllTypesAreIdenticalFrom(/*first_type_index=*/1); AddPrimitiveCaseWhenKernels(func, NumericTypes(), all_value_types_match); AddPrimitiveCaseWhenKernels(func, TemporalTypes(), all_value_types_match); AddPrimitiveCaseWhenKernels(func, IntervalTypes(), all_value_types_match); @@ -2954,9 +2928,9 @@ void RegisterScalarIfElse(FunctionRegistry* registry) { AddCoalesceKernel(func, Type::FIXED_SIZE_BINARY, CoalesceFunctor::Exec); AddCoalesceKernel(func, Type::DECIMAL128, CoalesceFunctor::Exec, - CoalesceFunction::DecimalMatchConstraint()); + AllTypesAreIdentical()); AddCoalesceKernel(func, Type::DECIMAL256, CoalesceFunctor::Exec, - CoalesceFunction::DecimalMatchConstraint()); + AllTypesAreIdentical()); for (const auto& ty : BaseBinaryTypes()) { AddCoalesceKernel(func, ty, GenerateTypeAgnosticVarBinaryBase(ty)); } diff --git a/cpp/src/arrow/compute/kernels/scalar_if_else_test.cc b/cpp/src/arrow/compute/kernels/scalar_if_else_test.cc index c9f0e7d7dc63..e05a1f081beb 100644 --- a/cpp/src/arrow/compute/kernels/scalar_if_else_test.cc +++ b/cpp/src/arrow/compute/kernels/scalar_if_else_test.cc @@ -3709,6 +3709,10 @@ TEST(TestCoalesce, DispatchBest) { {decimal256(3, 2), decimal256(3, 2)}); CheckDispatchBest("coalesce", {decimal256(3, 2), decimal128(3, 2)}, {decimal256(3, 2), decimal256(3, 2)}); + CheckDispatchBest("coalesce", {decimal256(4, 1), decimal128(3, 2)}, + {decimal256(5, 2), decimal256(5, 2)}); + CheckDispatchBest("coalesce", {decimal128(3, 2), decimal256(4, 1)}, + {decimal256(5, 2), decimal256(5, 2)}); CheckDispatchBest("coalesce", {timestamp(TimeUnit::SECOND), date32()}, {timestamp(TimeUnit::SECOND), timestamp(TimeUnit::SECOND)}); CheckDispatchBest("coalesce", {timestamp(TimeUnit::SECOND), timestamp(TimeUnit::MILLI)}, @@ -3735,6 +3739,8 @@ TEST(TestCoalesce, DispatchExact) { CheckDispatchExactFails("coalesce", {decimal128(4, 3), decimal128(3, 2)}); CheckDispatchExactFails("coalesce", {decimal128(3, 2), decimal256(3, 2)}); CheckDispatchExactFails("coalesce", {decimal256(3, 2), decimal128(3, 2)}); + CheckDispatchExactFails("coalesce", {decimal256(4, 1), decimal128(3, 2)}); + CheckDispatchExactFails("coalesce", {decimal128(3, 2), decimal256(4, 1)}); } template