diff --git a/cpp/src/arrow/compute/expression_test.cc b/cpp/src/arrow/compute/expression_test.cc index 5e1f3c093ee..b4ae405b35a 100644 --- a/cpp/src/arrow/compute/expression_test.cc +++ b/cpp/src/arrow/compute/expression_test.cc @@ -938,6 +938,70 @@ 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)), 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)), + 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); + 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) { + 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/kernel.cc b/cpp/src/arrow/compute/kernel.cc index addbb29edd2..dda1a5f8bda 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 0d4f9d6ff43..239a03a86bb 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 9317ae7a42d..5aad1effc82 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 1510dd9fc83..1d8e6c1f62e 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 @@ -2793,9 +2781,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; @@ -2911,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); @@ -2938,8 +2927,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, + AllTypesAreIdentical()); + AddCoalesceKernel(func, Type::DECIMAL256, CoalesceFunctor::Exec, + 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 a1ef82383e2..e05a1f081be 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,26 @@ 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", {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)}, @@ -3710,6 +3728,21 @@ 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)}); + CheckDispatchExactFails("coalesce", {decimal256(4, 1), decimal128(3, 2)}); + CheckDispatchExactFails("coalesce", {decimal128(3, 2), decimal256(4, 1)}); +} + template class TestChooseNumeric : public ::testing::Test {}; template