1/* @title: Red black tree */
2#pragma once
3#include <container_of.h>
4#include <stdbool.h>
5#include <stddef.h>
6#include <stdint.h>
7
8#define rbt_for_each_safe(pos, tmp, root) \
9 for (pos = rbt_first(root), tmp = rbt_next(pos); pos != NULL; \
10 pos = tmp, tmp = rbt_next(pos))
11
12#define rbt_for_each_entry_safe(pos, tmp, type, member, root) \
13 for (pos = rbt_entry(rbt_first(root), type, member), \
14 tmp = rbt_entry(rbt_next(&pos->member), type, member); \
15 &pos->member != NULL; \
16 pos = tmp, tmp = rbt_entry(rbt_next(&tmp->member), type, member))
17
18#define rbt_for_each_safe_reverse(pos, tmp, root) \
19 for (pos = rbt_last(root), tmp = rbt_prev(pos); pos != NULL; \
20 pos = tmp, tmp = rbt_prev(pos))
21
22#define rbt_for_each_entry_safe_reverse(pos, tmp, type, member, root) \
23 for (pos = rbt_entry(rbt_last(root), type, member), \
24 tmp = rbt_entry(rbt_prev(&pos->member), type, member); \
25 &pos->member != NULL; \
26 pos = tmp, tmp = rbt_entry(rbt_prev(&tmp->member), type, member))
27
28#define rbt_for_each(pos, root) \
29 for (pos = rbt_first(root); pos != NULL; pos = rbt_next(pos))
30
31#define rbt_for_each_entry(pos, type, member, root) \
32 for (pos = rbt_entry(rbt_first(root), type, member); &pos->member != NULL; \
33 pos = rbt_entry(rbt_next(&pos->member), type, member))
34
35#define rbt_for_each_reverse(pos, root) \
36 for (pos = rbt_last(root); pos != NULL; pos = rbt_prev(pos))
37
38#define rbt_for_each_entry_reverse(pos, type, member, root) \
39 for (pos = rbt_entry(rbt_last(root), type, member); &pos->member != NULL; \
40 pos = rbt_entry(rbt_prev(&pos->member), type, member))
41
42#define rbt_entry(ptr, type, member) container_of(ptr, type, member)
43#define rbt_parent(n) ((n)->parent)
44
45enum rbt_node_color { TREE_NODE_RED, TREE_NODE_BLACK };
46
47struct rbt_node {
48 enum rbt_node_color color;
49 struct rbt_node *left;
50 struct rbt_node *right;
51 struct rbt_node *parent;
52};
53
54typedef int32_t (*rbt_compare)(const struct rbt_node *a,
55 const struct rbt_node *b);
56typedef size_t (*rbt_get_data)(struct rbt_node *);
57
58struct rbt { /* TODO: stop using get_data. for now it works
59 * but in the future we may want to allow for rb-trees
60 * that are "backwards" or sorted by some other rule
61 * beyond integer field comparison */
62 rbt_get_data get_data;
63 rbt_compare compare;
64 struct rbt_node *root;
65};
66
67#define RBT_NODE_INIT \
68 (struct rbt_node) { \
69 .color = TREE_NODE_BLACK, .left = NULL, .right = NULL, .parent = NULL \
70 }
71
72static inline struct rbt_node *rbt_last(const struct rbt *root) {
73 struct rbt_node *node = root->root;
74 if (!node)
75 return NULL;
76 while (node->right)
77 node = node->right;
78 return node;
79}
80
81static inline struct rbt_node *rbt_prev(struct rbt_node *node) {
82 if (node->left) {
83 /* predecessor is rightmost node of left subtree */
84 node = node->left;
85 while (node->right)
86 node = node->right;
87 return node;
88 }
89
90 /* climb up until we come from the right */
91
92 struct rbt_node *parent = node->parent;
93 while (parent && node == parent->left) {
94 node = parent;
95 parent = parent->parent;
96 }
97 return parent;
98}
99
100static inline void rbt_init_node(struct rbt_node *n) {
101 n->color = TREE_NODE_BLACK;
102 n->left = n->right = n->parent = NULL;
103}
104
105static inline struct rbt_node *rbt_first(const struct rbt *root) {
106 struct rbt_node *node = root->root;
107 if (!node)
108 return NULL;
109 while (node->left)
110 node = node->left;
111 return node;
112}
113
114void rbt_link_node(struct rbt_node *node, struct rbt_node *parent,
115 struct rbt_node **link);
116void rbt_insert_color(struct rbt *tree, struct rbt_node *node);
117struct rbt *rbt_init(struct rbt *t, rbt_get_data get_data, rbt_compare compare);
118struct rbt *rbt_create(rbt_get_data get, rbt_compare compare);
119struct rbt_node *rbt_find_min(struct rbt_node *node);
120struct rbt_node *rbt_find_max(struct rbt_node *node);
121void rbt_delete(struct rbt *tree, struct rbt_node *z);
122struct rbt_node *rbt_search(struct rbt *tree, uint64_t data);
123void rbt_remove(struct rbt *tree, uint64_t data);
124void rbt_insert(struct rbt *tree, struct rbt_node *new_node);
125struct rbt_node *rbt_min(struct rbt *tree);
126struct rbt_node *rbt_max(struct rbt *tree);
127struct rbt_node *rbt_next(struct rbt_node *node);
128struct rbt_node *rbt_find_predecessor(struct rbt *tree, uint64_t data);
129struct rbt_node *rbt_find_successor(struct rbt *tree, uint64_t data);
130
131static inline bool rbt_empty(struct rbt *tree) {
132 return !tree->root;
133}
134
135bool rbt_has_node(struct rbt *tree, struct rbt_node *node);
136