Skip to main content

chia_datalayer/merkle/
iterators.rs

1use crate::merkle::error::Error;
2use crate::{BLOCK_SIZE, Block, Node, TreeIndex, try_get_block};
3use std::collections::{HashSet, VecDeque};
4
5struct LeftChildFirstIteratorItem {
6    visited: bool,
7    index: TreeIndex,
8}
9
10pub struct LeftChildFirstIterator<'a> {
11    blob: &'a [u8],
12    stack: Vec<LeftChildFirstIteratorItem>,
13    already_queued: HashSet<TreeIndex>,
14    predicate: Option<fn(&Block) -> bool>,
15    from_index: TreeIndex,
16}
17
18impl<'a> LeftChildFirstIterator<'a> {
19    pub fn new(blob: &'a [u8], from_index: Option<TreeIndex>) -> Self {
20        Self::new_with_block_predicate(blob, from_index, None)
21    }
22
23    pub fn new_with_block_predicate(
24        blob: &'a [u8],
25        from_index: Option<TreeIndex>,
26        predicate: Option<fn(&Block) -> bool>,
27    ) -> Self {
28        let mut stack = Vec::new();
29        let from_index = from_index.unwrap_or(TreeIndex(0));
30        if blob.len() / BLOCK_SIZE > 0 {
31            stack.push(LeftChildFirstIteratorItem {
32                visited: false,
33                index: from_index,
34            });
35        }
36
37        Self {
38            blob,
39            stack,
40            already_queued: HashSet::new(),
41            predicate,
42            from_index,
43        }
44    }
45}
46
47impl Iterator for LeftChildFirstIterator<'_> {
48    type Item = Result<(TreeIndex, Block), Error>;
49
50    fn next(&mut self) -> Option<Self::Item> {
51        // left sibling first, children before parents
52
53        loop {
54            let item = self.stack.pop()?;
55            let block = match try_get_block(self.blob, item.index) {
56                Ok(block) => block,
57                Err(e) => return Some(Err(e)),
58            };
59
60            if let Some(predicate) = self.predicate {
61                if !predicate(&block) {
62                    continue;
63                }
64            }
65
66            match block.node.parent().0 {
67                Some(index) => {
68                    if item.index == TreeIndex(0) {
69                        return Some(Err(Error::RootHasParent()));
70                    } else if item.index == self.from_index {
71                        match try_get_block(self.blob, index) {
72                            Ok(Block {
73                                node: Node::Internal(node),
74                                ..
75                            }) => {
76                                if item.index != node.left && item.index != node.right {
77                                    return Some(Err(Error::ParentDisagreesWithChild()));
78                                }
79                            }
80                            Ok(Block {
81                                node: Node::Leaf(_),
82                                ..
83                            }) => {
84                                return Some(Err(Error::LeafCannotBeParent()));
85                            }
86                            Err(Error::BlockIndexOutOfBounds(_)) => {
87                                return Some(Err(Error::ReferenceToUnknownParent()));
88                            }
89                            Err(e) => return Some(Err(e)),
90                        }
91                    } else if !self.already_queued.contains(&index) {
92                        return Some(Err(Error::ReferenceToUnknownParent()));
93                    }
94                }
95                None => {
96                    if item.index.0 != 0 {
97                        return Some(Err(Error::UnexpectedParentlessNode()));
98                    }
99                }
100            }
101
102            match block.node {
103                Node::Leaf(..) => {
104                    if block.metadata.dirty {
105                        return Some(Err(Error::DirtyLeaf(item.index)));
106                    }
107                    return Some(Ok((item.index, block)));
108                }
109                Node::Internal(ref node) => {
110                    if item.visited {
111                        return Some(Ok((item.index, block)));
112                    }
113
114                    if node.left == node.right
115                        || self.already_queued.contains(&node.left)
116                        || self.already_queued.contains(&node.right)
117                    {
118                        return Some(Err(Error::InvalidChildren()));
119                    }
120
121                    if self.already_queued.contains(&item.index) {
122                        return Some(Err(Error::CycleFound()));
123                    }
124                    self.already_queued.insert(item.index);
125
126                    self.stack.push(LeftChildFirstIteratorItem {
127                        visited: true,
128                        index: item.index,
129                    });
130                    self.stack.push(LeftChildFirstIteratorItem {
131                        visited: false,
132                        index: node.right,
133                    });
134                    self.stack.push(LeftChildFirstIteratorItem {
135                        visited: false,
136                        index: node.left,
137                    });
138                }
139            }
140        }
141    }
142}
143
144pub struct ParentFirstIterator<'a> {
145    blob: &'a [u8],
146    deque: VecDeque<TreeIndex>,
147    already_queued: HashSet<TreeIndex>,
148}
149
150impl<'a> ParentFirstIterator<'a> {
151    pub fn new(blob: &'a [u8], from_index: Option<TreeIndex>) -> Self {
152        let mut deque = VecDeque::new();
153        let from_index = from_index.unwrap_or(TreeIndex(0));
154        if blob.len() / BLOCK_SIZE > 0 {
155            deque.push_back(from_index);
156        }
157
158        Self {
159            blob,
160            deque,
161            already_queued: HashSet::new(),
162        }
163    }
164}
165
166impl Iterator for ParentFirstIterator<'_> {
167    type Item = Result<(TreeIndex, Block), Error>;
168
169    fn next(&mut self) -> Option<Self::Item> {
170        // left sibling first, parents before children
171
172        let index = self.deque.pop_front()?;
173        let block = match try_get_block(self.blob, index) {
174            Ok(block) => block,
175            Err(e) => return Some(Err(e)),
176        };
177
178        if let Node::Internal(ref node) = block.node {
179            if self.already_queued.contains(&index) {
180                return Some(Err(Error::CycleFound()));
181            }
182            self.already_queued.insert(index);
183
184            self.deque.push_back(node.left);
185            self.deque.push_back(node.right);
186        }
187
188        Some(Ok((index, block)))
189    }
190}
191
192pub struct BreadthFirstIterator<'a> {
193    blob: &'a [u8],
194    deque: VecDeque<TreeIndex>,
195    already_queued: HashSet<TreeIndex>,
196}
197
198impl<'a> BreadthFirstIterator<'a> {
199    #[allow(unused)]
200    pub fn new(blob: &'a [u8], from_index: Option<TreeIndex>) -> Self {
201        let mut deque = VecDeque::new();
202        let from_index = from_index.unwrap_or(TreeIndex(0));
203        if blob.len() / BLOCK_SIZE > 0 {
204            deque.push_back(from_index);
205        }
206
207        Self {
208            blob,
209            deque,
210            already_queued: HashSet::new(),
211        }
212    }
213}
214
215impl Iterator for BreadthFirstIterator<'_> {
216    type Item = Result<(TreeIndex, Block), Error>;
217
218    fn next(&mut self) -> Option<Self::Item> {
219        // left sibling first, parent depth before child depth
220
221        loop {
222            let index = self.deque.pop_front()?;
223            let block = match try_get_block(self.blob, index) {
224                Ok(block) => block,
225                Err(e) => return Some(Err(e)),
226            };
227
228            match block.node {
229                Node::Leaf(..) => return Some(Ok((index, block))),
230                Node::Internal(node) => {
231                    if self.already_queued.contains(&index) {
232                        return Some(Err(Error::CycleFound()));
233                    }
234                    self.already_queued.insert(index);
235
236                    self.deque.push_back(node.left);
237                    self.deque.push_back(node.right);
238                }
239            }
240        }
241    }
242}
243
244#[cfg(test)]
245mod tests {
246    use super::*;
247    use crate::merkle::test_util::open_dot;
248    use crate::merkle::test_util::traversal_blob;
249    use crate::{Hash, MerkleBlob, NodeType};
250    use expect_test::{Expect, expect};
251    use rstest::rstest;
252
253    fn iterator_test_reference(index: TreeIndex, block: &Block) -> (u32, NodeType, i64, i64, Hash) {
254        match block.node {
255            Node::Leaf(leaf) => (
256                index.0,
257                block.metadata.node_type,
258                leaf.key.0,
259                leaf.value.0,
260                block.node.hash(),
261            ),
262            Node::Internal(internal) => (
263                index.0,
264                block.metadata.node_type,
265                internal.left.0 as i64,
266                internal.right.0 as i64,
267                block.node.hash(),
268            ),
269        }
270    }
271
272    #[rstest]
273    // expect-test is adding them back
274    #[allow(clippy::needless_raw_string_hashes)]
275    #[case::left_child_first(
276        "left child first",
277        LeftChildFirstIterator::new,
278        Some(TreeIndex(0)),
279        expect![[r#"
280            [
281                (
282                    1,
283                    Leaf,
284                    2315169217770759719,
285                    3472611983179986487,
286                    Hash(
287                        0f980325ebe9426fa295f3f69cc38ef8fe6ce8f3b9f083556c0f927e67e56651,
288                    ),
289                ),
290                (
291                    3,
292                    Leaf,
293                    103,
294                    204,
295                    Hash(
296                        2d47301cff01acc863faa5f57e8fbc632114f1dc764772852ed0c29c0f248bd3,
297                    ),
298                ),
299                (
300                    5,
301                    Leaf,
302                    307,
303                    404,
304                    Hash(
305                        97148f80dd9289a1b67527c045fd47662d575ccdb594701a56c2255ac84f6113,
306                    ),
307                ),
308                (
309                    6,
310                    Internal,
311                    3,
312                    5,
313                    Hash(
314                        b946284149e4f4a0e767ef2feb397533fb112bf4d99c887348cec4438e38c1ce,
315                    ),
316                ),
317                (
318                    4,
319                    Internal,
320                    1,
321                    6,
322                    Hash(
323                        547b5bd537270427e570df6e43dda7c4ef23e6c3bec72cf19d912c3fe864f549,
324                    ),
325                ),
326                (
327                    2,
328                    Leaf,
329                    283686952306183,
330                    1157726452361532951,
331                    Hash(
332                        d8ddfc94e7201527a6a93ee04aed8c5c122ac38af6dbf6e5f1caefba2597230d,
333                    ),
334                ),
335                (
336                    0,
337                    Internal,
338                    4,
339                    2,
340                    Hash(
341                        cc7f12227cc5d96a631963804544872d67aef8b3a86ef9fbc798f7c5dfdbac2b,
342                    ),
343                ),
344            ]
345        "#]],
346    )]
347    #[allow(clippy::needless_raw_string_hashes)]
348    #[case::left_child_first(
349        "left child first - from non-root internal",
350        LeftChildFirstIterator::new,
351        Some(TreeIndex(4)),
352        expect![[r#"
353            [
354                (
355                    1,
356                    Leaf,
357                    2315169217770759719,
358                    3472611983179986487,
359                    Hash(
360                        0f980325ebe9426fa295f3f69cc38ef8fe6ce8f3b9f083556c0f927e67e56651,
361                    ),
362                ),
363                (
364                    3,
365                    Leaf,
366                    103,
367                    204,
368                    Hash(
369                        2d47301cff01acc863faa5f57e8fbc632114f1dc764772852ed0c29c0f248bd3,
370                    ),
371                ),
372                (
373                    5,
374                    Leaf,
375                    307,
376                    404,
377                    Hash(
378                        97148f80dd9289a1b67527c045fd47662d575ccdb594701a56c2255ac84f6113,
379                    ),
380                ),
381                (
382                    6,
383                    Internal,
384                    3,
385                    5,
386                    Hash(
387                        b946284149e4f4a0e767ef2feb397533fb112bf4d99c887348cec4438e38c1ce,
388                    ),
389                ),
390                (
391                    4,
392                    Internal,
393                    1,
394                    6,
395                    Hash(
396                        547b5bd537270427e570df6e43dda7c4ef23e6c3bec72cf19d912c3fe864f549,
397                    ),
398                ),
399            ]
400        "#]],
401    )]
402    #[allow(clippy::needless_raw_string_hashes)]
403    #[case::left_child_first(
404        "left child first - from non-root leaf",
405        LeftChildFirstIterator::new,
406        Some(TreeIndex(3)),
407        expect![[r#"
408            [
409                (
410                    3,
411                    Leaf,
412                    103,
413                    204,
414                    Hash(
415                        2d47301cff01acc863faa5f57e8fbc632114f1dc764772852ed0c29c0f248bd3,
416                    ),
417                ),
418            ]
419        "#]],
420    )]
421    // expect-test is adding them back
422    #[allow(clippy::needless_raw_string_hashes)]
423    #[case::parent_first(
424        "parent first",
425        ParentFirstIterator::new,
426        Some(TreeIndex(0)),
427        expect![[r#"
428            [
429                (
430                    0,
431                    Internal,
432                    4,
433                    2,
434                    Hash(
435                        cc7f12227cc5d96a631963804544872d67aef8b3a86ef9fbc798f7c5dfdbac2b,
436                    ),
437                ),
438                (
439                    4,
440                    Internal,
441                    1,
442                    6,
443                    Hash(
444                        547b5bd537270427e570df6e43dda7c4ef23e6c3bec72cf19d912c3fe864f549,
445                    ),
446                ),
447                (
448                    2,
449                    Leaf,
450                    283686952306183,
451                    1157726452361532951,
452                    Hash(
453                        d8ddfc94e7201527a6a93ee04aed8c5c122ac38af6dbf6e5f1caefba2597230d,
454                    ),
455                ),
456                (
457                    1,
458                    Leaf,
459                    2315169217770759719,
460                    3472611983179986487,
461                    Hash(
462                        0f980325ebe9426fa295f3f69cc38ef8fe6ce8f3b9f083556c0f927e67e56651,
463                    ),
464                ),
465                (
466                    6,
467                    Internal,
468                    3,
469                    5,
470                    Hash(
471                        b946284149e4f4a0e767ef2feb397533fb112bf4d99c887348cec4438e38c1ce,
472                    ),
473                ),
474                (
475                    3,
476                    Leaf,
477                    103,
478                    204,
479                    Hash(
480                        2d47301cff01acc863faa5f57e8fbc632114f1dc764772852ed0c29c0f248bd3,
481                    ),
482                ),
483                (
484                    5,
485                    Leaf,
486                    307,
487                    404,
488                    Hash(
489                        97148f80dd9289a1b67527c045fd47662d575ccdb594701a56c2255ac84f6113,
490                    ),
491                ),
492            ]
493        "#]])]
494    // expect-test is adding them back
495    #[allow(clippy::needless_raw_string_hashes)]
496    #[case::breadth_first(
497        "breadth first",
498        BreadthFirstIterator::new,
499        Some(TreeIndex(0)),
500        expect![[r#"
501            [
502                (
503                    2,
504                    Leaf,
505                    283686952306183,
506                    1157726452361532951,
507                    Hash(
508                        d8ddfc94e7201527a6a93ee04aed8c5c122ac38af6dbf6e5f1caefba2597230d,
509                    ),
510                ),
511                (
512                    1,
513                    Leaf,
514                    2315169217770759719,
515                    3472611983179986487,
516                    Hash(
517                        0f980325ebe9426fa295f3f69cc38ef8fe6ce8f3b9f083556c0f927e67e56651,
518                    ),
519                ),
520                (
521                    3,
522                    Leaf,
523                    103,
524                    204,
525                    Hash(
526                        2d47301cff01acc863faa5f57e8fbc632114f1dc764772852ed0c29c0f248bd3,
527                    ),
528                ),
529                (
530                    5,
531                    Leaf,
532                    307,
533                    404,
534                    Hash(
535                        97148f80dd9289a1b67527c045fd47662d575ccdb594701a56c2255ac84f6113,
536                    ),
537                ),
538            ]
539        "#]])]
540    fn test_iterators<'a, F, T>(
541        #[case] note: &str,
542        #[case] iterator_new: F,
543        #[case] from_index: Option<TreeIndex>,
544        #[case] expected: Expect,
545        #[by_ref] traversal_blob: &'a MerkleBlob,
546    ) where
547        F: Fn(&'a [u8], Option<TreeIndex>) -> T,
548        T: Iterator<Item = Result<(TreeIndex, Block), Error>>,
549    {
550        let mut dot_actual = traversal_blob.to_dot().unwrap();
551        dot_actual.set_note(note);
552
553        let mut actual = vec![];
554        {
555            let blob: &[u8] = &traversal_blob.blob;
556            for item in iterator_new(blob, from_index) {
557                let (index, block) = item.unwrap();
558                actual.push(iterator_test_reference(index, &block));
559                dot_actual.push_traversal(index);
560            }
561        }
562
563        traversal_blob.to_dot().unwrap();
564
565        open_dot(&mut dot_actual);
566
567        expected.assert_debug_eq(&actual);
568    }
569}