From 9a6c74f0cde6dea5a91cb3be360b3a4d32778bd7 Mon Sep 17 00:00:00 2001 From: Dana Jansens Date: Fri, 18 Apr 2025 10:17:48 -0400 Subject: [PATCH] Introduce FindIfOrNull() FindIfOrNone() and Contains() (#5322) `FindIfOrNull` returns a pointer to the element in the range if it's found, and nullptr otherwise. `FindIfOrNone` returns a copy of the element in the range if it's found, and `T::None` (for a range of elements of type `T`) otherwise. `Contains` returns a bool indicating whether the element in the range is found. These functions replace `llvm::find()` and `llvm::find_if()` when you want a single answer back instead of an iterator. This avoids the need to check against `end()`, allowing the return condition to be tested as a standard bool. We replace uses of `find()` and `find_if()` that did not require an iterator with these new helpers. Note that the return type of `FindIfOrNull` is a pointer since we can not write `optional`, which must be tested for null. If the null check is omitted, UB occurs and the resulting code may end up with an incorrect pointer (https://crbug.com/40153300) into the range (or elsewhere), rather than a null dereference. And this would be very confusing to debug. Hopefully debug builds and sanitizers keep this from being an issue we sink a bunch of time into debugging. --- common/BUILD | 22 +++++++ common/find.h | 92 +++++++++++++++++++++++++++++ common/find_test.cpp | 75 +++++++++++++++++++++++ testing/file_test/BUILD | 1 + testing/file_test/test_file.cpp | 12 ++-- toolchain/check/BUILD | 1 + toolchain/check/handle_function.cpp | 13 ++-- toolchain/parse/BUILD | 1 + toolchain/parse/extract.cpp | 3 +- toolchain/testing/BUILD | 1 + toolchain/testing/coverage_helper.h | 3 +- 11 files changed, 207 insertions(+), 17 deletions(-) create mode 100644 common/find.h create mode 100644 common/find_test.cpp diff --git a/common/BUILD b/common/BUILD index ed6056c32d50..b0b96fa8ea59 100644 --- a/common/BUILD +++ b/common/BUILD @@ -170,6 +170,28 @@ cc_test( ], ) +cc_library( + name = "find", + hdrs = ["find.h"], + deps = [ + ":check", + ":ostream", + ":raw_string_ostream", + "@llvm-project//llvm:Support", + ], +) + +cc_test( + name = "find_test", + size = "small", + srcs = ["find_test.cpp"], + deps = [ + ":find", + "//testing/base:gtest_main", + "@googletest//:gtest", + ], +) + cc_library( name = "hashing", srcs = ["hashing.cpp"], diff --git a/common/find.h b/common/find.h new file mode 100644 index 000000000000..d551cf1fab34 --- /dev/null +++ b/common/find.h @@ -0,0 +1,92 @@ +// Part of the Carbon Language project, under the Apache License v2.0 with LLVM +// Exceptions. See /LICENSE for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +#ifndef CARBON_COMMON_FIND_H_ +#define CARBON_COMMON_FIND_H_ + +#include +#include + +#include "llvm/ADT/STLExtras.h" + +namespace Carbon { + +namespace Internal { + +template +using RangePointerType = typename std::iterator_traits()))>::pointer; + +template +using RangeValueType = typename std::iterator_traits()))>::value_type; + +template +concept IsValidFindPredicate = + requires(const RangeValueType& elem, Pred pred) { + { pred(elem) } -> std::convertible_to; + }; + +template +concept IsComparable = requires(const A& a, const B& b) { + { a == b } -> std::convertible_to; +}; + +template +concept RangeValueHasNoneType = requires { + { RangeValueType::None } -> std::convertible_to>; +}; + +} // namespace Internal + +// Finds a value in the given `range` by testing the `predicate`. Returns a +// pointer to the value from the range on success, and nullptr if nothing is +// found. +// +// This is similar to `std::find_if()` but returns a pointer to the value +// instead of an iterator that must be tested against `end()`. +template + requires Internal::IsValidFindPredicate +constexpr auto FindIfOrNull(Range&& range, Pred predicate) + -> Internal::RangePointerType { + auto it = llvm::find_if(range, predicate); + if (it != range.end()) { + return std::addressof(*it); + } else { + return nullptr; + } +} + +// Finds a value in the given `range` by testing the `predicate` and returns a +// copy of it. If no match is found, returns `T::None` where the input range is +// over values of type `T`. +template + requires Internal::IsValidFindPredicate && + Internal::RangeValueHasNoneType && + std::copy_constructible> +constexpr auto FindIfOrNone(Range&& range, Pred predicate) + -> Internal::RangeValueType { + auto it = llvm::find_if(range, predicate); + if (it != range.end()) { + return *it; + } else { + return Internal::RangeValueType::None; + } +} + +// Finds a value in the given `range` by comparing to `query`. Returns a +// pointer to the value from the range on success, and nullptr if nothing is +// found. +// +// This is similar to `std::find_if()` but returns a pointer to the value +// instead of an iterator that must be tested against `end()`. +template > + requires Internal::IsComparable> +constexpr auto Contains(Range&& range, const Query& query) -> bool { + return llvm::find(range, query) != range.end(); +} + +} // namespace Carbon + +#endif // CARBON_COMMON_FIND_H_ diff --git a/common/find_test.cpp b/common/find_test.cpp new file mode 100644 index 000000000000..00de7a0fe5a0 --- /dev/null +++ b/common/find_test.cpp @@ -0,0 +1,75 @@ +// Part of the Carbon Language project, under the Apache License v2.0 with LLVM +// Exceptions. See /LICENSE for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +#include "common/find.h" + +#include + +#include + +namespace Carbon { +namespace { + +struct NoneType { + static const NoneType None; + int i; + + friend auto operator==(NoneType, NoneType) -> bool = default; +}; + +const NoneType NoneType::None = {.i = -1}; + +TEST(FindTest, ReturnType) { + const std::vector c; + std::vector m; + + auto pred = [](int) { return true; }; + static_assert(std::same_as); + static_assert(std::same_as); +} + +TEST(FindTest, FindIfOrNull) { + auto make_pred = [](int query) { + return [=](int elem) { return query == elem; }; + }; + + std::vector empty; + EXPECT_EQ(FindIfOrNull(empty, make_pred(0)), nullptr); + + std::vector range = {1, 2}; + EXPECT_EQ(FindIfOrNull(range, make_pred(0)), nullptr); + // NOLINTNEXTLINE(readability-container-data-pointer) + EXPECT_EQ(FindIfOrNull(range, make_pred(1)), &range[0]); + EXPECT_EQ(FindIfOrNull(range, make_pred(2)), &range[1]); + EXPECT_EQ(FindIfOrNull(range, make_pred(3)), nullptr); +} + +TEST(FindTest, FindIfOrNone) { + auto make_pred = [](NoneType query) { + return [=](NoneType elem) { return query == elem; }; + }; + + std::vector empty; + EXPECT_EQ(FindIfOrNone(empty, make_pred(NoneType{0})).i, -1); + + std::vector range = {NoneType{1}, NoneType{2}}; + EXPECT_EQ(FindIfOrNone(range, make_pred(NoneType{0})).i, -1); + EXPECT_EQ(FindIfOrNone(range, make_pred(NoneType{1})).i, 1); + EXPECT_EQ(FindIfOrNone(range, make_pred(NoneType{2})).i, 2); + EXPECT_EQ(FindIfOrNone(range, make_pred(NoneType{3})).i, -1); +} + +TEST(FindTest, Contains) { + std::vector empty; + EXPECT_EQ(Contains(empty, 0), false); + + std::vector range = {1, 2}; + EXPECT_EQ(Contains(range, 0), false); + EXPECT_EQ(Contains(range, 1), true); + EXPECT_EQ(Contains(range, 2), true); + EXPECT_EQ(Contains(range, 3), false); +} + +} // namespace +} // namespace Carbon diff --git a/testing/file_test/BUILD b/testing/file_test/BUILD index 9786e431f659..7d9c24896a96 100644 --- a/testing/file_test/BUILD +++ b/testing/file_test/BUILD @@ -43,6 +43,7 @@ cc_library( "//common:check", "//common:error", "//common:exe_path", + "//common:find", "//common:init_llvm", "//common:ostream", "//common:raw_string_ostream", diff --git a/testing/file_test/test_file.cpp b/testing/file_test/test_file.cpp index 92cacf4178d9..6775f77dda35 100644 --- a/testing/file_test/test_file.cpp +++ b/testing/file_test/test_file.cpp @@ -11,6 +11,7 @@ #include "common/check.h" #include "common/error.h" +#include "common/find.h" #include "common/raw_string_ostream.h" #include "llvm/ADT/StringExtras.h" #include "llvm/Support/JSON.h" @@ -132,14 +133,13 @@ static auto AutoFillDidOpenParams(llvm::json::Object* params, } CARBON_ASSIGN_OR_RETURN(auto file_path, ExtractFilePathFromUri(*uri)); - const auto* split_it = - llvm::find_if(splits, [&](const TestFile::Split& split) { - return split.filename == file_path; - }); - if (split_it == splits.end()) { + const auto* split = FindIfOrNull(splits, [&](const TestFile::Split& split) { + return split.filename == file_path; + }); + if (!split) { return ErrorBuilder() << "No split found for uri: " << *uri; } - attr_it->second = split_it->content; + attr_it->second = split->content; return Success(); } diff --git a/toolchain/check/BUILD b/toolchain/check/BUILD index 4c2f08e901d1..ad6e3a56a992 100644 --- a/toolchain/check/BUILD +++ b/toolchain/check/BUILD @@ -171,6 +171,7 @@ cc_library( ":pointer_dereference", "//common:check", "//common:error", + "//common:find", "//common:map", "//common:ostream", "//common:variant_helpers", diff --git a/toolchain/check/handle_function.cpp b/toolchain/check/handle_function.cpp index 4706a4bac688..31285bb11b0d 100644 --- a/toolchain/check/handle_function.cpp +++ b/toolchain/check/handle_function.cpp @@ -5,6 +5,7 @@ #include #include +#include "common/find.h" #include "toolchain/base/kind_switch.h" #include "toolchain/check/context.h" #include "toolchain/check/control_flow.h" @@ -90,15 +91,9 @@ static auto FindSelfPattern(Context& context, -> SemIR::InstId { auto implicit_param_patterns = context.inst_blocks().GetOrEmpty(implicit_param_patterns_id); - if (const auto* i = llvm::find_if(implicit_param_patterns, - [&](auto implicit_param_id) { - return SemIR::IsSelfPattern( - context.sem_ir(), implicit_param_id); - }); - i != implicit_param_patterns.end()) { - return *i; - } - return SemIR::InstId::None; + return FindIfOrNone(implicit_param_patterns, [&](auto implicit_param_id) { + return SemIR::IsSelfPattern(context.sem_ir(), implicit_param_id); + }); } // Diagnoses issues with the modifiers, removing modifiers that shouldn't be diff --git a/toolchain/parse/BUILD b/toolchain/parse/BUILD index de6f65b959b8..c5488815147d 100644 --- a/toolchain/parse/BUILD +++ b/toolchain/parse/BUILD @@ -153,6 +153,7 @@ cc_library( ":node_kind", "//common:check", "//common:error", + "//common:find", "//common:ostream", "//common:struct_reflection", "//toolchain/base:value_store", diff --git a/toolchain/parse/extract.cpp b/toolchain/parse/extract.cpp index bd7d68569783..a4b8d890073f 100644 --- a/toolchain/parse/extract.cpp +++ b/toolchain/parse/extract.cpp @@ -9,6 +9,7 @@ #include #include "common/error.h" +#include "common/find.h" #include "common/struct_reflection.h" #include "toolchain/parse/tree.h" #include "toolchain/parse/tree_and_subtrees.h" @@ -210,7 +211,7 @@ auto NodeExtractor::MatchesNodeIdOneOf( *trace_ << "\n"; } return false; - } else if (llvm::find(kinds, node_kind) == kinds.end()) { + } else if (!Contains(kinds, node_kind)) { if (trace_) { *trace_ << "NodeIdOneOf error: wrong kind " << node_kind << ", expected "; trace_kinds(); diff --git a/toolchain/testing/BUILD b/toolchain/testing/BUILD index 108c716d2a11..a6f50efae75e 100644 --- a/toolchain/testing/BUILD +++ b/toolchain/testing/BUILD @@ -49,6 +49,7 @@ cc_library( testonly = 1, hdrs = ["coverage_helper.h"], deps = [ + "//common:find", "//common:set", "@googletest//:gtest", "@llvm-project//llvm:Support", diff --git a/toolchain/testing/coverage_helper.h b/toolchain/testing/coverage_helper.h index cca5bc6ef310..668f7c51eca9 100644 --- a/toolchain/testing/coverage_helper.h +++ b/toolchain/testing/coverage_helper.h @@ -10,6 +10,7 @@ #include #include +#include "common/find.h" #include "common/set.h" #include "llvm/ADT/StringExtras.h" #include "re2/re2.h" @@ -52,7 +53,7 @@ auto TestKindCoverage(const std::string& manifest_path, llvm::SmallVector missing_kinds; for (auto kind : kinds) { - if (llvm::find(untested_kinds, kind) != untested_kinds.end()) { + if (Contains(untested_kinds, kind)) { EXPECT_FALSE(covered_kinds.Erase(kind.name())) << "Kind " << kind << " has coverage even though none was expected. If this has "