Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 64 additions & 0 deletions cpp/src/arrow/compute/expression_test.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -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<std::string>{"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")}),
Expand Down
16 changes: 16 additions & 0 deletions cpp/src/arrow/compute/kernel.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -519,6 +519,22 @@ std::shared_ptr<MatchConstraint> DecimalsHaveSameScale() {
return instance;
}

std::shared_ptr<MatchConstraint> AllTypesAreIdenticalFrom(size_t first_type_index) {
return MatchConstraint::Make(
[first_type_index](const std::vector<TypeHolder>& 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<MatchConstraint> AllTypesAreIdentical() {
static auto instance = AllTypesAreIdenticalFrom(/*first_type_index=*/0);
return instance;
}

// ----------------------------------------------------------------------
// KernelSignature

Expand Down
7 changes: 7 additions & 0 deletions cpp/src/arrow/compute/kernel.h
Original file line numberDiff line numberDiff line change
Expand Up@@ -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<MatchConstraint> DecimalsHaveSameScale();

/// \brief Constraint that all input types are identical.
ARROW_EXPORT std::shared_ptr<MatchConstraint> AllTypesAreIdentical();

/// \brief Constraint that all input types starting at first_type_index are identical.
ARROW_EXPORT std::shared_ptr<MatchConstraint> 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
Expand Down
17 changes: 17 additions & 0 deletions cpp/src/arrow/compute/kernel_test.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -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

Expand Down
25 changes: 8 additions & 17 deletions cpp/src/arrow/compute/kernels/scalar_if_else.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -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<MatchConstraint> AllValueTypesMatchConstraint() {
static auto constraint =
MatchConstraint::Make([](const std::vector<TypeHolder>& 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
Expand DownExpand Up@@ -2793,9 +2781,10 @@ void AddNestedCaseWhenKernels(const std::shared_ptr<CaseWhenFunction>& scalar_fu
}

void AddCoalesceKernel(const std::shared_ptr<ScalarFunction>& scalar_function,
detail::GetTypeId get_id, ArrayKernelExec exec) {
detail::GetTypeId get_id, ArrayKernelExec exec,
std::shared_ptr<MatchConstraint> 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;
Expand DownExpand Up@@ -2911,7 +2900,7 @@ void RegisterScalarIfElse(FunctionRegistry* registry) {
{
auto func = std::make_shared<CaseWhenFunction>(
"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);
Expand All@@ -2938,8 +2927,10 @@ void RegisterScalarIfElse(FunctionRegistry* registry) {
AddPrimitiveCoalesceKernels(func, {boolean(), null(), float16()});
AddCoalesceKernel(func, Type::FIXED_SIZE_BINARY,
CoalesceFunctor<FixedSizeBinaryType>::Exec);
AddCoalesceKernel(func, Type::DECIMAL128, CoalesceFunctor<FixedSizeBinaryType>::Exec);
AddCoalesceKernel(func, Type::DECIMAL256, CoalesceFunctor<FixedSizeBinaryType>::Exec);
AddCoalesceKernel(func, Type::DECIMAL128, CoalesceFunctor<FixedSizeBinaryType>::Exec,
AllTypesAreIdentical());
AddCoalesceKernel(func, Type::DECIMAL256, CoalesceFunctor<FixedSizeBinaryType>::Exec,
AllTypesAreIdentical());
for (const auto& ty : BaseBinaryTypes()) {
AddCoalesceKernel(func, ty, GenerateTypeAgnosticVarBinaryBase<CoalesceFunctor>(ty));
}
Expand Down
33 changes: 33 additions & 0 deletions cpp/src/arrow/compute/kernels/scalar_if_else_test.cc
Original file line numberDiff line numberDiff line change
Expand Up@@ -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)},

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

What about e.g. {decimal256(4, 1), decimal128(3, 2)}?

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good catch. I added this case in both argument orders. The common type is decimal256(5, 2), and the binding tests verify that both inputs are cast to it.

{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)},
Expand All@@ -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 <typename Type>
class TestChooseNumeric : public ::testing::Test {};
template <typename Type>
Expand Down
Loading