diff --git a/trie/database.go b/trie/database.go index b0fd787444..74c2f24bef 100644 --- a/trie/database.go +++ b/trie/database.go @@ -214,6 +214,36 @@ func gatherChildren(n node, children *[]common.Hash) { } } +// forChilds invokes the callback for all the tracked children of this node, +// both the implicit ones from inside the node as well as the explicit ones +//from outside the node. +func (n *cachedNode) forChilds(onChild func(hash common.Hash)) { + for child := range n.children { + onChild(child) + } + if _, ok := n.node.(rawNode); !ok { + forGatherChildren(n.node, onChild) + } +} + +// forGatherChildren traverses the node hierarchy of a collapsed storage node and +// invokes the callback for all the hashnode children. +func forGatherChildren(n node, onChild func(hash common.Hash)) { + switch n := n.(type) { + case *rawShortNode: + forGatherChildren(n.Val, onChild) + case rawFullNode: + for i := 0; i < 16; i++ { + forGatherChildren(n[i], onChild) + } + case hashNode: + onChild(common.BytesToHash(n)) + case valueNode, nil: + default: + panic(fmt.Sprintf("unknown node type: %T", n)) + } +} + // simplifyNode traverses the hierarchy of an expanded memory node and discards // all the internal caches, returning a node that only contains the raw data. func simplifyNode(n node) node { @@ -334,11 +364,11 @@ func (db *Database) insert(hash common.Hash, blob []byte, node node) { size: uint16(len(blob)), flushPrev: db.newest, } - for _, child := range entry.childs() { + entry.forChilds(func(child common.Hash) { if c := db.dirties[child]; c != nil { c.parents++ } - } + }) db.dirties[hash] = entry // Update the flush-list endpoints @@ -570,9 +600,9 @@ func (db *Database) dereference(child common.Hash, parent common.Hash) { db.dirties[node.flushNext].flushPrev = node.flushPrev } // Dereference all children and delete the node - for _, hash := range node.childs() { + node.forChilds(func(hash common.Hash) { db.dereference(hash, child) - } + }) delete(db.dirties, child) db.dirtiesSize -= common.StorageSize(common.HashLength + int(node.size)) if node.children != nil {