openmls/binary_tree/array_representation/
treemath.rs1use 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
9const MAX_TREE_INDEX: u32 = MAX_TREE_SIZE - 1;
12
13const MAX_LEAF: u32 = MAX_TREE_INDEX / 2;
15
16const MAX_PARENT: u32 = MAX_LEAF - 1;
19
20#[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 pub fn new(index: u32) -> Self {
48 LeafNodeIndex(index)
49 }
50
51 fn valid(&self) -> bool {
53 self.0 <= MAX_LEAF
54 }
55 pub fn u32(&self) -> u32 {
57 self.0
58 }
59
60 pub fn usize(&self) -> usize {
62 self.u32() as usize
63 }
64
65 fn to_tree_index(self) -> u32 {
67 self.0 * 2
68 }
69
70 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#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
79pub struct ParentNodeIndex(u32);
80
81impl ParentNodeIndex {
82 pub(crate) fn new(index: u32) -> Self {
84 ParentNodeIndex(index)
85 }
86
87 fn valid(&self) -> bool {
89 self.0 <= MAX_PARENT
90 }
91
92 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 fn to_tree_index(self) -> u32 {
103 self.0 * 2 + 1
104 }
105
106 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 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 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
144pub enum TreeNodeIndex {
145 Leaf(LeafNodeIndex),
146 Parent(ParentNodeIndex),
147}
148
149impl TreeNodeIndex {
150 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 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 #[cfg(any(feature = "test-utils", test))]
169 pub(crate) fn test_new(index: u32) -> Self {
170 Self::new(index)
171 }
172
173 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 #[cfg(any(feature = "test-utils", feature = "crypto-debug", test))]
183 pub(crate) fn test_u32(&self) -> u32 {
184 self.u32()
185 }
186
187 #[cfg(any(feature = "test-utils", test))]
189 fn usize(&self) -> usize {
190 self.u32() as usize
191 }
192
193 #[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 pub(crate) fn new(nodes: u32) -> Self {
220 let k = log2(nodes);
221 TreeSize((1 << (k + 1)) - 1)
222 }
223
224 #[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 pub(crate) fn leaf_count(&self) -> u32 {
232 (self.0 / 2) + 1
233 }
234
235 pub(crate) fn parent_count(&self) -> u32 {
237 self.0 / 2
238 }
239
240 pub(crate) fn u32(&self) -> u32 {
242 self.0
243 }
244
245 pub(crate) fn leaf_is_left(&self, leaf_index: LeafNodeIndex) -> bool {
248 leaf_index.u32() < self.leaf_count() / 2
249 }
250
251 pub(super) fn inc(&mut self) {
253 self.0 = self.0 * 2 + 1;
254 }
255
256 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]
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 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
337fn 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#[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#[cfg(any(feature = "test-utils", test))]
364pub(crate) fn test_sibling(index: TreeNodeIndex) -> TreeNodeIndex {
365 sibling(index)
366}
367
368pub(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
383pub(crate) fn copath(leaf_index: LeafNodeIndex, size: TreeSize) -> Vec<TreeNodeIndex> {
385 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 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
402pub(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
424pub(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}