| ... | ... | @@ -1,4 +1,4 @@ |
| 1 | | const std = @import("std.zig"); |
| 1 | const std = @import("std"); |
| 2 | 2 | const assert = std.debug.assert; |
| 3 | 3 | const testing = std.testing; |
| 4 | 4 | const Order = std.math.Order; |
| ... | ... | @@ -11,6 +11,7 @@ const Red = Color.Red; |
| 11 | 11 | const Black = Color.Black; |
| 12 | 12 | |
| 13 | 13 | const ReplaceError = error{NotEqual}; |
| 14 | const SortError = error{NotUnique}; // The new comparison function results in duplicates. |
| 14 | 15 | |
| 15 | 16 | /// Insert this into your struct that you want to add to a red-black tree. |
| 16 | 17 | /// Do not use a pointer. Turn the *rb.Node results of the functions in rb |
| ... | ... | @@ -132,7 +133,21 @@ pub const Node = struct { |
| 132 | 133 | |
| 133 | 134 | pub const Tree = struct { |
| 134 | 135 | root: ?*Node, |
| 135 | | compareFn: fn (*Node, *Node) Order, |
| 136 | compareFn: fn (*Node, *Node, *Tree) Order, |
| 137 | |
| 138 | /// Re-sorts a tree with a new compare function |
| 139 | pub fn sort(tree: *Tree, newCompareFn: fn (*Node, *Node, *Tree) Order) SortError!void { |
| 140 | var newTree = Tree.init(newCompareFn); |
| 141 | var node: *Node = undefined; |
| 142 | while (true) { |
| 143 | node = tree.first() orelse break; |
| 144 | tree.remove(node); |
| 145 | if (newTree.insert(node) != null) { |
| 146 | return error.NotUnique; // EEXISTS |
| 147 | } |
| 148 | } |
| 149 | tree.* = newTree; |
| 150 | } |
| 136 | 151 | |
| 137 | 152 | /// If you have a need for a version that caches this, please file a bug. |
| 138 | 153 | pub fn first(tree: *Tree) ?*Node { |
| ... | ... | @@ -244,6 +259,7 @@ pub const Tree = struct { |
| 244 | 259 | return doLookup(key, tree, &parent, &is_left); |
| 245 | 260 | } |
| 246 | 261 | |
| 262 | /// If node is not part of tree, behavior is undefined. |
| 247 | 263 | pub fn remove(tree: *Tree, nodeconst: *Node) void { |
| 248 | 264 | var node = nodeconst; |
| 249 | 265 | // as this has the same value as node, it is unsafe to access node after newnode |
| ... | ... | @@ -389,7 +405,7 @@ pub const Tree = struct { |
| 389 | 405 | var new = newconst; |
| 390 | 406 | |
| 391 | 407 | // I assume this can get optimized out if the caller already knows. |
| 392 | | if (tree.compareFn(old, new) != .eq) return ReplaceError.NotEqual; |
| 408 | if (tree.compareFn(old, new, tree) != .eq) return ReplaceError.NotEqual; |
| 393 | 409 | |
| 394 | 410 | if (old.getParent()) |parent| { |
| 395 | 411 | parent.setChild(new, parent.left == old); |
| ... | ... | @@ -404,9 +420,11 @@ pub const Tree = struct { |
| 404 | 420 | new.* = old.*; |
| 405 | 421 | } |
| 406 | 422 | |
| 407 | | pub fn init(tree: *Tree, f: fn (*Node, *Node) Order) void { |
| 408 | | tree.root = null; |
| 409 | | tree.compareFn = f; |
| 423 | pub fn init(f: fn (*Node, *Node, *Tree) Order) Tree { |
| 424 | return Tree{ |
| 425 | .root = null, |
| 426 | .compareFn = f, |
| 427 | }; |
| 410 | 428 | } |
| 411 | 429 | }; |
| 412 | 430 | |
| ... | ... | @@ -469,7 +487,7 @@ fn doLookup(key: *Node, tree: *Tree, pparent: *?*Node, is_left: *bool) ?*Node { |
| 469 | 487 | is_left.* = false; |
| 470 | 488 | |
| 471 | 489 | while (maybe_node) |node| { |
| 472 | | const res = tree.compareFn(node, key); |
| 490 | const res = tree.compareFn(node, key, tree); |
| 473 | 491 | if (res == .eq) { |
| 474 | 492 | return node; |
| 475 | 493 | } |
| ... | ... | @@ -498,7 +516,7 @@ fn testGetNumber(node: *Node) *testNumber { |
| 498 | 516 | return @fieldParentPtr(testNumber, "node", node); |
| 499 | 517 | } |
| 500 | 518 | |
| 501 | | fn testCompare(l: *Node, r: *Node) Order { |
| 519 | fn testCompare(l: *Node, r: *Node, contextIgnored: *Tree) Order { |
| 502 | 520 | var left = testGetNumber(l); |
| 503 | 521 | var right = testGetNumber(r); |
| 504 | 522 | |
| ... | ... | @@ -512,13 +530,17 @@ fn testCompare(l: *Node, r: *Node) Order { |
| 512 | 530 | unreachable; |
| 513 | 531 | } |
| 514 | 532 | |
| 533 | fn testCompareReverse(l: *Node, r: *Node, contextIgnored: *Tree) Order { |
| 534 | return testCompare(r, l, contextIgnored); |
| 535 | } |
| 536 | |
| 515 | 537 | test "rb" { |
| 516 | 538 | if (@import("builtin").arch == .aarch64) { |
| 517 | 539 | // TODO https://github.com/ziglang/zig/issues/3288 |
| 518 | 540 | return error.SkipZigTest; |
| 519 | 541 | } |
| 520 | 542 | |
| 521 | | var tree: Tree = undefined; |
| 543 | var tree = Tree.init(testCompare); |
| 522 | 544 | var ns: [10]testNumber = undefined; |
| 523 | 545 | ns[0].value = 42; |
| 524 | 546 | ns[1].value = 41; |
| ... | ... | @@ -534,7 +556,6 @@ test "rb" { |
| 534 | 556 | var dup: testNumber = undefined; |
| 535 | 557 | dup.value = 32345; |
| 536 | 558 | |
| 537 | | tree.init(testCompare); |
| 538 | 559 | _ = tree.insert(&ns[1].node); |
| 539 | 560 | _ = tree.insert(&ns[2].node); |
| 540 | 561 | _ = tree.insert(&ns[3].node); |
| ... | ... | @@ -557,8 +578,7 @@ test "rb" { |
| 557 | 578 | } |
| 558 | 579 | |
| 559 | 580 | test "inserting and looking up" { |
| 560 | | var tree: Tree = undefined; |
| 561 | | tree.init(testCompare); |
| 581 | var tree = Tree.init(testCompare); |
| 562 | 582 | var number: testNumber = undefined; |
| 563 | 583 | number.value = 1000; |
| 564 | 584 | _ = tree.insert(&number.node); |
| ... | ... | @@ -582,8 +602,7 @@ test "multiple inserts, followed by calling first and last" { |
| 582 | 602 | // TODO https://github.com/ziglang/zig/issues/3288 |
| 583 | 603 | return error.SkipZigTest; |
| 584 | 604 | } |
| 585 | | var tree: Tree = undefined; |
| 586 | | tree.init(testCompare); |
| 605 | var tree = Tree.init(testCompare); |
| 587 | 606 | var zeroth: testNumber = undefined; |
| 588 | 607 | zeroth.value = 0; |
| 589 | 608 | var first: testNumber = undefined; |
| ... | ... | @@ -601,4 +620,8 @@ test "multiple inserts, followed by calling first and last" { |
| 601 | 620 | var lookupNode: testNumber = undefined; |
| 602 | 621 | lookupNode.value = 3; |
| 603 | 622 | assert(tree.lookup(&lookupNode.node) == &third.node); |
| 623 | tree.sort(testCompareReverse) catch unreachable; |
| 624 | assert(testGetNumber(tree.first().?).value == 3); |
| 625 | assert(testGetNumber(tree.last().?).value == 0); |
| 626 | assert(tree.lookup(&lookupNode.node) == &third.node); |
| 604 | 627 | } |