Skip to main content

openmls/binary_tree/array_representation/
treemath.rs

1use std::cmp::Ordering;
2
3use serde::{Deserialize, Serialize};
4use tls_codec::{TlsDeserialize, TlsDeserializeBytes, TlsSerialize, TlsSize};
5
6pub(crate) const MAX_TREE_SIZE: u32 = (1 << 30) - 1;
7pub(crate) const MIN_TREE_SIZE: u32 = 1;
8
9/// Largest tree (node) index of any valid tree: node indices range over
10/// `0..=MAX_TREE_INDEX`. It is the last leaf of the maximal tree (even).
11const MAX_TREE_INDEX: u32 = MAX_TREE_SIZE - 1;
12
13/// Largest leaf payload: leaf `l` sits at tree index `2*l <= MAX_TREE_INDEX`.
14const MAX_LEAF: u32 = MAX_TREE_INDEX / 2;
15
16/// Largest parent payload: parent `p` sits at tree index `2*p + 1 < MAX_TREE_INDEX`
17/// (odd indices stop one short of the even maximum).
18const MAX_PARENT: u32 = MAX_LEAF - 1;
19
20/// LeafNodeIndex references a leaf node in a tree.
21#[derive(
22    Debug,
23    Clone,
24    Copy,
25    PartialEq,
26    Eq,
27    PartialOrd,
28    Ord,
29    Hash,
30    Serialize,
31    Deserialize,
32    TlsDeserialize,
33    TlsDeserializeBytes,
34    TlsSerialize,
35    TlsSize,
36)]
37pub struct LeafNodeIndex(u32);
38
39impl std::fmt::Display for LeafNodeIndex {
40    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41        f.write_fmt(format_args!("{:?}", self.0))
42    }
43}
44
45impl LeafNodeIndex {
46    /// Create a new `LeafNodeIndex` from a `u32`.
47    pub fn new(index: u32) -> Self {
48        LeafNodeIndex(index)
49    }
50
51    /// Checks that the wrapped index is valid.
52    fn valid(&self) -> bool {
53        self.0 <= MAX_LEAF
54    }
55    /// Return the inner value as `u32`.
56    pub fn u32(&self) -> u32 {
57        self.0
58    }
59
60    /// Return the inner value as `usize`.
61    pub fn usize(&self) -> usize {
62        self.u32() as usize
63    }
64
65    /// Return the index as a TreeNodeIndex value.
66    fn to_tree_index(self) -> u32 {
67        self.0 * 2
68    }
69
70    /// Warning: Only use when the node index represents a leaf node
71    fn from_tree_index(node_index: u32) -> Self {
72        debug_assert!(node_index.is_multiple_of(2));
73        LeafNodeIndex(node_index / 2)
74    }
75}
76
77/// ParentNodeIndex references a parent node in a tree.
78#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
79pub struct ParentNodeIndex(u32);
80
81impl ParentNodeIndex {
82    /// Create a new `ParentNodeIndex` from a `u32`.
83    pub(crate) fn new(index: u32) -> Self {
84        ParentNodeIndex(index)
85    }
86
87    /// Checks that the wrapped index is valid.
88    fn valid(&self) -> bool {
89        self.0 <= MAX_PARENT
90    }
91
92    /// Return the inner value as `u32`.
93    pub fn u32(&self) -> u32 {
94        self.0
95    }
96
97    pub(crate) fn usize(&self) -> usize {
98        self.0 as usize
99    }
100
101    /// Return the index as a TreeNodeIndex value.
102    fn to_tree_index(self) -> u32 {
103        self.0 * 2 + 1
104    }
105
106    /// Warning: Only use when the node index represents a parent node
107    fn from_tree_index(node_index: u32) -> Self {
108        debug_assert!(node_index > 0);
109        debug_assert!(node_index % 2 == 1);
110        ParentNodeIndex((node_index - 1) / 2)
111    }
112}
113
114#[cfg(test)]
115impl ParentNodeIndex {
116    /// Re-exported for testing.
117    pub(crate) fn test_from_tree_index(node_index: u32) -> Self {
118        Self::from_tree_index(node_index)
119    }
120}
121
122#[cfg(any(feature = "test-utils", test))]
123impl ParentNodeIndex {
124    /// Re-exported for testing.
125    pub(crate) fn test_to_tree_index(self) -> u32 {
126        self.to_tree_index()
127    }
128}
129
130impl From<LeafNodeIndex> for TreeNodeIndex {
131    fn from(leaf_index: LeafNodeIndex) -> Self {
132        TreeNodeIndex::Leaf(leaf_index)
133    }
134}
135
136impl From<ParentNodeIndex> for TreeNodeIndex {
137    fn from(parent_index: ParentNodeIndex) -> Self {
138        TreeNodeIndex::Parent(parent_index)
139    }
140}
141
142/// TreeNodeIndex references a node in a tree.
143#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
144pub enum TreeNodeIndex {
145    Leaf(LeafNodeIndex),
146    Parent(ParentNodeIndex),
147}
148
149impl TreeNodeIndex {
150    /// Create a new `TreeNodeIndex` from a `u32`.
151    fn new(index: u32) -> Self {
152        if index.is_multiple_of(2) {
153            TreeNodeIndex::Leaf(LeafNodeIndex::from_tree_index(index))
154        } else {
155            TreeNodeIndex::Parent(ParentNodeIndex::from_tree_index(index))
156        }
157    }
158
159    /// Checks that the wrapped index is valid.
160    fn valid(&self) -> bool {
161        match self {
162            TreeNodeIndex::Leaf(leaf_node_index) => leaf_node_index.valid(),
163            TreeNodeIndex::Parent(parent_node_index) => parent_node_index.valid(),
164        }
165    }
166
167    /// Re-exported for testing.
168    #[cfg(any(feature = "test-utils", test))]
169    pub(crate) fn test_new(index: u32) -> Self {
170        Self::new(index)
171    }
172
173    /// Return the inner value as `u32`.
174    fn u32(&self) -> u32 {
175        match self {
176            TreeNodeIndex::Leaf(index) => index.to_tree_index(),
177            TreeNodeIndex::Parent(index) => index.to_tree_index(),
178        }
179    }
180
181    /// Re-exported for testing.
182    #[cfg(any(feature = "test-utils", feature = "crypto-debug", test))]
183    pub(crate) fn test_u32(&self) -> u32 {
184        self.u32()
185    }
186
187    /// Return the inner value as `usize`.
188    #[cfg(any(feature = "test-utils", test))]
189    fn usize(&self) -> usize {
190        self.u32() as usize
191    }
192
193    /// Re-exported for testing.
194    #[cfg(any(feature = "test-utils", test))]
195    pub(crate) fn test_usize(&self) -> usize {
196        self.usize()
197    }
198}
199
200impl Ord for TreeNodeIndex {
201    fn cmp(&self, other: &TreeNodeIndex) -> Ordering {
202        self.u32().cmp(&other.u32())
203    }
204}
205
206impl PartialOrd for TreeNodeIndex {
207    fn partial_cmp(&self, other: &TreeNodeIndex) -> Option<Ordering> {
208        Some(self.cmp(other))
209    }
210}
211
212#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
213pub(crate) struct TreeSize(u32);
214
215impl TreeSize {
216    /// Create a new `TreeSize` from `nodes`, which will be rounded up to the
217    /// next power of 2. The tree size then reflects the smallest tree that can
218    /// contain the number of nodes.
219    pub(crate) fn new(nodes: u32) -> Self {
220        let k = log2(nodes);
221        TreeSize((1 << (k + 1)) - 1)
222    }
223
224    /// Creates a new `TreeSize` from a specific leaf count
225    #[cfg(any(feature = "test-utils", feature = "extensions-draft", test))]
226    pub(crate) fn from_leaf_count(leaf_count: u32) -> Self {
227        TreeSize::new(leaf_count * 2)
228    }
229
230    /// Return the number of leaf nodes in the tree.
231    pub(crate) fn leaf_count(&self) -> u32 {
232        (self.0 / 2) + 1
233    }
234
235    /// Return the number of parent nodes in the tree.
236    pub(crate) fn parent_count(&self) -> u32 {
237        self.0 / 2
238    }
239
240    /// Return the inner value as `u32`.
241    pub(crate) fn u32(&self) -> u32 {
242        self.0
243    }
244
245    /// Returns `true` if the leaf is in the left subtree and `false` otherwise.
246    /// If there is only one leaf in the tree, it returns `false`.
247    pub(crate) fn leaf_is_left(&self, leaf_index: LeafNodeIndex) -> bool {
248        leaf_index.u32() < self.leaf_count() / 2
249    }
250
251    /// Increase the size.
252    pub(super) fn inc(&mut self) {
253        self.0 = self.0 * 2 + 1;
254    }
255
256    /// Decrease the size.
257    pub(super) fn dec(&mut self) {
258        debug_assert!(self.0 >= 2);
259        if self.0 >= 2 {
260            self.0 = self.0.div_ceil(2) - 1;
261        } else {
262            self.0 = 0;
263        }
264    }
265}
266
267#[test]
268fn tree_size() {
269    assert_eq!(TreeSize::new(1).u32(), 1);
270    assert_eq!(TreeSize::new(3).u32(), 3);
271    assert_eq!(TreeSize::new(5).u32(), 7);
272    assert_eq!(TreeSize::new(7).u32(), 7);
273    assert_eq!(TreeSize::new(9).u32(), 15);
274    assert_eq!(TreeSize::new(11).u32(), 15);
275    assert_eq!(TreeSize::new(13).u32(), 15);
276    assert_eq!(TreeSize::new(15).u32(), 15);
277    assert_eq!(TreeSize::new(17).u32(), 31);
278}
279
280/// Test if the leaf is in the left subtree.
281#[test]
282fn test_leaf_is_left() {
283    assert!(!TreeSize::new(1).leaf_is_left(LeafNodeIndex::new(0)));
284
285    assert!(TreeSize::new(3).leaf_is_left(LeafNodeIndex::new(0)));
286    assert!(!TreeSize::new(3).leaf_is_left(LeafNodeIndex::new(1)));
287
288    assert!(TreeSize::new(5).leaf_is_left(LeafNodeIndex::new(0)));
289    assert!(TreeSize::new(5).leaf_is_left(LeafNodeIndex::new(1)));
290    assert!(!TreeSize::new(5).leaf_is_left(LeafNodeIndex::new(2)));
291    assert!(!TreeSize::new(5).leaf_is_left(LeafNodeIndex::new(3)));
292
293    assert!(TreeSize::new(15).leaf_is_left(LeafNodeIndex::new(0)));
294    assert!(TreeSize::new(15).leaf_is_left(LeafNodeIndex::new(1)));
295    assert!(TreeSize::new(15).leaf_is_left(LeafNodeIndex::new(2)));
296    assert!(TreeSize::new(15).leaf_is_left(LeafNodeIndex::new(3)));
297    assert!(!TreeSize::new(15).leaf_is_left(LeafNodeIndex::new(4)));
298    assert!(!TreeSize::new(15).leaf_is_left(LeafNodeIndex::new(5)));
299    assert!(!TreeSize::new(15).leaf_is_left(LeafNodeIndex::new(6)));
300    assert!(!TreeSize::new(15).leaf_is_left(LeafNodeIndex::new(7)));
301}
302
303fn log2(x: u32) -> usize {
304    if x == 0 {
305        return 0;
306    }
307    (31 - x.leading_zeros()) as usize
308}
309
310pub fn level(index: u32) -> usize {
311    // The cast is always valid, as there is at most 32 trailing ones
312    index.trailing_ones() as usize
313}
314
315pub(crate) fn root(size: TreeSize) -> TreeNodeIndex {
316    let size = size.u32();
317    debug_assert!(size > 0);
318    TreeNodeIndex::new((1 << log2(size)) - 1)
319}
320
321pub(crate) fn left(index: ParentNodeIndex) -> TreeNodeIndex {
322    let x = index.to_tree_index();
323    let k = level(x);
324    debug_assert!(k > 0);
325    let index = x ^ (0x01 << (k - 1));
326    TreeNodeIndex::new(index)
327}
328
329pub(crate) fn right(index: ParentNodeIndex) -> TreeNodeIndex {
330    let x = index.to_tree_index();
331    let k = level(x);
332    debug_assert!(k > 0);
333    let index = x ^ (0x03 << (k - 1));
334    TreeNodeIndex::new(index)
335}
336
337/// Warning: There is no check about the tree size and whether the parent is
338/// beyond the root
339fn parent(x: TreeNodeIndex) -> ParentNodeIndex {
340    let x = x.u32();
341    let k = level(x);
342    let b = (x >> (k + 1)) & 0x01;
343    let index = (x | (1 << k)) ^ (b << (k + 1));
344    ParentNodeIndex::from_tree_index(index)
345}
346
347/// Re-exported for testing.
348#[cfg(any(feature = "test-utils", test))]
349pub(crate) fn test_parent(index: TreeNodeIndex) -> ParentNodeIndex {
350    parent(index)
351}
352
353fn sibling(index: TreeNodeIndex) -> TreeNodeIndex {
354    let p = parent(index);
355    match index.u32().cmp(&p.to_tree_index()) {
356        Ordering::Less => right(p),
357        Ordering::Greater => left(p),
358        Ordering::Equal => left(p),
359    }
360}
361
362/// Re-exported for testing.
363#[cfg(any(feature = "test-utils", test))]
364pub(crate) fn test_sibling(index: TreeNodeIndex) -> TreeNodeIndex {
365    sibling(index)
366}
367
368/// Direct path from a node to the root.
369/// Does not include the node itself.
370pub(crate) fn direct_path(node_index: LeafNodeIndex, size: TreeSize) -> Vec<ParentNodeIndex> {
371    let r = root(size).u32();
372
373    let mut d = vec![];
374    let mut x = node_index.to_tree_index();
375    while x != r {
376        let parent = parent(TreeNodeIndex::new(x));
377        d.push(parent);
378        x = parent.to_tree_index();
379    }
380    d
381}
382
383/// Copath of a leaf node.
384pub(crate) fn copath(leaf_index: LeafNodeIndex, size: TreeSize) -> Vec<TreeNodeIndex> {
385    // Start with leaf
386    let mut full_path = vec![TreeNodeIndex::Leaf(leaf_index)];
387    let mut direct_path = direct_path(leaf_index, size);
388    if !direct_path.is_empty() {
389        // Remove root
390        direct_path.pop();
391    }
392    full_path.append(
393        &mut direct_path
394            .iter()
395            .map(|i| TreeNodeIndex::Parent(*i))
396            .collect(),
397    );
398
399    full_path.into_iter().map(sibling).collect()
400}
401
402/// Common ancestor of two leaf nodes, aka the node where their direct paths
403/// intersect.
404pub(super) fn lowest_common_ancestor(x: LeafNodeIndex, y: LeafNodeIndex) -> ParentNodeIndex {
405    let x = x.to_tree_index();
406    let y = y.to_tree_index();
407    let (lx, ly) = (level(x) + 1, level(y) + 1);
408    if (lx <= ly) && (x >> ly == y >> ly) {
409        return ParentNodeIndex::from_tree_index(y);
410    } else if (ly <= lx) && (x >> lx == y >> lx) {
411        return ParentNodeIndex::from_tree_index(x);
412    }
413
414    let (mut xn, mut yn) = (x, y);
415    let mut k = 0;
416    while xn != yn {
417        xn >>= 1;
418        yn >>= 1;
419        k += 1;
420    }
421    ParentNodeIndex::from_tree_index((xn << k) + (1 << (k - 1)) - 1)
422}
423
424/// The common direct path of two leaf nodes, i.e. the path from their common
425/// ancestor to the root.
426pub(crate) fn common_direct_path(
427    x: LeafNodeIndex,
428    y: LeafNodeIndex,
429    size: TreeSize,
430) -> Vec<ParentNodeIndex> {
431    let mut x_path = direct_path(x, size);
432    let mut y_path = direct_path(y, size);
433    x_path.reverse();
434    y_path.reverse();
435
436    let mut common_path = vec![];
437
438    for (x, y) in x_path.iter().zip(y_path.iter()) {
439        if x == y {
440            common_path.push(*x);
441        } else {
442            break;
443        }
444    }
445
446    common_path.reverse();
447    common_path
448}
449
450#[cfg(any(feature = "test-utils", test))]
451pub(crate) fn node_width(n: usize) -> usize {
452    if n == 0 {
453        0
454    } else {
455        2 * (n - 1) + 1
456    }
457}
458
459pub(crate) fn is_node_in_tree(node_index: TreeNodeIndex, size: TreeSize) -> bool {
460    node_index.valid() && node_index.u32() < size.u32()
461}
462
463#[test]
464fn test_node_in_tree() {
465    let tests = [(0u32, 3u32), (1, 3), (2, 5), (5, 7), (2, 11)];
466    for test in tests.iter() {
467        assert!(is_node_in_tree(
468            TreeNodeIndex::new(test.0),
469            TreeSize::new(test.1)
470        ));
471    }
472}
473
474#[test]
475fn test_node_not_in_tree() {
476    let tests = [(3u32, 1u32), (13, 7)];
477    for test in tests.iter() {
478        assert!(!is_node_in_tree(
479            TreeNodeIndex::new(test.0),
480            TreeSize::new(test.1)
481        ));
482    }
483}
484
485#[test]
486fn test_node_not_in_tree_wrapping() {
487    let tests = [1u32 << 31, u32::MAX];
488    for leaf in tests.iter() {
489        assert!(!is_node_in_tree(
490            TreeNodeIndex::Leaf(LeafNodeIndex::new(*leaf)),
491            TreeSize::new(3)
492        ));
493    }
494}