diff --git a/cpp/src/arrow/compute/kernels/scalar_nested.cc b/cpp/src/arrow/compute/kernels/scalar_nested.cc index e9c65aff1ce1..8990ec5b463e 100644 --- a/cpp/src/arrow/compute/kernels/scalar_nested.cc +++ b/cpp/src/arrow/compute/kernels/scalar_nested.cc @@ -20,6 +20,7 @@ #include #include "arrow/array/array_base.h" #include "arrow/array/builder_nested.h" +#include "arrow/array/builder_primitive.h" #include "arrow/compute/api_scalar.h" #include "arrow/compute/kernels/common_internal.h" #include "arrow/compute/registry_internal.h" @@ -28,8 +29,11 @@ #include "arrow/util/bit_block_counter.h" #include "arrow/util/bit_util.h" #include "arrow/util/bitmap_generate.h" +#include "arrow/util/bitmap_ops.h" +#include "arrow/util/float16.h" #include "arrow/util/logging_internal.h" #include "arrow/util/string.h" +#include "arrow/util/ubsan.h" #include "arrow/util/unreachable.h" namespace arrow { @@ -537,6 +541,233 @@ const FunctionDoc list_element_doc( "is emitted. Null values emit a null in the output."), {"lists", "index"}); +bool IsNaN(const Scalar& value) { + switch (value.type->id()) { + case Type::HALF_FLOAT: + return util::Float16::FromBits(checked_cast(value).value) + .is_nan(); + case Type::FLOAT: + return std::isnan(checked_cast(value).value); + case Type::DOUBLE: + return std::isnan(checked_cast(value).value); + default: + return false; + } +} + +// Null values are matched without "equal", but must be comparable all the same +Status CheckNullValueType(KernelContext* ctx, const DataType& values_type, + const DataType& value_type) { + if (value_type.id() == Type::NA || value_type.Equals(values_type)) { + return Status::OK(); + } + ARROW_ASSIGN_OR_RAISE(auto equal, + ctx->exec_context()->func_registry()->GetFunction("equal")); + std::vector types{&values_type, &value_type}; + return equal->DispatchBest(&types).status(); +} + +// Match `values` with a scalar `value`, or element-wise with an array of them. Like +// "is_in", and unlike "equal", a null value matches null values and a NaN value +// matches NaN values. +Result ListValuesMatch(KernelContext* ctx, const Datum& values, + const Datum& value) { + ExecContext* exec_ctx = ctx->exec_context(); + if (value.null_count() == value.length()) { + RETURN_NOT_OK(CheckNullValueType(ctx, *values.type(), *value.type())); + return CallFunction("is_null", {values}, exec_ctx); + } + if (value.is_scalar()) { + if (is_floating(values.type()->id()) && IsNaN(*value.scalar())) { + return CallFunction("is_nan", {values}, exec_ctx); + } + return CallFunction("equal", {values, value}, exec_ctx); + } + ARROW_ASSIGN_OR_RAISE(Datum match, CallFunction("equal", {values, value}, exec_ctx)); + ARROW_ASSIGN_OR_RAISE(Datum values_null, CallFunction("is_null", {values}, exec_ctx)); + ARROW_ASSIGN_OR_RAISE(Datum value_null, CallFunction("is_null", {value}, exec_ctx)); + ARROW_ASSIGN_OR_RAISE(Datum both_null, + CallFunction("and", {values_null, value_null}, exec_ctx)); + ARROW_ASSIGN_OR_RAISE(match, CallFunction("or_kleene", {match, both_null}, exec_ctx)); + if (is_floating(values.type()->id()) && is_floating(value.type()->id())) { + ARROW_ASSIGN_OR_RAISE(Datum values_nan, CallFunction("is_nan", {values}, exec_ctx)); + ARROW_ASSIGN_OR_RAISE(Datum value_nan, CallFunction("is_nan", {value}, exec_ctx)); + ARROW_ASSIGN_OR_RAISE(Datum both_nan, + CallFunction("and_kleene", {values_nan, value_nan}, exec_ctx)); + ARROW_ASSIGN_OR_RAISE(match, CallFunction("or_kleene", {match, both_nan}, exec_ctx)); + } + return match; +} + +// Returns a function giving the number of child values of list i, with none for null +// list views +template +auto GetListLengths(const ArraySpan& list) { + if constexpr (std::is_same_v) { + const int64_t width = checked_cast(*list.type).list_size(); + return [width](int64_t) { return width; }; + } else if constexpr (is_list_view_type::value) { + const auto* sizes = list.GetValues(2); + return + [&list, sizes](int64_t i) -> int64_t { return list.IsValid(i) ? sizes[i] : 0; }; + } else { + const auto* offsets = list.GetValues(1); + return [offsets](int64_t i) -> int64_t { return offsets[i + 1] - offsets[i]; }; + } +} + +// Returns the child values of all lists, one list after the other +template +Result> GetListValues(KernelContext* ctx, + const ArraySpan& list) { + if constexpr (std::is_same_v) { + const int64_t width = GetListLengths(list)(0); + return list.child_data[0].ToArrayData()->Slice(list.offset * width, + list.length * width); + } else if constexpr (is_list_view_type::value) { + // List views may reference child values in any order + typename TypeTraits::ArrayType list_view(list.ToArrayData()); + ARROW_ASSIGN_OR_RAISE(auto values, list_view.Flatten(ctx->memory_pool())); + return values->data(); + } else { + const auto* offsets = list.GetValues(1); + return list.child_data[0].ToArrayData()->Slice(offsets[0], + offsets[list.length] - offsets[0]); + } +} + +// Whether any of the `length` bits from `offset` is set in `bits`, a bitmap of `size` +// bytes +bool AnySet(const uint8_t* bits, int64_t size, int64_t offset, int64_t length) { + const int64_t byte_offset = offset / 8; + const int64_t bit_offset = offset % 8; + if (bit_offset + length <= 64 && byte_offset + 8 <= size) { + // Short ranges fit in a single word + const uint64_t word = + bit_util::FromLittleEndian(util::SafeLoadAs(bits + byte_offset)); + const uint64_t mask = length == 64 ? ~uint64_t{0} : (uint64_t{1} << length) - 1; + return ((word >> bit_offset) & mask) != 0; + } + arrow::internal::BitBlockCounter counter(bits, offset, length); + for (int64_t position = 0; position < length;) { + const auto block = counter.NextWord(); + if (block.popcount > 0) { + return true; + } + position += block.length; + } + return false; +} + +// Repeat the value of each list for each of its child values +template +Result RepeatListValues(KernelContext* ctx, const ArraySpan& list, + int64_t values_length, const ArraySpan& value) { + const auto list_length = GetListLengths(list); + Int64Builder indices(ctx->memory_pool()); + RETURN_NOT_OK(indices.Reserve(values_length)); + for (int64_t i = 0; i < list.length; ++i) { + const int64_t length = list_length(i); + for (int64_t j = 0; j < length; ++j) { + indices.UnsafeAppend(i); + } + } + ARROW_ASSIGN_OR_RAISE(auto repeat_indices, indices.Finish()); + return CallFunction("take", {value.ToArrayData(), repeat_indices}, ctx->exec_context()); +} + +// Match the child values of all lists with a single vectorized comparison, then check +// each list's range in the resulting bitmap. +template +Status ListContains(KernelContext* ctx, const ExecSpan& batch, ExecResult* out) { + if (batch.length == 0) { + return Status::OK(); + } + if (batch[0].is_scalar()) { + // A scalar list is only paired with an array of values here, so broadcast it + ARROW_ASSIGN_OR_RAISE(auto lists, MakeArrayFromScalar(*batch[0].scalar, batch.length, + ctx->memory_pool())); + ExecSpan array_batch = batch; + array_batch.values[0].SetArray(*lists->data()); + return ListContains(ctx, array_batch, out); + } + + const ArraySpan& list = batch[0].array; + ArraySpan* out_arr = out->array_span_mutable(); + + // Only null lists emit a null + if (list.MayHaveNulls()) { + arrow::internal::CopyBitmap(list.buffers[0].data, list.offset, list.length, + out_arr->buffers[0].data, out_arr->offset); + } else { + bit_util::SetBitsTo(out_arr->buffers[0].data, out_arr->offset, out_arr->length, true); + } + out_arr->null_count = list.null_count; + + ARROW_ASSIGN_OR_RAISE(auto values, GetListValues(ctx, list)); + Datum value; + if (batch[1].is_scalar()) { + value = batch[1].scalar->GetSharedPtr(); + } else if (batch[1].array.length == 1) { + // Also covers scalar inputs, which are promoted to arrays + ARROW_ASSIGN_OR_RAISE(value, batch[1].array.ToArray()->GetScalar(0)); + } else { + ARROW_ASSIGN_OR_RAISE( + value, RepeatListValues(ctx, list, values->length, batch[1].array)); + } + ARROW_ASSIGN_OR_RAISE(Datum match, ListValuesMatch(ctx, values, value)); + + // Null matches never count + const ArrayData& match_data = *match.array(); + std::shared_ptr matches = match_data.buffers[1]; + int64_t matches_offset = match_data.offset; + if (match_data.MayHaveNulls()) { + ARROW_ASSIGN_OR_RAISE( + matches, arrow::internal::BitmapAnd( + ctx->memory_pool(), matches->data(), match_data.offset, + match_data.buffers[0]->data(), match_data.offset, match_data.length, + /*out_offset=*/0)); + matches_offset = 0; + } + + // The output bits of null lists don't matter, so they are not special-cased + const auto list_length = GetListLengths(list); + int64_t start = matches_offset; + int64_t i = 0; + arrow::internal::GenerateBitsUnrolled( + out_arr->buffers[1].data, out_arr->offset, out_arr->length, [&] { + const int64_t length = list_length(i++); + const bool found = AnySet(matches->data(), matches->size(), start, length); + start += length; + return found; + }); + return Status::OK(); +} + +void AddListContainsKernels(ScalarFunction* func) { + auto add_kernel = [&](Type::type list_type_id, ArrayKernelExec exec) { + ScalarKernel kernel({InputType(list_type_id), InputType::Any()}, boolean(), exec); + // A null value is searched for rather than propagated + kernel.null_handling = NullHandling::COMPUTED_PREALLOCATE; + DCHECK_OK(func->AddKernel(std::move(kernel))); + }; + add_kernel(Type::LIST, ListContains); + add_kernel(Type::LARGE_LIST, ListContains); + add_kernel(Type::LIST_VIEW, ListContains); + add_kernel(Type::LARGE_LIST_VIEW, ListContains); + add_kernel(Type::FIXED_SIZE_LIST, ListContains); +} + +const FunctionDoc list_contains_doc( + "Check whether lists contain a given value", + ("`lists` must have a list-like type and `value` must be comparable\n" + "with the list value type.\n" + "For each list in `lists`, true is emitted if any of its values is equal\n" + "to the corresponding `value`, false otherwise. A null `value` matches\n" + "null list values, and a NaN `value` matches NaN list values; otherwise\n" + "null list values never match. Null lists emit a null in the output."), + {"lists", "value"}); + struct StructFieldFunctor { static Status Exec(KernelContext* ctx, const ExecSpan& batch, ExecResult* out) { const auto& options = OptionsWrapper::Get(ctx); @@ -960,6 +1191,11 @@ void RegisterScalarNested(FunctionRegistry* registry) { AddListElementKernels(list_element.get()); DCHECK_OK(registry->AddFunction(std::move(list_element))); + auto list_contains = std::make_shared("list_contains", Arity::Binary(), + list_contains_doc); + AddListContainsKernels(list_contains.get()); + DCHECK_OK(registry->AddFunction(std::move(list_contains))); + auto list_slice = std::make_shared("list_slice", Arity::Unary(), list_slice_doc); AddListSliceKernels(list_slice.get()); diff --git a/cpp/src/arrow/compute/kernels/scalar_nested_test.cc b/cpp/src/arrow/compute/kernels/scalar_nested_test.cc index b5a68d12cb0c..e4d6d69085ed 100644 --- a/cpp/src/arrow/compute/kernels/scalar_nested_test.cc +++ b/cpp/src/arrow/compute/kernels/scalar_nested_test.cc @@ -128,6 +128,237 @@ TEST(TestScalarNested, ListElementInvalid) { Raises(StatusCode::Invalid)); } +void CheckListContains(Datum lists, Datum value, const std::string& expected) { + CheckScalar("list_contains", {std::move(lists), std::move(value)}, + ArrayFromJSON(boolean(), expected)); +} + +TEST(TestScalarNested, ListContains) { + auto sample = "[[7, 5, 81], [6, null, 4, 7, 8], [], [5], null, [null]]"; + for (auto ty : NumericTypes()) { + for (auto list_type : + {list(ty), large_list(ty), list_view(ty), large_list_view(ty)}) { + auto input = ArrayFromJSON(list_type, sample); + CheckListContains(input, ScalarFromJSON(ty, "5"), + "[true, false, false, true, null, false]"); + CheckListContains(input, ScalarFromJSON(ty, "7"), + "[true, true, false, false, null, false]"); + CheckListContains(input, ScalarFromJSON(ty, "null"), + "[false, true, false, false, null, true]"); + CheckListContains(ArrayFromJSON(list_type, "[]"), ScalarFromJSON(ty, "5"), "[]"); + } + } + + auto input = + ArrayFromJSON(list(utf8()), R"([["a", "b"], ["a", "c"], ["b", "c", "d"]])"); + CheckListContains(input, ScalarFromJSON(utf8(), R"("a")"), "[true, true, false]"); + CheckListContains(input->Slice(1), ScalarFromJSON(utf8(), R"("b")"), "[false, true]"); + + // All-null lists contain no non-null value + CheckListContains(ArrayFromJSON(list(int64()), "[[null, null]]"), + ScalarFromJSON(int64(), "1"), "[false]"); +} + +TEST(TestScalarNested, ListContainsLongLists) { + // Matches and nulls beyond the first bitmap words, at unaligned offsets + std::string values = "["; + for (int i = 0; i < 300; ++i) { + values += i == 0 ? "" : ", "; + values += i == 100 ? "null" : i == 290 ? "7" : "1"; + } + values += "]"; + ASSERT_OK_AND_ASSIGN(auto input, + ListArray::FromArrays(*ArrayFromJSON(int32(), "[0, 3, 150, 300]"), + *ArrayFromJSON(int32(), values))); + CheckListContains(input, ScalarFromJSON(int32(), "7"), "[false, false, true]"); + CheckListContains(input, ScalarFromJSON(int32(), "null"), "[false, true, false]"); + CheckListContains(input, ArrayFromJSON(int32(), "[7, 7, null]"), + "[false, false, false]"); + CheckListContains(input, ArrayFromJSON(int32(), "[1, null, 7]"), "[true, true, true]"); +} + +TEST(TestScalarNested, ListContainsManyShortLists) { + // Short lists at every bit alignment, with null lists and null values + std::string lists = "["; + std::string contains_5 = "["; + std::string contains_null = "["; + for (int i = 0; i < 200; ++i) { + const std::string sep = i == 0 ? "" : ", "; + const bool is_null = i % 11 == 0; + lists += sep + (is_null ? "null" + : i % 7 == 0 ? "[5, 1]" + : i % 3 == 0 ? "[null, 2, 3]" + : "[1]"); + contains_5 += sep + (is_null ? "null" : i % 7 == 0 ? "true" : "false"); + contains_null += sep + (is_null ? "null" + : i % 7 != 0 && i % 3 == 0 ? "true" + : "false"); + } + auto input = ArrayFromJSON(list(int32()), lists + "]"); + CheckListContains(input, ScalarFromJSON(int32(), "5"), contains_5 + "]"); + CheckListContains(input, ScalarFromJSON(int32(), "null"), contains_null + "]"); +} + +TEST(TestScalarNested, ListContainsNull) { + // A null value matches lists holding a null + auto input = ArrayFromJSON( + list(int32()), "[[1, null], [null], [1, 2], [1, 1], [2, 1, null, 1], [], null]"); + for (auto value : {ScalarFromJSON(int32(), "null"), MakeNullScalar(null()), + ScalarFromJSON(int64(), "null")}) { + CheckListContains(input, value, "[true, true, false, false, true, false, null]"); + } + CheckListContains(ArrayFromJSON(list(int32()), "[[]]"), ScalarFromJSON(int32(), "null"), + "[false]"); + + input = ArrayFromJSON(list(utf8()), R"([["x", null], [null], ["x", "y"], [], null])"); + CheckListContains(input, ScalarFromJSON(utf8(), "null"), + "[true, true, false, false, null]"); + + input = ArrayFromJSON(list(list(int32())), + "[[[1], null], [null], [[1], [null]], [], null]"); + CheckListContains(input, MakeNullScalar(list(int32())), + "[true, true, false, false, null]"); + + input = ArrayFromJSON(list(null()), "[[null], [], null]"); + CheckListContains(input, MakeNullScalar(null()), "[true, false, null]"); +} + +TEST(TestScalarNested, ListContainsNaN) { + // Unlike "equal", a NaN value matches NaN list values + for (auto ty : {float32(), float64()}) { + auto input = ArrayFromJSON(list(ty), "[[1.5, null], [NaN], [1.5, NaN], [], null]"); + for (auto value_ty : {float32(), float64()}) { + CheckListContains(input, ScalarFromJSON(value_ty, "NaN"), + "[false, true, true, false, null]"); + CheckListContains(input, ScalarFromJSON(value_ty, "1.5"), + "[true, false, true, false, null]"); + } + } + CheckListContains(ArrayFromJSON(list(int64()), "[[1], []]"), + ScalarFromJSON(float64(), "NaN"), "[false, false]"); + + auto input = + ArrayFromJSON(list(float16()), "[[1.5, null], [NaN], [1.5, NaN], [], null]"); + CheckListContains(input, ScalarFromJSON(float16(), "NaN"), + "[false, true, true, false, null]"); + CheckListContains(ArrayFromJSON(list(float32()), "[[1.5], [NaN]]"), + ScalarFromJSON(float16(), "NaN"), "[false, true]"); +} + +TEST(TestScalarNested, ListContainsValueTypes) { + auto check = [](const std::shared_ptr& ty, const std::string& x, + const std::string& y) { + auto input = ArrayFromJSON( + list(ty), "[[" + x + ", null], [" + y + "], [" + x + ", " + y + "], [], null]"); + CheckListContains(input, ScalarFromJSON(ty, x), "[true, false, true, false, null]"); + CheckListContains(input, ScalarFromJSON(ty, y), "[false, true, true, false, null]"); + }; + check(utf8(), R"("x")", R"("y")"); + check(large_binary(), R"("x")", R"("y")"); + check(boolean(), "true", "false"); + check(date32(), "18262", "18628"); + check(decimal128(38, 2), R"("1.50")", R"("2.50")"); + check(timestamp(TimeUnit::NANO), "1", "2"); +} + +TEST(TestScalarNested, ListContainsImplicitCast) { + auto input = ArrayFromJSON(list(int64()), + "[[2, 2, 3, null, null], null, [], [null], [1], [0, -1]]"); + CheckListContains(input, ScalarFromJSON(float64(), "2.0"), + "[true, null, false, false, false, false]"); + CheckListContains(input, ScalarFromJSON(float64(), "1.5"), + "[false, null, false, false, false, false]"); + + // No overflow of the list values + CheckListContains(ArrayFromJSON(list(int8()), "[[44, 1], null, [], [null]]"), + ScalarFromJSON(int64(), "300"), "[false, null, false, false]"); + CheckListContains(ArrayFromJSON(list(float64()), "[[1.0, 2.0], [3.0], null, []]"), + ScalarFromJSON(int64(), "1"), "[true, false, null, false]"); + CheckListContains(ArrayFromJSON(list(decimal128(38, 2)), R"([["1.50"], ["2.50"]])"), + ScalarFromJSON(decimal128(3, 2), R"("1.50")"), "[true, false]"); +} + +TEST(TestScalarNested, ListContainsChunked) { + auto array = ArrayFromJSON(list(int32()), "[[1, 2], [3], null, [], [2, null]]"); + auto input = std::make_shared( + ArrayVector{array->Slice(3), array->Slice(0, 0), array->Slice(1, 2)}); + ASSERT_OK_AND_ASSIGN( + Datum result, CallFunction("list_contains", {input, ScalarFromJSON(int32(), "2")})); + AssertDatumsEqual( + ChunkedArrayFromJSON(boolean(), {"[false, true]", "[]", "[false, null]"}), result, + /*verbose=*/true); +} + +TEST(TestScalarNested, ListContainsFixedSizeList) { + auto input = ArrayFromJSON(fixed_size_list(int32(), 2), + "[[1, 2], [3, null], null, [2, 4], [null, null]]"); + CheckListContains(input, ScalarFromJSON(int32(), "2"), + "[true, false, null, true, false]"); + CheckListContains(input->Slice(2), ScalarFromJSON(int32(), "4"), "[null, true, false]"); +} + +TEST(TestScalarNested, ListContainsListViewOutOfOrder) { + // Views: [3, 4], [1, 2], [2, 3], [4], [] + ASSERT_OK_AND_ASSIGN( + auto input, ListViewArray::FromArrays(*ArrayFromJSON(int32(), "[2, 0, 1, 3, 0]"), + *ArrayFromJSON(int32(), "[2, 2, 2, 1, 0]"), + *ArrayFromJSON(int32(), "[1, 2, 3, 4]"))); + CheckListContains(input, ScalarFromJSON(int32(), "3"), + "[true, false, true, false, false]"); + CheckListContains(input, ScalarFromJSON(int32(), "1"), + "[false, true, false, false, false]"); + CheckListContains(input, ArrayFromJSON(int32(), "[3, 1, 1, 4, 1]"), + "[true, true, false, true, false]"); + // Only the child values referenced by the (sliced) views are compared + CheckListContains(input->Slice(2), ScalarFromJSON(int32(), "3"), + "[true, false, false]"); + CheckListContains(input->Slice(3), ScalarFromJSON(int32(), "1"), "[false, false]"); + CheckListContains(input->Slice(4), ScalarFromJSON(int32(), "1"), "[false]"); +} + +TEST(TestScalarNested, ListContainsArrayOfValues) { + // Each list is searched for the value at the same index + for (auto list_type : {list(int32()), large_list(int32()), list_view(int32()), + large_list_view(int32())}) { + CheckListContains( + ArrayFromJSON(list_type, "[[1, 2], [3, null], [], null, [4, 5], [null], [6]]"), + ArrayFromJSON(int32(), "[2, null, 1, 3, 4, null, 7]"), + "[true, true, false, null, true, true, false]"); + } + CheckListContains(ArrayFromJSON(fixed_size_list(int32(), 2), + "[[1, 2], [3, null], null, [4, 5], [null, 6]]"), + ArrayFromJSON(int32(), "[2, null, 1, 5, 7]"), + "[true, true, null, true, false]"); + CheckListContains(ArrayFromJSON(list(int64()), "[[1, 2], [3], [4]]"), + ArrayFromJSON(float64(), "[2.0, 3.5, 4.0]"), "[true, false, true]"); + CheckListContains(ArrayFromJSON(list(float64()), "[[1.5, NaN], [NaN], [1.5], [null]]"), + ArrayFromJSON(float64(), "[NaN, 1.5, NaN, NaN]"), + "[true, false, false, false]"); + CheckListContains(ArrayFromJSON(list(int32()), "[[1, null], [2], []]"), + ArrayFromJSON(null(), "[null, null, null]"), "[true, false, false]"); + + // A scalar list is searched for each value + CheckListContains(ScalarFromJSON(list(int32()), "[1, null]"), + ArrayFromJSON(int32(), "[1, 2, null]"), "[true, false, true]"); + CheckListContains(ScalarFromJSON(fixed_size_list(int32(), 2), "[1, 2]"), + ArrayFromJSON(int32(), "[2, 3]"), "[true, false]"); + CheckListContains(ScalarFromJSON(list(int32()), "null"), + ArrayFromJSON(int32(), "[1, 2]"), "[null, null]"); +} + +TEST(TestScalarNested, ListContainsInvalid) { + auto input = ArrayFromJSON(list(int32()), "[[1, 2], [3]]"); + // Null values must be comparable too + for (Datum value : + {Datum(ArrayFromJSON(utf8(), R"(["a", "b"])")), + Datum(ScalarFromJSON(utf8(), R"("a")")), Datum(ScalarFromJSON(boolean(), "true")), + Datum(ScalarFromJSON(utf8(), "null")), + Datum(ArrayFromJSON(utf8(), "[null, null]"))}) { + EXPECT_THAT(CallFunction("list_contains", {input, value}), + Raises(StatusCode::NotImplemented)); + } +} + using VarLenListLikeTypeFactory = std::shared_ptr (*)(std::shared_ptr); static const VarLenListLikeTypeFactory kVarLenListTypeFactories[] = { diff --git a/docs/source/cpp/compute.rst b/docs/source/cpp/compute.rst index f24fea20a9d8..5e29107a9f1c 100644 --- a/docs/source/cpp/compute.rst +++ b/docs/source/cpp/compute.rst @@ -1924,6 +1924,8 @@ Structural transforms +---------------------+------------+-------------------------------------+------------------+------------------------------+--------+ | Function name | Arity | Input types | Output type | Options class | Notes | +=====================+============+=====================================+==================+==============================+========+ +| list_contains | Binary | List-like (Arg 0), Any (Arg 1) | Boolean | | \(7) | ++---------------------+------------+-------------------------------------+------------------+------------------------------+--------+ | list_element | Binary | List-like (Arg 0), Integral (Arg 1) | List value type | | \(1) | +---------------------+------------+-------------------------------------+------------------+------------------------------+--------+ | list_flatten | Unary | List-like | List value type | | \(2) | @@ -1982,6 +1984,12 @@ Structural transforms index *n* and the type code at index *n* is 2. * The indices ``2`` and ``7`` are invalid. +* \(7) Output is true for each list containing a value equal to the second + argument. If the second argument is an array, each list is searched for the + value at the same index. Unlike ``equal``, a null second argument matches + null list values and a NaN second argument matches NaN list values; + otherwise null list values never match. Null lists emit a null in the output. + .. _cpp-compute-vector-replace-functions: Replace functions diff --git a/docs/source/python/api/compute.rst b/docs/source/python/api/compute.rst index 6a4b04468d4c..f83fce2be84e 100644 --- a/docs/source/python/api/compute.rst +++ b/docs/source/python/api/compute.rst @@ -569,6 +569,7 @@ Structural Transforms fill_null fill_null_backward fill_null_forward + list_contains list_element list_flatten list_parent_indices diff --git a/python/pyarrow/tests/test_compute.py b/python/pyarrow/tests/test_compute.py index 797fbc220ec3..042456a12439 100644 --- a/python/pyarrow/tests/test_compute.py +++ b/python/pyarrow/tests/test_compute.py @@ -3899,6 +3899,45 @@ def test_list_element(): assert result.equals(expected) +@pytest.mark.parametrize("list_type", [ + pa.list_(pa.int64()), pa.large_list(pa.int64()), + pa.list_view(pa.int64()), pa.large_list_view(pa.int64()), +]) +def test_list_contains(list_type): + lists = pa.array([[1, 2, None], [3], [], None], list_type) + assert pc.list_contains(lists, 2).to_pylist() == [True, False, False, None] + assert pc.list_contains(lists, 2.0).to_pylist() == [True, False, False, None] + assert pc.list_contains(lists, 4).to_pylist() == [False, False, False, None] + # A null value matches null list values + assert pc.list_contains(lists, None).to_pylist() == [True, False, False, None] + + +def test_list_contains_nan(): + lists = pa.array([[1.5, None], [float("nan")], []]) + result = pc.list_contains(lists, float("nan")) + assert result.to_pylist() == [False, True, False] + + +def test_list_contains_fixed_size_list(): + lists = pa.array([["a", "b"], ["c", None], None], pa.list_(pa.string(), 2)) + assert pc.list_contains(lists, "c").to_pylist() == [False, True, None] + + +def test_list_contains_array(): + lists = pa.array([[1, 2], [3, None], [4], None]) + values = pa.array([2, None, 5, 1]) + result = pc.list_contains(lists, values) + assert result.to_pylist() == [True, True, False, None] + + +def test_list_contains_invalid(): + lists = pa.array([[1, 2], [3]]) + with pytest.raises(pa.ArrowNotImplementedError): + pc.list_contains(lists, "a") + with pytest.raises(pa.ArrowNotImplementedError): + pc.list_contains(lists, pa.array(["a", "b"])) + + def test_count_distinct(): samples = [datetime.datetime(year=y, month=1, day=1) for y in range(1992, 2092)] arr = pa.array(samples, pa.timestamp("ns"))