Skip to main content

fbl/
wavl_tree.rs

1// Copyright 2026 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use crate::ptr_traits::{ManagedPtr, PtrTraits};
6use crate::sentinel::{is_sentinel_ptr, make_sentinel, make_sentinel_null, valid_sentinel_ptr};
7use crate::size_tracker::{NonTrackingSize, SizeTracker};
8use crate::tag::DefaultObjectTag;
9use core::borrow::Borrow;
10use core::cell::UnsafeCell;
11use core::marker::PhantomData;
12use core::pin::Pin;
13use pin_init::{PinInit, pin_data, pin_init, pinned_drop};
14
15/// Trait defining an observer for a `WavlTree`.
16///
17/// Observers are used by the test framework to record the number of insert,
18/// erase, rank-promote, rank-demote, and rotation operations performed during
19/// usage. The default implementation does nothing and is optimized away.
20///
21/// Observers may also be used to maintain additional application-specific per-node
22/// invariants. For example, maintaining subtree min/max values is useful for multikey
23/// partition searching.
24///
25/// Note: Records of promotions and demotions are used by tests to demonstrate
26/// that the computational complexity of insert/erase rebalancing is amortized
27/// constant. Promotions and demotions which are side effects of the rotation
28/// phase of rebalancing are considered to be part of the cost of rotation and
29/// are not tallied in the overall promote/demote accounting.
30pub trait WavlTreeObserver {
31    /// The type pointed to by the tree pointers.
32    type Target;
33
34    /// Invoked on the newly inserted node before rebalancing.
35    fn record_insert(&self, _node: *mut Self::Target) {}
36
37    /// Invoked on the node to be inserted and each ancestor node while traversing
38    /// the tree to find the initial insertion point.
39    fn record_insert_traverse(&self, _node: *mut Self::Target, _ancestor: *mut Self::Target) {}
40
41    /// Invoked on the node to be inserted and the colliding node with the same
42    /// key, during an insert-or-find operation. This method is mutually exclusive
43    /// with `record_insert_replace`, only one or the other is invoked during an
44    /// insert operation.
45    fn record_insert_collision(&self, _node: *mut Self::Target, _collision: *mut Self::Target) {}
46
47    /// Invoked on an existing node and its replacement, before swapping the
48    /// replacement into the tree, during an insert-or-replace operation. This
49    /// method is mutually exclusive with `record_insert_collision`, only one or the
50    /// other is invoked during an insert operation.
51    fn record_insert_replace(&self, _node: *mut Self::Target, _replacement: *mut Self::Target) {}
52
53    /// Invoked after each promotion during post-insert rebalancing.
54    fn record_insert_promote(&self) {}
55
56    /// Invoked after a single rotation during post-insert rebalancing.
57    fn record_insert_rotation(&self) {}
58
59    /// Invoked after a double rotation during post-insert rebalancing.
60    fn record_insert_double_rotation(&self) {}
61
62    /// Invoked on the pivot node, its parent, children, and sibling before a
63    /// rotation, just before updating the pointers in the relevant nodes. The
64    /// chirality of the children and sibling is relative to the direction of
65    /// rotation. The direction of rotation can be determined by comparing these
66    /// arguments with the values returned by the left and right child properties
67    /// of the pivot or parent arguments.
68    ///
69    /// The following diagrams the relationship of the nodes in a left rotation:
70    ///
71    /// ```text
72    ///             pivot                          parent                             |
73    ///            /     \                         /    \                             |
74    ///        parent  rl_child  <-----------  sibling  pivot                         |
75    ///        /    \                                   /   \                         |
76    ///   sibling  lr_child                       lr_child  rl_child                  |
77    /// ```
78    ///
79    /// In a right rotation, all of the relationships are reflected.
80    fn record_rotation(
81        &self,
82        _pivot: *mut Self::Target,
83        _lr_child: *mut Self::Target,
84        _rl_child: *mut Self::Target,
85        _parent: *mut Self::Target,
86        _sibling: *mut Self::Target,
87    ) {
88    }
89
90    /// Invoked on the node to be erased and the node in the tree where the
91    /// augmented invariants become invalid, leading up to the root. Called just
92    /// after updating the pointers in the relevant nodes, but before rebalancing.
93    ///
94    /// The following diagrams the relationship of the erased and invalidated
95    /// nodes:
96    ///
97    /// ```text
98    ///        root                                                                   |
99    ///       /    \                                                                  |
100    ///      A      B    <---- Invalidated starting here on up to the root            |
101    ///     / \    / \                                                                |
102    ///    C   D  E   F  <---- Erased node                                            |
103    /// ```
104    ///
105    /// When the node to be erased has two children, it is first swapped with the
106    /// leftmost child of the righthand subtree. In this case the invalidated node
107    /// is the parent of the original leftmost child of the righthand subtree, as
108    /// this is the deepest node to change after erasure.
109    ///
110    /// ```text
111    ///        root                       root                                        |
112    ///       /    \                     /    \                                       |
113    ///      A      B                   A      B                                      |
114    ///     / \    / \                 / \    / \                                     |
115    ///    C   D  E   F  <--+         C   D  E   H    <---- Invalidated starting here |
116    ///              / \    | Swap              / \                                   |
117    ///             G   H <-+                  G   F  <---- Erased node               |
118    /// ```
119    fn record_erase(&self, _node: *mut Self::Target, _invalidated: *mut Self::Target) {}
120
121    /// Invoked after each demotion during post-erase rebalancing.
122    fn record_erase_demote(&self) {}
123
124    /// Invoked after each single rotation during post-erase rebalancing.
125    fn record_erase_rotation(&self) {}
126
127    /// Invoked after each double rotation during post-erase rebalancing.
128    fn record_erase_double_rotation(&self) {}
129
130    /// Invoked during testing to verify WAVL tree rank rules for a given node.
131    fn verify_rank_rule(
132        &self,
133        _node: *mut Self::Target,
134        _left_most: *mut Self::Target,
135        _right_most: *mut Self::Target,
136        _sentinel: *mut Self::Target,
137    ) {
138    }
139
140    /// Invoked during testing to verify tree balance properties given the tree size and depth.
141    fn verify_balance(&self, _size: usize, _depth: usize) {}
142}
143
144pub struct DefaultWavlTreeObserver<T>(PhantomData<T>);
145impl<T> Default for DefaultWavlTreeObserver<T> {
146    fn default() -> Self {
147        Self(PhantomData)
148    }
149}
150impl<T> WavlTreeObserver for DefaultWavlTreeObserver<T> {
151    type Target = T;
152}
153
154/// Traits to implement to use `WavlTreeAugmentedInvariantObserver`.
155pub trait WavlTreeAugmentedInvariantObserverTraits {
156    /// The type of the node.
157    type Target;
158    /// The type of the invariant value.
159    type Value: Copy;
160
161    /// Returns the invariant value for the given node.
162    fn get_node_value(node: &Self::Target) -> Self::Value;
163
164    /// Returns the invariant value for the given node's subtree.
165    fn get_subtree_value(node: &Self::Target) -> Self::Value;
166
167    /// Combines subtree and node invariant values to produce a new subtree invariant value.
168    ///
169    /// Values are combined in in-order (left-to-right) sequence:
170    /// `left_subtree_value` is combined with `node_value`, and the result is combined
171    /// with `right_subtree_value`.
172    ///
173    /// The operation must be associative: `(a ∘ b) ∘ c == a ∘ (b ∘ c)`.
174    /// Non-commutative operations (such as string concatenation or sequence tracking)
175    /// are supported because child and node values are strictly combined in
176    /// in-order sequence across insertions, erasures, and rotations.
177    fn combine_values(a: Self::Value, b: Self::Value) -> Self::Value;
178
179    /// Sets the node's subtree invariant value.
180    fn set_subtree_value(node: &mut Self::Target, val: Self::Value);
181
182    /// Resets the node's subtree invariant value.
183    fn reset_subtree_value(node: &mut Self::Target);
184}
185
186/// A WAVL tree observer that maintains augmented invariants.
187pub struct WavlTreeAugmentedInvariantObserver<
188    Tag,
189    Traits,
190    const ALLOW_FIND_COLLISION: bool = true,
191    const ALLOW_REPLACE_COLLISION: bool = true,
192> {
193    _phantom: PhantomData<(Tag, Traits)>,
194}
195
196impl<Tag, Traits, const ALLOW_FIND_COLLISION: bool, const ALLOW_REPLACE_COLLISION: bool> Default
197    for WavlTreeAugmentedInvariantObserver<
198        Tag,
199        Traits,
200        ALLOW_FIND_COLLISION,
201        ALLOW_REPLACE_COLLISION,
202    >
203{
204    fn default() -> Self {
205        Self { _phantom: PhantomData }
206    }
207}
208
209impl<Tag, Traits, const ALLOW_FIND_COLLISION: bool, const ALLOW_REPLACE_COLLISION: bool>
210    WavlTreeAugmentedInvariantObserver<Tag, Traits, ALLOW_FIND_COLLISION, ALLOW_REPLACE_COLLISION>
211where
212    Traits: WavlTreeAugmentedInvariantObserverTraits,
213    Traits::Target: WavlTreeContainable<Traits::Target, Tag>,
214{
215    unsafe fn recompute_until_root(&self, mut current: *mut Traits::Target) {
216        while valid_sentinel_ptr(current) {
217            let current_ref = unsafe { &*current };
218            let node_value = Traits::get_node_value(current_ref);
219            let parent = current_ref.get_node().get_parent();
220            unsafe {
221                self.update_subtree_value(node_value, current);
222            }
223            current = parent;
224        }
225    }
226
227    unsafe fn update_subtree_value(&self, mut value: Traits::Value, node: *mut Traits::Target) {
228        let node_ref = unsafe { &*node };
229        let ns = node_ref.get_node();
230
231        let left = ns.get_left();
232        if valid_sentinel_ptr(left) {
233            let left_ref = unsafe { &*left };
234            let left_val = Traits::get_subtree_value(left_ref);
235            value = Traits::combine_values(left_val, value);
236        }
237
238        let right = ns.get_right();
239        if valid_sentinel_ptr(right) {
240            let right_ref = unsafe { &*right };
241            let right_val = Traits::get_subtree_value(right_ref);
242            value = Traits::combine_values(value, right_val);
243        }
244
245        let node_mut = unsafe { &mut *node };
246        Traits::set_subtree_value(node_mut, value);
247    }
248}
249
250#[allow(clippy::not_unsafe_ptr_arg_deref)]
251impl<Tag, Traits, const ALLOW_FIND_COLLISION: bool, const ALLOW_REPLACE_COLLISION: bool>
252    WavlTreeObserver
253    for WavlTreeAugmentedInvariantObserver<
254        Tag,
255        Traits,
256        ALLOW_FIND_COLLISION,
257        ALLOW_REPLACE_COLLISION,
258    >
259where
260    Traits: WavlTreeAugmentedInvariantObserverTraits,
261    Traits::Target: WavlTreeContainable<Traits::Target, Tag>,
262{
263    type Target = Traits::Target;
264
265    fn record_insert(&self, node: *mut Self::Target) {
266        let node_ref = unsafe { &*node };
267        let parent = node_ref.get_node().get_parent();
268        let val = Traits::get_node_value(node_ref);
269        let node_mut = unsafe { &mut *node };
270        Traits::set_subtree_value(node_mut, val);
271        if valid_sentinel_ptr(parent) {
272            unsafe { self.recompute_until_root(parent) };
273        }
274    }
275
276    fn record_insert_traverse(&self, _node: *mut Self::Target, _ancestor: *mut Self::Target) {}
277
278    fn record_insert_collision(&self, _node: *mut Self::Target, _collision: *mut Self::Target) {
279        debug_assert!(ALLOW_FIND_COLLISION);
280    }
281
282    fn record_insert_replace(&self, node: *mut Self::Target, replacement: *mut Self::Target) {
283        debug_assert!(ALLOW_REPLACE_COLLISION);
284        let replacement_ref = unsafe { &*replacement };
285
286        let parent = unsafe {
287            self.update_subtree_value(Traits::get_node_value(replacement_ref), node);
288            let ns = (*node).get_node();
289            ns.get_parent()
290        };
291        if valid_sentinel_ptr(parent) {
292            unsafe {
293                self.recompute_until_root(parent);
294            }
295        }
296
297        let replacement_mut = unsafe { &mut *replacement };
298        let node_mut = unsafe { &mut *node };
299        Traits::set_subtree_value(replacement_mut, Traits::get_subtree_value(node_mut));
300        Traits::reset_subtree_value(node_mut);
301    }
302
303    fn record_rotation(
304        &self,
305        pivot: *mut Self::Target,
306        lr_child: *mut Self::Target,
307        _rl_child: *mut Self::Target,
308        parent: *mut Self::Target,
309        sibling: *mut Self::Target,
310    ) {
311        let parent_ref = unsafe { &*parent };
312        let parent_subtree_val = Traits::get_subtree_value(parent_ref);
313        let parent_ns = parent_ref.get_node();
314
315        // Determine rotation direction before links are updated.
316        // If pivot is parent's right child, it is a left rotation; otherwise a right rotation.
317        let is_left_rotation = parent_ns.get_right() == pivot;
318        let (left_child, right_child) = if is_left_rotation {
319            // Left rotation: sibling remains left child, lr_child becomes right child
320            (sibling, lr_child)
321        } else {
322            // Right rotation: lr_child becomes left child, sibling remains right child
323            (lr_child, sibling)
324        };
325
326        let mut parent_value = Traits::get_node_value(parent_ref);
327
328        if valid_sentinel_ptr(left_child) {
329            let left_ref = unsafe { &*left_child };
330            let left_val = Traits::get_subtree_value(left_ref);
331            parent_value = Traits::combine_values(left_val, parent_value);
332        }
333
334        if valid_sentinel_ptr(right_child) {
335            let right_ref = unsafe { &*right_child };
336            let right_val = Traits::get_subtree_value(right_ref);
337            parent_value = Traits::combine_values(parent_value, right_val);
338        }
339
340        let pivot_mut = unsafe { &mut *pivot };
341        Traits::set_subtree_value(pivot_mut, parent_subtree_val);
342
343        let parent_mut = unsafe { &mut *parent };
344        Traits::set_subtree_value(parent_mut, parent_value);
345    }
346
347    fn record_erase(&self, node: *mut Self::Target, invalidated: *mut Self::Target) {
348        if valid_sentinel_ptr(invalidated) {
349            unsafe { self.recompute_until_root(invalidated) };
350        }
351        let node_mut = unsafe { &mut *node };
352        Traits::reset_subtree_value(node_mut);
353    }
354}
355
356/// Trait abstracting WAVL rank operations.
357pub trait WavlTreeRank: Copy {
358    /// The default rank value for a new node.
359    const DEFAULT: Self;
360    /// Returns the rank parity (true if odd, false if even).
361    fn rank_parity(rank: Self) -> bool;
362    /// Promotes the rank by 1.
363    fn promote_rank(rank: &mut Self);
364    /// Promotes the rank by 2.
365    fn double_promote_rank(rank: &mut Self);
366    /// Demotes the rank by 1.
367    fn demote_rank(rank: &mut Self);
368    /// Demotes the rank by 2.
369    fn double_demote_rank(rank: &mut Self);
370}
371
372impl WavlTreeRank for bool {
373    const DEFAULT: Self = false;
374    fn rank_parity(rank: Self) -> bool {
375        rank
376    }
377    fn promote_rank(rank: &mut Self) {
378        *rank = !*rank;
379    }
380    fn double_promote_rank(_rank: &mut Self) {} // no-op
381    fn demote_rank(rank: &mut Self) {
382        *rank = !*rank;
383    }
384    fn double_demote_rank(_rank: &mut Self) {} // no-op
385}
386
387impl WavlTreeRank for i32 {
388    const DEFAULT: Self = 0;
389    fn rank_parity(rank: Self) -> bool {
390        (rank & 1) != 0
391    }
392    fn promote_rank(rank: &mut Self) {
393        *rank += 1;
394    }
395    fn double_promote_rank(rank: &mut Self) {
396        *rank += 2;
397    }
398    fn demote_rank(rank: &mut Self) {
399        *rank -= 1;
400    }
401    fn double_demote_rank(rank: &mut Self) {
402        *rank -= 2;
403    }
404}
405
406/// A node in a Weak AVL (WAVL) Tree.
407#[repr(C)]
408pub struct WavlTreeNode<T, R: WavlTreeRank = bool> {
409    /// The parent element in the tree.
410    pub parent: UnsafeCell<*mut T>,
411    /// The left child element in the tree.
412    pub left: UnsafeCell<*mut T>,
413    /// The right child element in the tree.
414    pub right: UnsafeCell<*mut T>,
415    /// The integer rank of this node.
416    pub rank: UnsafeCell<R>,
417}
418
419impl<T, R: WavlTreeRank> WavlTreeNode<T, R> {
420    /// Creates a new, unlinked node.
421    pub const fn new() -> Self {
422        Self {
423            parent: UnsafeCell::new(core::ptr::null_mut()),
424            left: UnsafeCell::new(core::ptr::null_mut()),
425            right: UnsafeCell::new(core::ptr::null_mut()),
426            rank: UnsafeCell::new(R::DEFAULT),
427        }
428    }
429
430    /// Returns true if the node is currently in a tree.
431    pub fn in_container(&self) -> bool {
432        // SAFETY: Accessing parent pointer from UnsafeCell is safe because WavlTree coordinates
433        // exclusive mutations on containment states, ensuring no data races.
434        !unsafe { *self.parent.get() }.is_null()
435    }
436
437    fn get_parent(&self) -> *mut T {
438        // SAFETY: Accessing parent pointer from UnsafeCell is safe because it is only read sequentially
439        // or under logical exclusive containment borrow.
440        unsafe { *self.parent.get() }
441    }
442
443    fn set_parent(&self, parent: *mut T) {
444        // SAFETY: Mutating parent pointer in UnsafeCell is safe because the parent container holds
445        // exclusive mutable borrow of the containing tree structure.
446        unsafe {
447            *self.parent.get() = parent;
448        }
449    }
450
451    fn get_left(&self) -> *mut T {
452        // SAFETY: Accessing left pointer from UnsafeCell is safe because it is only read sequentially
453        // or under logical exclusive containment borrow.
454        unsafe { *self.left.get() }
455    }
456
457    fn set_left(&self, left: *mut T) {
458        // SAFETY: Mutating left pointer in UnsafeCell is safe because the parent container holds
459        // exclusive mutable borrow of the containing tree structure.
460        unsafe {
461            *self.left.get() = left;
462        }
463    }
464
465    fn get_right(&self) -> *mut T {
466        // SAFETY: Accessing right pointer from UnsafeCell is safe because it is only read sequentially
467        // or under logical exclusive containment borrow.
468        unsafe { *self.right.get() }
469    }
470
471    fn set_right(&self, right: *mut T) {
472        // SAFETY: Mutating right pointer in UnsafeCell is safe because the parent container holds
473        // exclusive mutable borrow of the containing tree structure.
474        unsafe {
475            *self.right.get() = right;
476        }
477    }
478
479    fn rank_parity(&self) -> bool {
480        // SAFETY: Reading rank from UnsafeCell is safe because it is only read sequentially or
481        // under logical exclusive container borrow.
482        unsafe { R::rank_parity(*self.rank.get()) }
483    }
484
485    /// Returns the rank value of this node.
486    pub fn rank(&self) -> R {
487        // SAFETY: Reading rank from UnsafeCell is safe because it is only read sequentially or
488        // under logical exclusive container borrow.
489        unsafe { *self.rank.get() }
490    }
491
492    fn promote_rank(&self) {
493        // SAFETY: Mutating rank in UnsafeCell is safe because the parent container holds
494        // exclusive mutable borrow of the containing tree structure.
495        unsafe {
496            R::promote_rank(&mut *self.rank.get());
497        }
498    }
499
500    fn double_promote_rank(&self) {
501        // SAFETY: Mutating rank in UnsafeCell is safe because the parent container holds
502        // exclusive mutable borrow of the containing tree structure.
503        unsafe {
504            R::double_promote_rank(&mut *self.rank.get());
505        }
506    }
507
508    fn demote_rank(&self) {
509        // SAFETY: Mutating rank in UnsafeCell is safe because the parent container holds
510        // exclusive mutable borrow of the containing tree structure.
511        unsafe {
512            R::demote_rank(&mut *self.rank.get());
513        }
514    }
515
516    fn double_demote_rank(&self) {
517        // SAFETY: Mutating rank in UnsafeCell is safe because the parent container holds
518        // exclusive mutable borrow of the containing tree structure.
519        unsafe {
520            R::double_demote_rank(&mut *self.rank.get());
521        }
522    }
523
524    /// Returns true if the node state invariants are currently valid.
525    pub fn is_valid(&self) -> bool {
526        let parent = self.get_parent();
527        let left = self.get_left();
528        let right = self.get_right();
529        !parent.is_null() || (parent.is_null() && left.is_null() && right.is_null())
530    }
531}
532
533impl<T, R: WavlTreeRank> core::fmt::Debug for WavlTreeNode<T, R> {
534    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
535        f.debug_struct("WavlTreeNode").field("in_container", &self.in_container()).finish()
536    }
537}
538
539impl<T, R: WavlTreeRank> Default for WavlTreeNode<T, R> {
540    fn default() -> Self {
541        Self::new()
542    }
543}
544
545impl<T, R: WavlTreeRank> Drop for WavlTreeNode<T, R> {
546    fn drop(&mut self) {
547        debug_assert!(!self.in_container(), "Object destroyed while still in container");
548    }
549}
550
551/// Trait that types must implement to be contained in a `WavlTree`.
552pub trait WavlTreeContainable<T, Tag = DefaultObjectTag> {
553    /// The rank type used by this node.
554    type Rank: WavlTreeRank;
555    /// Returns a reference to the tree node.
556    fn get_node(&self) -> &WavlTreeNode<T, Self::Rank>;
557}
558
559/// Trait that types must implement to expose a key for `WavlTree` sorting and lookup.
560pub trait WavlTreeKeyable<K: ?Sized> {
561    /// The type of key yielded by `get_key`, borrowing `K`.
562    ///
563    /// Implementations may return an owned/by-value key (e.g. `(u64, usize)` or `i32`)
564    /// or a borrowed/by-reference key (e.g. `&'a KeyStruct`).
565    type Key<'a>: Borrow<K>
566    where
567        Self: 'a;
568
569    /// Returns the key of this object.
570    fn get_key(&self) -> Self::Key<'_>;
571}
572
573#[allow(dead_code)]
574trait LrTraits {
575    type Inverse: LrTraits;
576
577    fn lr_child<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut T;
578    fn rl_child<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut T;
579
580    fn lr_child_ptr<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut *mut T;
581    fn rl_child_ptr<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut *mut T;
582
583    fn lr_most<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>) -> *mut P::Target
584    where
585        P: PtrTraits,
586        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
587        K: Ord,
588        S: SizeTracker,
589        O: WavlTreeObserver<Target = P::Target>;
590
591    fn rl_most<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>) -> *mut P::Target
592    where
593        P: PtrTraits,
594        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
595        K: Ord,
596        S: SizeTracker,
597        O: WavlTreeObserver<Target = P::Target>;
598
599    unsafe fn set_lr_most<K, P, Tag, S, O>(
600        tree: &mut WavlTree<K, P, Tag, S, O>,
601        val: *mut P::Target,
602    ) where
603        P: PtrTraits,
604        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
605        K: Ord,
606        S: SizeTracker,
607        O: WavlTreeObserver<Target = P::Target>;
608
609    unsafe fn set_rl_most<K, P, Tag, S, O>(
610        tree: &mut WavlTree<K, P, Tag, S, O>,
611        val: *mut P::Target,
612    ) where
613        P: PtrTraits,
614        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
615        K: Ord,
616        S: SizeTracker,
617        O: WavlTreeObserver<Target = P::Target>;
618}
619
620struct ForwardTraits;
621struct ReverseTraits;
622
623impl LrTraits for ForwardTraits {
624    type Inverse = ReverseTraits;
625
626    fn lr_child<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut T {
627        ns.get_left()
628    }
629    fn rl_child<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut T {
630        ns.get_right()
631    }
632
633    fn lr_child_ptr<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut *mut T {
634        ns.left.get()
635    }
636    fn rl_child_ptr<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut *mut T {
637        ns.right.get()
638    }
639
640    fn lr_most<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>) -> *mut P::Target
641    where
642        P: PtrTraits,
643        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
644        K: Ord,
645        S: SizeTracker,
646        O: WavlTreeObserver<Target = P::Target>,
647    {
648        tree.left_most
649    }
650
651    fn rl_most<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>) -> *mut P::Target
652    where
653        P: PtrTraits,
654        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
655        K: Ord,
656        S: SizeTracker,
657        O: WavlTreeObserver<Target = P::Target>,
658    {
659        tree.right_most
660    }
661
662    unsafe fn set_lr_most<K, P, Tag, S, O>(
663        tree: &mut WavlTree<K, P, Tag, S, O>,
664        val: *mut P::Target,
665    ) where
666        P: PtrTraits,
667        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
668        K: Ord,
669        S: SizeTracker,
670        O: WavlTreeObserver<Target = P::Target>,
671    {
672        tree.left_most = val;
673    }
674
675    unsafe fn set_rl_most<K, P, Tag, S, O>(
676        tree: &mut WavlTree<K, P, Tag, S, O>,
677        val: *mut P::Target,
678    ) where
679        P: PtrTraits,
680        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
681        K: Ord,
682        S: SizeTracker,
683        O: WavlTreeObserver<Target = P::Target>,
684    {
685        tree.right_most = val;
686    }
687}
688
689impl LrTraits for ReverseTraits {
690    type Inverse = ForwardTraits;
691
692    fn lr_child<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut T {
693        ns.get_right()
694    }
695    fn rl_child<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut T {
696        ns.get_left()
697    }
698
699    fn lr_child_ptr<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut *mut T {
700        ns.right.get()
701    }
702    fn rl_child_ptr<T, R: WavlTreeRank>(ns: &WavlTreeNode<T, R>) -> *mut *mut T {
703        ns.left.get()
704    }
705
706    fn lr_most<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>) -> *mut P::Target
707    where
708        P: PtrTraits,
709        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
710        K: Ord,
711        S: SizeTracker,
712        O: WavlTreeObserver<Target = P::Target>,
713    {
714        tree.right_most
715    }
716
717    fn rl_most<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>) -> *mut P::Target
718    where
719        P: PtrTraits,
720        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
721        K: Ord,
722        S: SizeTracker,
723        O: WavlTreeObserver<Target = P::Target>,
724    {
725        tree.left_most
726    }
727
728    unsafe fn set_lr_most<K, P, Tag, S, O>(
729        tree: &mut WavlTree<K, P, Tag, S, O>,
730        val: *mut P::Target,
731    ) where
732        P: PtrTraits,
733        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
734        K: Ord,
735        S: SizeTracker,
736        O: WavlTreeObserver<Target = P::Target>,
737    {
738        tree.right_most = val;
739    }
740
741    unsafe fn set_rl_most<K, P, Tag, S, O>(
742        tree: &mut WavlTree<K, P, Tag, S, O>,
743        val: *mut P::Target,
744    ) where
745        P: PtrTraits,
746        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
747        K: Ord,
748        S: SizeTracker,
749        O: WavlTreeObserver<Target = P::Target>,
750    {
751        tree.left_most = val;
752    }
753}
754
755/// A Weak AVL (WAVL) Tree associative container.
756///
757/// Implementation Notes:
758///
759/// WAVLTree<> is an implementation of a "Weak AVL" tree; a self
760/// balancing binary search tree whose rebalancing algorithm was
761/// originally described in
762///
763/// Bernhard Haeupler, Siddhartha Sen, and Robert E. Tarjan. 2015.
764/// Rank-Balanced Trees. ACM Trans. Algorithms 11, 4, Article 30 (June 2015), 26 pages.
765/// DOI=http://dx.doi.org/10.1145/2689412
766///
767/// See also
768/// https://en.wikipedia.org/wiki/WAVL_tree
769/// http://sidsen.azurewebsites.net/papers/rb-trees-talg.pdf
770///
771/// WAVLTree<>s, like HashTables, are associative containers and support all of
772/// the same key-centric operations (such as find() and insert_or_find()) that
773/// HashTables support.
774///
775/// Additionally, WAVLTree's are internally ordered by key (unlike HashTables
776/// which are un-ordered).  Iteration forwards or backwards runs in amortized
777/// constant time, but in O(log) time in an individual worst case.  Forward
778/// iteration will enumerate the elements in monotonically increasing order (as
779/// defined by the KeyTraits::LessThan operation).
780///
781/// Two additional operations are supported because of the ordered nature of a
782/// WAVLTree:
783/// upper_bound(key)        : Returns a cursor positioned at the first element (E) in the tree such that E.key > key.
784/// lower_bound(key)        : Returns a cursor positioned at the first element (E) in the tree such that E.key >= key.
785///
786/// The worst depth of a WAVL tree depends on whether or not the tree has ever
787/// been subject to erase operations.
788///
789/// ++ If the tree has seen only insert operations, the worst case depth of the
790///    tree is log_phi(N), where phi is the golden ratio.  This is the same bound
791///    as that of an AVL tree.
792/// ++ If the tree has seen erase operations in addition to insert operations,
793///    the worst case depth of the tree is 2*log_2(N).  This is the same bound as
794///    a Red-Black tree.
795///
796/// Insertion runs in O(log) time; finding the location takes O(log) time while
797/// post-insert rebalancing runs in amortized constant time.
798///
799/// Erase-by-key runs in O(log) time; finding the node to erase takes O(log) time
800/// while post-erase rebalancing runs in amortized constant time.
801///
802/// Because of the intrusive nature of the container, direct-erase operations
803/// (AKA, erase operations where the reference to the element to be erased is
804/// already known) run in amortized constant time.
805type TargetRank<P, Tag> =
806    <<P as PtrTraits>::Target as WavlTreeContainable<<P as PtrTraits>::Target, Tag>>::Rank;
807
808#[repr(C)]
809#[pin_data(PinnedDrop)]
810pub struct WavlTree<
811    K,
812    P,
813    Tag = DefaultObjectTag,
814    S = NonTrackingSize,
815    O = DefaultWavlTreeObserver<<P as PtrTraits>::Target>,
816> where
817    P: PtrTraits,
818    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
819    K: Ord,
820    S: SizeTracker,
821    O: WavlTreeObserver<Target = P::Target>,
822{
823    root: *mut P::Target,
824    left_most: *mut P::Target,
825    right_most: *mut P::Target,
826    size: S,
827    observer: O,
828    #[pin]
829    _pin: core::marker::PhantomPinned,
830    _phantom: core::marker::PhantomData<(K, P, Tag)>,
831}
832
833impl<K, P, Tag, S, O> WavlTree<K, P, Tag, S, O>
834where
835    P: PtrTraits,
836    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
837    K: Ord,
838    S: SizeTracker,
839    O: WavlTreeObserver<Target = P::Target>,
840{
841    /// Creates a new, empty tree with a custom observer.
842    pub fn new_with_observer(observer: O) -> impl PinInit<Self, core::convert::Infallible> {
843        pin_init!(&this in Self {
844            root: core::ptr::null_mut(),
845            left_most: make_sentinel(this.as_ptr()),
846            right_most: make_sentinel(this.as_ptr()),
847            size: S::INIT,
848            observer,
849            _pin: core::marker::PhantomPinned,
850            _phantom: core::marker::PhantomData,
851        })
852    }
853
854    /// Creates a new, empty tree.
855    pub fn new() -> impl PinInit<Self, core::convert::Infallible>
856    where
857        O: Default,
858    {
859        Self::new_with_observer(O::default())
860    }
861
862    fn get_sentinel(&self) -> *mut P::Target {
863        make_sentinel(self as *const Self as *mut Self)
864    }
865
866    /// Returns a reference to the node state of `ptr`.
867    ///
868    /// # Safety
869    ///
870    /// The caller must ensure that `ptr` is a valid, aligned, and dereferenceable pointer
871    /// to an initialized `P::Target` object that is alive for the lifetime `'a`.
872    unsafe fn get_node_ref<'a>(
873        ptr: *mut P::Target,
874    ) -> &'a WavlTreeNode<P::Target, TargetRank<P, Tag>> {
875        // SAFETY: The caller guarantees that `ptr` is valid, aligned, and dereferenceable.
876        unsafe { &(*ptr) }.get_node()
877    }
878
879    /// Returns true if the tree is empty.
880    pub fn is_empty(&self) -> bool {
881        self.root.is_null()
882    }
883
884    /// Returns a reference to the first (smallest) element of the tree, or `None` if it is empty.
885    pub fn front(&self) -> Option<&P::Target> {
886        // SAFETY: `self.left_most` is a valid pointer to a node in the tree or the sentinel.
887        // If `self.is_empty()` is false, it is guaranteed to be a valid, dereferenceable
888        // pointer to a node.
889        if self.is_empty() { None } else { unsafe { Some(&*self.left_most) } }
890    }
891
892    /// Returns a reference to the last (largest) element of the tree, or `None` if it is empty.
893    pub fn back(&self) -> Option<&P::Target> {
894        // SAFETY: `self.right_most` is a valid pointer to a node in the tree or the sentinel.
895        // If `self.is_empty()` is false, it is guaranteed to be a valid, dereferenceable
896        // pointer to a node.
897        if self.is_empty() { None } else { unsafe { Some(&*self.right_most) } }
898    }
899
900    /// Returns a mutable pointer to the pointer linking `node` into the tree.
901    ///
902    /// # Safety
903    ///
904    /// The caller must ensure that `node` is a valid, aligned, and dereferenceable pointer
905    /// to a node that is currently contained within this tree instance.
906    unsafe fn get_link_ptr_to_node(&mut self, node: *mut P::Target) -> *mut *mut P::Target {
907        debug_assert!(valid_sentinel_ptr(node));
908
909        // SAFETY: The caller guarantees that `node` is a valid pointer to a node currently in the tree.
910        let ns = unsafe { Self::get_node_ref(node) };
911        let parent = ns.get_parent();
912        if is_sentinel_ptr(parent) {
913            debug_assert_eq!(parent, self.get_sentinel());
914            debug_assert_eq!(self.root, node);
915            &mut self.root as *mut _
916        } else {
917            debug_assert!(!parent.is_null());
918            // SAFETY: `parent` is not a sentinel and is not null, so it must be a valid,
919            // dereferenceable pointer to a node in the tree.
920            let parent_ns = unsafe { Self::get_node_ref(parent) };
921            if parent_ns.get_left() == node {
922                parent_ns.left.get()
923            } else {
924                debug_assert_eq!(parent_ns.get_right(), node);
925                parent_ns.right.get()
926            }
927        }
928    }
929
930    /// Performs a tree rotation in the direction specified by `LR`.
931    ///
932    /// # Safety
933    ///
934    /// The caller must ensure that `node` and `parent` are valid, aligned, and dereferenceable
935    /// pointers to nodes currently contained within this tree instance, and that `node` is
936    /// a child of `parent` in the direction specified by `LR::Inverse`.
937    unsafe fn rotate_lr<LR: LrTraits>(&mut self, node: *mut P::Target, parent: *mut P::Target) {
938        // SAFETY: The caller guarantees that `node` and `parent` are valid pointers to nodes
939        // currently in the tree. The structural modifications (rotations) correctly permute the
940        // parent/child links while maintaining BST structural validity.
941        unsafe {
942            debug_assert!(valid_sentinel_ptr(node));
943            debug_assert!(valid_sentinel_ptr(parent));
944
945            let x = node;
946            let z = parent;
947
948            let x_ns = Self::get_node_ref(x);
949            let z_ns = Self::get_node_ref(z);
950
951            debug_assert_eq!(LR::rl_child(z_ns), x);
952
953            let x_link = LR::rl_child_ptr(z_ns);
954            let y_link = LR::lr_child_ptr(x_ns);
955            let z_link = self.get_link_ptr_to_node(z);
956
957            let g = z_ns.get_parent();
958            let y = *y_link;
959
960            debug_assert!(!is_sentinel_ptr(y));
961
962            // Permute the downstream links.
963            self.observer.record_rotation(x, y, LR::rl_child(x_ns), z, LR::lr_child(z_ns));
964            let tmp = *x_link;
965            *x_link = *y_link;
966            *y_link = *z_link;
967            *z_link = tmp;
968
969            // Update parent pointers.
970            x_ns.set_parent(g);
971            z_ns.set_parent(x);
972            if !y.is_null() {
973                Self::get_node_ref(y).set_parent(z);
974            }
975        }
976    }
977
978    /// Performs post-insertion balancing fixups.
979    ///
980    /// # Safety
981    ///
982    /// The caller must ensure that `node` and `parent` are valid, aligned, and dereferenceable
983    /// pointers to nodes currently contained within this tree instance, and that `node` is
984    /// a child of `parent` in the direction specified by `LR`.
985    unsafe fn post_insert_fixup_lr<LR: LrTraits>(
986        &mut self,
987        node: *mut P::Target,
988        parent: *mut P::Target,
989    ) {
990        type RL<LR> = <LR as LrTraits>::Inverse;
991
992        // SAFETY: The caller guarantees `node` and `parent` are valid pointers to nodes currently
993        // in the tree. The subsequent operations (including rotations and rank updates) correctly
994        // rebalance the tree according to WAVL insertion balance algorithms.
995        unsafe {
996            debug_assert!(valid_sentinel_ptr(node));
997            debug_assert!(valid_sentinel_ptr(parent));
998
999            let node_ns = Self::get_node_ref(node);
1000            let parent_ns = Self::get_node_ref(parent);
1001
1002            debug_assert_eq!(LR::lr_child(parent_ns), node);
1003
1004            let rl_child = LR::rl_child(node_ns);
1005            let rl_child_ns = if valid_sentinel_ptr(rl_child) {
1006                Some(Self::get_node_ref(rl_child))
1007            } else {
1008                None
1009            };
1010
1011            if rl_child_ns.is_none()
1012                || (rl_child_ns.unwrap().rank_parity() == node_ns.rank_parity())
1013            {
1014                // Case #1: single rotation.
1015                self.rotate_lr::<RL<LR>>(node, parent);
1016                parent_ns.demote_rank();
1017                self.observer.record_insert_rotation();
1018            } else {
1019                // Case #2: double rotation.
1020                let rl_child_ns = rl_child_ns.unwrap();
1021                self.rotate_lr::<LR>(rl_child, node);
1022                self.rotate_lr::<RL<LR>>(rl_child, parent);
1023
1024                rl_child_ns.promote_rank();
1025                node_ns.demote_rank();
1026                parent_ns.demote_rank();
1027                self.observer.record_insert_double_rotation();
1028            }
1029        }
1030    }
1031
1032    /// Rebalances the tree after an element is inserted.
1033    ///
1034    /// # Safety
1035    ///
1036    /// The caller must ensure that `node` is a valid, aligned, and dereferenceable pointer
1037    /// to a node that has just been inserted into this tree instance, and that the tree
1038    /// is structurally valid except for potential balance violations at `node` and its ancestors.
1039    unsafe fn balance_post_insert(&mut self, mut node: *mut P::Target) {
1040        // SAFETY: The caller guarantees `node` is a valid pointer to a node in the tree.
1041        // The loop climbs the tree, updating ranks and performing rotations as needed, maintaining
1042        // tree integrity.
1043        unsafe {
1044            let mut node_ns = Self::get_node_ref(node);
1045            debug_assert!(valid_sentinel_ptr(node_ns.get_parent()));
1046
1047            let mut parent = node_ns.get_parent();
1048            let mut parent_ns = Self::get_node_ref(parent);
1049
1050            if valid_sentinel_ptr(parent_ns.get_left()) && valid_sentinel_ptr(parent_ns.get_right())
1051            {
1052                return;
1053            }
1054
1055            let mut node_parity;
1056            let mut parent_parity;
1057            let mut sibling_parity;
1058            let mut is_left_child;
1059
1060            loop {
1061                // Promote.
1062                parent_ns.promote_rank();
1063                self.observer.record_insert_promote();
1064
1065                // Climb.
1066                node = parent;
1067                node_ns = Self::get_node_ref(node);
1068                parent = node_ns.get_parent();
1069
1070                if !valid_sentinel_ptr(parent) {
1071                    return;
1072                }
1073
1074                parent_ns = Self::get_node_ref(parent);
1075                is_left_child = parent_ns.get_left() == node;
1076                if is_left_child {
1077                    sibling_parity = if valid_sentinel_ptr(parent_ns.get_right()) {
1078                        Self::get_node_ref(parent_ns.get_right()).rank_parity()
1079                    } else {
1080                        true
1081                    };
1082                } else {
1083                    debug_assert_eq!(parent_ns.get_right(), node);
1084                    sibling_parity = if valid_sentinel_ptr(parent_ns.get_left()) {
1085                        Self::get_node_ref(parent_ns.get_left()).rank_parity()
1086                    } else {
1087                        true
1088                    };
1089                }
1090
1091                node_parity = node_ns.rank_parity();
1092                parent_parity = parent_ns.rank_parity();
1093
1094                if !((!node_parity && !parent_parity && sibling_parity)
1095                    || (node_parity && parent_parity && !sibling_parity))
1096                {
1097                    break;
1098                }
1099            }
1100
1101            if (node_parity != parent_parity) || (node_parity != sibling_parity) {
1102                return;
1103            }
1104
1105            if is_left_child {
1106                self.post_insert_fixup_lr::<ForwardTraits>(node, parent);
1107            } else {
1108                self.post_insert_fixup_lr::<ReverseTraits>(node, parent);
1109            }
1110        }
1111    }
1112
1113    /// Performs balance adjustments for a "2-2 leaf" node after erasure.
1114    ///
1115    /// # Safety
1116    ///
1117    /// The caller must ensure that `node` is a valid, aligned, and dereferenceable pointer
1118    /// to a node currently contained within this tree instance, and that the tree is structurally
1119    /// valid except for potential balance violations after an erasure at `node`.
1120    unsafe fn balance_post_erase_fix_22_leaf(&mut self, node: *mut P::Target) {
1121        // SAFETY: The caller guarantees `node` is a valid pointer to a node in the tree.
1122        // The function safely demotes the rank and propagates rebalancing up the tree.
1123        unsafe {
1124            debug_assert!(valid_sentinel_ptr(node));
1125
1126            let ns = Self::get_node_ref(node);
1127            if !ns.rank_parity()
1128                || valid_sentinel_ptr(ns.get_left())
1129                || valid_sentinel_ptr(ns.get_right())
1130            {
1131                return;
1132            }
1133
1134            ns.demote_rank();
1135            self.observer.record_erase_demote();
1136
1137            let parent = ns.get_parent();
1138            debug_assert!(!parent.is_null());
1139            if is_sentinel_ptr(parent) {
1140                return;
1141            }
1142
1143            let parent_ns = Self::get_node_ref(parent);
1144            let is_left_child = parent_ns.get_left() == node;
1145            debug_assert!(is_left_child || parent_ns.get_right() == node);
1146
1147            if is_left_child {
1148                self.balance_post_erase_fix_lr_3_child::<ForwardTraits>(parent);
1149            } else {
1150                self.balance_post_erase_fix_lr_3_child::<ReverseTraits>(parent);
1151            }
1152        }
1153    }
1154
1155    /// Rebalances the tree after an element is erased, resolving violations of the 3-child rule.
1156    ///
1157    /// # Safety
1158    ///
1159    /// The caller must ensure that `node` is a valid, aligned, and dereferenceable pointer
1160    /// to a node currently contained within this tree instance, and that `node` is the parent
1161    /// of a subtree that has just seen an erasure and violates the 3-child balance rule.
1162    unsafe fn balance_post_erase_fix_lr_3_child<LR: LrTraits>(&mut self, node: *mut P::Target) {
1163        type RL<LR> = <LR as LrTraits>::Inverse;
1164        // SAFETY: The caller guarantees `node` is a valid pointer to a node in the tree.
1165        // The loop walks up the tree performing rank adjustments and triggers rotations as required
1166        // by the WAVL erase balancing algorithms.
1167        unsafe {
1168            debug_assert!(valid_sentinel_ptr(node));
1169
1170            let mut z = node;
1171            let mut z_ns = Self::get_node_ref(z);
1172            let mut x = LR::lr_child(z_ns);
1173
1174            if valid_sentinel_ptr(x) != z_ns.rank_parity() {
1175                return;
1176            }
1177
1178            let mut x_is_lr_child = true;
1179            let mut y = LR::rl_child(z_ns);
1180
1181            loop {
1182                debug_assert!(valid_sentinel_ptr(y));
1183
1184                let y_ns = Self::get_node_ref(y);
1185                let y_is_2_child = y_ns.rank_parity() == z_ns.rank_parity();
1186
1187                if !y_is_2_child {
1188                    let y_is_22_node;
1189                    if y_ns.rank_parity() {
1190                        y_is_22_node = (!valid_sentinel_ptr(y_ns.get_left())
1191                            || Self::get_node_ref(y_ns.get_left()).rank_parity())
1192                            && (!valid_sentinel_ptr(y_ns.get_right())
1193                                || Self::get_node_ref(y_ns.get_right()).rank_parity());
1194                    } else {
1195                        y_is_22_node = valid_sentinel_ptr(y_ns.get_left())
1196                            && valid_sentinel_ptr(y_ns.get_right())
1197                            && !Self::get_node_ref(y_ns.get_left()).rank_parity()
1198                            && !Self::get_node_ref(y_ns.get_right()).rank_parity();
1199                    }
1200
1201                    if !y_is_22_node {
1202                        break;
1203                    }
1204                }
1205
1206                z_ns.demote_rank();
1207                self.observer.record_erase_demote();
1208                if !y_is_2_child {
1209                    y_ns.demote_rank();
1210                    self.observer.record_erase_demote();
1211                }
1212
1213                if !valid_sentinel_ptr(z_ns.get_parent()) {
1214                    return;
1215                }
1216
1217                let x_rank_parity = z_ns.rank_parity();
1218                x = z;
1219                z = z_ns.get_parent();
1220                z_ns = Self::get_node_ref(z);
1221
1222                if z_ns.rank_parity() == x_rank_parity {
1223                    return;
1224                }
1225
1226                x_is_lr_child = LR::lr_child(z_ns) == x;
1227                y = if x_is_lr_child { LR::rl_child(z_ns) } else { LR::lr_child(z_ns) };
1228            }
1229
1230            if x_is_lr_child {
1231                self.balance_post_erase_do_rotations::<LR>(y, z);
1232            } else {
1233                self.balance_post_erase_do_rotations::<RL<LR>>(y, z);
1234            }
1235        }
1236    }
1237
1238    /// Performs necessary rotations during post-erase rebalancing.
1239    ///
1240    /// # Safety
1241    ///
1242    /// The caller must ensure that `y` and `z` are valid, aligned, and dereferenceable
1243    /// pointers to nodes currently contained within this tree instance, and that `y` is the
1244    /// right child of `z` in the direction specified by `LR`.
1245    unsafe fn balance_post_erase_do_rotations<LR: LrTraits>(
1246        &mut self,
1247        y: *mut P::Target,
1248        z: *mut P::Target,
1249    ) {
1250        type RL<LR> = <LR as LrTraits>::Inverse;
1251        // SAFETY: The caller guarantees `y` and `z` are valid pointers to nodes in the tree.
1252        // The rotations correctly rebalance the tree at `z` and update ranks accordingly.
1253        unsafe {
1254            debug_assert!(valid_sentinel_ptr(y));
1255            debug_assert!(valid_sentinel_ptr(z));
1256
1257            let y_ns = Self::get_node_ref(y);
1258            let z_ns = Self::get_node_ref(z);
1259
1260            let w = LR::rl_child(y_ns);
1261            let w_rank_parity =
1262                if valid_sentinel_ptr(w) { Self::get_node_ref(w).rank_parity() } else { true };
1263
1264            if y_ns.rank_parity() != w_rank_parity {
1265                self.rotate_lr::<LR>(y, z);
1266                y_ns.promote_rank();
1267
1268                if !valid_sentinel_ptr(z_ns.get_left()) && !valid_sentinel_ptr(z_ns.get_right()) {
1269                    z_ns.double_demote_rank();
1270                } else {
1271                    z_ns.demote_rank();
1272                }
1273                self.observer.record_erase_rotation();
1274            } else {
1275                let v = LR::lr_child(y_ns);
1276                debug_assert!(valid_sentinel_ptr(v));
1277                let v_ns = Self::get_node_ref(v);
1278                debug_assert_ne!(v_ns.rank_parity(), y_ns.rank_parity());
1279
1280                self.rotate_lr::<RL<LR>>(v, y);
1281                self.rotate_lr::<LR>(v, z);
1282
1283                v_ns.double_promote_rank();
1284                y_ns.demote_rank();
1285                z_ns.double_demote_rank();
1286                self.observer.record_erase_double_rotation();
1287            }
1288        }
1289    }
1290
1291    /// Promotes the single child of `node` (in direction `LR`) to take `node`'s place in the tree.
1292    ///
1293    /// # Safety
1294    ///
1295    /// - `owner` must be a valid, aligned, dereferenceable pointer to a `*mut P::Target`.
1296    /// - `*owner` must be `null_mut()`.
1297    /// - `node` must be a valid, aligned, dereferenceable pointer to a node currently in the tree.
1298    /// - `node` must have exactly one child in the direction specified by `LR` (which must be
1299    ///   a valid non-null, non-sentinel node), and must NOT have a valid child in the opposite direction.
1300    unsafe fn promote_lr_child<LR: LrTraits>(
1301        &mut self,
1302        owner: *mut *mut P::Target,
1303        node: *mut P::Target,
1304    ) {
1305        // SAFETY: The caller must guarantee that `owner` is a valid, aligned, dereferenceable pointer
1306        // to `*mut P::Target` containing `null_mut()`, that `node` is a valid node in the tree,
1307        // and that the node has exactly one child in the `LR` direction to promote.
1308        unsafe {
1309            debug_assert!((*owner).is_null());
1310            debug_assert!(valid_sentinel_ptr(node));
1311
1312            let ns = Self::get_node_ref(node);
1313            let lr_child_ptr = LR::lr_child_ptr(ns);
1314            let rl_child_ptr = LR::rl_child_ptr(ns);
1315
1316            debug_assert!(valid_sentinel_ptr(*lr_child_ptr) && !valid_sentinel_ptr(*rl_child_ptr));
1317
1318            *owner = *lr_child_ptr;
1319            *lr_child_ptr = core::ptr::null_mut();
1320            Self::get_node_ref(*owner).set_parent(ns.get_parent());
1321
1322            let rl_most = LR::rl_most(self);
1323            debug_assert_eq!(rl_most == node, is_sentinel_ptr(*rl_child_ptr));
1324
1325            if is_sentinel_ptr(*rl_child_ptr) {
1326                let mut replacement = *owner;
1327                let mut next_rl_child_ptr;
1328
1329                loop {
1330                    let replacement_ns = Self::get_node_ref(replacement);
1331                    next_rl_child_ptr = LR::rl_child_ptr(replacement_ns);
1332
1333                    debug_assert!(!is_sentinel_ptr(*next_rl_child_ptr));
1334                    if (*next_rl_child_ptr).is_null() {
1335                        break;
1336                    }
1337                    replacement = *next_rl_child_ptr;
1338                }
1339
1340                LR::set_rl_most(self, replacement);
1341                *next_rl_child_ptr = self.get_sentinel();
1342                *rl_child_ptr = core::ptr::null_mut();
1343            }
1344
1345            ns.set_parent(core::ptr::null_mut());
1346            debug_assert!(ns.get_left().is_null());
1347            debug_assert!(ns.get_right().is_null());
1348        }
1349    }
1350
1351    /// Physically swaps the position of `node1` (pointed to by `ptr_ref1`) with `node2`
1352    /// (pointed to by `ptr_ref2`) in the tree's pointer structure.
1353    ///
1354    /// E.g. `node2` must be the leftmost descendant of the right child of `node1`.
1355    ///
1356    /// Returns the pointer to the slot originally containing `node2` (which now contains `node1`).
1357    ///
1358    /// # Safety
1359    ///
1360    /// - `ptr_ref1` must be a valid, aligned, dereferenceable pointer to `*mut P::Target`
1361    ///   which contains a valid, aligned, dereferenceable pointer to `node1`.
1362    /// - `ptr_ref2` must be a valid, aligned, dereferenceable pointer to `*mut P::Target`
1363    ///   which contains a valid, aligned, dereferenceable pointer to `node2`.
1364    /// - `node2` must be a descendant of `node1`'s right subtree.
1365    /// - Both `node1` and `node2` must reside within the same tree.
1366    unsafe fn swap_with_right_descendant(
1367        &mut self,
1368        ptr_ref1: *mut *mut P::Target,
1369        ptr_ref2: *mut *mut P::Target,
1370    ) -> *mut *mut P::Target {
1371        // SAFETY: The caller must guarantee that `ptr_ref1` and `ptr_ref2` are valid, aligned
1372        // pointers pointing to valid nodes in the tree, and that `node2` is a descendant in the
1373        // right subtree of `node1`. This method performs structural pointer manipulation to swap
1374        // the nodes physically, preserving local tree structure.
1375        unsafe {
1376            let node1 = *ptr_ref1;
1377            let node2 = *ptr_ref2;
1378
1379            let ns1 = Self::get_node_ref(node1);
1380            let ns2 = Self::get_node_ref(node2);
1381
1382            if ns1.get_right().is_null() {
1383                panic!("node1 right is NULL inside swap");
1384            }
1385
1386            let ns1_lp = if valid_sentinel_ptr(ns1.get_left()) {
1387                Self::get_node_ref(ns1.get_left()).parent.get()
1388            } else {
1389                core::ptr::null_mut()
1390            };
1391
1392            let ns2_lp = if valid_sentinel_ptr(ns2.get_left()) {
1393                Self::get_node_ref(ns2.get_left()).parent.get()
1394            } else {
1395                core::ptr::null_mut()
1396            };
1397
1398            let ns2_rp = if valid_sentinel_ptr(ns2.get_right()) {
1399                Self::get_node_ref(ns2.get_right()).parent.get()
1400            } else {
1401                core::ptr::null_mut()
1402            };
1403
1404            let r1 = ns1.get_right();
1405            if !valid_sentinel_ptr(r1) {
1406                if r1.is_null() {
1407                    panic!("ns1.get_right() is NULL");
1408                } else if is_sentinel_ptr(r1) {
1409                    panic!("ns1.get_right() is SENTINEL");
1410                } else {
1411                    panic!("ns1.get_right() is OTHER INVALID");
1412                }
1413            }
1414            let ns1_rp = Self::get_node_ref(ns1.get_right()).parent.get();
1415
1416            if node1 == self.left_most {
1417                self.left_most = node2;
1418            }
1419            if node2 == self.right_most {
1420                self.right_most = node1;
1421            }
1422
1423            // Swap parent.
1424            let parent_tmp = ns1.get_parent();
1425            ns1.set_parent(ns2.get_parent());
1426            ns2.set_parent(parent_tmp);
1427
1428            // Swap left.
1429            let left_tmp = ns1.get_left();
1430            ns1.set_left(ns2.get_left());
1431            ns2.set_left(left_tmp);
1432
1433            // Swap right.
1434            let right_tmp = ns1.get_right();
1435            ns1.set_right(ns2.get_right());
1436            ns2.set_right(right_tmp);
1437
1438            // Swap rank.
1439            let rank_tmp = *ns1.rank.get();
1440            *ns1.rank.get() = *ns2.rank.get();
1441            *ns2.rank.get() = rank_tmp;
1442
1443            if !ns1_lp.is_null() {
1444                *ns1_lp = node2;
1445            }
1446            if !ns2_lp.is_null() {
1447                *ns2_lp = node1;
1448            }
1449            if !ns2_rp.is_null() {
1450                *ns2_rp = node1;
1451            }
1452
1453            if ptr_ref2 != ns1.right.get() {
1454                #[allow(clippy::swap_ptr_to_ref)]
1455                core::mem::swap(&mut *ptr_ref1, &mut *ptr_ref2);
1456                *ns1_rp = node2;
1457                ptr_ref2
1458            } else {
1459                debug_assert_eq!(*ns1.parent.get(), node1);
1460                debug_assert_eq!(*ns2.right.get(), node2);
1461                #[allow(clippy::swap_ptr_to_ref)]
1462                core::mem::swap(&mut *ptr_ref1, &mut *ns2.right.get());
1463                *ns1.parent.get() = node2;
1464                ns2.right.get()
1465            }
1466        }
1467    }
1468
1469    /// Inserts a new node `ptr` into the WAVL tree.
1470    ///
1471    /// If a node with an identical key already exists, does not insert it, stores the colliding node's
1472    /// pointer in `collision`, and returns the original `ptr` as `Err(ptr)`.
1473    ///
1474    /// # Safety
1475    ///
1476    /// - `ptr` must wrap a valid, properly aligned, dereferenceable node.
1477    /// - The node must NOT be currently contained in this or any other intrusive container.
1478    /// - `collision` must be a valid, aligned, dereferenceable pointer to `*mut P::Target`.
1479    unsafe fn internal_insert(&mut self, ptr: P, collision: &mut *mut P::Target) -> Result<(), P> {
1480        // SAFETY: The caller guarantees that `ptr` represents a valid, unlinked node, and that
1481        // `collision` is a valid slot. Dereferencing pointers and mutating the parent/child links
1482        // preserves tree structure.
1483        unsafe {
1484            let raw = P::into_raw(ptr);
1485            debug_assert!(!raw.is_null());
1486
1487            let ns = Self::get_node_ref(raw);
1488            debug_assert!(ns.is_valid() && !ns.in_container());
1489
1490            *ns.rank.get() = <TargetRank<P, Tag>>::DEFAULT;
1491
1492            if self.root.is_null() {
1493                ns.set_parent(self.get_sentinel());
1494                ns.set_left(self.get_sentinel());
1495                ns.set_right(self.get_sentinel());
1496
1497                debug_assert!(is_sentinel_ptr(self.left_most) && is_sentinel_ptr(self.right_most));
1498                self.left_most = raw;
1499                self.right_most = raw;
1500
1501                self.root = raw;
1502                self.size.increment();
1503                self.observer.record_insert(raw);
1504                return Ok(());
1505            }
1506
1507            let key = (*raw).get_key();
1508            let key_borrow = key.borrow();
1509            let mut is_left_most = true;
1510            let mut is_right_most = true;
1511            let mut parent = self.root;
1512            let mut owner: *mut *mut P::Target;
1513
1514            loop {
1515                let parent_key = (*parent).get_key();
1516                let parent_key_borrow = parent_key.borrow();
1517                self.observer.record_insert_traverse(raw, parent);
1518
1519                if key_borrow == parent_key_borrow {
1520                    *collision = parent;
1521                    self.observer.record_insert_collision(raw, parent);
1522                    return Err(P::from_raw(raw));
1523                }
1524
1525                let parent_ns = Self::get_node_ref(parent);
1526
1527                if key_borrow < parent_key_borrow {
1528                    owner = parent_ns.left.get();
1529                    is_right_most = false;
1530                } else {
1531                    owner = parent_ns.right.get();
1532                    is_left_most = false;
1533                }
1534
1535                if !valid_sentinel_ptr(*owner) {
1536                    break;
1537                }
1538
1539                parent = *owner;
1540            }
1541
1542            debug_assert!(!is_left_most || !is_right_most);
1543
1544            if is_right_most {
1545                debug_assert!(is_sentinel_ptr(*owner));
1546                ns.set_right(self.get_sentinel());
1547                self.right_most = raw;
1548            } else if is_left_most {
1549                debug_assert!(is_sentinel_ptr(*owner));
1550                ns.set_left(self.get_sentinel());
1551                self.left_most = raw;
1552            }
1553
1554            debug_assert!(!valid_sentinel_ptr(*owner));
1555            ns.set_parent(parent);
1556
1557            *owner = raw;
1558            self.size.increment();
1559            self.observer.record_insert(raw);
1560
1561            self.balance_post_insert(*owner);
1562            Ok(())
1563        }
1564    }
1565
1566    /// Removes the node `ptr` from the WAVL tree, rebalancing if necessary.
1567    ///
1568    /// Returns the node wrapped in `P` if it was successfully erased and returned.
1569    ///
1570    /// # Safety
1571    ///
1572    /// - `ptr` must be a valid, properly aligned, dereferenceable raw pointer to a node.
1573    /// - If the node is not null or sentinel, it MUST be currently contained within this tree.
1574    unsafe fn internal_erase(&mut self, ptr: *mut P::Target) -> Option<P> {
1575        // SAFETY: The caller guarantees that `ptr` points to a valid node belonging to this tree.
1576        // Removing it involves swapping it out structurally, repairing BST/WAVL invariants,
1577        // and reclaiming ownership via `P::from_raw`.
1578        unsafe {
1579            if !valid_sentinel_ptr(ptr) {
1580                return None;
1581            }
1582
1583            let ns = Self::get_node_ref(ptr);
1584            let mut owner = self.get_link_ptr_to_node(ptr);
1585            debug_assert_eq!(*owner, ptr);
1586
1587            if valid_sentinel_ptr(ns.get_left()) && valid_sentinel_ptr(ns.get_right()) {
1588                let mut new_owner = ns.right.get();
1589                let mut new_ns = Self::get_node_ref(ns.get_right());
1590
1591                while !new_ns.get_left().is_null() {
1592                    debug_assert!(!is_sentinel_ptr(new_ns.get_left()));
1593                    new_owner = new_ns.left.get();
1594                    new_ns = Self::get_node_ref(*new_owner);
1595                }
1596
1597                owner = self.swap_with_right_descendant(owner, new_owner);
1598                debug_assert_eq!(*owner, ptr);
1599            }
1600
1601            let parent = ns.get_parent();
1602            let was_one_child;
1603            let was_left_child;
1604
1605            debug_assert!(!parent.is_null());
1606            if !is_sentinel_ptr(parent) {
1607                let parent_ns = Self::get_node_ref(parent);
1608                was_one_child = ns.rank_parity() != parent_ns.rank_parity();
1609                was_left_child = parent_ns.left.get() == owner;
1610            } else {
1611                was_one_child = false;
1612                was_left_child = false;
1613            }
1614
1615            *owner = core::ptr::null_mut();
1616
1617            let target = ptr;
1618            if valid_sentinel_ptr(ns.get_left()) {
1619                self.promote_lr_child::<ForwardTraits>(owner, target);
1620            } else if valid_sentinel_ptr(ns.get_right()) {
1621                self.promote_lr_child::<ReverseTraits>(owner, target);
1622            } else {
1623                debug_assert_eq!(is_sentinel_ptr(ns.get_left()), self.left_most == target);
1624                debug_assert_eq!(is_sentinel_ptr(ns.get_right()), self.right_most == target);
1625
1626                if is_sentinel_ptr(ns.get_left()) {
1627                    if is_sentinel_ptr(ns.get_right()) {
1628                        if S::IS_TRACKING {
1629                            debug_assert_eq!(self.size.get(), 1);
1630                        }
1631                        debug_assert!(is_sentinel_ptr(ns.get_parent()));
1632                        self.left_most = self.get_sentinel();
1633                        self.right_most = self.get_sentinel();
1634                        ns.set_left(core::ptr::null_mut());
1635                        ns.set_right(core::ptr::null_mut());
1636                    } else {
1637                        debug_assert!(valid_sentinel_ptr(ns.get_parent()));
1638                        debug_assert!(ns.get_right().is_null());
1639                        self.left_most = ns.get_parent();
1640                        *owner = ns.get_left();
1641                        ns.set_left(core::ptr::null_mut());
1642                    }
1643                } else if is_sentinel_ptr(ns.get_right()) {
1644                    debug_assert!(valid_sentinel_ptr(ns.get_parent()));
1645                    debug_assert!(ns.get_left().is_null());
1646                    self.right_most = ns.get_parent();
1647                    *owner = ns.get_right();
1648                    ns.set_right(core::ptr::null_mut());
1649                }
1650
1651                ns.set_parent(core::ptr::null_mut());
1652            }
1653
1654            debug_assert!(ns.is_valid() && !ns.in_container());
1655            self.observer.record_erase(target, parent);
1656
1657            self.size.decrement();
1658
1659            if !is_sentinel_ptr(parent) {
1660                if was_one_child {
1661                    self.balance_post_erase_fix_22_leaf(parent);
1662                } else {
1663                    if was_left_child {
1664                        self.balance_post_erase_fix_lr_3_child::<ForwardTraits>(parent);
1665                    } else {
1666                        self.balance_post_erase_fix_lr_3_child::<ReverseTraits>(parent);
1667                    }
1668                }
1669            }
1670
1671            Some(P::from_raw(target))
1672        }
1673    }
1674
1675    /// Replaces `old_node` with `new_node` in the tree's pointer structure.
1676    ///
1677    /// Returns the `old_node` wrapped in `P`.
1678    ///
1679    /// # Safety
1680    ///
1681    /// - `old_node` must be a valid, aligned, dereferenceable pointer to a node currently
1682    ///   contained within this tree.
1683    /// - `new_node` must be a valid node not currently contained in this or any other tree.
1684    /// - The key of `new_node` must exactly match the key of `old_node` to preserve the
1685    ///   Binary Search Tree (BST) ordering invariant.
1686    unsafe fn internal_swap(&mut self, old_node: *mut P::Target, new_node: P) -> Option<P> {
1687        // SAFETY: The caller must guarantee that `old_node` is in this tree, that `new_node`
1688        // is unlinked, and that their keys match. This method updates all parent and child
1689        // links to point to the new node, and reclaims ownership of `old_node`.
1690        unsafe {
1691            debug_assert!(!old_node.is_null());
1692            let new_raw = P::into_raw(new_node);
1693            debug_assert!(!new_raw.is_null());
1694            debug_assert!((*old_node).get_key().borrow() == (*new_raw).get_key().borrow());
1695
1696            let old_ns = Self::get_node_ref(old_node);
1697            let new_ns = Self::get_node_ref(new_raw);
1698
1699            debug_assert!(old_ns.in_container());
1700            debug_assert!(!new_ns.in_container());
1701            self.observer.record_insert_replace(old_node, new_raw);
1702
1703            if valid_sentinel_ptr(old_ns.get_left()) {
1704                Self::get_node_ref(old_ns.get_left()).set_parent(new_raw);
1705            } else {
1706                if is_sentinel_ptr(old_ns.get_left()) {
1707                    debug_assert_eq!(self.left_most, old_node);
1708                    self.left_most = new_raw;
1709                }
1710            }
1711            new_ns.set_left(old_ns.get_left());
1712            old_ns.set_left(core::ptr::null_mut());
1713
1714            if valid_sentinel_ptr(old_ns.get_right()) {
1715                Self::get_node_ref(old_ns.get_right()).set_parent(new_raw);
1716            } else {
1717                if is_sentinel_ptr(old_ns.get_right()) {
1718                    debug_assert_eq!(self.right_most, old_node);
1719                    self.right_most = new_raw;
1720                }
1721            }
1722            new_ns.set_right(old_ns.get_right());
1723            old_ns.set_right(core::ptr::null_mut());
1724
1725            *new_ns.rank.get() = *old_ns.rank.get();
1726
1727            *self.get_link_ptr_to_node(old_node) = new_raw;
1728            new_ns.set_parent(old_ns.get_parent());
1729            old_ns.set_parent(core::ptr::null_mut());
1730
1731            Some(P::from_raw(old_node))
1732        }
1733    }
1734
1735    /// Advances the node pointer `node` in-place to the next node in-order
1736    /// (according to the direction specified by `LR`).
1737    ///
1738    /// # Safety
1739    ///
1740    /// - `node` must be a valid, aligned, dereferenceable pointer to a raw pointer `*node`.
1741    /// - `*node` must point to a valid node currently contained within this tree.
1742    unsafe fn advance<LR: LrTraits>(node: &mut *mut P::Target) {
1743        // SAFETY: The caller must ensure that `*node` is a valid pointer to a node in this tree.
1744        // Traveling through parent/child links is safe as long as the tree structure is valid
1745        // and the node belongs to the tree.
1746        unsafe {
1747            debug_assert!(valid_sentinel_ptr(*node));
1748
1749            let mut ns = Self::get_node_ref(*node);
1750            let rl_child = LR::rl_child(ns);
1751            if !rl_child.is_null() {
1752                *node = rl_child;
1753
1754                if is_sentinel_ptr(*node) {
1755                    return;
1756                }
1757
1758                let mut lr_child = LR::lr_child(Self::get_node_ref(*node));
1759                while !lr_child.is_null() {
1760                    debug_assert!(!is_sentinel_ptr(lr_child));
1761                    *node = lr_child;
1762                    lr_child = LR::lr_child(Self::get_node_ref(*node));
1763                }
1764                return;
1765            }
1766
1767            let mut done;
1768            ns = Self::get_node_ref(*node);
1769            loop {
1770                debug_assert!(valid_sentinel_ptr(ns.get_parent()));
1771
1772                let parent_ns = Self::get_node_ref(ns.get_parent());
1773                done = LR::lr_child(parent_ns) == *node;
1774
1775                debug_assert!(done || LR::rl_child(parent_ns) == *node);
1776
1777                *node = ns.get_parent();
1778                ns = parent_ns;
1779
1780                if done {
1781                    break;
1782                }
1783            }
1784        }
1785    }
1786
1787    /// Inserts an element into the tree.
1788    ///
1789    /// For raw pointers, use [`insert_raw`] instead.
1790    pub fn insert(&mut self, ptr: P)
1791    where
1792        P: ManagedPtr,
1793    {
1794        // SAFETY: `P` is a `ManagedPtr`, which guarantees that the pointer is valid and that the
1795        // object will outlive its reference from this tree.
1796        unsafe { self.insert_raw(ptr) }
1797    }
1798
1799    /// Inserts an element into the tree.
1800    ///
1801    /// # Safety
1802    ///
1803    /// The caller must ensure that `ptr` is a valid pointer to a `T` and that the object outlives
1804    /// the reference from the tree.
1805    pub unsafe fn insert_raw(&mut self, ptr: P) {
1806        let mut collision = core::ptr::null_mut();
1807        // SAFETY: The caller guarantees `ptr` is valid and outlives the tree registration.
1808        let _ = unsafe { self.internal_insert(ptr, &mut collision) };
1809    }
1810
1811    /// Inserts the object pointed to by `ptr` if it is not already in the tree,
1812    /// or finds the object that `ptr` collided with instead.
1813    ///
1814    /// For raw pointers, use [`insert_or_find_raw`] instead.
1815    ///
1816    /// # Returns
1817    ///
1818    /// * `Ok(())` if there was no collision and the item was successfully inserted.
1819    ///   In this case, the tree takes ownership of `ptr`.
1820    /// * `Err((ptr, cursor))` if there was a collision. In this case, the
1821    ///   passed pointer `ptr` is returned back to the caller (not consumed), along with
1822    ///   a `CursorMut` positioned at the colliding node already in the tree.
1823    pub fn insert_or_find<'a>(
1824        &'a mut self,
1825        ptr: P,
1826    ) -> Result<(), (P, CursorMut<'a, K, P, Tag, S, O>)>
1827    where
1828        P: ManagedPtr,
1829    {
1830        // SAFETY: `P` is a `ManagedPtr`, which guarantees that the pointer is valid and that the
1831        // object will outlive its reference from this tree.
1832        unsafe { self.insert_or_find_raw(ptr) }
1833    }
1834
1835    /// Inserts the object pointed to by `ptr` if it is not already in the tree,
1836    /// or finds the object that `ptr` collided with instead.
1837    ///
1838    /// # Safety
1839    ///
1840    /// The caller must ensure that `ptr` is a valid pointer to a `T` and that the object outlives
1841    /// the reference from the tree.
1842    pub unsafe fn insert_or_find_raw<'a>(
1843        &'a mut self,
1844        ptr: P,
1845    ) -> Result<(), (P, CursorMut<'a, K, P, Tag, S, O>)> {
1846        let mut collision = core::ptr::null_mut();
1847        // SAFETY: The caller guarantees `ptr` is valid and outlives the tree registration.
1848        // If a collision occurs, `collision` is guaranteed to be a valid pointer to the colliding
1849        // node in this tree, which allows us to safely construct a `CursorMut` pointing to it.
1850        unsafe {
1851            match self.internal_insert(ptr, &mut collision) {
1852                Ok(()) => Ok(()),
1853                Err(ptr) => Err((ptr, CursorMut { tree: self, current: collision })),
1854            }
1855        }
1856    }
1857
1858    /// Finds the element in the tree with the same key as `*ptr` and replaces
1859    /// it with `ptr`, returning the element which was replaced.
1860    ///
1861    /// If no element in the tree shares a key with `*ptr`, simply adds `ptr` to
1862    /// the tree and returns `None`.
1863    ///
1864    /// In both cases, the input pointer `ptr` is consumed.
1865    ///
1866    /// For raw pointers, use [`insert_or_replace_raw`] instead.
1867    ///
1868    /// # Returns
1869    ///
1870    /// `Some(replaced)` containing the previous element if a collision occurred,
1871    /// or `None` if the element was newly inserted.
1872    pub fn insert_or_replace(&mut self, ptr: P) -> Option<P>
1873    where
1874        P: ManagedPtr,
1875    {
1876        // SAFETY: `P` is a `ManagedPtr`, which guarantees that the pointer is valid and that the
1877        // object will outlive its reference from this tree.
1878        unsafe { self.insert_or_replace_raw(ptr) }
1879    }
1880
1881    /// Finds the element in the tree with the same key as `*ptr` and replaces
1882    /// it with `ptr`, returning the element which was replaced.
1883    ///
1884    /// If no element in the tree shares a key with `*ptr`, simply adds `ptr` to
1885    /// the tree and returns `None`.
1886    ///
1887    /// # Safety
1888    ///
1889    /// The caller must ensure that `ptr` is a valid pointer to a `T` and that the object outlives
1890    /// the reference from the tree.
1891    pub unsafe fn insert_or_replace_raw(&mut self, ptr: P) -> Option<P> {
1892        let mut collision = core::ptr::null_mut();
1893        // SAFETY: The caller guarantees `ptr` is valid and outlives the tree registration.
1894        // If a collision occurs, `collision` points to a valid node in the tree sharing the same key,
1895        // making it safe to swap `collision` with `ptr` via `internal_swap`.
1896        unsafe {
1897            match self.internal_insert(ptr, &mut collision) {
1898                Ok(()) => None,
1899                Err(ptr) => self.internal_swap(collision, ptr),
1900            }
1901        }
1902    }
1903
1904    /// Removes and returns the first (smallest) element of the tree, or `None` if it is empty.
1905    pub fn pop_front(&mut self) -> Option<P> {
1906        if self.is_empty() {
1907            None
1908        } else {
1909            // SAFETY: If `self.is_empty()` is false, `self.left_most` is guaranteed to be a valid
1910            // pointer to a node currently contained in this tree instance.
1911            unsafe { self.internal_erase(self.left_most) }
1912        }
1913    }
1914
1915    /// Removes and returns the last (largest) element of the tree, or `None` if it is empty.
1916    pub fn pop_back(&mut self) -> Option<P> {
1917        if self.is_empty() {
1918            None
1919        } else {
1920            // SAFETY: If `self.is_empty()` is false, `self.right_most` is guaranteed to be a valid
1921            // pointer to a node currently contained in this tree instance.
1922            unsafe { self.internal_erase(self.right_most) }
1923        }
1924    }
1925
1926    /// Removes all elements from the tree.
1927    pub fn clear(&mut self) {
1928        while !self.is_empty() {
1929            self.pop_front();
1930        }
1931    }
1932
1933    /// Swaps the contents of this tree with another tree.
1934    ///
1935    /// This runs in O(1) time.
1936    pub fn swap(&mut self, other: &mut Self) {
1937        // Swap all fields except _pin and _phantom.
1938        core::mem::swap(&mut self.root, &mut other.root);
1939        core::mem::swap(&mut self.left_most, &mut other.left_most);
1940        core::mem::swap(&mut self.right_most, &mut other.right_most);
1941        core::mem::swap(&mut self.size, &mut other.size);
1942        core::mem::swap(&mut self.observer, &mut other.observer);
1943
1944        // Now repair the sentinel pointers which are self-referential.
1945        self.fix_sentinels_after_swap(other);
1946    }
1947
1948    fn fix_sentinels_after_swap(&mut self, other: &mut Self) {
1949        let self_sentinel = self.get_sentinel();
1950        let other_sentinel = other.get_sentinel();
1951
1952        // For `self` (which currently contains `other`'s old nodes):
1953        // The old sentinel in these nodes is `other_sentinel`. We update them to `self_sentinel`.
1954        if self.root.is_null() {
1955            self.left_most = self_sentinel;
1956            self.right_most = self_sentinel;
1957        } else {
1958            // SAFETY: Sentinels are verified to be valid and correspond to the correct node locations.
1959            unsafe {
1960                let root_ns = Self::get_node_ref(self.root);
1961                debug_assert_eq!(root_ns.get_parent(), other_sentinel);
1962                root_ns.set_parent(self_sentinel);
1963
1964                let left_ns = Self::get_node_ref(self.left_most);
1965                debug_assert_eq!(left_ns.get_left(), other_sentinel);
1966                left_ns.set_left(self_sentinel);
1967
1968                let right_ns = Self::get_node_ref(self.right_most);
1969                debug_assert_eq!(right_ns.get_right(), other_sentinel);
1970                right_ns.set_right(self_sentinel);
1971            }
1972        }
1973
1974        // For `other` (which currently contains `self`'s old nodes):
1975        // The old sentinel in these nodes is `self_sentinel`. We update them to `other_sentinel`.
1976        if other.root.is_null() {
1977            other.left_most = other_sentinel;
1978            other.right_most = other_sentinel;
1979        } else {
1980            // SAFETY: Sentinels are verified to be valid and correspond to the correct node locations.
1981            unsafe {
1982                let root_ns = Self::get_node_ref(other.root);
1983                debug_assert_eq!(root_ns.get_parent(), self_sentinel);
1984                root_ns.set_parent(other_sentinel);
1985
1986                let left_ns = Self::get_node_ref(other.left_most);
1987                debug_assert_eq!(left_ns.get_left(), self_sentinel);
1988                left_ns.set_left(other_sentinel);
1989
1990                let right_ns = Self::get_node_ref(other.right_most);
1991                debug_assert_eq!(right_ns.get_right(), self_sentinel);
1992                right_ns.set_right(other_sentinel);
1993            }
1994        }
1995    }
1996
1997    /// Traverses the tree to find the node with the given key.
1998    ///
1999    /// Returns a raw pointer to the matching node, or a sentinel pointer if not found.
2000    ///
2001    /// # Safety
2002    ///
2003    /// The returned raw pointer is only valid as long as the tree structure is not modified
2004    /// and no elements are deleted.
2005    unsafe fn find_raw(&self, key: &K) -> *mut P::Target {
2006        // SAFETY: Accessing node keys and traversing child pointers is safe since we only
2007        // traverse nodes contained inside this tree and ensure they are valid via `valid_sentinel_ptr`.
2008        unsafe {
2009            let mut node = self.root;
2010            while valid_sentinel_ptr(node) {
2011                let node_key = (*node).get_key();
2012                let b = node_key.borrow();
2013                if key == b {
2014                    return node;
2015                }
2016                let ns = Self::get_node_ref(node);
2017                node = if key < b { ns.get_left() } else { ns.get_right() };
2018            }
2019            self.get_sentinel()
2020        }
2021    }
2022
2023    /// Traverses the tree to find either the lower bound or upper bound node pointer.
2024    ///
2025    /// # Safety
2026    ///
2027    /// The returned raw pointer is only valid as long as the tree structure is not modified
2028    /// and no elements are deleted.
2029    unsafe fn bound_raw(&self, key: &K, strictly_greater: bool) -> *mut P::Target {
2030        // SAFETY: Accessing node keys and traversing child pointers is safe since we only
2031        // traverse nodes contained inside this tree and ensure they are valid via `valid_sentinel_ptr`.
2032        unsafe {
2033            let mut node = self.root;
2034            let mut found = self.get_sentinel();
2035
2036            while valid_sentinel_ptr(node) {
2037                let node_key = (*node).get_key();
2038                let b = node_key.borrow();
2039                let is_eligible = if strictly_greater { b > key } else { b >= key };
2040                if is_eligible {
2041                    found = node;
2042                    node = Self::get_node_ref(node).get_left();
2043                } else {
2044                    node = Self::get_node_ref(node).get_right();
2045                }
2046            }
2047            found
2048        }
2049    }
2050
2051    /// Finds an element in the tree by key.
2052    pub fn find(&self, key: &K) -> Option<&P::Target> {
2053        // SAFETY: find_raw returns either a sentinel pointer or a valid node in the tree.
2054        // If it is valid, returning a reference is safe for the lifetime of the borrow of `self`.
2055        unsafe {
2056            let node = self.find_raw(key);
2057            if valid_sentinel_ptr(node) { Some(&*node) } else { None }
2058        }
2059    }
2060
2061    /// Finds an element in the tree by key and returns a cursor positioned at it.
2062    ///
2063    /// If the key is not found, the returned cursor is positioned at the sentinel
2064    /// (i.e. `cursor.get()` will return `None`).
2065    pub fn find_cursor(&mut self, key: &K) -> CursorMut<'_, K, P, Tag, S, O> {
2066        // SAFETY: find_raw returns a valid node pointer or sentinel pointer belonging to this tree.
2067        let node = unsafe { self.find_raw(key) };
2068        CursorMut { tree: self, current: node }
2069    }
2070
2071    /// Returns a cursor positioned at the lower bound of the key (the first element
2072    /// in the tree whose key is greater than or equal to `key`).
2073    ///
2074    /// If no such element exists (e.g. all elements in the tree are smaller than `key`),
2075    /// the returned cursor is positioned at the sentinel.
2076    pub fn lower_bound(&mut self, key: &K) -> CursorMut<'_, K, P, Tag, S, O> {
2077        // SAFETY: bound_raw returns a valid node pointer or sentinel pointer in the tree.
2078        let node = unsafe { self.bound_raw(key, false) };
2079        CursorMut { tree: self, current: node }
2080    }
2081
2082    /// Returns a cursor positioned at the upper bound of the key (the first element
2083    /// in the tree whose key is strictly greater than `key`).
2084    ///
2085    /// If no such element exists (e.g. all elements in the tree are smaller than or
2086    /// equal to `key`), the returned cursor is positioned at the sentinel.
2087    pub fn upper_bound(&mut self, key: &K) -> CursorMut<'_, K, P, Tag, S, O> {
2088        // SAFETY: bound_raw returns a valid node pointer or sentinel pointer in the tree.
2089        let node = unsafe { self.bound_raw(key, true) };
2090        CursorMut { tree: self, current: node }
2091    }
2092
2093    /// Erases an element by key.
2094    pub fn erase(&mut self, key: &K) -> Option<P> {
2095        let mut cursor = self.find_cursor(key);
2096        cursor.erase()
2097    }
2098
2099    /// Erases an element by raw pointer.
2100    ///
2101    /// # Safety
2102    ///
2103    /// The caller must ensure that `obj` is a valid pointer to an object that is currently
2104    /// contained within this tree instance.
2105    pub unsafe fn erase_raw(&mut self, obj: *mut P::Target) -> Option<P> {
2106        // SAFETY: The caller guarantees that `obj` is currently contained in this WavlTree.
2107        // `internal_erase` is safe to execute on a contained pointer.
2108        unsafe {
2109            let node = (*obj).get_node();
2110            if !node.in_container() {
2111                return None;
2112            }
2113            self.internal_erase(obj)
2114        }
2115    }
2116
2117    /// Returns a cursor positioned at the front (smallest element) of the tree.
2118    pub fn cursor_mut(&mut self) -> CursorMut<'_, K, P, Tag, S, O> {
2119        let left_most = self.left_most;
2120        CursorMut { tree: self, current: left_most }
2121    }
2122
2123    /// Returns a read-only cursor positioned at the given element.
2124    ///
2125    /// # Safety
2126    ///
2127    /// The caller must ensure that `obj` is a member of this tree.
2128    /// It is undefined behavior to use the returned cursor if `obj` is not in the tree,
2129    /// or if it is in a different tree.
2130    pub unsafe fn cursor_at(&self, obj: &P::Target) -> Cursor<'_, K, P, Tag, S, O> {
2131        assert!(obj.get_node().in_container(), "Object must be in a container");
2132        Cursor { tree: self, current: obj as *const P::Target as *mut P::Target }
2133    }
2134
2135    /// Returns a mutable cursor positioned at the given element.
2136    ///
2137    /// # Safety
2138    ///
2139    /// The caller must ensure that `obj` is a member of this tree.
2140    /// It is undefined behavior to use the returned cursor if `obj` is not in the tree,
2141    /// or if it is in a different tree.
2142    pub unsafe fn cursor_mut_at(&mut self, obj: &P::Target) -> CursorMut<'_, K, P, Tag, S, O> {
2143        assert!(obj.get_node().in_container(), "Object must be in a container");
2144        CursorMut { tree: self, current: obj as *const P::Target as *mut P::Target }
2145    }
2146
2147    /// Returns an iterator over the elements of the tree.
2148    pub fn iter(&self) -> Iterator<'_, K, P, Tag, S, O> {
2149        Iterator::new(self)
2150    }
2151
2152    /// Returns a unidirectional forward iterator over the elements of the tree.
2153    pub fn forward_iter(&self) -> ForwardIterator<'_, K, P, Tag, S, O> {
2154        ForwardIterator::new(self.left_most)
2155    }
2156
2157    /// Returns a unidirectional reverse iterator over the elements of the tree.
2158    pub fn reverse_iter(&self) -> ReverseIterator<'_, K, P, Tag, S, O> {
2159        ReverseIterator::new(self.right_most)
2160    }
2161
2162    /// Returns a read-only cursor positioned at the root of the tree.
2163    pub fn root_cursor(&self) -> Cursor<'_, K, P, Tag, S, O> {
2164        Cursor { tree: self, current: self.root }
2165    }
2166
2167    /// Returns a read-only cursor positioned at the first (smallest) element.
2168    pub fn front_cursor(&self) -> Cursor<'_, K, P, Tag, S, O> {
2169        Cursor { tree: self, current: self.left_most }
2170    }
2171
2172    /// Returns a read-only cursor positioned at the last (largest) element.
2173    pub fn back_cursor(&self) -> Cursor<'_, K, P, Tag, S, O> {
2174        Cursor { tree: self, current: self.right_most }
2175    }
2176
2177    /// Returns the number of elements in the tree.
2178    pub fn len(&self) -> usize {
2179        self.size.get()
2180    }
2181}
2182
2183#[pinned_drop]
2184impl<K, P, Tag, S, O> PinnedDrop for WavlTree<K, P, Tag, S, O>
2185where
2186    P: PtrTraits,
2187    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2188    K: Ord,
2189    S: SizeTracker,
2190    O: WavlTreeObserver<Target = P::Target>,
2191{
2192    fn drop(self: Pin<&mut Self>) {
2193        if P::IS_MANAGED {
2194            let me = unsafe { self.get_unchecked_mut() };
2195            me.clear();
2196        } else {
2197            debug_assert!(self.is_empty(), "Tree must be empty on destruction");
2198            if S::IS_TRACKING {
2199                debug_assert_eq!(self.size.get(), 0, "Size must be zero on destruction");
2200            }
2201        }
2202    }
2203}
2204
2205/// A read-only cursor positioned in a `WavlTree`.
2206pub struct Cursor<
2207    'a,
2208    K,
2209    P,
2210    Tag = DefaultObjectTag,
2211    S = NonTrackingSize,
2212    O = DefaultWavlTreeObserver<<P as PtrTraits>::Target>,
2213> where
2214    P: PtrTraits,
2215    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2216    K: Ord,
2217    S: SizeTracker,
2218    O: WavlTreeObserver<Target = P::Target>,
2219{
2220    tree: &'a WavlTree<K, P, Tag, S, O>,
2221    current: *mut P::Target,
2222}
2223
2224impl<'a, K, P, Tag, S, O> Clone for Cursor<'a, K, P, Tag, S, O>
2225where
2226    P: PtrTraits,
2227    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2228    K: Ord,
2229    S: SizeTracker,
2230    O: WavlTreeObserver<Target = P::Target>,
2231{
2232    fn clone(&self) -> Self {
2233        *self
2234    }
2235}
2236
2237impl<'a, K, P, Tag, S, O> Copy for Cursor<'a, K, P, Tag, S, O>
2238where
2239    P: PtrTraits,
2240    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2241    K: Ord,
2242    S: SizeTracker,
2243    O: WavlTreeObserver<Target = P::Target>,
2244{
2245}
2246
2247impl<'a, K, P, Tag, S, O> PartialEq for Cursor<'a, K, P, Tag, S, O>
2248where
2249    P: PtrTraits,
2250    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2251    K: Ord,
2252    S: SizeTracker,
2253    O: WavlTreeObserver<Target = P::Target>,
2254{
2255    fn eq(&self, other: &Self) -> bool {
2256        self.current == other.current
2257    }
2258}
2259
2260impl<'a, K, P, Tag, S, O> core::fmt::Debug for Cursor<'a, K, P, Tag, S, O>
2261where
2262    P: PtrTraits,
2263    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2264    K: Ord,
2265    S: SizeTracker,
2266    O: WavlTreeObserver<Target = P::Target>,
2267{
2268    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
2269        f.debug_struct("Cursor").field("current", &self.current).finish()
2270    }
2271}
2272
2273impl<'a, K, P, Tag, S, O> Cursor<'a, K, P, Tag, S, O>
2274where
2275    P: PtrTraits,
2276    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2277    K: Ord,
2278    S: SizeTracker,
2279    O: WavlTreeObserver<Target = P::Target>,
2280{
2281    /// Returns a reference to the current element, or `None` if the cursor is at the sentinel.
2282    pub fn get(&self) -> Option<&'a P::Target> {
2283        if is_sentinel_ptr(self.current) {
2284            None
2285        } else {
2286            // SAFETY: `self.current` is checked to be non-sentinel.
2287            // The lifetime `'a` is tied to the `WavlTree` borrow.
2288            unsafe { Some(&*self.current) }
2289        }
2290    }
2291
2292    /// Returns true if the cursor is positioned at a valid element (not the sentinel).
2293    pub fn is_valid(&self) -> bool {
2294        valid_sentinel_ptr(self.current)
2295    }
2296
2297    /// Returns a cursor positioned at the left child of the current element.
2298    /// If the current element is the sentinel, returns a cursor at the sentinel.
2299    pub fn left(&self) -> Self {
2300        if !self.is_valid() {
2301            *self
2302        } else {
2303            // SAFETY: `self.current` is checked to be non-sentinel (which implies it is a valid,
2304            // non-null node in the tree because `current` is never null). Accessing the tree node state is safe.
2305            let ns = unsafe { WavlTree::<K, P, Tag, S, O>::get_node_ref(self.current) };
2306            Self { tree: self.tree, current: ns.get_left() }
2307        }
2308    }
2309
2310    /// Returns a cursor positioned at the right child of the current element.
2311    pub fn right(&self) -> Self {
2312        if !self.is_valid() {
2313            *self
2314        } else {
2315            // SAFETY: `self.current` is checked to be non-sentinel (which implies it is a valid,
2316            // non-null node in the tree because `current` is never null). Accessing the tree node state is safe.
2317            let ns = unsafe { WavlTree::<K, P, Tag, S, O>::get_node_ref(self.current) };
2318            Self { tree: self.tree, current: ns.get_right() }
2319        }
2320    }
2321
2322    /// Returns a cursor positioned at the parent of the current element.
2323    pub fn parent(&self) -> Self {
2324        if !self.is_valid() {
2325            *self
2326        } else {
2327            // SAFETY: `self.current` is checked to be non-sentinel (which implies it is a valid,
2328            // non-null node in the tree because `current` is never null). Accessing the tree node state is safe.
2329            let ns = unsafe { WavlTree::<K, P, Tag, S, O>::get_node_ref(self.current) };
2330            let parent = ns.get_parent();
2331            Self { tree: self.tree, current: parent }
2332        }
2333    }
2334
2335    /// Returns the raw pointer to the current element.
2336    /// This may be a sentinel pointer if the cursor is at the sentinel.
2337    pub fn as_raw_ptr(&self) -> *mut P::Target {
2338        self.current
2339    }
2340}
2341
2342/// A cursor over elements in a `WavlTree`.
2343pub struct CursorMut<
2344    'a,
2345    K,
2346    P,
2347    Tag = DefaultObjectTag,
2348    S = NonTrackingSize,
2349    O = DefaultWavlTreeObserver<<P as PtrTraits>::Target>,
2350> where
2351    P: PtrTraits,
2352    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2353    K: Ord,
2354    S: SizeTracker,
2355    O: WavlTreeObserver<Target = P::Target>,
2356{
2357    tree: &'a mut WavlTree<K, P, Tag, S, O>,
2358    current: *mut P::Target,
2359}
2360
2361impl<'a, K, P, Tag, S, O> CursorMut<'a, K, P, Tag, S, O>
2362where
2363    P: PtrTraits,
2364    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2365    K: Ord,
2366    S: SizeTracker,
2367    O: WavlTreeObserver<Target = P::Target>,
2368{
2369    /// Returns a reference to the current element.
2370    pub fn get(&self) -> Option<&P::Target> {
2371        if is_sentinel_ptr(self.current) {
2372            None
2373        } else {
2374            // SAFETY: `self.current` is checked to be non-sentinel (which implies it is a valid,
2375            // non-null node in the tree because `current` is never null). Since `CursorMut` mutably
2376            // borrows the tree, the node is guaranteed to be valid and dereferenceable for the
2377            // lifetime of the reference.
2378            unsafe { Some(&*self.current) }
2379        }
2380    }
2381
2382    /// Moves the cursor to the next (larger) element.
2383    pub fn move_next(&mut self) {
2384        if valid_sentinel_ptr(self.current) {
2385            // SAFETY: `self.current` is verified to be a valid node in the tree. Advancing through
2386            // tree pointers is safe.
2387            unsafe {
2388                WavlTree::<K, P, Tag, S, O>::advance::<ForwardTraits>(&mut self.current);
2389            }
2390        }
2391    }
2392
2393    /// Moves the cursor to the previous (smaller) element.
2394    pub fn move_prev(&mut self) {
2395        if valid_sentinel_ptr(self.current) {
2396            // SAFETY: `self.current` is verified to be a valid node in the tree. Advancing through
2397            // tree pointers is safe.
2398            unsafe {
2399                WavlTree::<K, P, Tag, S, O>::advance::<ReverseTraits>(&mut self.current);
2400            }
2401        } else if is_sentinel_ptr(self.current) {
2402            self.current = self.tree.right_most;
2403        }
2404    }
2405
2406    /// Erases the current element and moves the cursor to the next element.
2407    pub fn erase(&mut self) -> Option<P> {
2408        if !valid_sentinel_ptr(self.current) {
2409            return None;
2410        }
2411
2412        let to_erase = self.current;
2413        // SAFETY: `to_erase` is verified to be a valid, non-sentinel node in the tree.
2414        // `advance` moves the cursor to a safe position before `internal_erase` physically removes
2415        // `to_erase` from the tree.
2416        unsafe {
2417            WavlTree::<K, P, Tag, S, O>::advance::<ForwardTraits>(&mut self.current);
2418            self.tree.internal_erase(to_erase)
2419        }
2420    }
2421}
2422
2423/// A unidirectional forward iterator over the elements of a `WavlTree`.
2424pub struct ForwardIterator<
2425    'a,
2426    K,
2427    P,
2428    Tag = DefaultObjectTag,
2429    S = NonTrackingSize,
2430    O = DefaultWavlTreeObserver<<P as PtrTraits>::Target>,
2431> where
2432    P: PtrTraits,
2433    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2434    K: Ord,
2435    S: SizeTracker,
2436    O: WavlTreeObserver<Target = P::Target>,
2437{
2438    current: *mut P::Target,
2439    _phantom: core::marker::PhantomData<&'a WavlTree<K, P, Tag, S, O>>,
2440}
2441
2442impl<'a, K, P, Tag, S, O> ForwardIterator<'a, K, P, Tag, S, O>
2443where
2444    P: PtrTraits,
2445    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2446    K: Ord,
2447    S: SizeTracker,
2448    O: WavlTreeObserver<Target = P::Target>,
2449{
2450    fn new(current: *mut P::Target) -> Self {
2451        Self { current, _phantom: core::marker::PhantomData }
2452    }
2453
2454    /// Creates an iterator starting from a specific element.
2455    ///
2456    /// # Panics
2457    ///
2458    /// Panics if the object is not in a container.
2459    pub fn from_element(obj: &'a P::Target) -> Self {
2460        assert!(obj.get_node().in_container(), "Object must be in a container");
2461        Self { current: obj as *const _ as *mut _, _phantom: core::marker::PhantomData }
2462    }
2463
2464    fn get_current(&self) -> Option<&'a P::Target> {
2465        if is_sentinel_ptr(self.current) {
2466            None
2467        } else {
2468            // SAFETY: `self.current` is checked to be non-sentinel (which implies it is a valid,
2469            // non-null node in the tree because `current` is never null). Since the iterator
2470            // holds a lifetime borrow of the `WavlTree`, and the tree remains unmodified,
2471            // the node is guaranteed to be valid and dereferenceable for `'a`.
2472            unsafe { Some(&*self.current) }
2473        }
2474    }
2475}
2476
2477impl<'a, K, P, Tag, S, O> Clone for ForwardIterator<'a, K, P, Tag, S, O>
2478where
2479    P: PtrTraits,
2480    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2481    K: Ord,
2482    S: SizeTracker,
2483    O: WavlTreeObserver<Target = P::Target>,
2484{
2485    fn clone(&self) -> Self {
2486        Self { current: self.current, _phantom: core::marker::PhantomData }
2487    }
2488}
2489
2490impl<'a, K, P, Tag, S, O> core::iter::Iterator for ForwardIterator<'a, K, P, Tag, S, O>
2491where
2492    P: PtrTraits,
2493    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2494    K: Ord,
2495    S: SizeTracker,
2496    O: WavlTreeObserver<Target = P::Target>,
2497{
2498    type Item = &'a P::Target;
2499
2500    fn next(&mut self) -> Option<Self::Item> {
2501        let current = self.get_current()?;
2502        // SAFETY: `self.current` is validated as non-sentinel by `get_current`.
2503        // Moving through tree pointers is safe.
2504        unsafe {
2505            WavlTree::<K, P, Tag, S, O>::advance::<ForwardTraits>(&mut self.current);
2506        }
2507        Some(current)
2508    }
2509}
2510
2511/// A unidirectional reverse iterator over the elements of a `WavlTree`.
2512pub struct ReverseIterator<
2513    'a,
2514    K,
2515    P,
2516    Tag = DefaultObjectTag,
2517    S = NonTrackingSize,
2518    O = DefaultWavlTreeObserver<<P as PtrTraits>::Target>,
2519> where
2520    P: PtrTraits,
2521    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2522    K: Ord,
2523    S: SizeTracker,
2524    O: WavlTreeObserver<Target = P::Target>,
2525{
2526    current: *mut P::Target,
2527    _phantom: core::marker::PhantomData<&'a WavlTree<K, P, Tag, S, O>>,
2528}
2529
2530impl<'a, K, P, Tag, S, O> ReverseIterator<'a, K, P, Tag, S, O>
2531where
2532    P: PtrTraits,
2533    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2534    K: Ord,
2535    S: SizeTracker,
2536    O: WavlTreeObserver<Target = P::Target>,
2537{
2538    fn new(current: *mut P::Target) -> Self {
2539        Self { current, _phantom: core::marker::PhantomData }
2540    }
2541
2542    /// Creates an iterator starting from a specific element.
2543    ///
2544    /// # Panics
2545    ///
2546    /// Panics if the object is not in a container.
2547    pub fn from_element(obj: &'a P::Target) -> Self {
2548        assert!(obj.get_node().in_container(), "Object must be in a container");
2549        Self { current: obj as *const _ as *mut _, _phantom: core::marker::PhantomData }
2550    }
2551
2552    fn get_current(&self) -> Option<&'a P::Target> {
2553        if is_sentinel_ptr(self.current) {
2554            None
2555        } else {
2556            // SAFETY: `self.current` is checked to be non-sentinel (which implies it is a valid,
2557            // non-null node in the tree because `current` is never null). Since the iterator
2558            // holds a lifetime borrow of the `WavlTree`, and the tree remains unmodified,
2559            // the node is guaranteed to be valid and dereferenceable for `'a`.
2560            unsafe { Some(&*self.current) }
2561        }
2562    }
2563}
2564
2565impl<'a, K, P, Tag, S, O> Clone for ReverseIterator<'a, K, P, Tag, S, O>
2566where
2567    P: PtrTraits,
2568    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2569    K: Ord,
2570    S: SizeTracker,
2571    O: WavlTreeObserver<Target = P::Target>,
2572{
2573    fn clone(&self) -> Self {
2574        Self { current: self.current, _phantom: core::marker::PhantomData }
2575    }
2576}
2577
2578impl<'a, K, P, Tag, S, O> core::iter::Iterator for ReverseIterator<'a, K, P, Tag, S, O>
2579where
2580    P: PtrTraits,
2581    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2582    K: Ord,
2583    S: SizeTracker,
2584    O: WavlTreeObserver<Target = P::Target>,
2585{
2586    type Item = &'a P::Target;
2587
2588    fn next(&mut self) -> Option<Self::Item> {
2589        let current = self.get_current()?;
2590        // SAFETY: `self.current` is validated as non-sentinel by `get_current`.
2591        // Moving through tree pointers is safe.
2592        unsafe {
2593            WavlTree::<K, P, Tag, S, O>::advance::<ReverseTraits>(&mut self.current);
2594        }
2595        Some(current)
2596    }
2597}
2598
2599/// An iterator over the elements of a `WavlTree`.
2600pub struct Iterator<
2601    'a,
2602    K,
2603    P,
2604    Tag = DefaultObjectTag,
2605    S = NonTrackingSize,
2606    O = DefaultWavlTreeObserver<<P as PtrTraits>::Target>,
2607> where
2608    P: PtrTraits,
2609    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2610    K: Ord,
2611    S: SizeTracker,
2612    O: WavlTreeObserver<Target = P::Target>,
2613{
2614    front: ForwardIterator<'a, K, P, Tag, S, O>,
2615    back: ReverseIterator<'a, K, P, Tag, S, O>,
2616}
2617
2618impl<'a, K, P, Tag, S, O> Iterator<'a, K, P, Tag, S, O>
2619where
2620    P: PtrTraits,
2621    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2622    K: Ord,
2623    S: SizeTracker,
2624    O: WavlTreeObserver<Target = P::Target>,
2625{
2626    fn new(tree: &'a WavlTree<K, P, Tag, S, O>) -> Self {
2627        if tree.is_empty() {
2628            Self {
2629                front: ForwardIterator::new(make_sentinel_null()),
2630                back: ReverseIterator::new(make_sentinel_null()),
2631            }
2632        } else {
2633            Self {
2634                front: ForwardIterator::new(tree.left_most),
2635                back: ReverseIterator::new(tree.right_most),
2636            }
2637        }
2638    }
2639}
2640
2641impl<'a, K, P, Tag, S, O> Clone for Iterator<'a, K, P, Tag, S, O>
2642where
2643    P: PtrTraits,
2644    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2645    K: Ord,
2646    S: SizeTracker,
2647    O: WavlTreeObserver<Target = P::Target>,
2648{
2649    fn clone(&self) -> Self {
2650        Self { front: self.front.clone(), back: self.back.clone() }
2651    }
2652}
2653
2654impl<'a, K, P, Tag, S, O> core::iter::Iterator for Iterator<'a, K, P, Tag, S, O>
2655where
2656    P: PtrTraits,
2657    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2658    K: Ord,
2659    S: SizeTracker,
2660    O: WavlTreeObserver<Target = P::Target>,
2661{
2662    type Item = &'a P::Target;
2663
2664    fn next(&mut self) -> Option<Self::Item> {
2665        let met = self.front.current == self.back.current;
2666        let item = self.front.next();
2667        if item.is_some() {
2668            if met {
2669                self.front.current = make_sentinel_null();
2670                self.back.current = make_sentinel_null();
2671            }
2672        }
2673        item
2674    }
2675}
2676
2677impl<'a, K, P, Tag, S, O> core::iter::DoubleEndedIterator for Iterator<'a, K, P, Tag, S, O>
2678where
2679    P: PtrTraits,
2680    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
2681    K: Ord,
2682    S: SizeTracker,
2683    O: WavlTreeObserver<Target = P::Target>,
2684{
2685    fn next_back(&mut self) -> Option<Self::Item> {
2686        let met = self.front.current == self.back.current;
2687        let item = self.back.next();
2688        if item.is_some() {
2689            if met {
2690                self.front.current = make_sentinel_null();
2691                self.back.current = make_sentinel_null();
2692            }
2693        }
2694        item
2695    }
2696}
2697
2698impl<K, T, Tag, S, O> WavlTree<K, *mut T, Tag, S, O>
2699where
2700    T: WavlTreeContainable<T, Tag> + WavlTreeKeyable<K>,
2701    K: Ord,
2702    S: SizeTracker,
2703    O: WavlTreeObserver<Target = T>,
2704{
2705    /// Unsafely removes all elements from the tree without modifying node memory.
2706    ///
2707    /// This method resets the tree's internal pointers, effectively emptying it, but does
2708    /// NOT modify the node state of the elements that were in the tree.
2709    ///
2710    /// # Safety
2711    ///
2712    /// Because the nodes are not modified, they will still believe they are in a container
2713    /// (i.e. `in_container()` will return `true` for them). If these elements are subsequently
2714    /// dropped, they will trigger a `debug_assert` panic (as `WavlTreeNode` asserts on drop
2715    /// that it is not in a container).
2716    ///
2717    /// The caller is responsible for manually clearing the node state of the elements, or
2718    /// ensuring they are never dropped while in this "dirty" state.
2719    ///
2720    /// Only usable with containers of unmanaged pointers. Think carefully before calling this!
2721    pub fn clear_unsafe(&mut self) {
2722        self.root = core::ptr::null_mut();
2723        self.left_most = self.get_sentinel();
2724        self.right_most = self.get_sentinel();
2725        self.size.set(0);
2726    }
2727}
2728
2729impl<K, P, Tag, S, O> core::fmt::Debug for WavlTree<K, P, Tag, S, O>
2730where
2731    P: PtrTraits,
2732    P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K> + core::fmt::Debug,
2733    K: Ord,
2734    S: SizeTracker,
2735    O: WavlTreeObserver<Target = P::Target>,
2736{
2737    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
2738        f.debug_list().entries(self.iter()).finish()
2739    }
2740}
2741
2742#[cfg(test)]
2743mod tests {
2744    use super::*;
2745    use crate::intrusive_container_test_support::*;
2746    use crate::ref_counted::HasRefCount;
2747    use crate::ref_ptr::RefPtr;
2748    use crate::size_tracker::TrackingSize;
2749    use crate::unique_ptr::UniquePtr;
2750    use crate::{Recyclable, WavlTreeContainable};
2751    use core::ffi::c_void;
2752    use pin_init::stack_pin_init;
2753
2754    trait AsTargetRef {
2755        type Target;
2756        unsafe fn as_target_ref(&self) -> &Self::Target;
2757    }
2758
2759    impl<T> AsTargetRef for *mut T {
2760        type Target = T;
2761        unsafe fn as_target_ref(&self) -> &T {
2762            unsafe { &**self }
2763        }
2764    }
2765
2766    impl<T: Recyclable> AsTargetRef for UniquePtr<T> {
2767        type Target = T;
2768        unsafe fn as_target_ref(&self) -> &T {
2769            &**self
2770        }
2771    }
2772
2773    impl<T: HasRefCount + Recyclable> AsTargetRef for RefPtr<T> {
2774        type Target = T;
2775        unsafe fn as_target_ref(&self) -> &T {
2776            &**self
2777        }
2778    }
2779
2780    #[derive(crate::WavlTreeContainable, crate::Recyclable)]
2781    struct TestObject {
2782        value: i32,
2783        #[wavl_node]
2784        node: WavlTreeNode<TestObject>,
2785    }
2786
2787    impl TestObject {
2788        fn new(value: i32) -> Self {
2789            Self { value, node: WavlTreeNode::new() }
2790        }
2791    }
2792
2793    impl WavlTreeKeyable<i32> for TestObject {
2794        type Key<'a> = i32;
2795        fn get_key(&self) -> i32 {
2796            self.value
2797        }
2798    }
2799
2800    impl TestValue for TestObject {
2801        fn new(value: i32) -> Self {
2802            Self::new(value)
2803        }
2804    }
2805
2806    ::zr::static_assert!(
2807        core::mem::size_of::<WavlTree<i32, *mut TestObject>>()
2808            == 3 * core::mem::size_of::<*mut TestObject>()
2809    );
2810    ::zr::static_assert!(
2811        core::mem::align_of::<WavlTree<i32, *mut TestObject>>()
2812            == core::mem::align_of::<*mut TestObject>()
2813    );
2814
2815    ::zr::static_assert!(
2816        core::mem::size_of::<WavlTree<i32, *mut TestObject, DefaultObjectTag, TrackingSize>>()
2817            == 4 * core::mem::size_of::<*mut TestObject>()
2818    );
2819    ::zr::static_assert!(
2820        core::mem::align_of::<WavlTree<i32, *mut TestObject, DefaultObjectTag, TrackingSize>>()
2821            == core::mem::align_of::<*mut TestObject>()
2822    );
2823
2824    #[derive(crate::WavlTreeContainable, crate::Recyclable)]
2825    struct UniqueTestObject {
2826        value: i32,
2827        #[wavl_node]
2828        node: WavlTreeNode<UniqueTestObject>,
2829    }
2830
2831    impl UniqueTestObject {
2832        fn new(value: i32) -> Self {
2833            Self { value, node: WavlTreeNode::new() }
2834        }
2835    }
2836
2837    impl WavlTreeKeyable<i32> for UniqueTestObject {
2838        type Key<'a> = i32;
2839        fn get_key(&self) -> i32 {
2840            self.value
2841        }
2842    }
2843
2844    impl TestValue for UniqueTestObject {
2845        fn new(value: i32) -> Self {
2846            Self::new(value)
2847        }
2848    }
2849
2850    #[fbl::ref_counted]
2851    #[derive(crate::WavlTreeContainable, crate::Recyclable)]
2852    #[repr(C)]
2853    pub struct RefTestObject {
2854        value: i32,
2855        #[wavl_node]
2856        node: WavlTreeNode<RefTestObject>,
2857    }
2858
2859    impl WavlTreeKeyable<i32> for RefTestObject {
2860        type Key<'a> = i32;
2861        fn get_key(&self) -> i32 {
2862            self.value
2863        }
2864    }
2865
2866    impl TestValue for RefTestObject {
2867        fn new_ref_counted(value: i32) -> RefPtr<Self> {
2868            crate::make_ref_counted!(RefTestObject { value: value, node: WavlTreeNode::new() })
2869                .unwrap()
2870        }
2871    }
2872
2873    macro_rules! generate_tree_tests {
2874        ($mod_name:ident, $ptr_type:ty, $factory_type:ty, $get_val:expr, $insert:expr, $insert_or_find:expr, $insert_or_replace:expr) => {
2875            mod $mod_name {
2876                use super::*;
2877
2878                #[test]
2879                fn test_basic_sorting() {
2880                    let mut factory = <$factory_type>::new();
2881                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
2882                    let tree = unsafe { tree.get_unchecked_mut() };
2883                    assert!(tree.is_empty());
2884
2885                    // Insert in scrambled order
2886                    $insert(tree, factory.create(3));
2887                    $insert(tree, factory.create(1));
2888                    $insert(tree, factory.create(4));
2889                    $insert(tree, factory.create(2));
2890
2891                    assert!(!tree.is_empty());
2892
2893                    // Iteration should be sorted
2894                    let mut iter = tree.iter();
2895                    assert_eq!($get_val(iter.next().unwrap()), 1);
2896                    assert_eq!($get_val(iter.next().unwrap()), 2);
2897                    assert_eq!($get_val(iter.next().unwrap()), 3);
2898                    assert_eq!($get_val(iter.next().unwrap()), 4);
2899                    assert!(iter.next().is_none());
2900
2901                    tree.clear();
2902                    assert!(tree.is_empty());
2903                }
2904
2905                #[test]
2906                fn test_double_ended_iterator() {
2907                    let mut factory = <$factory_type>::new();
2908                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
2909                    let tree = unsafe { tree.get_unchecked_mut() };
2910                    $insert(tree, factory.create(30));
2911                    $insert(tree, factory.create(10));
2912                    $insert(tree, factory.create(20));
2913
2914                    let mut iter = tree.iter();
2915                    assert_eq!($get_val(iter.next().unwrap()), 10);
2916                    assert_eq!($get_val(iter.next_back().unwrap()), 30);
2917                    assert_eq!($get_val(iter.next().unwrap()), 20);
2918                    assert!(iter.next().is_none());
2919                    assert!(iter.next_back().is_none());
2920
2921                    tree.clear();
2922                }
2923
2924                #[test]
2925                fn test_find() {
2926                    let mut factory = <$factory_type>::new();
2927                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
2928                    let tree = unsafe { tree.get_unchecked_mut() };
2929                    $insert(tree, factory.create(3));
2930                    $insert(tree, factory.create(1));
2931                    $insert(tree, factory.create(2));
2932
2933                    assert!(tree.find(&2).is_some());
2934                    assert_eq!($get_val(tree.find(&2).unwrap()), 2);
2935                    assert!(tree.find(&4).is_none());
2936
2937                    tree.clear();
2938                }
2939
2940                #[test]
2941                fn test_bounds() {
2942                    let mut factory = <$factory_type>::new();
2943                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
2944                    let tree = unsafe { tree.get_unchecked_mut() };
2945                    $insert(tree, factory.create(10));
2946                    $insert(tree, factory.create(30));
2947                    $insert(tree, factory.create(20));
2948
2949                    // lower_bound(>=)
2950                    assert_eq!($get_val(tree.lower_bound(&15).get().unwrap()), 20);
2951                    assert_eq!($get_val(tree.lower_bound(&20).get().unwrap()), 20);
2952                    assert!(tree.lower_bound(&35).get().is_none());
2953
2954                    // upper_bound(>)
2955                    assert_eq!($get_val(tree.upper_bound(&15).get().unwrap()), 20);
2956                    assert_eq!($get_val(tree.upper_bound(&20).get().unwrap()), 30);
2957                    assert!(tree.upper_bound(&30).get().is_none());
2958
2959                    tree.clear();
2960                }
2961
2962                #[test]
2963                fn test_pops() {
2964                    let mut factory = <$factory_type>::new();
2965                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
2966                    let tree = unsafe { tree.get_unchecked_mut() };
2967                    $insert(tree, factory.create(10));
2968                    $insert(tree, factory.create(30));
2969                    $insert(tree, factory.create(20));
2970
2971                    let popped = tree.pop_front();
2972                    assert!(popped.is_some());
2973                    let val = popped.unwrap();
2974                    assert_eq!($get_val(unsafe { val.as_target_ref() }), 10);
2975
2976                    let popped = tree.pop_back();
2977                    assert!(popped.is_some());
2978                    let val = popped.unwrap();
2979                    assert_eq!($get_val(unsafe { val.as_target_ref() }), 30);
2980
2981                    let popped = tree.pop_front();
2982                    assert!(popped.is_some());
2983                    let val = popped.unwrap();
2984                    assert_eq!($get_val(unsafe { val.as_target_ref() }), 20);
2985
2986                    assert!(tree.pop_front().is_none());
2987                }
2988
2989                #[test]
2990                fn test_erase_cursor() {
2991                    let mut factory = <$factory_type>::new();
2992                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
2993                    let tree = unsafe { tree.get_unchecked_mut() };
2994                    $insert(tree, factory.create(10));
2995                    $insert(tree, factory.create(30));
2996                    $insert(tree, factory.create(20));
2997
2998                    let mut cursor = tree.find_cursor(&20);
2999                    let erased = cursor.erase();
3000                    assert!(erased.is_some());
3001                    let val = erased.unwrap();
3002                    assert_eq!($get_val(unsafe { val.as_target_ref() }), 20);
3003
3004                    // Cursor should advance to next element (30)
3005                    assert_eq!($get_val(cursor.get().unwrap()), 30);
3006
3007                    let mut iter = tree.iter();
3008                    assert_eq!($get_val(iter.next().unwrap()), 10);
3009                    assert_eq!($get_val(iter.next().unwrap()), 30);
3010                    assert!(iter.next().is_none());
3011
3012                    tree.clear();
3013                }
3014
3015                #[test]
3016                fn test_insert_or_find() {
3017                    let mut factory = <$factory_type>::new();
3018                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
3019                    let tree = unsafe { tree.get_unchecked_mut() };
3020                    $insert(tree, factory.create(10));
3021
3022                    let new_item = factory.create(10); // Duplicate key
3023                    let res = $insert_or_find(tree, new_item);
3024                    assert!(res.is_err());
3025                    let (failed_ptr, collision) = res.err().unwrap();
3026                    assert_eq!($get_val(unsafe { failed_ptr.as_target_ref() }), 10);
3027                    assert_eq!($get_val(collision.get().unwrap()), 10);
3028
3029                    let ok_item = factory.create(20);
3030                    assert!($insert_or_find(tree, ok_item).is_ok());
3031
3032                    tree.clear();
3033                }
3034
3035                #[test]
3036                fn test_insert_or_replace() {
3037                    let mut factory = <$factory_type>::new();
3038                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
3039                    let tree = unsafe { tree.get_unchecked_mut() };
3040                    $insert(tree, factory.create(10));
3041
3042                    let replacement = factory.create(10);
3043                    let res = $insert_or_replace(tree, replacement);
3044                    assert!(res.is_some());
3045                    let old_item = res.unwrap();
3046                    assert_eq!($get_val(unsafe { old_item.as_target_ref() }), 10);
3047
3048                    let found = tree.find(&10);
3049                    assert!(found.is_some());
3050                    assert_eq!($get_val(found.unwrap()), 10);
3051
3052                    tree.clear();
3053                }
3054
3055                #[test]
3056                fn test_from_element() {
3057                    let mut factory = <$factory_type>::new();
3058                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
3059                    let tree = unsafe { tree.get_unchecked_mut() };
3060
3061                    $insert(tree, factory.create(10));
3062                    $insert(tree, factory.create(20));
3063                    $insert(tree, factory.create(30));
3064
3065                    let target_ref = tree.find(&20).unwrap();
3066
3067                    let mut forward_iter: ForwardIterator<'_, i32, $ptr_type> = ForwardIterator::from_element(target_ref);
3068                    assert_eq!($get_val(forward_iter.next().unwrap()), 20);
3069                    assert_eq!($get_val(forward_iter.next().unwrap()), 30);
3070                    assert!(forward_iter.next().is_none());
3071
3072                    let mut reverse_iter: ReverseIterator<'_, i32, $ptr_type> = ReverseIterator::from_element(target_ref);
3073                    assert_eq!($get_val(reverse_iter.next().unwrap()), 20);
3074                    assert_eq!($get_val(reverse_iter.next().unwrap()), 10);
3075                    assert!(reverse_iter.next().is_none());
3076
3077                    tree.clear();
3078                }
3079
3080                #[test]
3081                fn test_cursor_at() {
3082                    let mut factory = <$factory_type>::new();
3083                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
3084                    let tree = unsafe { tree.get_unchecked_mut() };
3085
3086                    $insert(tree, factory.create(10));
3087                    $insert(tree, factory.create(20));
3088                    $insert(tree, factory.create(30));
3089
3090                    let target_ptr = tree.find(&20).unwrap() as *const <$ptr_type as PtrTraits>::Target;
3091                    // SAFETY: target_ptr is a valid pointer to an object currently in the tree.
3092                    // Using a raw pointer bypasses the borrow checker, allowing us to obtain an unbound reference
3093                    // and borrow the tree mutably afterward.
3094                    let target_ref = unsafe { &*target_ptr };
3095
3096                    // Test read-only cursor_at
3097                    let cursor = unsafe { tree.cursor_at(target_ref) };
3098                    assert!(cursor.is_valid());
3099                    assert_eq!($get_val(cursor.get().unwrap()), 20);
3100                    assert_eq!($get_val(cursor.left().get().unwrap()), 10);
3101                    assert_eq!($get_val(cursor.right().get().unwrap()), 30);
3102
3103                    // Test mutable cursor_mut_at
3104                    let mut cursor_mut = unsafe { tree.cursor_mut_at(target_ref) };
3105                    assert_eq!($get_val(cursor_mut.get().unwrap()), 20);
3106
3107                    // Verify we can erase using the cursor returned by cursor_mut_at
3108                    let erased = cursor_mut.erase();
3109                    assert!(erased.is_some());
3110                    let val = erased.unwrap();
3111                    assert_eq!($get_val(unsafe { val.as_target_ref() }), 20);
3112
3113                    // Cursor should now be at the next element (30)
3114                    assert_eq!($get_val(cursor_mut.get().unwrap()), 30);
3115
3116                    tree.clear();
3117                }
3118
3119                #[test]
3120                fn test_iterator_clone() {
3121                    let mut factory = <$factory_type>::new();
3122                    stack_pin_init!(let tree = WavlTree::<i32, $ptr_type>::new());
3123                    let tree = unsafe { tree.get_unchecked_mut() };
3124                    $insert(tree, factory.create(10));
3125                    $insert(tree, factory.create(20));
3126                    $insert(tree, factory.create(30));
3127
3128                    let mut iter = tree.iter();
3129                    assert_eq!($get_val(iter.next().unwrap()), 10);
3130
3131                    let mut cloned_iter = iter.clone();
3132
3133                    assert_eq!($get_val(iter.next().unwrap()), 20);
3134                    assert_eq!($get_val(iter.next().unwrap()), 30);
3135                    assert!(iter.next().is_none());
3136
3137                    assert_eq!($get_val(cloned_iter.next().unwrap()), 20);
3138                    assert_eq!($get_val(cloned_iter.next().unwrap()), 30);
3139                    assert!(cloned_iter.next().is_none());
3140
3141                    tree.clear();
3142                }
3143
3144                #[test]
3145                fn test_swap() {
3146                    let mut factory = <$factory_type>::new();
3147                    stack_pin_init!(let tree1 = WavlTree::<i32, $ptr_type>::new());
3148                    let tree1 = unsafe { tree1.get_unchecked_mut() };
3149                    stack_pin_init!(let tree2 = WavlTree::<i32, $ptr_type>::new());
3150                    let tree2 = unsafe { tree2.get_unchecked_mut() };
3151
3152                    $insert(tree1, factory.create(1));
3153                    $insert(tree1, factory.create(3));
3154
3155                    $insert(tree2, factory.create(2));
3156                    $insert(tree2, factory.create(4));
3157
3158                    tree1.swap(tree2);
3159
3160                    let mut iter1 = tree1.iter();
3161                    assert_eq!($get_val(iter1.next().unwrap()), 2);
3162                    assert_eq!($get_val(iter1.next().unwrap()), 4);
3163                    assert!(iter1.next().is_none());
3164
3165                    let mut iter2 = tree2.iter();
3166                    assert_eq!($get_val(iter2.next().unwrap()), 1);
3167                    assert_eq!($get_val(iter2.next().unwrap()), 3);
3168                    assert!(iter2.next().is_none());
3169
3170                    tree1.clear();
3171                    tree2.clear();
3172                }
3173            }
3174        };
3175    }
3176
3177    generate_tree_tests!(
3178        raw_ptr_tests,
3179        *mut TestObject,
3180        RawFactory<TestObject>,
3181        |p: &TestObject| p.value,
3182        |tree, obj| unsafe { WavlTree::<i32, *mut TestObject>::insert_raw(tree, obj) },
3183        |tree, obj| unsafe { WavlTree::<i32, *mut TestObject>::insert_or_find_raw(tree, obj) },
3184        |tree, obj| unsafe { WavlTree::<i32, *mut TestObject>::insert_or_replace_raw(tree, obj) }
3185    );
3186
3187    generate_tree_tests!(
3188        unique_ptr_tests,
3189        UniquePtr<UniqueTestObject>,
3190        UniqueFactory<UniqueTestObject>,
3191        |p: &UniqueTestObject| p.value,
3192        |tree, obj| WavlTree::<i32, UniquePtr<UniqueTestObject>>::insert(tree, obj),
3193        |tree, obj| WavlTree::<i32, UniquePtr<UniqueTestObject>>::insert_or_find(tree, obj),
3194        |tree, obj| WavlTree::<i32, UniquePtr<UniqueTestObject>>::insert_or_replace(tree, obj)
3195    );
3196
3197    generate_tree_tests!(
3198        ref_ptr_tests,
3199        RefPtr<RefTestObject>,
3200        RefFactory<RefTestObject>,
3201        |p: &RefTestObject| p.value,
3202        |tree, obj| WavlTree::<i32, RefPtr<RefTestObject>>::insert(tree, obj),
3203        |tree, obj| WavlTree::<i32, RefPtr<RefTestObject>>::insert_or_find(tree, obj),
3204        |tree, obj| WavlTree::<i32, RefPtr<RefTestObject>>::insert_or_replace(tree, obj)
3205    );
3206
3207    #[test]
3208    fn test_erase_by_raw_pointer() {
3209        stack_pin_init!(let tree = WavlTree::<i32, *mut TestObject, DefaultObjectTag, TrackingSize>::new());
3210        let tree = unsafe { tree.get_unchecked_mut() };
3211        let mut obj1 = TestObject::new(10);
3212        let mut obj2 = TestObject::new(20);
3213        let mut obj3 = TestObject::new(30);
3214
3215        unsafe {
3216            tree.insert_raw(core::ptr::addr_of_mut!(obj1));
3217            tree.insert_raw(core::ptr::addr_of_mut!(obj2));
3218            tree.insert_raw(core::ptr::addr_of_mut!(obj3));
3219        }
3220
3221        assert_eq!(tree.len(), 3);
3222
3223        // Erase obj2 directly
3224        let erased = unsafe { tree.erase_raw(core::ptr::addr_of_mut!(obj2)) };
3225        assert!(erased.is_some());
3226        assert_eq!(unsafe { &*erased.unwrap() }.value, 20);
3227        assert_eq!(tree.len(), 2);
3228
3229        let mut iter = tree.iter();
3230        assert_eq!(iter.next().unwrap().value, 10);
3231        assert_eq!(iter.next().unwrap().value, 30);
3232        assert!(iter.next().is_none());
3233
3234        tree.clear();
3235    }
3236
3237    #[test]
3238    fn test_clear_unsafe() {
3239        stack_pin_init!(let tree = WavlTree::<i32, *mut TestObject, DefaultObjectTag, TrackingSize>::new());
3240        let tree = unsafe { tree.get_unchecked_mut() };
3241        let mut obj1 = TestObject::new(10);
3242        let mut obj2 = TestObject::new(20);
3243        let mut obj3 = TestObject::new(30);
3244
3245        unsafe {
3246            tree.insert_raw(core::ptr::addr_of_mut!(obj1));
3247            tree.insert_raw(core::ptr::addr_of_mut!(obj2));
3248            tree.insert_raw(core::ptr::addr_of_mut!(obj3));
3249        }
3250
3251        assert_eq!(tree.len(), 3);
3252        assert!(!tree.is_empty());
3253
3254        tree.clear_unsafe();
3255
3256        assert_eq!(tree.len(), 0);
3257        assert!(tree.is_empty());
3258
3259        // Clean up the nodes manually so that they can be safely dropped.
3260        unsafe {
3261            (*obj1.get_node().parent.get()) = core::ptr::null_mut();
3262            (*obj1.get_node().left.get()) = core::ptr::null_mut();
3263            (*obj1.get_node().right.get()) = core::ptr::null_mut();
3264
3265            (*obj2.get_node().parent.get()) = core::ptr::null_mut();
3266            (*obj2.get_node().left.get()) = core::ptr::null_mut();
3267            (*obj2.get_node().right.get()) = core::ptr::null_mut();
3268
3269            (*obj3.get_node().parent.get()) = core::ptr::null_mut();
3270            (*obj3.get_node().left.get()) = core::ptr::null_mut();
3271            (*obj3.get_node().right.get()) = core::ptr::null_mut();
3272        }
3273    }
3274
3275    #[test]
3276    fn test_tracking_size() {
3277        stack_pin_init!(let tree = WavlTree::<i32, UniquePtr<UniqueTestObject>, DefaultObjectTag, TrackingSize>::new());
3278        let tree = unsafe { tree.get_unchecked_mut() };
3279
3280        assert_eq!(tree.len(), 0);
3281        tree.insert(UniquePtr::try_new(UniqueTestObject::new(10)).unwrap());
3282        assert_eq!(tree.len(), 1);
3283        tree.insert(UniquePtr::try_new(UniqueTestObject::new(20)).unwrap());
3284        assert_eq!(tree.len(), 2);
3285        tree.pop_front();
3286        assert_eq!(tree.len(), 1);
3287        tree.clear();
3288        assert_eq!(tree.len(), 0);
3289    }
3290
3291    struct Tag2;
3292
3293    #[fbl::ref_counted]
3294    #[derive(crate::WavlTreeContainable, crate::Recyclable)]
3295    #[repr(C)]
3296    struct MultiTreeObject {
3297        value: i32,
3298        #[wavl_node]
3299        node1: WavlTreeNode<MultiTreeObject>,
3300        #[wavl_node(tag = Tag2)]
3301        node2: WavlTreeNode<MultiTreeObject>,
3302    }
3303
3304    impl WavlTreeKeyable<i32> for MultiTreeObject {
3305        type Key<'a> = i32;
3306        fn get_key(&self) -> i32 {
3307            self.value
3308        }
3309    }
3310
3311    #[test]
3312    fn test_multiple_containers() {
3313        stack_pin_init!(let tree1 = WavlTree::<i32, RefPtr<MultiTreeObject>, DefaultObjectTag>::new());
3314        let tree1 = unsafe { tree1.get_unchecked_mut() };
3315        stack_pin_init!(let tree2 = WavlTree::<i32, RefPtr<MultiTreeObject>, Tag2>::new());
3316        let tree2 = unsafe { tree2.get_unchecked_mut() };
3317
3318        let obj1 = fbl::make_ref_counted!(MultiTreeObject {
3319            value: 10,
3320            node1: WavlTreeNode::new(),
3321            node2: WavlTreeNode::new(),
3322        })
3323        .unwrap();
3324
3325        let obj2 = fbl::make_ref_counted!(MultiTreeObject {
3326            value: 20,
3327            node1: WavlTreeNode::new(),
3328            node2: WavlTreeNode::new(),
3329        })
3330        .unwrap();
3331
3332        tree1.insert(obj1.clone());
3333        tree1.insert(obj2.clone());
3334
3335        tree2.insert(obj1.clone());
3336        tree2.insert(obj2.clone());
3337
3338        let mut iter1 = tree1.iter();
3339        assert_eq!(iter1.next().unwrap().value, 10);
3340        assert_eq!(iter1.next().unwrap().value, 20);
3341
3342        let mut iter2 = tree2.iter();
3343        assert_eq!(iter2.next().unwrap().value, 10);
3344        assert_eq!(iter2.next().unwrap().value, 20);
3345
3346        tree1.clear();
3347        tree2.clear();
3348    }
3349
3350    extern crate alloc;
3351    use alloc::boxed::Box;
3352    use alloc::sync::Arc;
3353    use alloc::vec::Vec;
3354    use core::sync::atomic::{AtomicUsize, Ordering};
3355
3356    struct Lfsr {
3357        core: u64,
3358    }
3359
3360    impl Lfsr {
3361        fn new(initial_core: u64) -> Self {
3362            Self { core: initial_core }
3363        }
3364
3365        fn set_core(&mut self, val: u64) {
3366            self.core = val;
3367        }
3368
3369        fn get_next(&mut self) -> u64 {
3370            let mut ret = 0u64;
3371            let mut flag = 1u64;
3372            let generator = 0xD800000000000000u64;
3373
3374            for _ in 0..(core::mem::size_of::<usize>() * 8) {
3375                let bit = (self.core & 1) != 0;
3376                self.core >>= 1;
3377                if bit {
3378                    self.core ^= generator;
3379                    ret |= flag;
3380                }
3381                flag <<= 1;
3382            }
3383
3384            ret
3385        }
3386    }
3387
3388    struct OpCounts {
3389        insert_ops: AtomicUsize,
3390        insert_promotes: AtomicUsize,
3391        insert_rotations: AtomicUsize,
3392        insert_double_rotations: AtomicUsize,
3393        insert_collisions: AtomicUsize,
3394        insert_replacements: AtomicUsize,
3395        insert_traversals: AtomicUsize,
3396        inspected_rotations: AtomicUsize,
3397        erase_ops: AtomicUsize,
3398        erase_demotes: AtomicUsize,
3399        erase_rotations: AtomicUsize,
3400        erase_double_rotations: AtomicUsize,
3401    }
3402
3403    impl OpCounts {
3404        const fn new() -> Self {
3405            Self {
3406                insert_ops: AtomicUsize::new(0),
3407                insert_promotes: AtomicUsize::new(0),
3408                insert_rotations: AtomicUsize::new(0),
3409                insert_double_rotations: AtomicUsize::new(0),
3410                insert_collisions: AtomicUsize::new(0),
3411                insert_replacements: AtomicUsize::new(0),
3412                insert_traversals: AtomicUsize::new(0),
3413                inspected_rotations: AtomicUsize::new(0),
3414                erase_ops: AtomicUsize::new(0),
3415                erase_demotes: AtomicUsize::new(0),
3416                erase_rotations: AtomicUsize::new(0),
3417                erase_double_rotations: AtomicUsize::new(0),
3418            }
3419        }
3420    }
3421
3422    #[derive(crate::WavlTreeContainable)]
3423    #[repr(C)]
3424    struct BalanceTestObj {
3425        key: u64,
3426        min_subtree_key: u64,
3427        max_subtree_key: u64,
3428        erase_deck_ptr: core::cell::Cell<*mut BalanceTestObj>,
3429        #[wavl_node(rank = i32)]
3430        node: WavlTreeNode<BalanceTestObj, i32>,
3431    }
3432
3433    impl BalanceTestObj {
3434        fn new(key: u64) -> Self {
3435            Self {
3436                key,
3437                min_subtree_key: 0,
3438                max_subtree_key: 0,
3439                erase_deck_ptr: core::cell::Cell::new(core::ptr::null_mut()),
3440                node: WavlTreeNode::new(),
3441            }
3442        }
3443
3444        fn swap_erase_deck_ptr(a: &BalanceTestObj, b: &BalanceTestObj) {
3445            let tmp = a.erase_deck_ptr.get();
3446            a.erase_deck_ptr.set(b.erase_deck_ptr.get());
3447            b.erase_deck_ptr.set(tmp);
3448        }
3449    }
3450
3451    impl WavlTreeKeyable<u64> for BalanceTestObj {
3452        type Key<'a> = u64;
3453        fn get_key(&self) -> u64 {
3454            self.key
3455        }
3456    }
3457
3458    struct WavlBalanceTestObserver {
3459        op_counts: Arc<OpCounts>,
3460    }
3461    impl WavlTreeObserver for WavlBalanceTestObserver {
3462        type Target = BalanceTestObj;
3463
3464        fn record_insert(&self, node: *mut BalanceTestObj) {
3465            self.op_counts.insert_ops.fetch_add(1, Ordering::Relaxed);
3466            unsafe {
3467                (*node).min_subtree_key = (*node).key;
3468                (*node).max_subtree_key = (*node).key;
3469            }
3470        }
3471
3472        fn record_insert_traverse(&self, node: *mut BalanceTestObj, ancestor: *mut BalanceTestObj) {
3473            self.op_counts.insert_traversals.fetch_add(1, Ordering::Relaxed);
3474            unsafe {
3475                (*ancestor).min_subtree_key =
3476                    core::cmp::min((*ancestor).min_subtree_key, (*node).key);
3477                (*ancestor).max_subtree_key =
3478                    core::cmp::max((*ancestor).max_subtree_key, (*node).key);
3479            }
3480        }
3481
3482        fn record_insert_collision(
3483            &self,
3484            _node: *mut BalanceTestObj,
3485            _collision: *mut BalanceTestObj,
3486        ) {
3487            self.op_counts.insert_collisions.fetch_add(1, Ordering::Relaxed);
3488        }
3489
3490        fn record_insert_replace(
3491            &self,
3492            node: *mut BalanceTestObj,
3493            replacement: *mut BalanceTestObj,
3494        ) {
3495            self.op_counts.insert_replacements.fetch_add(1, Ordering::Relaxed);
3496            unsafe {
3497                (*replacement).min_subtree_key = (*node).min_subtree_key;
3498                (*replacement).max_subtree_key = (*node).max_subtree_key;
3499            }
3500        }
3501
3502        fn record_insert_promote(&self) {
3503            self.op_counts.insert_promotes.fetch_add(1, Ordering::Relaxed);
3504        }
3505
3506        fn record_insert_rotation(&self) {
3507            self.op_counts.insert_rotations.fetch_add(1, Ordering::Relaxed);
3508        }
3509
3510        fn record_insert_double_rotation(&self) {
3511            self.op_counts.insert_double_rotations.fetch_add(1, Ordering::Relaxed);
3512        }
3513
3514        fn record_rotation(
3515            &self,
3516            pivot: *mut BalanceTestObj,
3517            lr_child: *mut BalanceTestObj,
3518            _rl_child: *mut BalanceTestObj,
3519            parent: *mut BalanceTestObj,
3520            sibling: *mut BalanceTestObj,
3521        ) {
3522            self.op_counts.inspected_rotations.fetch_add(1, Ordering::Relaxed);
3523            unsafe {
3524                (*pivot).min_subtree_key = (*parent).min_subtree_key;
3525                (*pivot).max_subtree_key = (*parent).max_subtree_key;
3526
3527                (*parent).min_subtree_key = (*parent).key;
3528                (*parent).max_subtree_key = (*parent).key;
3529
3530                if valid_sentinel_ptr(sibling) {
3531                    (*parent).min_subtree_key =
3532                        core::cmp::min((*parent).min_subtree_key, (*sibling).min_subtree_key);
3533                    (*parent).max_subtree_key =
3534                        core::cmp::max((*parent).max_subtree_key, (*sibling).max_subtree_key);
3535                }
3536                if valid_sentinel_ptr(lr_child) {
3537                    (*parent).min_subtree_key =
3538                        core::cmp::min((*parent).min_subtree_key, (*lr_child).min_subtree_key);
3539                    (*parent).max_subtree_key =
3540                        core::cmp::max((*parent).max_subtree_key, (*lr_child).max_subtree_key);
3541                }
3542            }
3543        }
3544
3545        fn record_erase(&self, _node: *mut BalanceTestObj, invalidated: *mut BalanceTestObj) {
3546            self.op_counts.erase_ops.fetch_add(1, Ordering::Relaxed);
3547            unsafe {
3548                let mut current = invalidated;
3549                while valid_sentinel_ptr(current) {
3550                    (*current).min_subtree_key = (*current).key;
3551                    (*current).max_subtree_key = (*current).key;
3552
3553                    let ns = (*current).get_node();
3554                    let left = ns.get_left();
3555                    if valid_sentinel_ptr(left) {
3556                        (*current).min_subtree_key =
3557                            core::cmp::min((*current).min_subtree_key, (*left).min_subtree_key);
3558                        (*current).max_subtree_key =
3559                            core::cmp::max((*current).max_subtree_key, (*left).max_subtree_key);
3560                    }
3561                    let right = ns.get_right();
3562                    if valid_sentinel_ptr(right) {
3563                        (*current).min_subtree_key =
3564                            core::cmp::min((*current).min_subtree_key, (*right).min_subtree_key);
3565                        (*current).max_subtree_key =
3566                            core::cmp::max((*current).max_subtree_key, (*right).max_subtree_key);
3567                    }
3568                    current = ns.get_parent();
3569                }
3570            }
3571        }
3572
3573        fn record_erase_demote(&self) {
3574            self.op_counts.erase_demotes.fetch_add(1, Ordering::Relaxed);
3575        }
3576
3577        fn record_erase_rotation(&self) {
3578            self.op_counts.erase_rotations.fetch_add(1, Ordering::Relaxed);
3579        }
3580
3581        fn record_erase_double_rotation(&self) {
3582            self.op_counts.erase_double_rotations.fetch_add(1, Ordering::Relaxed);
3583        }
3584
3585        fn verify_rank_rule(
3586            &self,
3587            node: *mut BalanceTestObj,
3588            _left_most: *mut BalanceTestObj,
3589            _right_most: *mut BalanceTestObj,
3590            _sentinel: *mut BalanceTestObj,
3591        ) {
3592            unsafe {
3593                let ns = (*node).get_node();
3594                let rank = ns.rank();
3595                assert!(rank >= 0, "All ranks must be non-negative.");
3596
3597                let left = ns.get_left();
3598                let right = ns.get_right();
3599
3600                if !valid_sentinel_ptr(left) && !valid_sentinel_ptr(right) {
3601                    assert_eq!(rank, 0i32, "Leaf nodes must have rank 0!");
3602                } else {
3603                    if valid_sentinel_ptr(left) {
3604                        let left_ns = (*left).get_node();
3605                        let delta = rank - left_ns.rank();
3606                        assert!(
3607                            delta >= 1 && delta <= 2,
3608                            "Left hand rank difference not in range [1, 2]"
3609                        );
3610                    }
3611
3612                    if valid_sentinel_ptr(right) {
3613                        let right_ns = (*right).get_node();
3614                        let delta = rank - right_ns.rank();
3615                        assert!(
3616                            delta >= 1 && delta <= 2,
3617                            "Right hand rank difference not in range [1, 2]"
3618                        );
3619                    }
3620                }
3621            }
3622        }
3623
3624        fn verify_balance(&self, size: usize, depth: usize) {
3625            if size > 0 {
3626                let log2_n = (size as f64).log2();
3627                let erase_ops = self.op_counts.erase_ops.load(Ordering::Relaxed);
3628                let scale = if erase_ops > 0 { 2.0 } else { 1.4404200904125564 };
3629                let max_depth = (log2_n * scale) as usize + 1;
3630                assert!(
3631                    max_depth >= depth,
3632                    "Depth bound exceeded! max_depth: {}, actual depth: {}",
3633                    max_depth,
3634                    depth
3635                );
3636
3637                let insert_rotations = self.op_counts.insert_rotations.load(Ordering::Relaxed);
3638                let insert_double_rotations =
3639                    self.op_counts.insert_double_rotations.load(Ordering::Relaxed);
3640                let insert_promotes = self.op_counts.insert_promotes.load(Ordering::Relaxed);
3641                let insert_ops = self.op_counts.insert_ops.load(Ordering::Relaxed);
3642
3643                let total_insert_rotations = insert_rotations + insert_double_rotations;
3644                assert!(
3645                    insert_promotes <= (3 * insert_ops) + (2 * erase_ops),
3646                    "#insert promotes must be <= (3 * #inserts) + (2 * #erases)"
3647                );
3648                assert!(
3649                    total_insert_rotations <= insert_ops,
3650                    "#insert_rotations must be <= #inserts"
3651                );
3652
3653                let erase_demotes = self.op_counts.erase_demotes.load(Ordering::Relaxed);
3654                let erase_rotations = self.op_counts.erase_rotations.load(Ordering::Relaxed);
3655                let erase_double_rotations =
3656                    self.op_counts.erase_double_rotations.load(Ordering::Relaxed);
3657
3658                let total_erase_rotations = erase_rotations + erase_double_rotations;
3659                assert!(erase_demotes <= erase_ops, "#erase demotes must be <= #erases");
3660                assert!(total_erase_rotations <= erase_ops, "#erase_rotations must be <= #erases");
3661
3662                let inspected_rotations =
3663                    self.op_counts.inspected_rotations.load(Ordering::Relaxed);
3664                let total_inspected_rotations = insert_rotations
3665                    + erase_rotations
3666                    + 2 * insert_double_rotations
3667                    + 2 * erase_double_rotations;
3668                assert_eq!(
3669                    total_inspected_rotations, inspected_rotations,
3670                    "#inspected rotations must be == #rotations"
3671                );
3672            }
3673        }
3674    }
3675
3676    struct WavlTreeChecker;
3677    impl WavlTreeChecker {
3678        fn verify_parent_back_links<K, P, Tag, S, O>(cursor: Cursor<'_, K, P, Tag, S, O>)
3679        where
3680            P: PtrTraits,
3681            P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
3682            K: Ord,
3683            S: SizeTracker,
3684            O: WavlTreeObserver<Target = P::Target>,
3685        {
3686            assert!(cursor.is_valid());
3687            let left = cursor.left();
3688            if left.is_valid() {
3689                assert_eq!(
3690                    cursor.as_raw_ptr(),
3691                    left.parent().as_raw_ptr(),
3692                    "Corrupt left-side parent back-link!"
3693                );
3694            }
3695
3696            let right = cursor.right();
3697            if right.is_valid() {
3698                assert_eq!(
3699                    cursor.as_raw_ptr(),
3700                    right.parent().as_raw_ptr(),
3701                    "Corrupt right-side parent back-link!"
3702                );
3703            }
3704        }
3705
3706        fn sanity_check<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>)
3707        where
3708            P: PtrTraits,
3709            P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
3710            K: Ord,
3711            S: SizeTracker,
3712            O: WavlTreeObserver<Target = P::Target>,
3713        {
3714            let is_empty = tree.is_empty();
3715            let root = tree.root_cursor();
3716            let front = tree.front_cursor();
3717            let back = tree.back_cursor();
3718
3719            let sentinel_ptr =
3720                if is_empty { front.as_raw_ptr() } else { front.left().as_raw_ptr() };
3721
3722            if is_empty {
3723                assert!(!root.is_valid());
3724                assert!(!front.is_valid());
3725                assert!(!back.is_valid());
3726                if S::IS_TRACKING {
3727                    assert_eq!(tree.len(), 0);
3728                }
3729            } else {
3730                assert!(root.is_valid());
3731                assert!(front.is_valid());
3732                assert!(back.is_valid());
3733                assert!(!front.left().is_valid());
3734                assert!(!back.right().is_valid());
3735                if S::IS_TRACKING {
3736                    assert!(tree.len() > 0);
3737                }
3738            }
3739
3740            let mut cur_depth = 0;
3741            let mut depth = 0;
3742            let mut size = 0;
3743
3744            let mut cursor = root;
3745
3746            while cursor.is_valid() {
3747                Self::verify_parent_back_links(cursor);
3748                cur_depth += 1;
3749
3750                let left = cursor.left();
3751                if !left.is_valid() {
3752                    break;
3753                }
3754                cursor = left;
3755            }
3756
3757            while cursor.is_valid() {
3758                if depth < cur_depth {
3759                    depth = cur_depth;
3760                }
3761                size += 1;
3762
3763                Self::verify_parent_back_links(cursor);
3764                tree.observer.verify_rank_rule(
3765                    cursor.as_raw_ptr(),
3766                    front.as_raw_ptr(),
3767                    back.as_raw_ptr(),
3768                    sentinel_ptr,
3769                );
3770
3771                let right = cursor.right();
3772                if right.is_valid() {
3773                    cur_depth += 1;
3774                    cursor = right;
3775                    Self::verify_parent_back_links(cursor);
3776
3777                    loop {
3778                        let left = cursor.left();
3779                        if !left.is_valid() {
3780                            break;
3781                        }
3782                        cur_depth += 1;
3783                        cursor = left;
3784                        Self::verify_parent_back_links(cursor);
3785                    }
3786                    continue;
3787                }
3788
3789                let mut parent = cursor.parent();
3790                let mut keep_going = false;
3791                while parent.is_valid() {
3792                    let is_left = parent.left() == cursor;
3793                    let is_right = parent.right() == cursor;
3794
3795                    assert!(is_left != is_right);
3796                    assert!(is_left || is_right);
3797
3798                    cursor = parent;
3799                    cur_depth -= 1;
3800
3801                    if is_left {
3802                        keep_going = true;
3803                        break;
3804                    }
3805
3806                    parent = parent.parent();
3807                }
3808
3809                if !keep_going {
3810                    break;
3811                }
3812            }
3813
3814            if S::IS_TRACKING {
3815                assert_eq!(tree.len(), size);
3816            }
3817            tree.observer.verify_balance(size, depth);
3818        }
3819    }
3820
3821    fn shuffle_erase_deck(objects: &[Box<BalanceTestObj>], rng: &mut Lfsr, size: usize) {
3822        for i in (2..size).rev() {
3823            let ndx = (rng.get_next() as usize) % i;
3824            if ndx != i {
3825                BalanceTestObj::swap_erase_deck_ptr(&objects[i], &objects[ndx]);
3826            }
3827        }
3828    }
3829
3830    fn check_augmented_invariants<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>)
3831    where
3832        P: PtrTraits,
3833        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
3834        K: Ord,
3835        S: SizeTracker,
3836        O: WavlTreeObserver<Target = P::Target>,
3837    {
3838        if tree.is_empty() {
3839            return;
3840        }
3841        let root = tree.root_cursor().as_raw_ptr() as *mut BalanceTestObj;
3842        let left = tree.front_cursor().as_raw_ptr() as *mut BalanceTestObj;
3843        let right = tree.back_cursor().as_raw_ptr() as *mut BalanceTestObj;
3844
3845        unsafe {
3846            assert_eq!((*root).min_subtree_key, (*left).key, "Min subtree key invariant violated!");
3847            assert_eq!(
3848                (*root).max_subtree_key,
3849                (*right).key,
3850                "Max subtree key invariant violated!"
3851            );
3852        }
3853    }
3854
3855    fn check_iterators<K, P, Tag, S, O>(tree: &WavlTree<K, P, Tag, S, O>)
3856    where
3857        P: PtrTraits,
3858        P::Target: WavlTreeContainable<P::Target, Tag> + WavlTreeKeyable<K>,
3859        K: Ord,
3860        S: SizeTracker,
3861        O: WavlTreeObserver<Target = P::Target>,
3862    {
3863        if tree.is_empty() {
3864            return;
3865        }
3866        let root = tree.root_cursor();
3867        let left_most = tree.front_cursor();
3868        let right_most = tree.back_cursor();
3869
3870        let mut left_cursor = root;
3871        let mut right_cursor = root;
3872        let mut i = 0;
3873
3874        let limit = if S::IS_TRACKING { tree.len() } else { 10000 };
3875
3876        while (left_cursor != left_most || right_cursor != right_most) && i < limit {
3877            assert!(left_cursor.is_valid());
3878            if left_cursor == left_most {
3879                assert!(!left_cursor.left().is_valid());
3880            } else {
3881                left_cursor = left_cursor.left();
3882            }
3883
3884            assert!(right_cursor.is_valid());
3885            if right_cursor == right_most {
3886                assert!(!right_cursor.right().is_valid());
3887            } else {
3888                right_cursor = right_cursor.right();
3889            }
3890
3891            i += 1;
3892        }
3893
3894        assert_eq!(left_cursor, left_most);
3895        assert_eq!(right_cursor, right_most);
3896
3897        let limit = i;
3898        left_cursor = left_most;
3899        right_cursor = right_most;
3900        i = 0;
3901
3902        while (left_cursor != root || right_cursor != root) && i < limit {
3903            assert!(left_cursor.is_valid());
3904            if left_cursor == root {
3905                assert!(!left_cursor.parent().is_valid());
3906            } else {
3907                left_cursor = left_cursor.parent();
3908            }
3909
3910            assert!(right_cursor.is_valid());
3911            if right_cursor == root {
3912                assert!(!right_cursor.parent().is_valid());
3913            } else {
3914                right_cursor = right_cursor.parent();
3915            }
3916
3917            i += 1;
3918        }
3919
3920        assert_eq!(left_cursor, root);
3921        assert_eq!(right_cursor, root);
3922    }
3923
3924    #[test]
3925    fn test_balance_and_invariants() {
3926        let seeds = [0xe87e1062fc1f4f80u64, 0x03d6bffb124b4918u64, 0x8f7d83e8d10b4765u64];
3927        let test_size = 128;
3928        let replacement_count = test_size / 8;
3929        let mut rng = Lfsr::new(1);
3930
3931        for seed_ndx in 0..seeds.len() {
3932            let seed = seeds[seed_ndx];
3933            rng.set_core(seed);
3934
3935            let op_counts = Arc::new(OpCounts::new());
3936            let observer = WavlBalanceTestObserver { op_counts: Arc::clone(&op_counts) };
3937
3938            stack_pin_init!(let tree = WavlTree::<u64, *mut BalanceTestObj, DefaultObjectTag, TrackingSize, WavlBalanceTestObserver>::new_with_observer(observer));
3939            let tree = unsafe { tree.get_unchecked_mut() };
3940
3941            let mut objects = Vec::with_capacity(test_size);
3942            let mut replacements = Vec::with_capacity(replacement_count);
3943
3944            match seed_ndx {
3945                0 => {
3946                    for i in 0..test_size {
3947                        let obj = Box::new(BalanceTestObj::new(i as u64));
3948                        let raw = &*obj as *const BalanceTestObj as *mut BalanceTestObj;
3949                        obj.erase_deck_ptr.set(raw);
3950                        objects.push(obj);
3951
3952                        if i < replacement_count {
3953                            let rep = Box::new(BalanceTestObj::new(i as u64));
3954                            let raw = &*rep as *const BalanceTestObj as *mut BalanceTestObj;
3955                            rep.erase_deck_ptr.set(raw);
3956                            replacements.push(rep);
3957                        }
3958                    }
3959                }
3960                1 => {
3961                    for i in 0..test_size {
3962                        let obj = Box::new(BalanceTestObj::new((test_size - i) as u64));
3963                        let raw = &*obj as *const BalanceTestObj as *mut BalanceTestObj;
3964                        obj.erase_deck_ptr.set(raw);
3965                        objects.push(obj);
3966
3967                        if i < replacement_count {
3968                            let rep = Box::new(BalanceTestObj::new((test_size - i) as u64));
3969                            let raw = &*rep as *const BalanceTestObj as *mut BalanceTestObj;
3970                            rep.erase_deck_ptr.set(raw);
3971                            replacements.push(rep);
3972                        }
3973                    }
3974                }
3975                _ => {
3976                    for i in 0..test_size {
3977                        let val = rng.get_next();
3978                        let obj = Box::new(BalanceTestObj::new(val));
3979                        let raw = &*obj as *const BalanceTestObj as *mut BalanceTestObj;
3980                        obj.erase_deck_ptr.set(raw);
3981                        objects.push(obj);
3982
3983                        if i < replacement_count {
3984                            let rep = Box::new(BalanceTestObj::new(val));
3985                            let raw = &*rep as *const BalanceTestObj as *mut BalanceTestObj;
3986                            rep.erase_deck_ptr.set(raw);
3987                            replacements.push(rep);
3988                        }
3989                    }
3990                }
3991            }
3992
3993            // 1. Insert all objects
3994            for i in 0..test_size {
3995                unsafe {
3996                    check_augmented_invariants(tree);
3997                    WavlTreeChecker::sanity_check(tree);
3998                    let raw = &mut *objects[i] as *mut BalanceTestObj;
3999                    tree.insert_raw(raw);
4000                    check_augmented_invariants(tree);
4001                    WavlTreeChecker::sanity_check(tree);
4002                }
4003            }
4004
4005            check_iterators(tree);
4006
4007            // 2. Collide replacements
4008            for i in 0..replacement_count {
4009                unsafe {
4010                    check_augmented_invariants(tree);
4011                    WavlTreeChecker::sanity_check(tree);
4012                    let raw = &mut *replacements[i] as *mut BalanceTestObj;
4013                    assert!(tree.insert_or_find_raw(raw).is_err());
4014                    check_augmented_invariants(tree);
4015                    WavlTreeChecker::sanity_check(tree);
4016                }
4017            }
4018
4019            // 3. Replace original nodes with replacements
4020            for i in 0..replacement_count {
4021                unsafe {
4022                    check_augmented_invariants(tree);
4023                    WavlTreeChecker::sanity_check(tree);
4024                    let raw = &mut *replacements[i] as *mut BalanceTestObj;
4025                    assert!(tree.insert_or_replace_raw(raw).is_some());
4026                    check_augmented_invariants(tree);
4027                    WavlTreeChecker::sanity_check(tree);
4028                }
4029            }
4030
4031            check_iterators(tree);
4032
4033            // 4. Swap them back
4034            for i in 0..replacement_count {
4035                unsafe {
4036                    check_augmented_invariants(tree);
4037                    WavlTreeChecker::sanity_check(tree);
4038                    let raw = &mut *objects[i] as *mut BalanceTestObj;
4039                    assert!(tree.insert_or_replace_raw(raw).is_some());
4040                    check_augmented_invariants(tree);
4041                    WavlTreeChecker::sanity_check(tree);
4042                }
4043            }
4044
4045            check_iterators(tree);
4046
4047            // Shuffle erase deck
4048            shuffle_erase_deck(&objects, &mut rng, test_size);
4049
4050            // 5. Erase half the elements
4051            for i in 0..(test_size / 2) {
4052                unsafe {
4053                    check_augmented_invariants(tree);
4054                    WavlTreeChecker::sanity_check(tree);
4055                    let raw_target = objects[i].erase_deck_ptr.get();
4056                    let erased = tree.erase_raw(raw_target);
4057                    assert!(erased.is_some());
4058                    assert_eq!(erased.unwrap(), raw_target);
4059                    check_augmented_invariants(tree);
4060                    WavlTreeChecker::sanity_check(tree);
4061                }
4062            }
4063
4064            check_iterators(tree);
4065
4066            // 6. Put them back
4067            for i in 0..(test_size / 2) {
4068                unsafe {
4069                    check_augmented_invariants(tree);
4070                    WavlTreeChecker::sanity_check(tree);
4071                    let raw_target = objects[i].erase_deck_ptr.get();
4072                    tree.insert_raw(raw_target);
4073                    check_augmented_invariants(tree);
4074                    WavlTreeChecker::sanity_check(tree);
4075                }
4076            }
4077
4078            check_iterators(tree);
4079
4080            // Shuffle erase deck again
4081            shuffle_erase_deck(&objects, &mut rng, test_size);
4082
4083            // 7. Erase everything
4084            for i in 0..test_size {
4085                unsafe {
4086                    check_augmented_invariants(tree);
4087                    WavlTreeChecker::sanity_check(tree);
4088                    let raw_target = objects[i].erase_deck_ptr.get();
4089                    let erased = tree.erase_raw(raw_target);
4090                    assert!(erased.is_some());
4091                    assert_eq!(erased.unwrap(), raw_target);
4092                    check_augmented_invariants(tree);
4093                    WavlTreeChecker::sanity_check(tree);
4094                }
4095            }
4096
4097            check_iterators(tree);
4098            assert_eq!(tree.size.get(), 0);
4099
4100            assert!(op_counts.insert_ops.load(Ordering::Relaxed) > 0);
4101            assert!(op_counts.insert_promotes.load(Ordering::Relaxed) > 0);
4102            assert!(op_counts.insert_rotations.load(Ordering::Relaxed) > 0);
4103            assert!(op_counts.insert_traversals.load(Ordering::Relaxed) > 0);
4104            assert!(op_counts.erase_ops.load(Ordering::Relaxed) > 0);
4105            assert!(op_counts.erase_demotes.load(Ordering::Relaxed) > 0);
4106            assert!(op_counts.erase_rotations.load(Ordering::Relaxed) > 0);
4107        }
4108    }
4109
4110    // WavlTree FFI Declarations
4111    unsafe extern "C" {
4112        // UniqueTree Helpers
4113        fn cpp_create_unique_tree() -> *mut c_void;
4114        fn cpp_destroy_unique_tree(tree: *mut c_void);
4115        fn cpp_unique_tree_insert(tree: *mut c_void, item: *mut c_void);
4116        fn cpp_unique_tree_erase(tree: *mut c_void, key: i32) -> *mut c_void;
4117        fn cpp_unique_tree_find(tree: *mut c_void, key: i32) -> *mut c_void;
4118        fn cpp_unique_tree_is_empty(tree: *mut c_void) -> bool;
4119
4120        // RefTree Helpers
4121        fn cpp_create_ref_tree() -> *mut c_void;
4122        fn cpp_destroy_ref_tree(tree: *mut c_void);
4123        fn cpp_ref_tree_insert(tree: *mut c_void, item: *mut c_void);
4124        fn cpp_ref_tree_erase(tree: *mut c_void, key: i32) -> *mut c_void;
4125        fn cpp_ref_tree_find(tree: *mut c_void, key: i32) -> *mut c_void;
4126        fn cpp_ref_tree_is_empty(tree: *mut c_void) -> bool;
4127
4128        // SharedUniqueObject Helpers (Defined in intrusive_container_test_support.cc)
4129        fn cpp_create_unique_object(value: i32, destruction_flag: *mut bool) -> *mut c_void;
4130        fn cpp_get_unique_object_value(obj: *mut c_void) -> i32;
4131
4132        // SharedRefObject Helpers (Defined in intrusive_container_test_support.cc)
4133        fn cpp_create_ref_object(value: i32, destruction_flag: *mut bool) -> *mut c_void;
4134        fn cpp_get_ref_object_value(obj: *mut c_void) -> i32;
4135    }
4136
4137    #[test]
4138    fn test_interop_rust_tree_cpp_unique_objects() {
4139        use core::sync::atomic::{AtomicBool, Ordering};
4140
4141        let destroyed1 = AtomicBool::new(false);
4142        let destroyed2 = AtomicBool::new(false);
4143
4144        unsafe {
4145            stack_pin_init!(let tree = WavlTree::<i32, UniquePtr<SharedUniqueObject>>::new());
4146            let tree = tree.get_unchecked_mut();
4147
4148            let cpp_raw1 = cpp_create_unique_object(10, destroyed1.as_ptr() as *mut bool);
4149            let cpp_raw2 = cpp_create_unique_object(20, destroyed2.as_ptr() as *mut bool);
4150
4151            let obj1 = UniquePtr::from_raw(cpp_raw1 as *mut SharedUniqueObject);
4152            let obj2 = UniquePtr::from_raw(cpp_raw2 as *mut SharedUniqueObject);
4153
4154            tree.insert(obj1);
4155            tree.insert(obj2);
4156
4157            assert!(!destroyed1.load(Ordering::Relaxed));
4158            assert!(!destroyed2.load(Ordering::Relaxed));
4159
4160            // Find one
4161            let found = tree.find(&10);
4162            assert!(found.is_some());
4163            assert_eq!(found.unwrap().value, 10);
4164
4165            // Erase one
4166            let popped = tree.erase(&20);
4167            assert!(popped.is_some());
4168            assert_eq!(popped.as_ref().unwrap().value, 20);
4169
4170            // Drop popped -> should destroy in C++!
4171            drop(popped);
4172            assert!(!destroyed1.load(Ordering::Relaxed));
4173            assert!(destroyed2.load(Ordering::Relaxed));
4174
4175            // Drop tree -> should destroy remaining in C++!
4176        }
4177        assert!(destroyed1.load(Ordering::Relaxed));
4178    }
4179
4180    #[test]
4181    fn test_interop_cpp_tree_rust_unique_objects() {
4182        use alloc::sync::Arc;
4183        use core::sync::atomic::{AtomicBool, Ordering};
4184
4185        let destroyed1 = Arc::new(AtomicBool::new(false));
4186        let destroyed2 = Arc::new(AtomicBool::new(false));
4187
4188        unsafe {
4189            let cpp_tree = cpp_create_unique_tree();
4190            assert!(cpp_unique_tree_is_empty(cpp_tree));
4191
4192            let obj1 = UniquePtr::try_new(SharedUniqueObject::new(10)).unwrap();
4193            let obj2 = UniquePtr::try_new(SharedUniqueObject::new(20)).unwrap();
4194
4195            // Set destruction flags
4196            let raw1 = UniquePtr::as_ptr(&obj1) as *mut SharedUniqueObject;
4197            (*raw1).destruction_flag = destroyed1.as_ptr() as *mut bool;
4198            let raw2 = UniquePtr::as_ptr(&obj2) as *mut SharedUniqueObject;
4199            (*raw2).destruction_flag = destroyed2.as_ptr() as *mut bool;
4200
4201            // Push to C++ tree (transfers ownership)
4202            cpp_unique_tree_insert(cpp_tree, UniquePtr::into_raw(obj1) as *mut c_void);
4203            cpp_unique_tree_insert(cpp_tree, UniquePtr::into_raw(obj2) as *mut c_void);
4204
4205            assert!(!destroyed1.load(Ordering::Relaxed));
4206            assert!(!destroyed2.load(Ordering::Relaxed));
4207
4208            // Find in C++
4209            let found = cpp_unique_tree_find(cpp_tree, 10);
4210            assert!(!found.is_null());
4211            assert_eq!(cpp_get_unique_object_value(found), 10);
4212
4213            // Erase one from C++
4214            let popped = cpp_unique_tree_erase(cpp_tree, 20);
4215            assert!(!popped.is_null());
4216            assert_eq!(cpp_get_unique_object_value(popped), 20);
4217
4218            // Convert back to Rust UniquePtr and drop -> should free in Rust!
4219            let popped_rust = UniquePtr::from_raw(popped as *mut SharedUniqueObject);
4220            drop(popped_rust);
4221            assert!(!destroyed1.load(Ordering::Relaxed));
4222            assert!(destroyed2.load(Ordering::Relaxed));
4223
4224            // Destroy C++ tree -> should destroy remaining in Rust!
4225            cpp_destroy_unique_tree(cpp_tree);
4226        }
4227        assert!(destroyed1.load(Ordering::Relaxed));
4228    }
4229
4230    #[test]
4231    fn test_interop_rust_tree_cpp_ref_objects() {
4232        use core::sync::atomic::{AtomicBool, Ordering};
4233
4234        let destroyed1 = AtomicBool::new(false);
4235        let destroyed2 = AtomicBool::new(false);
4236
4237        unsafe {
4238            stack_pin_init!(let tree = WavlTree::<i32, RefPtr<SharedRefObject>>::new());
4239            let tree = tree.get_unchecked_mut();
4240
4241            let cpp_raw1 = cpp_create_ref_object(10, destroyed1.as_ptr() as *mut bool);
4242            let cpp_raw2 = cpp_create_ref_object(20, destroyed2.as_ptr() as *mut bool);
4243
4244            let obj1 = RefPtr::from_raw(cpp_raw1 as *mut SharedRefObject);
4245            let obj2 = RefPtr::from_raw(cpp_raw2 as *mut SharedRefObject);
4246
4247            tree.insert(obj1);
4248            tree.insert(obj2);
4249
4250            assert!(!destroyed1.load(Ordering::Relaxed));
4251            assert!(!destroyed2.load(Ordering::Relaxed));
4252
4253            // Find one
4254            let found = tree.find(&10);
4255            assert!(found.is_some());
4256            assert_eq!(found.unwrap().value, 10);
4257
4258            // Erase one
4259            let popped = tree.erase(&20);
4260            assert!(popped.is_some());
4261            assert_eq!(popped.as_ref().unwrap().value, 20);
4262
4263            // Drop popped -> should destroy in C++!
4264            drop(popped);
4265            assert!(!destroyed1.load(Ordering::Relaxed));
4266            assert!(destroyed2.load(Ordering::Relaxed));
4267
4268            // Drop tree -> should destroy remaining in C++!
4269        }
4270        assert!(destroyed1.load(Ordering::Relaxed));
4271    }
4272
4273    #[test]
4274    fn test_interop_cpp_tree_rust_ref_objects() {
4275        use alloc::sync::Arc;
4276        use core::sync::atomic::{AtomicBool, Ordering};
4277
4278        let destroyed1 = Arc::new(AtomicBool::new(false));
4279        let destroyed2 = Arc::new(AtomicBool::new(false));
4280
4281        unsafe {
4282            let cpp_tree = cpp_create_ref_tree();
4283            assert!(cpp_ref_tree_is_empty(cpp_tree));
4284
4285            let obj1 = SharedRefObject::new_ref_counted(10);
4286            let obj2 = SharedRefObject::new_ref_counted(20);
4287
4288            // Set destruction flags
4289            let raw1 = RefPtr::as_ptr(&obj1) as *mut SharedRefObject;
4290            (*raw1).destruction_flag = destroyed1.as_ptr() as *mut bool;
4291            let raw2 = RefPtr::as_ptr(&obj2) as *mut SharedRefObject;
4292            (*raw2).destruction_flag = destroyed2.as_ptr() as *mut bool;
4293
4294            // Insert to C++ tree (transfers ownership)
4295            cpp_ref_tree_insert(
4296                cpp_tree,
4297                RefPtr::into_raw(obj1) as *mut SharedRefObject as *mut c_void,
4298            );
4299            cpp_ref_tree_insert(
4300                cpp_tree,
4301                RefPtr::into_raw(obj2) as *mut SharedRefObject as *mut c_void,
4302            );
4303
4304            assert!(!destroyed1.load(Ordering::Relaxed));
4305            assert!(!destroyed2.load(Ordering::Relaxed));
4306
4307            // Find in C++
4308            let found = cpp_ref_tree_find(cpp_tree, 10);
4309            assert!(!found.is_null());
4310            assert_eq!(cpp_get_ref_object_value(found), 10);
4311
4312            // Erase one from C++
4313            let popped = cpp_ref_tree_erase(cpp_tree, 20);
4314            assert!(!popped.is_null());
4315            assert_eq!(cpp_get_ref_object_value(popped), 20);
4316
4317            // Convert back to Rust RefPtr and drop -> should free in Rust!
4318            let popped_rust = RefPtr::from_raw(popped as *mut SharedRefObject);
4319            drop(popped_rust);
4320            assert!(!destroyed1.load(Ordering::Relaxed));
4321            assert!(destroyed2.load(Ordering::Relaxed));
4322
4323            // Destroy C++ tree -> should destroy remaining in Rust!
4324            cpp_destroy_ref_tree(cpp_tree);
4325        }
4326        assert!(destroyed1.load(Ordering::Relaxed));
4327    }
4328
4329    #[derive(PartialEq, Eq, PartialOrd, Ord, Debug)]
4330    struct LargeKey {
4331        name: [u8; 32],
4332    }
4333
4334    #[derive(crate::WavlTreeContainable, crate::Recyclable)]
4335    struct TestRefKeyObject {
4336        key: LargeKey,
4337        #[wavl_node]
4338        node: WavlTreeNode<TestRefKeyObject>,
4339    }
4340
4341    impl TestRefKeyObject {
4342        fn new(name: [u8; 32]) -> Self {
4343            Self { key: LargeKey { name }, node: WavlTreeNode::new() }
4344        }
4345    }
4346
4347    impl WavlTreeKeyable<LargeKey> for TestRefKeyObject {
4348        type Key<'a> = &'a LargeKey;
4349        fn get_key(&self) -> &LargeKey {
4350            &self.key
4351        }
4352    }
4353
4354    #[test]
4355    fn test_by_reference_key() {
4356        use crate::UniquePtr;
4357
4358        type TestTree = WavlTree<LargeKey, UniquePtr<TestRefKeyObject>>;
4359        stack_pin_init!(let tree = TestTree::new());
4360        let tree = unsafe { tree.get_unchecked_mut() };
4361
4362        let mut key1 = [0u8; 32];
4363        key1[0] = 10;
4364        let mut key2 = [0u8; 32];
4365        key2[0] = 20;
4366
4367        let obj1 = UniquePtr::try_new(TestRefKeyObject::new(key1)).unwrap();
4368        let obj2 = UniquePtr::try_new(TestRefKeyObject::new(key2)).unwrap();
4369        tree.insert(obj1);
4370        tree.insert(obj2);
4371
4372        let query = LargeKey { name: key1 };
4373        let found = tree.find(&query);
4374        assert!(found.is_some());
4375        assert_eq!(found.unwrap().key, query);
4376
4377        let erased = tree.erase(&query);
4378        assert!(erased.is_some());
4379        assert_eq!(erased.unwrap().key, query);
4380    }
4381
4382    #[derive(WavlTreeContainable, Recyclable)]
4383    struct TestAugmentedObject {
4384        value: i32,
4385        subtree_sum: i32,
4386        #[wavl_node]
4387        node: WavlTreeNode<TestAugmentedObject>,
4388    }
4389
4390    impl TestAugmentedObject {
4391        fn new(value: i32) -> Self {
4392            Self { value, subtree_sum: 0, node: WavlTreeNode::new() }
4393        }
4394    }
4395
4396    impl WavlTreeKeyable<i32> for TestAugmentedObject {
4397        type Key<'a> = i32;
4398        fn get_key(&self) -> i32 {
4399            self.value
4400        }
4401    }
4402
4403    struct SubtreeSumTraits;
4404    impl WavlTreeAugmentedInvariantObserverTraits for SubtreeSumTraits {
4405        type Target = TestAugmentedObject;
4406        type Value = i32;
4407
4408        fn get_node_value(node: &Self::Target) -> Self::Value {
4409            node.value
4410        }
4411
4412        fn get_subtree_value(node: &Self::Target) -> Self::Value {
4413            node.subtree_sum
4414        }
4415
4416        fn combine_values(a: Self::Value, b: Self::Value) -> Self::Value {
4417            a + b
4418        }
4419
4420        fn set_subtree_value(node: &mut Self::Target, val: Self::Value) {
4421            node.subtree_sum = val;
4422        }
4423
4424        fn reset_subtree_value(node: &mut Self::Target) {
4425            node.subtree_sum = 0;
4426        }
4427    }
4428
4429    #[test]
4430    fn test_augmented_invariant_observer() {
4431        type TestObserver = WavlTreeAugmentedInvariantObserver<DefaultObjectTag, SubtreeSumTraits>;
4432        type TestTree = WavlTree<
4433            i32,
4434            UniquePtr<TestAugmentedObject>,
4435            DefaultObjectTag,
4436            TrackingSize,
4437            TestObserver,
4438        >;
4439
4440        stack_pin_init!(let tree = TestTree::new_with_observer(TestObserver::default()));
4441        let tree = unsafe { tree.get_unchecked_mut() };
4442
4443        // Helper to verify subtree sums recursively
4444        fn verify_subtree_sums(node: *mut TestAugmentedObject) -> i32 {
4445            if !valid_sentinel_ptr(node) {
4446                return 0;
4447            }
4448            let node_ref = unsafe { &*node };
4449            let ns = node_ref.get_node();
4450            let left_sum = verify_subtree_sums(ns.get_left());
4451            let right_sum = verify_subtree_sums(ns.get_right());
4452            let expected_sum = node_ref.value + left_sum + right_sum;
4453            assert_eq!(
4454                node_ref.subtree_sum, expected_sum,
4455                "Subtree sum mismatch at node {}",
4456                node_ref.value
4457            );
4458            expected_sum
4459        }
4460
4461        // Insert some nodes
4462        let values = [10, 20, 5, 15, 25, 2, 7];
4463        for &val in &values {
4464            let obj = UniquePtr::try_new(TestAugmentedObject::new(val)).unwrap();
4465            tree.insert(obj);
4466            // Verify after each insert
4467            if !tree.root.is_null() {
4468                verify_subtree_sums(tree.root);
4469            }
4470        }
4471
4472        // Verify final sum at root
4473        let expected_total_sum: i32 = values.iter().sum();
4474        assert_eq!(unsafe { &*tree.root }.subtree_sum, expected_total_sum);
4475
4476        // Erase some nodes
4477        let to_erase = [15, 5, 20];
4478        let mut current_sum = expected_total_sum;
4479        for &val in &to_erase {
4480            let erased = tree.erase(&val);
4481            assert!(erased.is_some());
4482            assert_eq!(erased.as_ref().unwrap().value, val);
4483            current_sum -= val;
4484            if !tree.root.is_null() {
4485                assert_eq!(unsafe { &*tree.root }.subtree_sum, current_sum);
4486                verify_subtree_sums(tree.root);
4487            }
4488        }
4489    }
4490
4491    #[derive(Debug, PartialEq, Eq, Clone, Copy, Default)]
4492    struct NonCommutativeVal {
4493        hash: u64,
4494        len: u32,
4495    }
4496
4497    impl NonCommutativeVal {
4498        fn single(val: i32) -> Self {
4499            Self { hash: (val as u64) & 0xffff, len: 1 }
4500        }
4501
4502        fn combine(left: Self, right: Self) -> Self {
4503            if left.len == 0 {
4504                return right;
4505            }
4506            if right.len == 0 {
4507                return left;
4508            }
4509            let mut power = 1u64;
4510            for _ in 0..right.len {
4511                power = power.wrapping_mul(31);
4512            }
4513            Self {
4514                hash: left.hash.wrapping_mul(power).wrapping_add(right.hash),
4515                len: left.len + right.len,
4516            }
4517        }
4518    }
4519
4520    #[derive(WavlTreeContainable, Recyclable)]
4521    struct TestNonCommutativeObject {
4522        value: i32,
4523        subtree_val: NonCommutativeVal,
4524        #[wavl_node]
4525        node: WavlTreeNode<TestNonCommutativeObject>,
4526    }
4527
4528    impl TestNonCommutativeObject {
4529        fn new(value: i32) -> Self {
4530            Self { value, subtree_val: NonCommutativeVal::default(), node: WavlTreeNode::new() }
4531        }
4532    }
4533
4534    impl WavlTreeKeyable<i32> for TestNonCommutativeObject {
4535        type Key<'a> = i32;
4536        fn get_key(&self) -> i32 {
4537            self.value
4538        }
4539    }
4540
4541    struct NonCommutativeTraits;
4542    impl WavlTreeAugmentedInvariantObserverTraits for NonCommutativeTraits {
4543        type Target = TestNonCommutativeObject;
4544        type Value = NonCommutativeVal;
4545
4546        fn get_node_value(node: &Self::Target) -> Self::Value {
4547            NonCommutativeVal::single(node.value)
4548        }
4549
4550        fn get_subtree_value(node: &Self::Target) -> Self::Value {
4551            node.subtree_val
4552        }
4553
4554        fn combine_values(a: Self::Value, b: Self::Value) -> Self::Value {
4555            NonCommutativeVal::combine(a, b)
4556        }
4557
4558        fn set_subtree_value(node: &mut Self::Target, val: Self::Value) {
4559            node.subtree_val = val;
4560        }
4561
4562        fn reset_subtree_value(node: &mut Self::Target) {
4563            node.subtree_val = NonCommutativeVal::default();
4564        }
4565    }
4566
4567    #[test]
4568    fn test_augmented_invariant_observer_non_commutative() {
4569        type TestObserver =
4570            WavlTreeAugmentedInvariantObserver<DefaultObjectTag, NonCommutativeTraits>;
4571        type TestTree = WavlTree<
4572            i32,
4573            UniquePtr<TestNonCommutativeObject>,
4574            DefaultObjectTag,
4575            TrackingSize,
4576            TestObserver,
4577        >;
4578
4579        stack_pin_init!(let tree = TestTree::new_with_observer(TestObserver::default()));
4580        let tree = unsafe { tree.get_unchecked_mut() };
4581
4582        // Helper to verify non-commutative subtree values recursively
4583        fn verify_subtree_values(node: *mut TestNonCommutativeObject) -> NonCommutativeVal {
4584            if !valid_sentinel_ptr(node) {
4585                return NonCommutativeVal::default();
4586            }
4587            let node_ref = unsafe { &*node };
4588            let ns = node_ref.get_node();
4589            let left_val = verify_subtree_values(ns.get_left());
4590            let node_val = NonCommutativeVal::single(node_ref.value);
4591            let right_val = verify_subtree_values(ns.get_right());
4592            let expected = NonCommutativeVal::combine(
4593                NonCommutativeVal::combine(left_val, node_val),
4594                right_val,
4595            );
4596            assert_eq!(
4597                node_ref.subtree_val, expected,
4598                "Subtree non-commutative value mismatch at node {}",
4599                node_ref.value
4600            );
4601            expected
4602        }
4603
4604        // Helper to compute expected in-order value from a list of sorted keys
4605        fn expected_in_order(keys: &[i32]) -> NonCommutativeVal {
4606            let mut sorted: alloc::vec::Vec<i32> = keys.iter().copied().collect();
4607            sorted.sort();
4608            let mut acc = NonCommutativeVal::default();
4609            for &k in &sorted {
4610                acc = NonCommutativeVal::combine(acc, NonCommutativeVal::single(k));
4611            }
4612            acc
4613        }
4614
4615        // Insert nodes in an order that triggers single and double rotations in both directions
4616        let values = [10, 20, 5, 15, 25, 2, 7, 1, 3, 6, 8, 12, 17, 22, 30];
4617        for (i, &val) in values.iter().enumerate() {
4618            let obj = UniquePtr::try_new(TestNonCommutativeObject::new(val)).unwrap();
4619            tree.insert(obj);
4620            if !tree.root.is_null() {
4621                verify_subtree_values(tree.root);
4622                let expected = expected_in_order(&values[..=i]);
4623                assert_eq!(unsafe { &*tree.root }.subtree_val, expected);
4624            }
4625        }
4626
4627        // Erase nodes causing rebalancings/rotations
4628        let to_erase = [5, 20, 15, 2, 25];
4629        let mut remaining: alloc::vec::Vec<i32> = values.iter().copied().collect();
4630        for &val in &to_erase {
4631            let erased = tree.erase(&val);
4632            assert!(erased.is_some());
4633            assert_eq!(erased.as_ref().unwrap().value, val);
4634            remaining.retain(|&x| x != val);
4635            if !tree.root.is_null() {
4636                verify_subtree_values(tree.root);
4637                let expected = expected_in_order(&remaining);
4638                assert_eq!(unsafe { &*tree.root }.subtree_val, expected);
4639            }
4640        }
4641    }
4642}