Add AVL balancing

This commit is contained in:
2026-07-28 14:12:49 +01:00
parent 7503561c40
commit 9ac1407907
5 changed files with 161 additions and 23 deletions
+6 -6
View File
@@ -4,12 +4,9 @@
#include "vase/vase.h"
int main() {
char *text;
uint32_t len;
int s = read_file("./flake.nix", &text, &len);
if (!s)
return 1;
const char *text_o = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
uint32_t len = strlen(text_o);
char *text = strdup(text_o);
Vase vase = Vase(text, len);
@@ -18,6 +15,9 @@ int main() {
std::cout << "\n->\n\n";
vase.insert(14, "gr\ntt", 5);
vase.insert(200, "gr\ntt", 5);
vase.insert(100, "gr\ntt", 5);
vase.insert(20, "gr\ntt", 5);
print_shard(vase.root.ptr);
+130 -9
View File
@@ -1,5 +1,104 @@
#include "vase/shard.h"
int height(Shard *n) {
return n ? n->height : 0;
}
int balance_factor(Shard *n) {
Branch *b = (Branch *)n;
return height(b->left.ptr) - height(b->right.ptr);
}
ShardPtr rotate_right(Branch *z) {
Branch *y = (Branch *)z->left.ptr;
return new Branch(
y->left.ptr,
new Branch(y->right.ptr, z->right.ptr)
);
}
ShardPtr rotate_left(Branch *z) {
Branch *y = (Branch *)z->right.ptr;
return new Branch(
new Branch(z->left.ptr, y->left.ptr),
y->right.ptr
);
}
ShardPtr balance(Shard *node) {
if (!node || node->kind == Shard::ShardKind::Petal)
return ShardPtr(node);
Branch *b = (Branch *)node;
int bf = balance_factor(node);
// left heavy
if (bf > 1) {
Branch *left = (Branch *)b->left.ptr;
// Left-right case
if (balance_factor(left) < 0) {
auto new_left = rotate_left(left);
auto rebuilt = new Branch(
new_left.ptr,
b->right.ptr
);
return rotate_right((Branch *)rebuilt);
}
// Left-left case
return rotate_right(b);
}
// right heavy
if (bf < -1) {
Branch *right = (Branch *)b->right.ptr;
// Right-left case
if (balance_factor(right) > 0) {
auto new_right = rotate_right(right);
auto rebuilt = new Branch(
b->left.ptr,
new_right.ptr
);
return rotate_left((Branch *)rebuilt);
}
// Right-right case
return rotate_left(b);
}
return ShardPtr(node);
}
ShardPtr merge(Shard *a, Shard *b) {
if (!a)
return ShardPtr(b);
if (!b)
return ShardPtr(a);
if (a->height > b->height + 1) {
Branch *ba = (Branch *)a;
auto r = merge(ba->right.ptr, b);
return balance(new Branch(ba->left.ptr, r.ptr));
}
if (b->height > a->height + 1) {
Branch *bb = (Branch *)b;
auto l = merge(a, bb->left.ptr);
return balance(new Branch(l.ptr, bb->right.ptr));
}
return balance(new Branch(a, b));
}
std::pair<ShardPtr, ShardPtr> split_shard(Shard *n, uint32_t offset) {
if (!n)
return {nullptr, nullptr};
@@ -10,17 +109,15 @@ std::pair<ShardPtr, ShardPtr> split_shard(Shard *n, uint32_t offset) {
if (n->kind == Shard::ShardKind::Branch) {
Branch *b = (Branch *)n;
if (offset < b->left.ptr->length) {
auto [a, b2] = split_shard(b->left.ptr, offset);
return {a, ShardPtr(new Branch(b2.ptr, b->right.ptr))};
return {a, merge(b2.ptr, b->right.ptr)};
} else {
auto [a, b2] = split_shard(b->right.ptr, offset - b->left.ptr->length);
return {ShardPtr(new Branch(b->left.ptr, a.ptr)), b2};
return {merge(b->left.ptr, a.ptr), b2};
}
} else {
Petal *p = (Petal *)n;
auto left = new Petal(
offset,
p->source->count_lines(0, offset),
@@ -38,12 +135,35 @@ std::pair<ShardPtr, ShardPtr> split_shard(Shard *n, uint32_t offset) {
}
}
ShardPtr merge_leaves(Shard *a, Shard *b) {
if (a->kind != Shard::ShardKind::Petal || b->kind != Shard::ShardKind::Petal)
return merge(a, b);
Petal *pa = (Petal *)a;
Petal *pb = (Petal *)b;
if (!(pa->source == pb->source && pa->pos + pa->length == pb->pos))
return merge(a, b);
return new Petal(
pa->length + pb->length,
pa->lines + pb->lines,
pa->source,
pa->pos
);
}
ShardPtr append_leaf(Shard *root, Shard *leaf) {
if (!root)
return ShardPtr(leaf);
if (root->kind == Shard::ShardKind::Petal)
return merge_leaves(root, leaf);
Branch *b = (Branch *)root;
auto new_right = append_leaf(b->right.ptr, leaf);
return balance(new Branch(b->left.ptr, new_right.ptr));
}
ShardPtr concat_shard(ShardPtr left, ShardPtr right) {
if (!left.ptr)
return right;
if (!right.ptr)
return left;
return new Branch(left.ptr, right.ptr);
return merge(left.ptr, right.ptr);
}
void print_shard(const Shard *shard, int depth) {
@@ -63,6 +183,7 @@ void print_shard(const Shard *shard, int depth) {
<< " @" << shard
<< " len=" << shard->length
<< " lines=" << shard->lines
<< " height=" << (int)shard->height
<< " refs=" << shard->refs.load()
<< "\n";