Skip to content
Merged
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 number Diff line number Diff 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 number Diff line number Diff 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 number Diff line number Diff 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 number Diff line number Diff 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 number Diff line number Diff 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 Down Expand 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 Down Expand 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 number Diff line number Diff 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)},
Comment thread
pitrou marked this conversation as resolved.
{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