diff --git a/base/BUILD b/base/BUILD index 600889182..9dba8296e 100644 --- a/base/BUILD +++ b/base/BUILD @@ -38,15 +38,28 @@ cc_library( "@com_google_absl//absl/base:no_destructor", "@com_google_absl//absl/base:nullability", "@com_google_absl//absl/container:btree", + "@com_google_absl//absl/functional:overload", + "@com_google_absl//absl/log:absl_check", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_absl//absl/types:optional", + "@com_google_absl//absl/types:optional_ref", "@com_google_absl//absl/types:span", "@com_google_absl//absl/types:variant", ], ) +cc_test( + name = "attribute_test", + srcs = ["attribute_test.cc"], + deps = [ + ":attributes", + "//internal:testing", + "@com_google_absl//absl/types:optional", + ], +) + cc_library( name = "kind", hdrs = ["kind.h"], diff --git a/base/attribute.cc b/base/attribute.cc index f750a1850..e2b41b242 100644 --- a/base/attribute.cc +++ b/base/attribute.cc @@ -17,12 +17,15 @@ #include #include #include +#include #include "absl/base/macros.h" #include "absl/base/nullability.h" +#include "absl/log/absl_check.h" #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/str_cat.h" +#include "absl/strings/string_view.h" #include "absl/types/variant.h" #include "base/kind.h" #include "internal/status_macros.h" @@ -36,14 +39,10 @@ class AttributeStringPrinter { public: // String representation for the given qualifier is appended to output. // output must be non-null. - explicit AttributeStringPrinter(std::string* output, Kind type) - : output_(*output), type_(type) {} + explicit AttributeStringPrinter(std::string* output) : output_(*output) {} - absl::Status operator()(const Kind& ignored) const { - // Attributes are represented as a variant, with illegal attribute - // qualifiers represented with their type as the first alternative. - return absl::InvalidArgumentError( - absl::StrCat("Unsupported attribute qualifier ", KindToString(type_))); + absl::Status operator()(absl::monostate) const { + return absl::InvalidArgumentError("bad attribute qualifier"); } absl::Status operator()(int64_t index) { @@ -68,22 +67,19 @@ class AttributeStringPrinter { private: std::string& output_; - Kind type_; }; // Visitor for appending string representation for different qualifier kinds. class AttributeQualifierStringPrinter { public: // String representation for the given qualifier is appended to output. - explicit AttributeQualifierStringPrinter(std::string* absl_nonnull output, - Kind type) - : output_(*output), type_(type) {} + explicit AttributeQualifierStringPrinter(std::string* absl_nonnull output) + : output_(*output) {} - absl::Status operator()(const Kind& ignored) const { + absl::Status operator()(absl::monostate) const { // Attributes are represented as a variant, with illegal attribute // qualifiers represented with their type as the first alternative. - return absl::InvalidArgumentError( - absl::StrCat("Unsupported attribute qualifier ", KindToString(type_))); + return absl::InvalidArgumentError("bad attribute qualifier"); } absl::Status operator()(int64_t index) { @@ -108,142 +104,26 @@ class AttributeQualifierStringPrinter { private: std::string& output_; - Kind type_; }; struct AttributeQualifierTypeVisitor final { - Kind operator()(const Kind& type) const { return type; } - - Kind operator()(int64_t ignored) const { - static_cast(ignored); - return Kind::kInt64; - } - - Kind operator()(uint64_t ignored) const { - static_cast(ignored); - return Kind::kUint64; - } - - Kind operator()(const std::string& ignored) const { - static_cast(ignored); - return Kind::kString; - } - - Kind operator()(bool ignored) const { - static_cast(ignored); - return Kind::kBool; - } -}; - -struct AttributeQualifierTypeComparator final { - const Kind lhs; - - bool operator()(const Kind& rhs) const { - return static_cast(lhs) < static_cast(rhs); - } - - bool operator()(int64_t) const { return false; } - - bool operator()(uint64_t other) const { return false; } - - bool operator()(const std::string&) const { return false; } - - bool operator()(bool other) const { return false; } -}; - -struct AttributeQualifierIntComparator final { - const int64_t lhs; - - bool operator()(const Kind&) const { return true; } - - bool operator()(int64_t rhs) const { return lhs < rhs; } + Kind operator()(absl::monostate) const { return Kind::kNull; } - bool operator()(uint64_t) const { return true; } + Kind operator()(int64_t) const { return Kind::kInt64; } - bool operator()(const std::string&) const { return true; } + Kind operator()(uint64_t) const { return Kind::kUint64; } - bool operator()(bool) const { return false; } -}; - -struct AttributeQualifierUintComparator final { - const uint64_t lhs; - - bool operator()(const Kind&) const { return true; } - - bool operator()(int64_t) const { return false; } - - bool operator()(uint64_t rhs) const { return lhs < rhs; } - - bool operator()(const std::string&) const { return true; } - - bool operator()(bool) const { return false; } -}; - -struct AttributeQualifierStringComparator final { - const std::string& lhs; - - bool operator()(const Kind&) const { return true; } - - bool operator()(int64_t) const { return false; } - - bool operator()(uint64_t) const { return false; } + Kind operator()(const std::string&) const { return Kind::kString; } - bool operator()(const std::string& rhs) const { return lhs < rhs; } - - bool operator()(bool) const { return false; } -}; - -struct AttributeQualifierBoolComparator final { - const bool lhs; - - bool operator()(const Kind&) const { return true; } - - bool operator()(int64_t) const { return true; } - - bool operator()(uint64_t) const { return true; } - - bool operator()(const std::string&) const { return true; } - - bool operator()(bool rhs) const { return lhs < rhs; } + Kind operator()(bool) const { return Kind::kBool; } }; } // namespace -struct AttributeQualifier::ComparatorVisitor final { - const AttributeQualifier::Variant& rhs; - - bool operator()(const Kind& lhs) const { - return absl::visit(AttributeQualifierTypeComparator{lhs}, rhs); - } - - bool operator()(int64_t lhs) const { - return absl::visit(AttributeQualifierIntComparator{lhs}, rhs); - } - - bool operator()(uint64_t lhs) const { - return absl::visit(AttributeQualifierUintComparator{lhs}, rhs); - } - - bool operator()(const std::string& lhs) const { - return absl::visit(AttributeQualifierStringComparator{lhs}, rhs); - } - - bool operator()(bool lhs) const { - return absl::visit(AttributeQualifierBoolComparator{lhs}, rhs); - } -}; - Kind AttributeQualifier::kind() const { return absl::visit(AttributeQualifierTypeVisitor{}, value_); } -bool AttributeQualifier::operator<(const AttributeQualifier& other) const { - // The order is not publicly documented because it is subject to change. - // Currently we sort in the following order, with each type being sorted - // against itself: bool, int, uint, string, type. - return absl::visit(ComparatorVisitor{other.value_}, value_); -} - bool Attribute::operator==(const Attribute& other) const { // We cannot check pointer equality as a short circuit because we have to // treat all invalid AttributeQualifier as not equal to each other. @@ -296,7 +176,7 @@ bool Attribute::operator<(const Attribute& other) const { return false; } -const absl::StatusOr Attribute::AsString() const { +absl::StatusOr Attribute::AsString() const { if (variable_name().empty()) { return absl::InvalidArgumentError( "Only ident rooted attributes are supported."); @@ -305,26 +185,335 @@ const absl::StatusOr Attribute::AsString() const { std::string result = std::string(variable_name()); for (const auto& qualifier : qualifier_path()) { - CEL_RETURN_IF_ERROR(absl::visit( - AttributeStringPrinter(&result, qualifier.kind()), qualifier.value_)); + CEL_RETURN_IF_ERROR(absl::visit(AttributeStringPrinter(&result), + common_internal::AsVariant(qualifier))); } return result; } +std::string AttributeQualifier::ToString() const { + std::string result; + absl::Status status = + absl::visit(AttributeQualifierStringPrinter{&result}, value_); + ABSL_DCHECK_OK(status) << "bad attribute qualifier"; + status.IgnoreError(); + return result; +} + bool AttributeQualifier::IsMatch(const AttributeQualifier& other) const { - if (absl::holds_alternative(value_) || - absl::holds_alternative(other.value_)) { + ABSL_DCHECK(*this) << "bad attribute qualifier"; + ABSL_DCHECK(other) << "bad attribute qualifier"; + if (absl::holds_alternative(value_) || + absl::holds_alternative(other.value_)) { return false; } - return value_ == other.value_; + return *this == other; } -absl::StatusOr AttributeQualifier::AsString() const { - std::string result; - CEL_RETURN_IF_ERROR( - absl::visit(AttributeQualifierStringPrinter(&result, kind()), value_)); - return result; +bool AttributeQualifier::IsMatch(absl::string_view other_key) const { + ABSL_DCHECK(*this) << "bad attribute qualifier"; + if (auto string = AsString(); string.has_value()) { + return *string == other_key; + } + return false; +} + +bool AttributeQualifierPattern::IsMatch( + const AttributeQualifier& qualifier) const { + ABSL_DCHECK(*this) << "bad attribute qualifier pattern"; + ABSL_DCHECK(qualifier) << "bad attribute qualifier"; + if (!*this || !qualifier) { + return false; + } + if (IsWildcard()) { + return true; + } + return *this == qualifier; +} + +bool AttributeQualifierPattern::IsMatch(absl::string_view other_key) const { + ABSL_DCHECK(*this) << "bad attribute qualifier pattern"; + if (IsWildcard()) { + return true; + } + if (auto string = AsString(); string.has_value()) { + return *string == other_key; + } + return false; +} + +namespace { + +struct AttributeQualifierEqualTo { + bool operator()(const absl::monostate&, const absl::monostate&) const { + return false; + } + + template + std::enable_if_t, absl::monostate>, + bool> + operator()(const absl::monostate&, const T&) const { + return false; + } + + template + std::enable_if_t, absl::monostate>, + bool> + operator()(const T&, const absl::monostate&) const { + return false; + } + + bool operator()(bool lhs, bool rhs) const { return lhs == rhs; } + + bool operator()(bool, int64_t) const { return false; } + + bool operator()(bool, uint64_t) const { return false; } + + bool operator()(bool, absl::string_view) const { return false; } + + bool operator()(int64_t, bool) const { return false; } + + bool operator()(int64_t lhs, int64_t rhs) const { return lhs == rhs; } + + bool operator()(int64_t, uint64_t) const { return false; } + + bool operator()(int64_t, absl::string_view) const { return false; } + + bool operator()(uint64_t, bool) const { return false; } + + bool operator()(uint64_t, int64_t) const { return false; } + + bool operator()(uint64_t lhs, uint64_t rhs) const { return lhs == rhs; } + + bool operator()(uint64_t, absl::string_view) const { return false; } + + bool operator()(absl::string_view, bool) const { return false; } + + bool operator()(absl::string_view, int64_t) const { return false; } + + bool operator()(absl::string_view, uint64_t) const { return false; } + + bool operator()(absl::string_view lhs, absl::string_view rhs) const { + return lhs == rhs; + } + + bool operator()(const common_internal::WildcardType&, + const common_internal::WildcardType&) const { + return true; + } + + template + std::enable_if_t< + !std::is_same_v, common_internal::WildcardType> && + !std::is_same_v, absl::monostate>, + bool> + operator()(const common_internal::WildcardType&, const T&) const { + return false; + } + + template + std::enable_if_t< + !std::is_same_v, common_internal::WildcardType> && + !std::is_same_v, absl::monostate>, + bool> + operator()(const T&, const common_internal::WildcardType&) const { + return false; + } +}; + +struct AttributeQualifierLess { + bool operator()(const absl::monostate&, const absl::monostate&) const { + return false; + } + + template + std::enable_if_t, absl::monostate>, + bool> + operator()(const absl::monostate&, const T&) const { + return true; + } + + template + std::enable_if_t, absl::monostate>, + bool> + operator()(const T&, const absl::monostate&) const { + return false; + } + + bool operator()(bool lhs, bool rhs) const { return lhs < rhs; } + + bool operator()(bool, int64_t) const { return false; } + + bool operator()(bool, uint64_t) const { return false; } + + bool operator()(bool, absl::string_view) const { return false; } + + bool operator()(int64_t, bool) const { return true; } + + bool operator()(int64_t lhs, int64_t rhs) const { return lhs < rhs; } + + bool operator()(int64_t, uint64_t) const { return true; } + + bool operator()(int64_t, absl::string_view) const { return true; } + + bool operator()(uint64_t, bool) const { return true; } + + bool operator()(uint64_t, int64_t) const { return false; } + + bool operator()(uint64_t lhs, uint64_t rhs) const { return lhs < rhs; } + + bool operator()(uint64_t, absl::string_view) const { return true; } + + bool operator()(absl::string_view, bool) const { return true; } + + bool operator()(absl::string_view, int64_t) const { return false; } + + bool operator()(absl::string_view, uint64_t) const { return false; } + + bool operator()(absl::string_view lhs, absl::string_view rhs) const { + return lhs < rhs; + } + + bool operator()(const common_internal::WildcardType&, + const common_internal::WildcardType&) const { + return false; + } + + template + std::enable_if_t< + !std::is_same_v, common_internal::WildcardType> && + !std::is_same_v, absl::monostate>, + bool> + operator()(const common_internal::WildcardType&, const T&) const { + return false; + } + + template + std::enable_if_t< + !std::is_same_v, common_internal::WildcardType> && + !std::is_same_v, absl::monostate>, + bool> + operator()(const T&, const common_internal::WildcardType&) const { + return true; + } +}; + +} // namespace + +bool operator==(const AttributeQualifier& lhs, const AttributeQualifier& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator==(const AttributeQualifierView& lhs, + const AttributeQualifierView& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator==(const AttributeQualifierPattern& lhs, + const AttributeQualifierPattern& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator==(const AttributeQualifier& lhs, + const AttributeQualifierView& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator==(const AttributeQualifier& lhs, + const AttributeQualifierPattern& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator==(const AttributeQualifierView& lhs, + const AttributeQualifier& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator==(const AttributeQualifierView& lhs, + const AttributeQualifierPattern& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator==(const AttributeQualifierPattern& lhs, + const AttributeQualifier& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator==(const AttributeQualifierPattern& lhs, + const AttributeQualifierView& rhs) { + return absl::visit(AttributeQualifierEqualTo{}, + common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifier& lhs, const AttributeQualifier& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifierView& lhs, + const AttributeQualifierView& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifierPattern& lhs, + const AttributeQualifierPattern& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifier& lhs, + const AttributeQualifierView& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifier& lhs, + const AttributeQualifierPattern& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifierView& lhs, + const AttributeQualifier& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifierView& lhs, + const AttributeQualifierPattern& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifierPattern& lhs, + const AttributeQualifier& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); +} + +bool operator<(const AttributeQualifierPattern& lhs, + const AttributeQualifierView& rhs) { + return absl::visit(AttributeQualifierLess{}, common_internal::AsVariant(lhs), + common_internal::AsVariant(rhs)); } } // namespace cel diff --git a/base/attribute.h b/base/attribute.h index 69dcaf161..aef6f49ca 100644 --- a/base/attribute.h +++ b/base/attribute.h @@ -15,6 +15,7 @@ #ifndef THIRD_PARTY_CEL_CPP_BASE_ATTRIBUTE_H_ #define THIRD_PARTY_CEL_CPP_BASE_ATTRIBUTE_H_ +#include #include #include #include @@ -22,24 +23,63 @@ #include #include +#include "absl/base/attributes.h" +#include "absl/base/macros.h" +#include "absl/functional/overload.h" +#include "absl/log/absl_check.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" #include "absl/types/optional.h" +#include "absl/types/optional_ref.h" #include "absl/types/span.h" #include "absl/types/variant.h" #include "base/kind.h" namespace cel { -// AttributeQualifier represents a segment in -// attribute resolutuion path. A segment can be qualified by values of +namespace common_internal { +class AttributeMatcherNode; +struct WildcardType {}; +using AttributeQualifierVariant = + absl::variant; +using AttributeQualifierPatternVariant = + absl::variant; +using AttributeQualifierViewVariant = + absl::variant; +} // namespace common_internal + +class AttributeQualifier; +class AttributeQualifierPattern; +class Attribute; +class AttributePattern; +class AttributeQualifierView; + +namespace common_internal { +[[nodiscard]] +const AttributeQualifierVariant& AsVariant( + const AttributeQualifier& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND); +[[nodiscard]] +AttributeQualifierVariant&& AsVariant( + AttributeQualifier&& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND); +[[nodiscard]] +const AttributeQualifierPatternVariant& AsVariant( + const AttributeQualifierPattern& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND); +[[nodiscard]] +AttributeQualifierPatternVariant&& AsVariant( + AttributeQualifierPattern&& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND); +[[nodiscard]] +const AttributeQualifierViewVariant& AsVariant( + const AttributeQualifierView& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND); +[[nodiscard]] +AttributeQualifierViewVariant&& AsVariant( + AttributeQualifierView&& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND); +} // namespace common_internal + +// AttributeQualifier represents a segment in the +// attribute resolution path. A segment can be qualified by values of // following types: string/int64_t/uint64_t/bool. -class AttributeQualifier final { - private: - struct ComparatorVisitor; - - using Variant = absl::variant; - +class AttributeQualifier { public: static AttributeQualifier OfInt(int64_t value) { return AttributeQualifier(absl::in_place_type, std::move(value)); @@ -49,11 +89,20 @@ class AttributeQualifier final { return AttributeQualifier(absl::in_place_type, std::move(value)); } + static AttributeQualifier OfString(const char* value) { + return OfString(absl::string_view(value)); + } + static AttributeQualifier OfString(std::string value) { return AttributeQualifier(absl::in_place_type, std::move(value)); } + static AttributeQualifier OfString(absl::string_view value) { + return AttributeQualifier(absl::in_place_type, + std::string(value)); + } + static AttributeQualifier OfBool(bool value) { return AttributeQualifier(absl::in_place_type, std::move(value)); } @@ -68,55 +117,126 @@ class AttributeQualifier final { Kind kind() const; - // Family of Get... methods. Return values if requested type matches the - // stored one. - absl::optional GetInt64Key() const { - return absl::holds_alternative(value_) - ? absl::optional(absl::get<1>(value_)) - : absl::nullopt; + [[nodiscard]] + std::string ToString() const; + + ABSL_DEPRECATE_AND_INLINE() + absl::optional GetInt64Key() const { return AsInt(); } + + ABSL_DEPRECATE_AND_INLINE() + absl::optional GetUint64Key() const { return AsUint(); } + + ABSL_DEPRECATED("Use AsString") + absl::optional GetStringKey() const { + if (auto string = AsString(); string.has_value()) { + return *string; + } + return absl::nullopt; } - absl::optional GetUint64Key() const { - return absl::holds_alternative(value_) - ? absl::optional(absl::get<2>(value_)) - : absl::nullopt; + ABSL_DEPRECATE_AND_INLINE() + absl::optional GetBoolKey() const { return AsBool(); } + + explicit operator bool() const { + return !absl::holds_alternative(value_); } - absl::optional GetStringKey() const { - return absl::holds_alternative(value_) - ? absl::optional(absl::get<3>(value_)) - : absl::nullopt; + [[nodiscard]] + bool IsBool() const { + return absl::holds_alternative(value_); } - absl::optional GetBoolKey() const { - return absl::holds_alternative(value_) - ? absl::optional(absl::get<4>(value_)) - : absl::nullopt; + [[nodiscard]] + bool IsInt() const { + return absl::holds_alternative(value_); } - bool operator==(const AttributeQualifier& other) const { - return IsMatch(other); + [[nodiscard]] + bool IsUint() const { + return absl::holds_alternative(value_); } - bool operator<(const AttributeQualifier& other) const; + [[nodiscard]] + bool IsString() const { + return absl::holds_alternative(value_); + } - bool IsMatch(absl::string_view other_key) const { - absl::optional key = GetStringKey(); - return (key.has_value() && key.value() == other_key); + [[nodiscard]] + bool GetBool() const { + ABSL_DCHECK(IsBool()); + return absl::get(value_); } - absl::StatusOr AsString() const; + [[nodiscard]] + int64_t GetInt() const { + ABSL_DCHECK(IsInt()); + return absl::get(value_); + } - private: - friend class Attribute; - friend struct ComparatorVisitor; + [[nodiscard]] + uint64_t GetUint() const { + ABSL_DCHECK(IsUint()); + return absl::get(value_); + } - template - AttributeQualifier(absl::in_place_type_t in_place_type, T&& value) - : value_(in_place_type, std::forward(value)) {} + [[nodiscard]] + const std::string& GetString() const { + ABSL_DCHECK(IsString()); + return absl::get(value_); + } + + [[nodiscard]] + absl::optional AsBool() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; + } + + [[nodiscard]] + absl::optional AsInt() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; + } + [[nodiscard]] + absl::optional AsUint() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; + } + + [[nodiscard]] + absl::optional_ref AsString() const { + if (const auto* value = absl::get_if(&value_); + value != nullptr) { + return *value; + } + return absl::nullopt; + } + + [[nodiscard]] bool IsMatch(const AttributeQualifier& other) const; + [[nodiscard]] + bool IsMatch(absl::string_view other_key) const; + + private: + friend const common_internal::AttributeQualifierVariant& + common_internal::AsVariant(const AttributeQualifier& qualifier); + friend common_internal::AttributeQualifierVariant&& + common_internal::AsVariant(AttributeQualifier&& qualifier); + + template + explicit AttributeQualifier(absl::in_place_type_t in_place_type, + Args&&... args) + : value_(in_place_type, std::forward(args)...) {} + + using Variant = common_internal::AttributeQualifierVariant; + // The previous implementation of Attribute preserved all value // instances, regardless of whether they are supported in this context or not. // We represented unsupported types by using the first alternative and thus @@ -124,58 +244,582 @@ class AttributeQualifier final { Variant value_; }; -// AttributeQualifierPattern matches a segment in -// attribute resolutuion path. AttributeQualifierPattern is capable of -// matching path elements of types string/int64/uint64/bool. -class AttributeQualifierPattern final { +class AttributeQualifierView { + public: + static AttributeQualifierView OfInt(int64_t value) { + return AttributeQualifierView(absl::in_place_type, + std::move(value)); + } + + static AttributeQualifierView OfUint(uint64_t value) { + return AttributeQualifierView(absl::in_place_type, + std::move(value)); + } + + static AttributeQualifierView OfString(const char* value) { + return OfString(absl::string_view(value)); + } + + static AttributeQualifierView OfString(absl::string_view value) { + return AttributeQualifierView(absl::in_place_type, + std::move(value)); + } + + static AttributeQualifierView OfString(std::string&&) = delete; + + static AttributeQualifierView OfBool(bool value) { + return AttributeQualifierView(absl::in_place_type, std::move(value)); + } + + AttributeQualifierView() = default; + AttributeQualifierView(const AttributeQualifierView&) = default; + AttributeQualifierView& operator=(const AttributeQualifierView&) = default; + + // NOLINTNEXTLINE(google-explicit-constructor) + AttributeQualifierView( + const AttributeQualifier& other ABSL_ATTRIBUTE_LIFETIME_BOUND) + : value_(absl::visit( + absl::Overload( + [](absl::monostate value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](bool value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](int64_t value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](uint64_t value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](const std::string& value) -> Variant { + return Variant(absl::in_place_type, value); + }), + common_internal::AsVariant(other))) {} + + // NOLINTNEXTLINE(google-explicit-constructor) + AttributeQualifierView& operator=( + const AttributeQualifier& other ABSL_ATTRIBUTE_LIFETIME_BOUND) { + return *this = AttributeQualifierView(other); + } + + AttributeQualifierView& operator=(AttributeQualifier&&) = delete; + + ABSL_DEPRECATE_AND_INLINE() + absl::optional GetInt64Key() const { return AsInt(); } + + ABSL_DEPRECATE_AND_INLINE() + absl::optional GetUint64Key() const { return AsUint(); } + + ABSL_DEPRECATE_AND_INLINE() + absl::optional GetStringKey() const { return AsString(); } + + ABSL_DEPRECATE_AND_INLINE() + absl::optional GetBoolKey() const { return AsBool(); } + + [[nodiscard]] + bool IsBool() const { + return absl::holds_alternative(value_); + } + + [[nodiscard]] + bool IsInt() const { + return absl::holds_alternative(value_); + } + + [[nodiscard]] + bool IsUint() const { + return absl::holds_alternative(value_); + } + + [[nodiscard]] + bool IsString() const { + return absl::holds_alternative(value_); + } + + [[nodiscard]] + bool GetBool() const { + ABSL_DCHECK(IsBool()); + return absl::get(value_); + } + + [[nodiscard]] + int64_t GetInt() const { + ABSL_DCHECK(IsInt()); + return absl::get(value_); + } + + [[nodiscard]] + uint64_t GetUint() const { + ABSL_DCHECK(IsUint()); + return absl::get(value_); + } + + [[nodiscard]] + absl::string_view GetString() const { + ABSL_DCHECK(IsString()); + return absl::get(value_); + } + + [[nodiscard]] + absl::optional AsBool() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; + } + + [[nodiscard]] + absl::optional AsInt() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; + } + + [[nodiscard]] + absl::optional AsUint() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; + } + + [[nodiscard]] + absl::optional AsString() const { + if (const auto* value = absl::get_if(&value_); + value != nullptr) { + return *value; + } + return absl::nullopt; + } + private: - // Qualifier value. If not set, treated as wildcard. - std::optional value_; + friend const common_internal::AttributeQualifierViewVariant& + common_internal::AsVariant(const AttributeQualifierView& qualifier); + friend common_internal::AttributeQualifierViewVariant&& + common_internal::AsVariant(AttributeQualifierView&& qualifier); - explicit AttributeQualifierPattern(std::optional value) - : value_(std::move(value)) {} + using Variant = common_internal::AttributeQualifierViewVariant; + template + AttributeQualifierView(absl::in_place_type_t in_place_type, T&& value) + : value_(in_place_type, std::forward(value)) {} + + Variant value_; +}; + +// AttributeQualifierPattern matches a segment in +// attribute resolutuion path. AttributeQualifierPattern is capable of +// matching path elements of types string/int64/uint64/bool. +class AttributeQualifierPattern { public: static AttributeQualifierPattern OfInt(int64_t value) { - return AttributeQualifierPattern(AttributeQualifier::OfInt(value)); + return AttributeQualifierPattern(absl::in_place_type, value); } static AttributeQualifierPattern OfUint(uint64_t value) { - return AttributeQualifierPattern(AttributeQualifier::OfUint(value)); + return AttributeQualifierPattern(absl::in_place_type, value); + } + + static AttributeQualifierPattern OfString(const char* value) { + return OfString(absl::string_view(value)); } static AttributeQualifierPattern OfString(std::string value) { - return AttributeQualifierPattern( - AttributeQualifier::OfString(std::move(value))); + return AttributeQualifierPattern(absl::in_place_type, + std::move(value)); + } + + static AttributeQualifierPattern OfString(absl::string_view value) { + return AttributeQualifierPattern(absl::in_place_type, + std::string(value)); } static AttributeQualifierPattern OfBool(bool value) { - return AttributeQualifierPattern(AttributeQualifier::OfBool(value)); + return AttributeQualifierPattern(absl::in_place_type, value); + } + + ABSL_DEPRECATE_AND_INLINE() + static AttributeQualifierPattern CreateWildcard() { return Wildcard(); } + + static AttributeQualifierPattern Wildcard() { + return AttributeQualifierPattern(absl::in_place_type); + } + + // NOLINTNEXTLINE(google-explicit-constructor) + AttributeQualifierPattern(const AttributeQualifier& value) + : value_(absl::visit( + absl::Overload( + [](absl::monostate value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](bool value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](int64_t value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](uint64_t value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](const std::string& value) -> Variant { + return Variant(absl::in_place_type, value); + }), + common_internal::AsVariant(value))) {} + + // NOLINTNEXTLINE(google-explicit-constructor) + AttributeQualifierPattern(AttributeQualifier&& value) + : value_(absl::visit( + absl::Overload( + [](absl::monostate value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](bool value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](int64_t value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](uint64_t value) -> Variant { + return Variant(absl::in_place_type, value); + }, + [](std::string&& value) -> Variant { + return Variant(absl::in_place_type, + std::move(value)); + }), + common_internal::AsVariant(std::move(value)))) {} + + explicit operator bool() const { + return !absl::holds_alternative(value_); } - static AttributeQualifierPattern CreateWildcard() { - return AttributeQualifierPattern(std::nullopt); + [[nodiscard]] + bool IsBool() const { + return absl::holds_alternative(value_); } - explicit AttributeQualifierPattern(AttributeQualifier qualifier) - : AttributeQualifierPattern( - std::optional(std::move(qualifier))) {} + [[nodiscard]] + bool IsInt() const { + return absl::holds_alternative(value_); + } + + [[nodiscard]] + bool IsUint() const { + return absl::holds_alternative(value_); + } + + [[nodiscard]] + bool IsString() const { + return absl::holds_alternative(value_); + } + + [[nodiscard]] + bool IsWildcard() const { + return absl::holds_alternative(value_); + } + + [[nodiscard]] + bool GetBool() const { + ABSL_DCHECK(IsBool()); + return absl::get(value_); + } + + [[nodiscard]] + int64_t GetInt() const { + ABSL_DCHECK(IsInt()); + return absl::get(value_); + } + + [[nodiscard]] + uint64_t GetUint() const { + ABSL_DCHECK(IsUint()); + return absl::get(value_); + } + + [[nodiscard]] + const std::string& GetString() const { + ABSL_DCHECK(IsString()); + return absl::get(value_); + } + + [[nodiscard]] + absl::optional AsBool() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; + } - bool IsWildcard() const { return !value_.has_value(); } + [[nodiscard]] + absl::optional AsInt() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; + } - bool IsMatch(const AttributeQualifier& qualifier) const { - if (IsWildcard()) return true; - return value_.value() == qualifier; + [[nodiscard]] + absl::optional AsUint() const { + if (const auto* value = absl::get_if(&value_); value != nullptr) { + return *value; + } + return absl::nullopt; } - bool IsMatch(absl::string_view other_key) const { - if (!value_.has_value()) return true; - return value_->IsMatch(other_key); + [[nodiscard]] + absl::optional_ref AsString() const { + if (const auto* value = absl::get_if(&value_); + value != nullptr) { + return *value; + } + return absl::nullopt; } + + [[nodiscard]] + absl::optional ToQualifier() const { + return absl::visit( + absl::Overload( + [](absl::monostate) -> absl::optional { + return AttributeQualifier(); + }, + [](bool value) -> absl::optional { + return AttributeQualifier::OfBool(value); + }, + [](int64_t value) -> absl::optional { + return AttributeQualifier::OfInt(value); + }, + [](uint64_t value) -> absl::optional { + return AttributeQualifier::OfUint(value); + }, + [](const std::string& value) -> absl::optional { + return AttributeQualifier::OfString(value); + }, + [](common_internal::WildcardType value) + -> absl::optional { + return absl::nullopt; + }), + value_); + } + + [[nodiscard]] + absl::optional ToQualifierView() const + ABSL_ATTRIBUTE_LIFETIME_BOUND { + return absl::visit( + absl::Overload( + [](absl::monostate) -> absl::optional { + return AttributeQualifierView(); + }, + [](bool value) -> absl::optional { + return AttributeQualifierView::OfBool(value); + }, + [](int64_t value) -> absl::optional { + return AttributeQualifierView::OfInt(value); + }, + [](uint64_t value) -> absl::optional { + return AttributeQualifierView::OfUint(value); + }, + [](const std::string& value) + -> absl::optional { + return AttributeQualifierView::OfString(value); + }, + [](common_internal::WildcardType value) + -> absl::optional { + return absl::nullopt; + }), + value_); + } + + [[nodiscard]] + bool IsMatch(const AttributeQualifier& qualifier) const; + + [[nodiscard]] + bool IsMatch(absl::string_view other_key) const; + + private: + friend const common_internal::AttributeQualifierPatternVariant& + common_internal::AsVariant(const AttributeQualifierPattern& qualifier); + friend common_internal::AttributeQualifierPatternVariant&& + common_internal::AsVariant(AttributeQualifierPattern&& qualifier); + + using Variant = common_internal::AttributeQualifierPatternVariant; + using WildcardType = common_internal::WildcardType; + + template + explicit AttributeQualifierPattern(absl::in_place_type_t in_place_type, + Args&&... args) + : value_(in_place_type, std::forward(args)...) {} + + // Qualifier value. If not set, treated as wildcard. + common_internal::AttributeQualifierPatternVariant value_; }; +[[nodiscard]] +bool operator==(const AttributeQualifier& lhs, const AttributeQualifier& rhs); + +[[nodiscard]] +bool operator==(const AttributeQualifierView& lhs, + const AttributeQualifierView& rhs); + +[[nodiscard]] +bool operator==(const AttributeQualifierPattern& lhs, + const AttributeQualifierPattern& rhs); + +[[nodiscard]] +bool operator==(const AttributeQualifier& lhs, + const AttributeQualifierView& rhs); + +[[nodiscard]] +bool operator==(const AttributeQualifier& lhs, + const AttributeQualifierPattern& rhs); + +[[nodiscard]] +bool operator==(const AttributeQualifierView& lhs, + const AttributeQualifier& rhs); + +[[nodiscard]] +bool operator==(const AttributeQualifierView& lhs, + const AttributeQualifierPattern& rhs); + +[[nodiscard]] +bool operator==(const AttributeQualifierPattern& lhs, + const AttributeQualifier& rhs); + +[[nodiscard]] +bool operator==(const AttributeQualifierPattern& lhs, + const AttributeQualifierView& rhs); + +[[nodiscard]] +inline bool operator!=(const AttributeQualifier& lhs, + const AttributeQualifier& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +inline bool operator!=(const AttributeQualifierView& lhs, + const AttributeQualifierView& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +inline bool operator!=(const AttributeQualifierPattern& lhs, + const AttributeQualifierPattern& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +inline bool operator!=(const AttributeQualifier& lhs, + const AttributeQualifierView& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +inline bool operator!=(const AttributeQualifier& lhs, + const AttributeQualifierPattern& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +inline bool operator!=(const AttributeQualifierView& lhs, + const AttributeQualifier& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +inline bool operator!=(const AttributeQualifierView& lhs, + const AttributeQualifierPattern& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +inline bool operator!=(const AttributeQualifierPattern& lhs, + const AttributeQualifier& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +inline bool operator!=(const AttributeQualifierPattern& lhs, + const AttributeQualifierView& rhs) { + return !operator==(lhs, rhs); +} + +[[nodiscard]] +bool operator<(const AttributeQualifier& lhs, const AttributeQualifier& rhs); + +[[nodiscard]] +bool operator<(const AttributeQualifierView& lhs, + const AttributeQualifierView& rhs); + +[[nodiscard]] +bool operator<(const AttributeQualifierPattern& lhs, + const AttributeQualifierPattern& rhs); + +[[nodiscard]] +bool operator<(const AttributeQualifier& lhs, + const AttributeQualifierView& rhs); + +[[nodiscard]] +bool operator<(const AttributeQualifier& lhs, + const AttributeQualifierPattern& rhs); + +[[nodiscard]] +bool operator<(const AttributeQualifierView& lhs, + const AttributeQualifier& rhs); + +[[nodiscard]] +bool operator<(const AttributeQualifierView& lhs, + const AttributeQualifierPattern& rhs); + +[[nodiscard]] +bool operator<(const AttributeQualifierPattern& lhs, + const AttributeQualifier& rhs); + +[[nodiscard]] +bool operator<(const AttributeQualifierPattern& lhs, + const AttributeQualifierView& rhs); + +namespace common_internal { + +[[nodiscard]] +inline const AttributeQualifierVariant& AsVariant( + const AttributeQualifier& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND) { + return qualifier.value_; +} + +[[nodiscard]] +inline AttributeQualifierVariant&& AsVariant( + AttributeQualifier&& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND) { + return std::move(qualifier.value_); +} + +[[nodiscard]] +inline const AttributeQualifierPatternVariant& AsVariant( + const AttributeQualifierPattern& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND) { + return qualifier.value_; +} + +[[nodiscard]] +inline AttributeQualifierPatternVariant&& AsVariant( + AttributeQualifierPattern&& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND) { + return std::move(qualifier.value_); +} + +[[nodiscard]] +inline const AttributeQualifierViewVariant& AsVariant( + const AttributeQualifierView& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND) { + return qualifier.value_; +} + +[[nodiscard]] +inline AttributeQualifierViewVariant&& AsVariant( + AttributeQualifierView&& qualifier ABSL_ATTRIBUTE_LIFETIME_BOUND) { + return std::move(qualifier.value_); +} + +} // namespace common_internal + // Attribute represents resolved attribute path. -class Attribute final { +class Attribute { public: explicit Attribute(std::string variable_name) : Attribute(std::move(variable_name), {}) {} @@ -197,7 +841,7 @@ class Attribute final { bool operator<(const Attribute& other) const; - const absl::StatusOr AsString() const; + absl::StatusOr AsString() const; private: struct Impl final { @@ -218,7 +862,7 @@ class Attribute final { // - field selection; // - map lookup by key; // - list access by index. -class AttributePattern final { +class AttributePattern { public: // MatchType enum specifies how closely pattern is matching the attribute: enum class MatchType { diff --git a/base/attribute_test.cc b/base/attribute_test.cc new file mode 100644 index 000000000..9ecf43060 --- /dev/null +++ b/base/attribute_test.cc @@ -0,0 +1,78 @@ +// Copyright 2022 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "base/attribute.h" + +#include "absl/types/optional.h" +#include "internal/testing.h" + +namespace cel { +namespace { + +using ::testing::Optional; + +TEST(AttributeQualifierView, Bool) { + AttributeQualifierView qualifier = AttributeQualifierView::OfBool(true); + EXPECT_TRUE(qualifier.IsBool()); + EXPECT_FALSE(qualifier.IsInt()); + EXPECT_FALSE(qualifier.IsUint()); + EXPECT_FALSE(qualifier.IsString()); + EXPECT_TRUE(qualifier.GetBool()); + EXPECT_THAT(qualifier.AsBool(), Optional(true)); + EXPECT_THAT(qualifier.AsInt(), absl::nullopt); + EXPECT_THAT(qualifier.AsUint(), absl::nullopt); + EXPECT_THAT(qualifier.AsString(), absl::nullopt); +} + +TEST(AttributeQualifierView, Int) { + AttributeQualifierView qualifier = AttributeQualifierView::OfInt(1); + EXPECT_FALSE(qualifier.IsBool()); + EXPECT_TRUE(qualifier.IsInt()); + EXPECT_FALSE(qualifier.IsUint()); + EXPECT_FALSE(qualifier.IsString()); + EXPECT_EQ(qualifier.GetInt(), 1); + EXPECT_THAT(qualifier.AsBool(), absl::nullopt); + EXPECT_THAT(qualifier.AsInt(), Optional(1)); + EXPECT_THAT(qualifier.AsUint(), absl::nullopt); + EXPECT_THAT(qualifier.AsString(), absl::nullopt); +} + +TEST(AttributeQualifierView, Uint) { + AttributeQualifierView qualifier = AttributeQualifierView::OfUint(1); + EXPECT_FALSE(qualifier.IsBool()); + EXPECT_FALSE(qualifier.IsInt()); + EXPECT_TRUE(qualifier.IsUint()); + EXPECT_FALSE(qualifier.IsString()); + EXPECT_EQ(qualifier.GetUint(), 1); + EXPECT_THAT(qualifier.AsBool(), absl::nullopt); + EXPECT_THAT(qualifier.AsInt(), absl::nullopt); + EXPECT_THAT(qualifier.AsUint(), Optional(1)); + EXPECT_THAT(qualifier.AsString(), absl::nullopt); +} + +TEST(AttributeQualifierView, String) { + AttributeQualifierView qualifier = AttributeQualifierView::OfString("foo"); + EXPECT_FALSE(qualifier.IsBool()); + EXPECT_FALSE(qualifier.IsInt()); + EXPECT_FALSE(qualifier.IsUint()); + EXPECT_TRUE(qualifier.IsString()); + EXPECT_EQ(qualifier.GetString(), "foo"); + EXPECT_THAT(qualifier.AsBool(), absl::nullopt); + EXPECT_THAT(qualifier.AsInt(), absl::nullopt); + EXPECT_THAT(qualifier.AsUint(), absl::nullopt); + EXPECT_THAT(qualifier.AsString(), Optional("foo")); +} + +} // namespace +} // namespace cel diff --git a/eval/public/cel_attribute_test.cc b/eval/public/cel_attribute_test.cc index b72189332..424dae92f 100644 --- a/eval/public/cel_attribute_test.cc +++ b/eval/public/cel_attribute_test.cc @@ -52,7 +52,7 @@ TEST(CelAttributeQualifierTest, TestBoolAccess) { EXPECT_FALSE(qualifier.GetUint64Key().has_value()); EXPECT_TRUE(qualifier.GetBoolKey().has_value()); EXPECT_THAT(qualifier.GetBoolKey().value(), Eq(true)); - EXPECT_THAT(qualifier.AsString(), IsOkAndHolds("true")); + EXPECT_THAT(qualifier.ToString(), "true"); } TEST(CelAttributeQualifierTest, TestInt64Access) { @@ -64,7 +64,7 @@ TEST(CelAttributeQualifierTest, TestInt64Access) { EXPECT_TRUE(qualifier.GetInt64Key().has_value()); EXPECT_THAT(qualifier.GetInt64Key().value(), Eq(-1)); - EXPECT_THAT(qualifier.AsString(), IsOkAndHolds("-1")); + EXPECT_THAT(qualifier.ToString(), "-1"); } TEST(CelAttributeQualifierTest, TestUint64Access) { @@ -76,7 +76,7 @@ TEST(CelAttributeQualifierTest, TestUint64Access) { EXPECT_TRUE(qualifier.GetUint64Key().has_value()); EXPECT_THAT(qualifier.GetUint64Key().value(), Eq(1UL)); - EXPECT_THAT(qualifier.AsString(), IsOkAndHolds("1")); + EXPECT_THAT(qualifier.ToString(), "1"); } TEST(CelAttributeQualifierTest, TestStringAccess) { @@ -89,7 +89,7 @@ TEST(CelAttributeQualifierTest, TestStringAccess) { EXPECT_TRUE(qualifier.GetStringKey().has_value()); EXPECT_THAT(qualifier.GetStringKey().value(), Eq("test")); - EXPECT_THAT(qualifier.AsString(), IsOkAndHolds("test")); + EXPECT_THAT(qualifier.ToString(), "test"); } void TestAllInequalities(const CelAttributeQualifier& qualifier) { diff --git a/extensions/protobuf/internal/qualify.cc b/extensions/protobuf/internal/qualify.cc index 37ad30011..1d53dac0d 100644 --- a/extensions/protobuf/internal/qualify.cc +++ b/extensions/protobuf/internal/qualify.cc @@ -339,11 +339,7 @@ absl::StatusOr ProtoQualifyState::CheckMapIn qualifier)); if (!value_ref.has_value()) { - std::string key_string; - absl::StatusOr key_string_or = qualifier.AsString(); - if (key_string_or.ok()) { - key_string = *key_string_or; - } + std::string key_string = qualifier.ToString(); return runtime_internal::CreateNoSuchKeyError(key_string); } return std::move(value_ref).value(); diff --git a/internal/BUILD b/internal/BUILD index 189853323..5b44efbad 100644 --- a/internal/BUILD +++ b/internal/BUILD @@ -42,6 +42,26 @@ cc_test( ], ) +cc_library( + name = "arena_tree", + srcs = ["arena_tree.cc"], + hdrs = ["arena_tree.h"], + deps = [ + "@com_google_absl//absl/base:nullability", + "@com_google_absl//absl/log:absl_check", + "@com_google_protobuf//:protobuf", + ], +) + +cc_test( + name = "arena_tree_test", + srcs = ["arena_tree_test.cc"], + deps = [ + ":arena_tree", + ":testing", + ], +) + cc_library( name = "new", srcs = ["new.cc"], diff --git a/internal/arena_tree.cc b/internal/arena_tree.cc new file mode 100644 index 000000000..9bd2922c3 --- /dev/null +++ b/internal/arena_tree.cc @@ -0,0 +1,352 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "internal/arena_tree.h" + +#include "absl/base/nullability.h" + +namespace cel::internal { + +namespace { + +void ArenaTreeRotateLeft(ArenaTreeNodeBase** head, ArenaTreeNodeBase* elem) { + ArenaTreeNodeBase* tmp = ArenaTreeNodeGetRight(elem); + if (ArenaTreeNodeSetRight(elem, ArenaTreeNodeGetLeft(tmp)) != nullptr) { + ArenaTreeNodeSetParent(ArenaTreeNodeGetLeft(tmp), elem); + } + ArenaTreeNodeBase* parent = ArenaTreeNodeGetParent(elem); + ArenaTreeNodeSetParent(tmp, parent); + if (parent != nullptr) { + if (elem == ArenaTreeNodeGetLeft(parent)) { + ArenaTreeNodeSetLeft(parent, tmp); + } else { + ArenaTreeNodeSetRight(parent, tmp); + } + } else { + *head = tmp; + } + ArenaTreeNodeSetLeft(tmp, elem); + ArenaTreeNodeSetParent(elem, tmp); +} + +void ArenaTreeRotateRight(ArenaTreeNodeBase** head, ArenaTreeNodeBase* elem) { + ArenaTreeNodeBase* tmp = ArenaTreeNodeGetLeft(elem); + if (ArenaTreeNodeSetLeft(elem, ArenaTreeNodeGetRight(tmp)) != nullptr) { + ArenaTreeNodeSetParent(ArenaTreeNodeGetRight(tmp), elem); + } + ArenaTreeNodeBase* parent = ArenaTreeNodeGetParent(elem); + ArenaTreeNodeSetParent(tmp, parent); + if (parent != nullptr) { + if (elem == ArenaTreeNodeGetLeft(parent)) { + ArenaTreeNodeSetLeft(parent, tmp); + } else { + ArenaTreeNodeSetRight(parent, tmp); + } + } else { + *head = tmp; + } + ArenaTreeNodeSetRight(tmp, elem); + ArenaTreeNodeSetParent(elem, tmp); +} + +void ArenaTreeRemoveColor(ArenaTreeNodeBase** head, ArenaTreeNodeBase* parent, + ArenaTreeNodeBase* elem) { + ArenaTreeNodeBase* tmp; + while ((elem == nullptr || + ArenaTreeNodeGetColor(elem) == ArenaTreeNodeColor::kBlack) && + elem != *head) { + if (ArenaTreeNodeGetLeft(parent) == elem) { + tmp = ArenaTreeNodeGetRight(parent); + if (ArenaTreeNodeGetColor(tmp) == ArenaTreeNodeColor::kRed) { + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeColor::kBlack); + ArenaTreeNodeSetColor(parent, ArenaTreeNodeColor::kRed); + ArenaTreeRotateLeft(head, parent); + tmp = ArenaTreeNodeGetRight(parent); + } + if ((ArenaTreeNodeGetLeft(tmp) == nullptr || + ArenaTreeNodeGetColor(ArenaTreeNodeGetLeft(tmp)) == + ArenaTreeNodeColor::kBlack) && + (ArenaTreeNodeGetRight(tmp) == nullptr || + ArenaTreeNodeGetColor(ArenaTreeNodeGetRight(tmp)) == + ArenaTreeNodeColor::kBlack)) { + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeColor::kRed); + elem = parent; + parent = ArenaTreeNodeGetParent(elem); + } else { + if (ArenaTreeNodeGetRight(tmp) == nullptr || + ArenaTreeNodeGetColor(ArenaTreeNodeGetRight(tmp)) == + ArenaTreeNodeColor::kBlack) { + ArenaTreeNodeBase* left; + if ((left = ArenaTreeNodeGetLeft(tmp)) != nullptr) { + ArenaTreeNodeSetColor(left, ArenaTreeNodeColor::kBlack); + } + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeColor::kRed); + ArenaTreeRotateRight(head, tmp); + tmp = ArenaTreeNodeGetRight(parent); + } + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeGetColor(parent)); + ArenaTreeNodeSetColor(parent, ArenaTreeNodeColor::kBlack); + if (ArenaTreeNodeGetRight(tmp) != nullptr) { + ArenaTreeNodeSetColor(ArenaTreeNodeGetRight(tmp), + ArenaTreeNodeColor::kBlack); + } + ArenaTreeRotateLeft(head, parent); + elem = *head; + break; + } + } else { + tmp = ArenaTreeNodeGetLeft(parent); + if (ArenaTreeNodeGetColor(tmp) == ArenaTreeNodeColor::kRed) { + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeColor::kBlack); + ArenaTreeNodeSetColor(parent, ArenaTreeNodeColor::kRed); + ArenaTreeRotateRight(head, parent); + tmp = ArenaTreeNodeGetLeft(parent); + } + if ((ArenaTreeNodeGetLeft(tmp) == nullptr || + ArenaTreeNodeGetColor(ArenaTreeNodeGetLeft(tmp)) == + ArenaTreeNodeColor::kBlack) && + (ArenaTreeNodeGetRight(tmp) == nullptr || + ArenaTreeNodeGetColor(ArenaTreeNodeGetRight(tmp)) == + ArenaTreeNodeColor::kBlack)) { + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeColor::kRed); + elem = parent; + parent = ArenaTreeNodeGetParent(elem); + } else { + if (ArenaTreeNodeGetLeft(tmp) == nullptr || + ArenaTreeNodeGetColor(ArenaTreeNodeGetLeft(tmp)) == + ArenaTreeNodeColor::kBlack) { + ArenaTreeNodeBase* right; + if ((right = ArenaTreeNodeGetRight(tmp)) != nullptr) { + ArenaTreeNodeSetColor(right, ArenaTreeNodeColor::kBlack); + } + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeColor::kRed); + ArenaTreeRotateLeft(head, tmp); + tmp = ArenaTreeNodeGetLeft(parent); + } + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeGetColor(parent)); + ArenaTreeNodeSetColor(parent, ArenaTreeNodeColor::kBlack); + if (ArenaTreeNodeGetLeft(tmp) != nullptr) { + ArenaTreeNodeSetColor(ArenaTreeNodeGetLeft(tmp), + ArenaTreeNodeColor::kBlack); + } + ArenaTreeRotateRight(head, parent); + elem = *head; + break; + } + } + } + if (elem != nullptr) { + ArenaTreeNodeSetColor(elem, ArenaTreeNodeColor::kBlack); + } +} + +} // namespace + +const ArenaTreeNodeBase* absl_nullable ArenaTreeNext( + const ArenaTreeNodeBase* absl_nullable node) { + if (node != nullptr) { + const ArenaTreeNodeBase* right = ArenaTreeNodeGetRight(node); + if (right != nullptr) { + node = right; + const ArenaTreeNodeBase* left; + while ((left = ArenaTreeNodeGetLeft(node)) != nullptr) { + node = left; + } + } else { + const ArenaTreeNodeBase* parent = ArenaTreeNodeGetParent(node); + if (parent != nullptr && node == ArenaTreeNodeGetLeft(parent)) { + node = parent; + } else { + while (parent != nullptr && ArenaTreeNodeGetRight(parent) == node) { + node = ArenaTreeNodeGetParent(node); + parent = ArenaTreeNodeGetParent(node); + } + node = parent; + } + } + } + return node; +} + +const ArenaTreeNodeBase* absl_nullable ArenaTreePrev( + const ArenaTreeNodeBase* absl_nullable node) { + if (node != nullptr) { + const ArenaTreeNodeBase* left = ArenaTreeNodeGetLeft(node); + if (left != nullptr) { + node = left; + const ArenaTreeNodeBase* right; + while ((right = ArenaTreeNodeGetRight(node)) != nullptr) { + node = right; + } + } else { + const ArenaTreeNodeBase* parent = ArenaTreeNodeGetParent(node); + if (parent != nullptr && node == ArenaTreeNodeGetRight(parent)) { + node = parent; + } else { + while (parent != nullptr && ArenaTreeNodeGetLeft(parent) == node) { + node = ArenaTreeNodeGetParent(node); + parent = ArenaTreeNodeGetParent(node); + } + node = parent; + } + } + } + return node; +} + +const ArenaTreeNodeBase* absl_nullable ArenaTreeMin( + const ArenaTreeNodeBase* absl_nullable node) { + const ArenaTreeNodeBase* tmp = node; + const ArenaTreeNodeBase* parent = nullptr; + while (tmp != nullptr) { + parent = tmp; + tmp = ArenaTreeNodeGetLeft(tmp); + } + return parent; +} + +const ArenaTreeNodeBase* absl_nullable ArenaTreeMax( + const ArenaTreeNodeBase* absl_nullable node) { + const ArenaTreeNodeBase* tmp = node; + const ArenaTreeNodeBase* parent = nullptr; + while (tmp != nullptr) { + parent = tmp; + tmp = ArenaTreeNodeGetRight(tmp); + } + return parent; +} + +void ArenaTreeRemove(ArenaTreeNodeBase* absl_nullable* absl_nonnull head, + ArenaTreeNodeBase* absl_nonnull elem) { + ArenaTreeNodeBase* child; + ArenaTreeNodeBase* parent; + ArenaTreeNodeBase* const old = elem; + ArenaTreeNodeColor color; + if (ArenaTreeNodeGetLeft(elem) == nullptr) { + child = ArenaTreeNodeGetRight(elem); + } else if (ArenaTreeNodeGetRight(elem) == nullptr) { + child = ArenaTreeNodeGetLeft(elem); + } else { + ArenaTreeNodeBase* left; + elem = ArenaTreeNodeGetRight(elem); + while ((left = ArenaTreeNodeGetLeft(elem)) != nullptr) { + elem = left; + } + child = ArenaTreeNodeGetRight(elem); + parent = ArenaTreeNodeGetParent(elem); + color = ArenaTreeNodeGetColor(elem); + if (child != nullptr) { + ArenaTreeNodeSetParent(child, parent); + } + if (parent != nullptr) { + if (ArenaTreeNodeGetLeft(parent) == elem) { + ArenaTreeNodeSetLeft(parent, child); + } else { + ArenaTreeNodeSetRight(parent, child); + } + } else { + *head = child; + } + if (ArenaTreeNodeGetParent(elem) == old) { + parent = elem; + } + ArenaTreeNodeBase* old_parent = ArenaTreeNodeGetParent(old); + if (old_parent != nullptr) { + if (ArenaTreeNodeGetLeft(old_parent) == old) { + ArenaTreeNodeSetLeft(old_parent, elem); + } else { + ArenaTreeNodeSetRight(old_parent, elem); + } + } else { + *head = elem; + } + ArenaTreeNodeSetParent(ArenaTreeNodeGetLeft(old), elem); + if (ArenaTreeNodeGetRight(old) != nullptr) { + ArenaTreeNodeSetParent(ArenaTreeNodeGetRight(old), elem); + } + goto color; + } + parent = ArenaTreeNodeGetParent(elem); + color = ArenaTreeNodeGetColor(elem); + if (child != nullptr) { + ArenaTreeNodeSetParent(child, parent); + } + if (parent != nullptr) { + if (ArenaTreeNodeGetLeft(parent) == elem) { + ArenaTreeNodeSetLeft(parent, child); + } else { + ArenaTreeNodeSetRight(parent, child); + } + } else { + *head = child; + } +color: + if (color == ArenaTreeNodeColor::kBlack) { + ArenaTreeRemoveColor(head, parent, child); + } + ArenaTreeNodeClear(old); +} + +void ArenaTreeInsertColor(ArenaTreeNodeBase* absl_nullable* absl_nonnull head, + ArenaTreeNodeBase* absl_nonnull elem) { + ArenaTreeNodeBase* parent; + ArenaTreeNodeBase* grandparent; + ArenaTreeNodeBase* tmp; + while ((parent = ArenaTreeNodeGetParent(elem)) != nullptr && + ArenaTreeNodeGetColor(parent) == ArenaTreeNodeColor::kRed) { + grandparent = ArenaTreeNodeGetParent(parent); + if (parent == ArenaTreeNodeGetLeft(grandparent)) { + tmp = ArenaTreeNodeGetRight(grandparent); + if (tmp != nullptr && + ArenaTreeNodeGetColor(tmp) == ArenaTreeNodeColor::kRed) { + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeColor::kBlack); + ArenaTreeNodeSetColor(parent, ArenaTreeNodeColor::kBlack); + ArenaTreeNodeSetColor(grandparent, ArenaTreeNodeColor::kRed); + elem = grandparent; + continue; + } + if (ArenaTreeNodeGetRight(parent) == elem) { + ArenaTreeRotateLeft(head, parent); + tmp = parent; + parent = elem; + elem = tmp; + } + ArenaTreeNodeSetColor(parent, ArenaTreeNodeColor::kBlack); + ArenaTreeNodeSetColor(grandparent, ArenaTreeNodeColor::kRed); + ArenaTreeRotateRight(head, grandparent); + } else { + tmp = ArenaTreeNodeGetLeft(grandparent); + if (tmp != nullptr && + ArenaTreeNodeGetColor(tmp) == ArenaTreeNodeColor::kRed) { + ArenaTreeNodeSetColor(tmp, ArenaTreeNodeColor::kBlack); + ArenaTreeNodeSetColor(parent, ArenaTreeNodeColor::kBlack); + ArenaTreeNodeSetColor(grandparent, ArenaTreeNodeColor::kRed); + elem = grandparent; + continue; + } + if (ArenaTreeNodeGetLeft(parent) == elem) { + ArenaTreeRotateRight(head, parent); + tmp = parent; + parent = elem; + elem = tmp; + } + ArenaTreeNodeSetColor(parent, ArenaTreeNodeColor::kBlack); + ArenaTreeNodeSetColor(grandparent, ArenaTreeNodeColor::kRed); + ArenaTreeRotateLeft(head, grandparent); + } + } + ArenaTreeNodeSetColor(*head, ArenaTreeNodeColor::kBlack); +} + +} // namespace cel::internal diff --git a/internal/arena_tree.h b/internal/arena_tree.h new file mode 100644 index 000000000..9e76bd84a --- /dev/null +++ b/internal/arena_tree.h @@ -0,0 +1,398 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// ArenaTree is a low level implementation of an RBTree tailored for use with +// google::protobuf::Arena. + +#ifndef THIRD_PARTY_CEL_CPP_INTERNAL_ARENA_TREE_H_ +#define THIRD_PARTY_CEL_CPP_INTERNAL_ARENA_TREE_H_ + +#include +#include +#include +#include + +#include "absl/base/nullability.h" +#include "absl/log/absl_check.h" +#include "google/protobuf/arena.h" + +namespace cel::internal { + +enum class ArenaTreeNodeColor : uintptr_t { + kBlack = 0, + kRed = 1, +}; + +struct ArenaTreeNodeBase; + +[[nodiscard]] +const ArenaTreeNodeBase* absl_nullable ArenaTreeNext( + const ArenaTreeNodeBase* absl_nullable node); + +[[nodiscard]] +const ArenaTreeNodeBase* absl_nullable ArenaTreePrev( + const ArenaTreeNodeBase* absl_nullable node); + +[[nodiscard]] +const ArenaTreeNodeBase* absl_nullable ArenaTreeMin( + const ArenaTreeNodeBase* absl_nullable node); + +[[nodiscard]] +const ArenaTreeNodeBase* absl_nullable ArenaTreeMax( + const ArenaTreeNodeBase* absl_nullable node); + +template +using IsDerivedFromArenaTreeNodeBase = std::conjunction< + std::is_base_of, + std::negation>>>; + +template +constexpr bool kIsDerivedFromArenaTreeNodeBase = + IsDerivedFromArenaTreeNodeBase::value; + +template +using EnableIfDerivedFromArenaTreeNodeBase = + std::enable_if_t, U>; + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreeNext( + T* absl_nullable node) { + return static_cast(const_cast( + (ArenaTreeNext)(static_cast(node)))); +} + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreePrev( + T* absl_nullable node) { + return static_cast(const_cast( + (ArenaTreePrev)(static_cast(node)))); +} + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreeMin( + T* absl_nullable node) { + return static_cast(const_cast( + (ArenaTreeMin)(static_cast(node)))); +} + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreeMax( + T* absl_nullable node) { + return static_cast(const_cast( + (ArenaTreeMax)(static_cast(node)))); +} + +[[nodiscard]] +ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetParent( + const ArenaTreeNodeBase* absl_nonnull node); + +void ArenaTreeNodeSetParent(ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown parent); + +[[nodiscard]] +ArenaTreeNodeColor ArenaTreeNodeGetColor( + const ArenaTreeNodeBase* absl_nonnull node); + +void ArenaTreeNodeSetColor(ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeColor color); + +void ArenaTreeNodeSet(ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullable parent); + +void ArenaTreeNodeClear(ArenaTreeNodeBase* absl_nonnull node); + +[[nodiscard]] +ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetLeft( + const ArenaTreeNodeBase* absl_nonnull node); + +ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeSetLeft( + ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown left); + +[[nodiscard]] +ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetRight( + const ArenaTreeNodeBase* absl_nonnull node); + +ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeSetRight( + ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown right); + +struct ArenaTreeNodeBase { + private: + uintptr_t parent_and_color = 0; + ArenaTreeNodeBase* absl_nullable left = nullptr; + ArenaTreeNodeBase* absl_nullable right = nullptr; + + friend ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetParent( + const ArenaTreeNodeBase* absl_nonnull node); + friend void ArenaTreeNodeSetParent( + ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown parent); + friend ArenaTreeNodeColor ArenaTreeNodeGetColor( + const ArenaTreeNodeBase* absl_nonnull node); + friend void ArenaTreeNodeSetColor(ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeColor color); + friend void ArenaTreeNodeSet(ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullable parent); + friend void ArenaTreeNodeClear(ArenaTreeNodeBase* absl_nonnull node); + friend ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetLeft( + const ArenaTreeNodeBase* absl_nonnull node); + friend ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeSetLeft( + ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown left); + friend ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetRight( + const ArenaTreeNodeBase* absl_nonnull node); + friend ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeSetRight( + ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown right); +}; + +[[nodiscard]] +inline ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetParent( + const ArenaTreeNodeBase* absl_nonnull node) { + return reinterpret_cast(node->parent_and_color & + ~uintptr_t{1}); +} + +inline void ArenaTreeNodeSetParent( + ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown parent) { + node->parent_and_color = static_cast(ArenaTreeNodeGetColor(node)) | + reinterpret_cast(parent); +} + +[[nodiscard]] +inline ArenaTreeNodeColor ArenaTreeNodeGetColor( + const ArenaTreeNodeBase* absl_nonnull node) { + return static_cast(node->parent_and_color & uintptr_t{1}); +} + +inline void ArenaTreeNodeSetColor(ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeColor color) { + node->parent_and_color = + reinterpret_cast(ArenaTreeNodeGetParent(node)) | + static_cast(color); +} + +inline void ArenaTreeNodeSet(ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullable parent) { + node->parent_and_color = reinterpret_cast(parent) | + static_cast(ArenaTreeNodeColor::kRed); + node->left = node->right = nullptr; +} + +inline void ArenaTreeNodeClear(ArenaTreeNodeBase* absl_nonnull node) { + node->parent_and_color = 0; + node->left = node->right = nullptr; +} + +[[nodiscard]] +inline ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetLeft( + const ArenaTreeNodeBase* absl_nonnull node) { + return node->left; +} + +inline ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeSetLeft( + ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown left) { + return node->left = left; +} + +[[nodiscard]] +inline ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeGetRight( + const ArenaTreeNodeBase* absl_nonnull node) { + return node->right; +} + +inline ArenaTreeNodeBase* absl_nullability_unknown ArenaTreeNodeSetRight( + ArenaTreeNodeBase* absl_nonnull node, + ArenaTreeNodeBase* absl_nullability_unknown right) { + return node->right = right; +} + +void ArenaTreeRemove(ArenaTreeNodeBase* absl_nullable* absl_nonnull head, + ArenaTreeNodeBase* absl_nonnull elem); + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreeRemove( + T* absl_nullable* absl_nonnull head, T* absl_nonnull elem) { + ArenaTreeNodeBase* head_base = *head; + (ArenaTreeRemove)(&head_base, static_cast(elem)); + *head = static_cast(head_base); +} + +void ArenaTreeInsertColor(ArenaTreeNodeBase* absl_nullable* absl_nonnull head, + ArenaTreeNodeBase* absl_nonnull elem); + +template +struct ArenaTreeNode; + +template +[[nodiscard]] +const T& ArenaTreeNodeGetValue(const ArenaTreeNode* absl_nonnull node); + +template +struct ArenaTreeNode : ArenaTreeNodeBase { + template + explicit ArenaTreeNode(Args&&... args) + : ArenaTreeNodeBase(), value(std::forward(args)...) {} + + private: + template + friend const U& ArenaTreeNodeGetValue( + const ArenaTreeNode* absl_nonnull node); + + T value; +}; + +template +[[nodiscard]] inline const T& ArenaTreeNodeGetValue( + const ArenaTreeNode* absl_nonnull node) { + return node->value; +} + +template +struct ArenaTreeNodeCrtp : ArenaTreeNodeBase { + using ArenaTreeNodeBase::ArenaTreeNodeBase; +}; + +template +[[nodiscard]] inline const T& ArenaTreeNodeGetValue( + const ArenaTreeNodeCrtp* absl_nonnull node) { + return *static_cast(node); +} + +template +[[nodiscard]] +T* absl_nonnull ArenaTreeInsert(T* absl_nullable* absl_nonnull head, + T* absl_nonnull elem, const Compare& compare) { + T* tmp = *head; + T* parent = nullptr; + int diff = 0; + while (tmp != nullptr) { + parent = tmp; + diff = std::invoke(compare, (ArenaTreeNodeGetValue)(elem), + (ArenaTreeNodeGetValue)(parent)); + if (diff < 0) { + tmp = static_cast((ArenaTreeNodeGetLeft)(tmp)); + } else if (diff > 0) { + tmp = static_cast((ArenaTreeNodeGetRight)(tmp)); + } else { + return tmp; + } + } + (ArenaTreeNodeSet)(elem, parent); + if (parent != nullptr) { + if (diff < 0) { + (ArenaTreeNodeSetLeft)(parent, elem); + } else { + (ArenaTreeNodeSetRight)(parent, elem); + } + } else { + *head = elem; + } + ArenaTreeNodeBase* head_base = *head; + (ArenaTreeInsertColor)(&head_base, elem); + *head = static_cast(head_base); + return elem; +} + +template +struct ArenaTreeNodeConstructor { + template + void operator()(Args&&... args) const { + ABSL_DCHECK(*out == nullptr); + *out = google::protobuf::Arena::Create(arena, std::forward(args)...); + } + + google::protobuf::Arena* const absl_nonnull arena; + T** out; +}; + +template +[[nodiscard]] +std::pair ArenaTreeLazyEmplace( + google::protobuf::Arena* absl_nonnull arena, T* absl_nullable* absl_nonnull head, + const K& key, const Compare& compare, Emplacer&& emplacer) { + T* tmp = *head; + T* parent = nullptr; + int diff = 0; + while (tmp != nullptr) { + parent = tmp; + diff = std::invoke(compare, key, (ArenaTreeNodeGetValue)(parent)); + if (diff < 0) { + tmp = static_cast((ArenaTreeNodeGetLeft)(tmp)); + } else if (diff > 0) { + tmp = static_cast((ArenaTreeNodeGetRight)(tmp)); + } else { + return {tmp, false}; + } + } + T* elem = nullptr; + ArenaTreeNodeConstructor constructor{ + .arena = arena, + .out = &elem, + }; + std::invoke(std::forward(emplacer), + static_cast&>(constructor)); + ABSL_DCHECK(elem != nullptr); + (ArenaTreeNodeSet)(elem, parent); + if (parent != nullptr) { + if (diff < 0) { + (ArenaTreeNodeSetLeft)(parent, elem); + } else { + (ArenaTreeNodeSetRight)(parent, elem); + } + } else { + *head = elem; + } + ArenaTreeNodeBase* head_base = *head; + (ArenaTreeInsertColor)(&head_base, elem); + *head = static_cast(head_base); + return {elem, true}; +} + +template +[[nodiscard]] +const T* absl_nullable ArenaTreeFind(const T* absl_nullable head, const K& key, + const Compare& compare) { + const T* tmp = head; + while (tmp != nullptr) { + int diff = std::invoke(compare, key, (ArenaTreeNodeGetValue)(tmp)); + if (diff < 0) { + tmp = static_cast((ArenaTreeNodeGetLeft)(tmp)); + } else if (diff > 0) { + tmp = static_cast((ArenaTreeNodeGetRight)(tmp)); + } else { + return tmp; + } + } + return nullptr; +} + +template +[[nodiscard]] +T* absl_nullable ArenaTreeFind(T* absl_nullable head, const K& key, + const Compare& compare) { + return (ArenaTreeFind)(static_cast(head), key, compare); +} + +} // namespace cel::internal + +#endif // THIRD_PARTY_CEL_CPP_INTERNAL_ARENA_TREE_H_ diff --git a/internal/arena_tree_test.cc b/internal/arena_tree_test.cc new file mode 100644 index 000000000..ddc6b9885 --- /dev/null +++ b/internal/arena_tree_test.cc @@ -0,0 +1,223 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "internal/arena_tree.h" + +#include + +#include "internal/testing.h" + +namespace cel::internal { +namespace { + +using ::testing::IsNull; + +using TestNode = ArenaTreeNode; + +struct TestNodeCompare { + int operator()(int lhs, int rhs) const { + if (lhs < rhs) { + return -1; + } + if (lhs > rhs) { + return 1; + } + return 0; + } +}; + +struct TestNodeCrtp : ArenaTreeNodeCrtp { + explicit TestNodeCrtp(int value) : ArenaTreeNodeCrtp(), value(value) {} + + int value; +}; + +struct TestNodeCrtpCompare { + int operator()(int lhs, int rhs) const { + if (lhs < rhs) { + return -1; + } + if (lhs > rhs) { + return 1; + } + return 0; + } + + int operator()(int lhs, const TestNodeCrtp& rhs) const { + return (*this)(lhs, rhs.value); + } + + int operator()(const TestNodeCrtp& lhs, int rhs) const { + return (*this)(lhs.value, rhs); + } + + int operator()(const TestNodeCrtp& lhs, const TestNodeCrtp& rhs) const { + return (*this)(lhs.value, rhs.value); + } +}; + +TEST(ArenaTree, Empty) { + TestNode* head = nullptr; + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); +} + +TEST(ArenaTree, Single) { + TestNode* head = nullptr; + std::unique_ptr node1 = std::make_unique(1); + EXPECT_EQ(ArenaTreeInsert(&head, node1.get(), TestNodeCompare{}), + node1.get()); + EXPECT_EQ(ArenaTreePrev(head), nullptr); + EXPECT_EQ(ArenaTreeNext(head), nullptr); + EXPECT_EQ(ArenaTreeMin(head), node1.get()); + EXPECT_EQ(ArenaTreeMax(head), node1.get()); + ArenaTreeRemove(&head, node1.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); +} + +TEST(ArenaTree, Couple) { + TestNode* head = nullptr; + std::unique_ptr node1 = std::make_unique(1); + std::unique_ptr node2 = std::make_unique(2); + + EXPECT_EQ(ArenaTreeInsert(&head, node1.get(), TestNodeCompare{}), + node1.get()); + EXPECT_EQ(ArenaTreeInsert(&head, node2.get(), TestNodeCompare{}), + node2.get()); + EXPECT_EQ(ArenaTreePrev(node1.get()), nullptr); + EXPECT_EQ(ArenaTreePrev(node2.get()), node1.get()); + EXPECT_EQ(ArenaTreeNext(node1.get()), node2.get()); + EXPECT_EQ(ArenaTreeNext(node2.get()), nullptr); + EXPECT_EQ(ArenaTreeMin(head), node1.get()); + EXPECT_EQ(ArenaTreeMax(head), node2.get()); + + ArenaTreeRemove(&head, node1.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_EQ(ArenaTreeMin(head), node2.get()); + EXPECT_EQ(ArenaTreeMax(head), node2.get()); + + ArenaTreeRemove(&head, node2.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); + + EXPECT_EQ(ArenaTreeInsert(&head, node2.get(), TestNodeCompare{}), + node2.get()); + EXPECT_EQ(ArenaTreeInsert(&head, node1.get(), TestNodeCompare{}), + node1.get()); + EXPECT_EQ(ArenaTreePrev(node1.get()), nullptr); + EXPECT_EQ(ArenaTreePrev(node2.get()), node1.get()); + EXPECT_EQ(ArenaTreeNext(node1.get()), node2.get()); + EXPECT_EQ(ArenaTreeNext(node2.get()), nullptr); + EXPECT_EQ(ArenaTreeMin(head), node1.get()); + EXPECT_EQ(ArenaTreeMax(head), node2.get()); + + ArenaTreeRemove(&head, node2.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_EQ(ArenaTreeMin(head), node1.get()); + EXPECT_EQ(ArenaTreeMax(head), node1.get()); + + ArenaTreeRemove(&head, node1.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); +} + +TEST(ArenaTree, CrtpEmpty) { + TestNodeCrtp* head = nullptr; + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); +} + +TEST(ArenaTree, CrtpSingle) { + TestNodeCrtp* head = nullptr; + std::unique_ptr node1 = std::make_unique(1); + EXPECT_EQ(ArenaTreeInsert(&head, node1.get(), TestNodeCrtpCompare{}), + node1.get()); + EXPECT_EQ(ArenaTreePrev(head), nullptr); + EXPECT_EQ(ArenaTreeNext(head), nullptr); + EXPECT_EQ(ArenaTreeMin(head), node1.get()); + EXPECT_EQ(ArenaTreeMax(head), node1.get()); + ArenaTreeRemove(&head, node1.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); +} + +TEST(ArenaTree, CrtpCouple) { + TestNodeCrtp* head = nullptr; + std::unique_ptr node1 = std::make_unique(1); + std::unique_ptr node2 = std::make_unique(2); + + EXPECT_EQ(ArenaTreeInsert(&head, node1.get(), TestNodeCrtpCompare{}), + node1.get()); + EXPECT_EQ(ArenaTreeInsert(&head, node2.get(), TestNodeCrtpCompare{}), + node2.get()); + EXPECT_EQ(ArenaTreePrev(node1.get()), nullptr); + EXPECT_EQ(ArenaTreePrev(node2.get()), node1.get()); + EXPECT_EQ(ArenaTreeNext(node1.get()), node2.get()); + EXPECT_EQ(ArenaTreeNext(node2.get()), nullptr); + EXPECT_EQ(ArenaTreeMin(head), node1.get()); + EXPECT_EQ(ArenaTreeMax(head), node2.get()); + + ArenaTreeRemove(&head, node1.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_EQ(ArenaTreeMin(head), node2.get()); + EXPECT_EQ(ArenaTreeMax(head), node2.get()); + + ArenaTreeRemove(&head, node2.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); + + EXPECT_EQ(ArenaTreeInsert(&head, node2.get(), TestNodeCrtpCompare{}), + node2.get()); + EXPECT_EQ(ArenaTreeInsert(&head, node1.get(), TestNodeCrtpCompare{}), + node1.get()); + EXPECT_EQ(ArenaTreePrev(node1.get()), nullptr); + EXPECT_EQ(ArenaTreePrev(node2.get()), node1.get()); + EXPECT_EQ(ArenaTreeNext(node1.get()), node2.get()); + EXPECT_EQ(ArenaTreeNext(node2.get()), nullptr); + EXPECT_EQ(ArenaTreeMin(head), node1.get()); + EXPECT_EQ(ArenaTreeMax(head), node2.get()); + + ArenaTreeRemove(&head, node2.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_EQ(ArenaTreeMin(head), node1.get()); + EXPECT_EQ(ArenaTreeMax(head), node1.get()); + + ArenaTreeRemove(&head, node1.get()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); +} + +} // namespace +} // namespace cel::internal