From d21b527e54caec85f134dd166a926fd57f970116 Mon Sep 17 00:00:00 2001 From: Nick Johnson Date: Thu, 2 Feb 2017 18:17:37 +0000 Subject: [PATCH] trie: Implement differenceIterator --- trie/iterator.go | 140 +++++++++++++++++++++++++++++++++++++----- trie/iterator_test.go | 57 +++++++++++++++++ 2 files changed, 182 insertions(+), 15 deletions(-) diff --git a/trie/iterator.go b/trie/iterator.go index cbfab3dd41..1ff13226e0 100644 --- a/trie/iterator.go +++ b/trie/iterator.go @@ -16,11 +16,13 @@ package trie -import "github.com/ethereum/go-ethereum/common" +import ( + "bytes" + "github.com/ethereum/go-ethereum/common" +) // Iterator is a key-value trie iterator that traverses a Trie. type Iterator struct { - trie *Trie nodeIt NodeIterator Key []byte // Current data key on which the iterator is positioned on @@ -30,12 +32,19 @@ type Iterator struct { // NewIterator creates a new key-value iterator. func NewIterator(trie *Trie) *Iterator { return &Iterator{ - trie: trie, nodeIt: NewNodeIterator(trie), Key: nil, } } +// FromNodeIterator creates a new key-value iterator from a node iterator +func FromNodeIterator(it NodeIterator) *Iterator { + return &Iterator{ + nodeIt: it, + Key: nil, + } +} + // Next moves the iterator forward one key-value entry. func (it *Iterator) Next() bool { for it.nodeIt.Next(true) { @@ -179,6 +188,7 @@ func (it *nodeIterator) step(children bool) error { } // Continue iteration to the next child +outer: for { if len(it.stack) == 0 { it.trie = nil @@ -191,24 +201,28 @@ func (it *nodeIterator) step(children bool) error { } if node, ok := parent.node.(*fullNode); ok { // Full node, iterate over children - parent.child++ - if parent.child < len(node.Children) { - it.stack = append(it.stack, &nodeIteratorState{ - hash: common.BytesToHash(node.flags.hash), - node: node.Children[parent.child], - parent: ancestor, - child: -1, - pathlen: len(it.path), - }) - it.path = append(it.path, byte(parent.child)) - break + for parent.child++; parent.child < len(node.Children); parent.child++ { + child := node.Children[parent.child] + if child != nil { + hash, _ := child.cache() + it.stack = append(it.stack, &nodeIteratorState{ + hash: common.BytesToHash(hash), + node: child, + parent: ancestor, + child: -1, + pathlen: len(it.path), + }) + it.path = append(it.path, byte(parent.child)) + break outer + } } } else if node, ok := parent.node.(*shortNode); ok { // Short node, return the pointer singleton child if parent.child < 0 { parent.child++ + hash, _ := node.Val.cache() it.stack = append(it.stack, &nodeIteratorState{ - hash: common.BytesToHash(node.flags.hash), + hash: common.BytesToHash(hash), node: node.Val, parent: ancestor, child: -1, @@ -244,3 +258,99 @@ func (it *nodeIterator) step(children bool) error { } return nil } + +type differenceIterator struct { + a, b NodeIterator + eof bool + count int // Number of nodes scanned on either trie +} + +// NewDifferenceIterator constructs a NodeIterator that iterates over elements in b that +// are not in a. +func NewDifferenceIterator(a, b NodeIterator) (NodeIterator, *int) { + a.Next(true) + it := &differenceIterator{ + a: a, + b: b, + eof: false, + count: 0, + } + return it, &it.count +} + +func (it *differenceIterator) Hash() common.Hash { + return it.b.Hash() +} + +func (it *differenceIterator) Parent() common.Hash { + return it.b.Parent() +} + +func (it *differenceIterator) Leaf() bool { + return it.b.Leaf() +} + +func (it *differenceIterator) LeafBlob() []byte { + return it.b.LeafBlob() +} + +func (it *differenceIterator) Path() []byte { + return it.b.Path() +} + +func (it *differenceIterator) Next(bool) bool { + // Invariants: + // - We always advance at least one element in b. + // - At the start of this function, a's path is lexically greater than b's. + if !it.b.Next(true) { + return false + } + it.count += 1 + + if it.eof { + // a has reached eof, so we just return all elements from b + return true + } + + for { + apath, bpath := it.a.Path(), it.b.Path() + cmp := bytes.Compare(apath, bpath) + if cmp < 0 { + // b jumped past a; advance a + if !it.a.Next(true) { + it.eof = true + return true + } + it.count += 1 + } else if cmp > 0 { + // b is before a + return true + } else { + if it.a.Hash() != it.b.Hash() || it.a.Leaf() != it.b.Leaf() { + // Keys are identical, but hashes or leaf status differs + return true + } + if it.a.Leaf() && it.b.Leaf() && !bytes.Equal(it.a.LeafBlob(), it.b.LeafBlob()) { + // Both are leaf nodes, but with different values + return true + } + + // a and b are identical; skip this whole subtree if the nodes have hashes + if !it.b.Next(it.a.Hash() == common.Hash{}) { + return false + } + if !it.a.Next(it.a.Hash() == common.Hash{}) { + it.eof = true + return true + } + it.count += 2 + } + } +} + +func (it *differenceIterator) Error() error { + if err := it.a.Error(); err != nil { + return err + } + return it.b.Error() +} diff --git a/trie/iterator_test.go b/trie/iterator_test.go index adfc30d460..51ad065cb7 100644 --- a/trie/iterator_test.go +++ b/trie/iterator_test.go @@ -116,3 +116,60 @@ func TestNodeIteratorCoverage(t *testing.T) { } } } + +func TestDifferenceIterator(t *testing.T) { + triea := newEmpty() + valsa := []struct{ k, v string }{ + {"bar", "b"}, + {"barb", "ba"}, + {"bars", "bb"}, + {"bard", "bc"}, + {"fab", "z"}, + {"foo", "a"}, + {"food", "ab"}, + {"foos", "aa"}, + } + for _, val := range valsa { + triea.Update([]byte(val.k), []byte(val.v)) + } + triea.Commit() + + trieb := newEmpty() + valsb := []struct{ k, v string }{ + {"aardvark", "c"}, + {"bar", "b"}, + {"barb", "bd"}, + {"bars", "be"}, + {"fab", "z"}, + {"foo", "a"}, + {"foos", "aa"}, + {"food", "ab"}, + {"jars", "d"}, + } + for _, val := range valsb { + trieb.Update([]byte(val.k), []byte(val.v)) + } + trieb.Commit() + + found := make(map[string]string) + di, _ := NewDifferenceIterator(NewNodeIterator(triea), NewNodeIterator(trieb)) + it := FromNodeIterator(di) + for it.Next() { + found[string(it.Key)] = string(it.Value) + } + + all := []struct{ k, v string }{ + {"aardvark", "c"}, + {"barb", "bd"}, + {"bars", "be"}, + {"jars", "d"}, + } + for _, item := range all { + if found[item.k] != item.v { + t.Errorf("iterator value mismatch for %s: got %q want %q", item.k, found[item.k], item.v) + } + } + if len(found) != len(all) { + t.Errorf("iterator count mismatch: got %d values, want %d", len(found), len(all)) + } +}