From 2ff1fbd238f09bf625f705f38c1cb365dae90e41 Mon Sep 17 00:00:00 2001 From: Justin King Date: Tue, 29 Sep 2026 10:20:57 -0700 Subject: [PATCH] Introduce RBTree implementation tailored for `proto2::Arena` PiperOrigin-RevId: 990385693 --- internal/BUILD | 21 ++ internal/arena_tree.cc | 345 ++++++++++++++++++++ internal/arena_tree.h | 390 ++++++++++++++++++++++ internal/arena_tree_test.cc | 627 ++++++++++++++++++++++++++++++++++++ 4 files changed, 1383 insertions(+) create mode 100644 internal/arena_tree.cc create mode 100644 internal/arena_tree.h create mode 100644 internal/arena_tree_test.cc diff --git a/internal/BUILD b/internal/BUILD index 189853323..d6781b538 100644 --- a/internal/BUILD +++ b/internal/BUILD @@ -42,6 +42,27 @@ 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", + "@com_google_protobuf//:protobuf", + ], +) + 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..8db8fde6c --- /dev/null +++ b/internal/arena_tree.cc @@ -0,0 +1,345 @@ +// 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" + +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* ArenaTreeNext(const ArenaTreeNodeBase* 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* ArenaTreePrev(const ArenaTreeNodeBase* 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* ArenaTreeMin(const ArenaTreeNodeBase* node) { + const ArenaTreeNodeBase* tmp = node; + const ArenaTreeNodeBase* parent = nullptr; + while (tmp != nullptr) { + parent = tmp; + tmp = ArenaTreeNodeGetLeft(tmp); + } + return parent; +} + +const ArenaTreeNodeBase* ArenaTreeMax(const ArenaTreeNodeBase* node) { + const ArenaTreeNodeBase* tmp = node; + const ArenaTreeNodeBase* parent = nullptr; + while (tmp != nullptr) { + parent = tmp; + tmp = ArenaTreeNodeGetRight(tmp); + } + return parent; +} + +void ArenaTreeRemove(ArenaTreeNodeBase** head, ArenaTreeNodeBase* 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; + } + ArenaTreeNodeCopy(elem, old); + 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** head, ArenaTreeNodeBase* 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..05f0e74bd --- /dev/null +++ b/internal/arena_tree.h @@ -0,0 +1,390 @@ +// 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. It avoids needing to deal with using a custom Allocator with +// std::set, std::map, absl::btree_set, or absl::btree_map, where the options +// are manually construct on the arena to avoid registering the destructor +// (technically undefined behavior) or registering the destructor and running it +// even though it is a no-op at runtime. Using absl::btree_set or +// absl::btree_map also does not preserve pointer stability and they have a +// higher overhead for small containers. +// +// Note that you should not use this unless you are really sure you need to. In +// most cases you should prefer std::set, std::map, absl::btree_set, or +// absl::btree_map. + +#ifndef THIRD_PARTY_CEL_CPP_INTERNAL_ARENA_TREE_H_ +#define THIRD_PARTY_CEL_CPP_INTERNAL_ARENA_TREE_H_ + +#include +#include +#include +#include + +#include "absl/log/absl_check.h" +#include "google/protobuf/arena.h" + +namespace cel::internal { + +enum class ArenaTreeNodeColor : uintptr_t { + kBlack = 0, + kRed = 1, +}; + +class ArenaTreeNodeBase; + +[[nodiscard]] +const ArenaTreeNodeBase* ArenaTreeNext(const ArenaTreeNodeBase* node); + +[[nodiscard]] +const ArenaTreeNodeBase* ArenaTreePrev(const ArenaTreeNodeBase* node); + +[[nodiscard]] +const ArenaTreeNodeBase* ArenaTreeMin(const ArenaTreeNodeBase* node); + +[[nodiscard]] +const ArenaTreeNodeBase* ArenaTreeMax(const ArenaTreeNodeBase* 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* node) { + return static_cast(const_cast( + (ArenaTreeNext)(static_cast(node)))); +} + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreePrev(T* node) { + return static_cast(const_cast( + (ArenaTreePrev)(static_cast(node)))); +} + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreeMin(T* node) { + return static_cast(const_cast( + (ArenaTreeMin)(static_cast(node)))); +} + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreeMax(T* node) { + return static_cast(const_cast( + (ArenaTreeMax)(static_cast(node)))); +} + +[[nodiscard]] +ArenaTreeNodeBase* ArenaTreeNodeGetParent(const ArenaTreeNodeBase* node); + +void ArenaTreeNodeSetParent(ArenaTreeNodeBase* node, ArenaTreeNodeBase* parent); + +[[nodiscard]] +ArenaTreeNodeColor ArenaTreeNodeGetColor(const ArenaTreeNodeBase* node); + +void ArenaTreeNodeSetColor(ArenaTreeNodeBase* node, ArenaTreeNodeColor color); + +void ArenaTreeNodeSet(ArenaTreeNodeBase* node, ArenaTreeNodeBase* parent); + +void ArenaTreeNodeCopy(ArenaTreeNodeBase* dst, const ArenaTreeNodeBase* src); + +void ArenaTreeNodeClear(ArenaTreeNodeBase* node); + +[[nodiscard]] +ArenaTreeNodeBase* ArenaTreeNodeGetLeft(const ArenaTreeNodeBase* node); + +ArenaTreeNodeBase* ArenaTreeNodeSetLeft(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* left); + +[[nodiscard]] +ArenaTreeNodeBase* ArenaTreeNodeGetRight(const ArenaTreeNodeBase* node); + +ArenaTreeNodeBase* ArenaTreeNodeSetRight(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* right); + +class ArenaTreeNodeBase { + private: + uintptr_t parent_and_color_ = 0; + ArenaTreeNodeBase* left_ = nullptr; + ArenaTreeNodeBase* right_ = nullptr; + + friend ArenaTreeNodeBase* ArenaTreeNodeGetParent( + const ArenaTreeNodeBase* node); + friend void ArenaTreeNodeSetParent(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* parent); + friend ArenaTreeNodeColor ArenaTreeNodeGetColor( + const ArenaTreeNodeBase* node); + friend void ArenaTreeNodeSetColor(ArenaTreeNodeBase* node, + ArenaTreeNodeColor color); + friend void ArenaTreeNodeSet(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* parent); + friend void ArenaTreeNodeClear(ArenaTreeNodeBase* node); + friend void ArenaTreeNodeCopy(ArenaTreeNodeBase* dst, + const ArenaTreeNodeBase* src); + friend ArenaTreeNodeBase* ArenaTreeNodeGetLeft(const ArenaTreeNodeBase* node); + friend ArenaTreeNodeBase* ArenaTreeNodeSetLeft(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* left); + friend ArenaTreeNodeBase* ArenaTreeNodeGetRight( + const ArenaTreeNodeBase* node); + friend ArenaTreeNodeBase* ArenaTreeNodeSetRight(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* right); +}; + +[[nodiscard]] +inline ArenaTreeNodeBase* ArenaTreeNodeGetParent( + const ArenaTreeNodeBase* node) { + return reinterpret_cast(node->parent_and_color_ & + ~uintptr_t{1}); +} + +inline void ArenaTreeNodeSetParent(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* parent) { + node->parent_and_color_ = + static_cast(ArenaTreeNodeGetColor(node)) | + reinterpret_cast(parent); +} + +[[nodiscard]] +inline ArenaTreeNodeColor ArenaTreeNodeGetColor(const ArenaTreeNodeBase* node) { + return static_cast(node->parent_and_color_ & + uintptr_t{1}); +} + +inline void ArenaTreeNodeSetColor(ArenaTreeNodeBase* node, + ArenaTreeNodeColor color) { + node->parent_and_color_ = + reinterpret_cast(ArenaTreeNodeGetParent(node)) | + static_cast(color); +} + +inline void ArenaTreeNodeSet(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* parent) { + node->parent_and_color_ = reinterpret_cast(parent) | + static_cast(ArenaTreeNodeColor::kRed); + node->left_ = node->right_ = nullptr; +} + +inline void ArenaTreeNodeClear(ArenaTreeNodeBase* node) { + node->parent_and_color_ = 0; + node->left_ = node->right_ = nullptr; +} + +inline void ArenaTreeNodeCopy(ArenaTreeNodeBase* dst, + const ArenaTreeNodeBase* src) { + dst->parent_and_color_ = src->parent_and_color_; + dst->left_ = src->left_; + dst->right_ = src->right_; +} + +[[nodiscard]] +inline ArenaTreeNodeBase* ArenaTreeNodeGetLeft(const ArenaTreeNodeBase* node) { + return node->left_; +} + +inline ArenaTreeNodeBase* ArenaTreeNodeSetLeft(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* left) { + return node->left_ = left; +} + +[[nodiscard]] +inline ArenaTreeNodeBase* ArenaTreeNodeGetRight(const ArenaTreeNodeBase* node) { + return node->right_; +} + +inline ArenaTreeNodeBase* ArenaTreeNodeSetRight(ArenaTreeNodeBase* node, + ArenaTreeNodeBase* right) { + return node->right_ = right; +} + +void ArenaTreeRemove(ArenaTreeNodeBase** head, ArenaTreeNodeBase* elem); + +template +[[nodiscard]] +inline EnableIfDerivedFromArenaTreeNodeBase ArenaTreeRemove(T** head, + T* elem) { + ArenaTreeNodeBase* head_base = *head; + (ArenaTreeRemove)(&head_base, static_cast(elem)); + *head = static_cast(head_base); +} + +void ArenaTreeInsertColor(ArenaTreeNodeBase** head, ArenaTreeNodeBase* elem); + +template +class ArenaTreeNode; + +template +[[nodiscard]] +const T& ArenaTreeNodeGetValue(const ArenaTreeNode* node); + +template +class ArenaTreeNode : public ArenaTreeNodeBase { + public: + template + explicit ArenaTreeNode(Args&&... args) + : ArenaTreeNodeBase(), value_(std::forward(args)...) {} + + private: + template + friend const U& ArenaTreeNodeGetValue(const ArenaTreeNode* node); + + T value_; +}; + +template +[[nodiscard]] inline const T& ArenaTreeNodeGetValue( + const ArenaTreeNode* node) { + return node->value_; +} + +template +class ArenaTreeNodeCrtp : public ArenaTreeNodeBase { + public: + using ArenaTreeNodeBase::ArenaTreeNodeBase; +}; + +template +[[nodiscard]] inline const T& ArenaTreeNodeGetValue( + const ArenaTreeNodeCrtp* node) { + return *static_cast(node); +} + +template +[[nodiscard]] +T* ArenaTreeInsert(T** head, T* 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 arena; + T** out; +}; + +template +[[nodiscard]] +std::pair ArenaTreeLazyEmplace(google::protobuf::Arena* arena, T** 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* ArenaTreeFind(const T* 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* ArenaTreeFind(T* head, const K& key, const Compare& compare) { + return const_cast( + (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..9f6dfdbf7 --- /dev/null +++ b/internal/arena_tree_test.cc @@ -0,0 +1,627 @@ +// 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 +#include +#include +#include + +#include "internal/testing.h" +#include "google/protobuf/arena.h" + +namespace cel::internal { +namespace { + +using ::testing::ElementsAre; +using ::testing::IsEmpty; +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); + } +}; + +template +int VerifySubtree(const T* node, const T* expected_parent, const T* lower_bound, + const T* upper_bound, const Compare& compare) { + if (node == nullptr) { + return 1; + } + EXPECT_EQ(ArenaTreeNodeGetParent(node), expected_parent); + if (lower_bound != nullptr) { + EXPECT_LT(compare(ArenaTreeNodeGetValue(lower_bound), + ArenaTreeNodeGetValue(node)), + 0); + } + if (upper_bound != nullptr) { + EXPECT_LT(compare(ArenaTreeNodeGetValue(node), + ArenaTreeNodeGetValue(upper_bound)), + 0); + } + const auto* left = static_cast(ArenaTreeNodeGetLeft(node)); + const auto* right = static_cast(ArenaTreeNodeGetRight(node)); + if (ArenaTreeNodeGetColor(node) == ArenaTreeNodeColor::kRed) { + if (left != nullptr) { + EXPECT_EQ(ArenaTreeNodeGetColor(left), ArenaTreeNodeColor::kBlack); + } + if (right != nullptr) { + EXPECT_EQ(ArenaTreeNodeGetColor(right), ArenaTreeNodeColor::kBlack); + } + } + int left_black_height = VerifySubtree(left, node, lower_bound, node, compare); + int right_black_height = + VerifySubtree(right, node, node, upper_bound, compare); + EXPECT_EQ(left_black_height, right_black_height); + return left_black_height + + (ArenaTreeNodeGetColor(node) == ArenaTreeNodeColor::kBlack ? 1 : 0); +} + +template +std::vector VerifyTree(T* head, const Compare& compare) { + if (head == nullptr) { + EXPECT_THAT(ArenaTreeMin(head), IsNull()); + EXPECT_THAT(ArenaTreeMax(head), IsNull()); + EXPECT_THAT(ArenaTreeNext(head), IsNull()); + EXPECT_THAT(ArenaTreePrev(head), IsNull()); + return {}; + } + EXPECT_EQ(ArenaTreeNodeGetColor(head), ArenaTreeNodeColor::kBlack); + VerifySubtree(static_cast(head), static_cast(nullptr), + static_cast(nullptr), static_cast(nullptr), + compare); + + std::vector forward; + for (const T* curr = ArenaTreeMin(static_cast(head)); + curr != nullptr; curr = static_cast(ArenaTreeNext(curr))) { + forward.push_back(curr); + EXPECT_EQ(ArenaTreeFind(static_cast(head), + ArenaTreeNodeGetValue(curr), compare), + curr); + EXPECT_EQ(ArenaTreeFind(head, ArenaTreeNodeGetValue(curr), compare), curr); + } + EXPECT_FALSE(forward.empty()); + EXPECT_EQ(forward.front(), ArenaTreeMin(head)); + EXPECT_EQ(forward.back(), ArenaTreeMax(head)); + + std::vector backward; + for (T* curr = ArenaTreeMax(head); curr != nullptr; + curr = ArenaTreePrev(curr)) { + backward.push_back(curr); + } + std::reverse(backward.begin(), backward.end()); + EXPECT_EQ(forward.size(), backward.size()); + for (size_t i = 0; i < forward.size(); ++i) { + EXPECT_EQ(forward[i], backward[i]); + } + + std::vector values; + values.reserve(forward.size()); + for (const T* node : forward) { + if constexpr (std::is_same_v) { + values.push_back(ArenaTreeNodeGetValue(node)); + } else { + values.push_back(ArenaTreeNodeGetValue(node).value); + } + } + return values; +} + +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; + auto 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()); + EXPECT_EQ(ArenaTreeFind(head, 1, TestNodeCompare{}), 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()); + EXPECT_THAT(ArenaTreeFind(head, 1, TestNodeCompare{}), IsNull()); +} + +TEST(ArenaTree, Couple) { + TestNode* head = nullptr; + auto node1 = std::make_unique(1); + auto 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; + auto 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()); + EXPECT_EQ(ArenaTreeFind(head, 1, TestNodeCrtpCompare{}), 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()); + EXPECT_THAT(ArenaTreeFind(head, 1, TestNodeCrtpCompare{}), IsNull()); +} + +TEST(ArenaTree, CrtpCouple) { + TestNodeCrtp* head = nullptr; + auto node1 = std::make_unique(1); + auto 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()); +} + +TEST(ArenaTree, InsertDuplicateReturnsExistingNode) { + TestNode* head = nullptr; + auto node2 = std::make_unique(2); + auto node1 = std::make_unique(1); + auto node3 = std::make_unique(3); + auto dup2 = std::make_unique(2); + auto dup1 = std::make_unique(1); + auto dup3 = std::make_unique(3); + + EXPECT_EQ(ArenaTreeInsert(&head, node2.get(), TestNodeCompare{}), + node2.get()); + EXPECT_EQ(ArenaTreeInsert(&head, node1.get(), TestNodeCompare{}), + node1.get()); + EXPECT_EQ(ArenaTreeInsert(&head, node3.get(), TestNodeCompare{}), + node3.get()); + + EXPECT_EQ(ArenaTreeInsert(&head, dup2.get(), TestNodeCompare{}), node2.get()); + EXPECT_EQ(ArenaTreeInsert(&head, dup1.get(), TestNodeCompare{}), node1.get()); + EXPECT_EQ(ArenaTreeInsert(&head, dup3.get(), TestNodeCompare{}), node3.get()); + EXPECT_THAT(VerifyTree(head, TestNodeCompare{}), ElementsAre(1, 2, 3)); +} + +TEST(ArenaTree, LazyEmplace) { + google::protobuf::Arena arena; + TestNode* head = nullptr; + + int emplace_calls = 0; + auto emplace_val = [&](int key) { + return ArenaTreeLazyEmplace(&arena, &head, key, TestNodeCompare{}, + [&](const auto& ctor) { + ++emplace_calls; + ctor(key); + }); + }; + + auto [node2, inserted2] = emplace_val(2); + EXPECT_TRUE(inserted2); + EXPECT_EQ(ArenaTreeNodeGetValue(node2), 2); + EXPECT_EQ(emplace_calls, 1); + + auto [node1, inserted1] = emplace_val(1); + EXPECT_TRUE(inserted1); + EXPECT_EQ(ArenaTreeNodeGetValue(node1), 1); + EXPECT_EQ(emplace_calls, 2); + + auto [node3, inserted3] = emplace_val(3); + EXPECT_TRUE(inserted3); + EXPECT_EQ(ArenaTreeNodeGetValue(node3), 3); + EXPECT_EQ(emplace_calls, 3); + + // Duplicate keys should not invoke the emplacer. + auto [dup2, dup_inserted2] = emplace_val(2); + EXPECT_FALSE(dup_inserted2); + EXPECT_EQ(dup2, node2); + EXPECT_EQ(emplace_calls, 3); + + auto [dup1, dup_inserted1] = emplace_val(1); + EXPECT_FALSE(dup_inserted1); + EXPECT_EQ(dup1, node1); + EXPECT_EQ(emplace_calls, 3); + + auto [dup3, dup_inserted3] = emplace_val(3); + EXPECT_FALSE(dup_inserted3); + EXPECT_EQ(dup3, node3); + EXPECT_EQ(emplace_calls, 3); + + EXPECT_THAT(VerifyTree(head, TestNodeCompare{}), ElementsAre(1, 2, 3)); +} + +TEST(ArenaTree, CrtpLazyEmplace) { + google::protobuf::Arena arena; + TestNodeCrtp* head = nullptr; + + int emplace_calls = 0; + auto emplace_val = [&](int key) { + return ArenaTreeLazyEmplace(&arena, &head, key, TestNodeCrtpCompare{}, + [&](const auto& ctor) { + ++emplace_calls; + ctor(key); + }); + }; + + auto [node2, inserted2] = emplace_val(2); + EXPECT_TRUE(inserted2); + EXPECT_EQ(ArenaTreeNodeGetValue(node2).value, 2); + EXPECT_EQ(emplace_calls, 1); + + auto [node1, inserted1] = emplace_val(1); + EXPECT_TRUE(inserted1); + EXPECT_EQ(ArenaTreeNodeGetValue(node1).value, 1); + EXPECT_EQ(emplace_calls, 2); + + auto [node3, inserted3] = emplace_val(3); + EXPECT_TRUE(inserted3); + EXPECT_EQ(ArenaTreeNodeGetValue(node3).value, 3); + EXPECT_EQ(emplace_calls, 3); + + auto [dup2, dup_inserted2] = emplace_val(2); + EXPECT_FALSE(dup_inserted2); + EXPECT_EQ(dup2, node2); + EXPECT_EQ(emplace_calls, 3); + + EXPECT_THAT(VerifyTree(head, TestNodeCrtpCompare{}), ElementsAre(1, 2, 3)); +} + +TEST(ArenaTree, FindLeftRightAndMissing) { + google::protobuf::Arena arena; + TestNode* head = nullptr; + for (int key : {8, 4, 12, 2, 6, 10, 14}) { + auto [node, inserted] = + ArenaTreeLazyEmplace(&arena, &head, key, TestNodeCompare{}, + [key](const auto& ctor) { ctor(key); }); + EXPECT_TRUE(inserted); + } + + const TestNode* const_head = head; + for (int key : {2, 4, 6, 8, 10, 12, 14}) { + TestNode* found = ArenaTreeFind(head, key, TestNodeCompare{}); + ASSERT_NE(found, nullptr); + EXPECT_EQ(ArenaTreeNodeGetValue(found), key); + EXPECT_EQ(ArenaTreeFind(const_head, key, TestNodeCompare{}), found); + } + for (int missing : {1, 3, 5, 7, 9, 11, 13, 15}) { + EXPECT_THAT(ArenaTreeFind(head, missing, TestNodeCompare{}), IsNull()); + EXPECT_THAT(ArenaTreeFind(const_head, missing, TestNodeCompare{}), + IsNull()); + } +} + +TEST(ArenaTree, InsertAndRemoveAscendingAndDescending) { + constexpr int kCount = 64; + std::vector> nodes; + nodes.reserve(kCount); + for (int i = 0; i < kCount; ++i) { + nodes.push_back(std::make_unique(i)); + } + + // Insert ascending, remove ascending. + TestNode* head = nullptr; + for (int i = 0; i < kCount; ++i) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[i].get(), TestNodeCompare{}), + nodes[i].get()); + VerifyTree(head, TestNodeCompare{}); + } + for (int i = 0; i < kCount; ++i) { + ArenaTreeRemove(&head, nodes[i].get()); + VerifyTree(head, TestNodeCompare{}); + } + EXPECT_THAT(head, IsNull()); + + // Insert descending, remove descending. + for (int i = kCount - 1; i >= 0; --i) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[i].get(), TestNodeCompare{}), + nodes[i].get()); + VerifyTree(head, TestNodeCompare{}); + } + for (int i = kCount - 1; i >= 0; --i) { + ArenaTreeRemove(&head, nodes[i].get()); + VerifyTree(head, TestNodeCompare{}); + } + EXPECT_THAT(head, IsNull()); + + // Insert ascending, remove descending. + for (int i = 0; i < kCount; ++i) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[i].get(), TestNodeCompare{}), + nodes[i].get()); + } + VerifyTree(head, TestNodeCompare{}); + for (int i = kCount - 1; i >= 0; --i) { + ArenaTreeRemove(&head, nodes[i].get()); + VerifyTree(head, TestNodeCompare{}); + } + EXPECT_THAT(head, IsNull()); + + // Insert descending, remove ascending. + for (int i = kCount - 1; i >= 0; --i) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[i].get(), TestNodeCompare{}), + nodes[i].get()); + } + VerifyTree(head, TestNodeCompare{}); + for (int i = 0; i < kCount; ++i) { + ArenaTreeRemove(&head, nodes[i].get()); + VerifyTree(head, TestNodeCompare{}); + } + EXPECT_THAT(head, IsNull()); +} + +TEST(ArenaTree, InsertZigZagAndRemoveInternalNodes) { + constexpr int kCount = 63; + std::vector> nodes; + nodes.reserve(kCount); + for (int i = 0; i < kCount; ++i) { + nodes.push_back(std::make_unique(i)); + } + + // Insert in a zig-zag pattern (alternating low and high values) to exercise + // Left-Right and Right-Left triangle rotations in ArenaTreeInsertColor. + TestNode* head = nullptr; + int lo = 0; + int hi = kCount - 1; + while (lo <= hi) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[lo].get(), TestNodeCompare{}), + nodes[lo].get()); + VerifyTree(head, TestNodeCompare{}); + if (lo < hi) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[hi].get(), TestNodeCompare{}), + nodes[hi].get()); + VerifyTree(head, TestNodeCompare{}); + } + ++lo; + --hi; + } + + // Repeatedly remove left child of root, right child of root, and root itself + // to cover 2-child removals where old_parent is null, old is a left child, + // and old is a right child. + int step = 0; + while (head != nullptr) { + TestNode* target = head; + if (step % 3 == 0 && ArenaTreeNodeGetLeft(head) != nullptr) { + target = static_cast(ArenaTreeNodeGetLeft(head)); + } else if (step % 3 == 1 && ArenaTreeNodeGetRight(head) != nullptr) { + target = static_cast(ArenaTreeNodeGetRight(head)); + } + ArenaTreeRemove(&head, target); + VerifyTree(head, TestNodeCompare{}); + ++step; + } +} + +TEST(ArenaTree, RemoveSingleLeftAndSingleRightChildCases) { + // Specifically exercise ArenaTreeRemove when the removed node has only a left + // child or only a right child (as root, left child of parent, and right child + // of parent). + std::vector> nodes; + for (int i = 0; i < 10; ++i) { + nodes.push_back(std::make_unique(i)); + } + + // Case 1: Node with only a left child as left child of parent and right child + // of parent. + TestNode* head = nullptr; + for (int idx : {5, 2, 8, 1, 7}) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[idx].get(), TestNodeCompare{}), + nodes[idx].get()); + } + // Node 2 has only left child 1; Node 8 has only left child 7. + ArenaTreeRemove(&head, nodes[2].get()); + VerifyTree(head, TestNodeCompare{}); + ArenaTreeRemove(&head, nodes[8].get()); + VerifyTree(head, TestNodeCompare{}); + while (head != nullptr) { + ArenaTreeRemove(&head, head); + } + + // Case 2: Node with only a right child as left child of parent and right + // child of parent. + for (int idx : {5, 2, 8, 3, 9}) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[idx].get(), TestNodeCompare{}), + nodes[idx].get()); + } + // Node 2 has only right child 3; Node 8 has only right child 9. + ArenaTreeRemove(&head, nodes[2].get()); + VerifyTree(head, TestNodeCompare{}); + ArenaTreeRemove(&head, nodes[8].get()); + VerifyTree(head, TestNodeCompare{}); + while (head != nullptr) { + ArenaTreeRemove(&head, head); + } +} + +TEST(ArenaTree, AllPermutationsInsertAndRemove) { + constexpr int kPermSize = 6; + std::vector> nodes; + nodes.reserve(kPermSize); + std::vector insert_order(kPermSize); + for (int i = 0; i < kPermSize; ++i) { + nodes.push_back(std::make_unique(i)); + insert_order[i] = i; + } + + do { + TestNode* head = nullptr; + for (int idx : insert_order) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[idx].get(), TestNodeCompare{}), + nodes[idx].get()); + } + EXPECT_THAT(VerifyTree(head, TestNodeCompare{}), + ElementsAre(0, 1, 2, 3, 4, 5)); + + // Remove in insertion order. + for (int idx : insert_order) { + ArenaTreeRemove(&head, nodes[idx].get()); + } + EXPECT_THAT(VerifyTree(head, TestNodeCompare{}), IsEmpty()); + + // Re-insert and remove in reverse insertion order. + for (int idx : insert_order) { + EXPECT_EQ(ArenaTreeInsert(&head, nodes[idx].get(), TestNodeCompare{}), + nodes[idx].get()); + } + for (int i = kPermSize - 1; i >= 0; --i) { + ArenaTreeRemove(&head, nodes[insert_order[i]].get()); + } + EXPECT_THAT(VerifyTree(head, TestNodeCompare{}), IsEmpty()); + } while (std::next_permutation(insert_order.begin(), insert_order.end())); +} + +} // namespace +} // namespace cel::internal