diff --git a/trie/stacktrie.go b/trie/stacktrie.go index 0d65ee75e0..5c4cd58453 100644 --- a/trie/stacktrie.go +++ b/trie/stacktrie.go @@ -17,11 +17,7 @@ package trie import ( - "bufio" - "bytes" - "encoding/gob" "errors" - "io" "sync" "github.com/ethereum/go-ethereum/common" @@ -29,186 +25,50 @@ import ( "github.com/ethereum/go-ethereum/log" ) -var ErrCommitDisabled = errors.New("no database for committing") - -var stPool = sync.Pool{ - New: func() interface{} { - return NewStackTrie(nil) - }, -} +var ( + ErrCommitDisabled = errors.New("no database for committing") + stPool = sync.Pool{New: func() any { return new(stNode) }} + _ = types.TrieHasher((*StackTrie)(nil)) +) // NodeWriteFunc is used to provide all information of a dirty node for committing // so that callers can flush nodes into database with desired scheme. type NodeWriteFunc = func(owner common.Hash, path []byte, hash common.Hash, blob []byte) -func stackTrieFromPool(writeFn NodeWriteFunc, owner common.Hash) *StackTrie { - st := stPool.Get().(*StackTrie) - st.owner = owner - st.writeFn = writeFn - return st -} - -func returnToPool(st *StackTrie) { - st.Reset() - stPool.Put(st) -} - // StackTrie is a trie implementation that expects keys to be inserted // in order. Once it determines that a subtree will no longer be inserted // into, it will hash it and free up the memory it uses. type StackTrie struct { - owner common.Hash // the owner of the trie - nodeType uint8 // node type (as in branch, ext, leaf) - val []byte // value contained by this node if it's a leaf - key []byte // key chunk covered by this (leaf|ext) node - children [16]*StackTrie // list of children (for branch and exts) - writeFn NodeWriteFunc // function for committing nodes, can be nil + owner common.Hash // the owner of the trie + writeFn NodeWriteFunc // function for committing nodes, can be nil + root *stNode + h *hasher } // NewStackTrie allocates and initializes an empty trie. func NewStackTrie(writeFn NodeWriteFunc) *StackTrie { return &StackTrie{ - nodeType: emptyNode, - writeFn: writeFn, + writeFn: writeFn, + root: stPool.Get().(*stNode), + h: newHasher(false), } } // NewStackTrieWithOwner allocates and initializes an empty trie, but with // the additional owner field. func NewStackTrieWithOwner(writeFn NodeWriteFunc, owner common.Hash) *StackTrie { - return &StackTrie{ - owner: owner, - nodeType: emptyNode, - writeFn: writeFn, - } + stack := NewStackTrie(writeFn) + stack.owner = owner + return stack } -// NewFromBinary initialises a serialized stacktrie with the given db. -func NewFromBinary(data []byte, writeFn NodeWriteFunc) (*StackTrie, error) { - var st StackTrie - if err := st.UnmarshalBinary(data); err != nil { - return nil, err - } - // If a database is used, we need to recursively add it to every child - if writeFn != nil { - st.setWriter(writeFn) - } - return &st, nil -} - -// MarshalBinary implements encoding.BinaryMarshaler -func (st *StackTrie) MarshalBinary() (data []byte, err error) { - var ( - b bytes.Buffer - w = bufio.NewWriter(&b) - ) - if err := gob.NewEncoder(w).Encode(struct { - Owner common.Hash - NodeType uint8 - Val []byte - Key []byte - }{ - st.owner, - st.nodeType, - st.val, - st.key, - }); err != nil { - return nil, err - } - for _, child := range st.children { - if child == nil { - w.WriteByte(0) - continue - } - w.WriteByte(1) - if childData, err := child.MarshalBinary(); err != nil { - return nil, err - } else { - w.Write(childData) - } - } - w.Flush() - return b.Bytes(), nil -} - -// UnmarshalBinary implements encoding.BinaryUnmarshaler -func (st *StackTrie) UnmarshalBinary(data []byte) error { - r := bytes.NewReader(data) - return st.unmarshalBinary(r) -} - -func (st *StackTrie) unmarshalBinary(r io.Reader) error { - var dec struct { - Owner common.Hash - NodeType uint8 - Val []byte - Key []byte - } - if err := gob.NewDecoder(r).Decode(&dec); err != nil { - return err - } - st.owner = dec.Owner - st.nodeType = dec.NodeType - st.val = dec.Val - st.key = dec.Key - - var hasChild = make([]byte, 1) - for i := range st.children { - if _, err := r.Read(hasChild); err != nil { - return err - } else if hasChild[0] == 0 { - continue - } - var child StackTrie - if err := child.unmarshalBinary(r); err != nil { - return err - } - st.children[i] = &child - } - return nil -} - -func (st *StackTrie) setWriter(writeFn NodeWriteFunc) { - st.writeFn = writeFn - for _, child := range st.children { - if child != nil { - child.setWriter(writeFn) - } - } -} - -func newLeaf(owner common.Hash, key, val []byte, writeFn NodeWriteFunc) *StackTrie { - st := stackTrieFromPool(writeFn, owner) - st.nodeType = leafNode - st.key = append(st.key, key...) - st.val = val - return st -} - -func newExt(owner common.Hash, key []byte, child *StackTrie, writeFn NodeWriteFunc) *StackTrie { - st := stackTrieFromPool(writeFn, owner) - st.nodeType = extNode - st.key = append(st.key, key...) - st.children[0] = child - return st -} - -// List all values that StackTrie#nodeType can hold -const ( - emptyNode = iota - branchNode - extNode - leafNode - hashedNode -) - // Update inserts a (key, value) pair into the stack trie. -func (st *StackTrie) Update(key, value []byte) error { +func (stack *StackTrie) Update(key, value []byte) error { k := keybytesToHex(key) if len(value) == 0 { panic("deletion not supported") } - st.insert(k[:len(k)-1], value, nil) + stack.insert(stack.root, k[:len(k)-1], value, nil) return nil } @@ -220,21 +80,59 @@ func (st *StackTrie) MustUpdate(key, value []byte) { } } -func (st *StackTrie) Reset() { - st.owner = common.Hash{} - st.writeFn = nil +func (stack *StackTrie) Reset() { + stack.owner = (common.Hash{}) + stack.writeFn = nil + stack.root = stPool.Get().(*stNode) +} + +// stNode represents a node within a StackTrie +type stNode struct { + nodeType uint8 // node type (as in branch, ext, leaf) + val []byte // value contained by this node if it's a leaf + key []byte // key chunk covered by this (leaf|ext) node + children [16]*stNode // list of children (for branch and exts) +} + +func newLeaf(key, val []byte) *stNode { + st := stPool.Get().(*stNode) + st.nodeType = leafNode + st.key = append(st.key, key...) + st.val = val + return st +} + +func newExt(key []byte, child *stNode) *stNode { + st := stPool.Get().(*stNode) + st.nodeType = extNode + st.key = append(st.key, key...) + st.children[0] = child + return st +} + +// List all values that stNode#nodeType can hold +const ( + emptyNode = iota + branchNode + extNode + leafNode + hashedNode +) + +func (st *stNode) Reset() *stNode { st.key = st.key[:0] st.val = nil for i := range st.children { st.children[i] = nil } st.nodeType = emptyNode + return st } // Helper function that, given a full key, determines the index // at which the chunk pointed by st.keyOffset is different from // the same chunk in the full key. -func (st *StackTrie) getDiffIndex(key []byte) int { +func (st *stNode) getDiffIndex(key []byte) int { for idx, nibble := range st.key { if nibble != key[idx] { return idx @@ -245,7 +143,7 @@ func (st *StackTrie) getDiffIndex(key []byte) int { // Helper function to that inserts a (key, value) pair into // the trie. -func (st *StackTrie) insert(key, value []byte, prefix []byte) { +func (stack *StackTrie) insert(st *stNode, key, value []byte, prefix []byte) { switch st.nodeType { case branchNode: /* Branch */ idx := int(key[0]) @@ -254,7 +152,7 @@ func (st *StackTrie) insert(key, value []byte, prefix []byte) { for i := idx - 1; i >= 0; i-- { if st.children[i] != nil { if st.children[i].nodeType != hashedNode { - st.children[i].hash(append(prefix, byte(i))) + stack.hash(st.children[i], append(prefix, byte(i))) } break } @@ -262,9 +160,9 @@ func (st *StackTrie) insert(key, value []byte, prefix []byte) { // Add new child if st.children[idx] == nil { - st.children[idx] = newLeaf(st.owner, key[1:], value, st.writeFn) + st.children[idx] = newLeaf(key[1:], value) } else { - st.children[idx].insert(key[1:], value, append(prefix, key[0])) + stack.insert(st.children[idx], key[1:], value, append(prefix, key[0])) } case extNode: /* Ext */ @@ -279,29 +177,29 @@ func (st *StackTrie) insert(key, value []byte, prefix []byte) { if diffidx == len(st.key) { // Ext key and key segment are identical, recurse into // the child node. - st.children[0].insert(key[diffidx:], value, append(prefix, key[:diffidx]...)) + stack.insert(st.children[0], key[diffidx:], value, append(prefix, key[:diffidx]...)) return } // Save the original part. Depending if the break is // at the extension's last byte or not, create an // intermediate extension or use the extension's child // node directly. - var n *StackTrie + var n *stNode if diffidx < len(st.key)-1 { // Break on the non-last byte, insert an intermediate // extension. The path prefix of the newly-inserted // extension should also contain the different byte. - n = newExt(st.owner, st.key[diffidx+1:], st.children[0], st.writeFn) - n.hash(append(prefix, st.key[:diffidx+1]...)) + n = newExt(st.key[diffidx+1:], st.children[0]) + stack.hash(n, append(prefix, st.key[:diffidx+1]...)) } else { // Break on the last byte, no need to insert // an extension node: reuse the current node. // The path prefix of the original part should // still be same. n = st.children[0] - n.hash(append(prefix, st.key...)) + stack.hash(n, append(prefix, st.key...)) } - var p *StackTrie + var p *stNode if diffidx == 0 { // the break is on the first byte, so // the current node is converted into @@ -313,12 +211,12 @@ func (st *StackTrie) insert(key, value []byte, prefix []byte) { // the common prefix is at least one byte // long, insert a new intermediate branch // node. - st.children[0] = stackTrieFromPool(st.writeFn, st.owner) + st.children[0] = stPool.Get().(*stNode) st.children[0].nodeType = branchNode p = st.children[0] } // Create a leaf for the inserted part - o := newLeaf(st.owner, key[diffidx+1:], value, st.writeFn) + o := newLeaf(key[diffidx+1:], value) // Insert both child leaves where they belong: origIdx := st.key[diffidx] @@ -344,7 +242,7 @@ func (st *StackTrie) insert(key, value []byte, prefix []byte) { // Check if the split occurs at the first nibble of the // chunk. In that case, no prefix extnode is necessary. // Otherwise, create that - var p *StackTrie + var p *stNode if diffidx == 0 { // Convert current leaf into a branch st.nodeType = branchNode @@ -354,7 +252,7 @@ func (st *StackTrie) insert(key, value []byte, prefix []byte) { // Convert current node into an ext, // and insert a child branch node. st.nodeType = extNode - st.children[0] = NewStackTrieWithOwner(st.writeFn, st.owner) + st.children[0] = stPool.Get().(*stNode) st.children[0].nodeType = branchNode p = st.children[0] } @@ -363,11 +261,11 @@ func (st *StackTrie) insert(key, value []byte, prefix []byte) { // value and another containing the new value. The child leaf // is hashed directly in order to free up some memory. origIdx := st.key[diffidx] - p.children[origIdx] = newLeaf(st.owner, st.key[diffidx+1:], st.val, st.writeFn) - p.children[origIdx].hash(append(prefix, st.key[:diffidx+1]...)) + p.children[origIdx] = newLeaf(st.key[diffidx+1:], st.val) + stack.hash(p.children[origIdx], append(prefix, st.key[:diffidx+1]...)) newIdx := key[diffidx] - p.children[newIdx] = newLeaf(st.owner, key[diffidx+1:], value, st.writeFn) + p.children[newIdx] = newLeaf(key[diffidx+1:], value) // Finally, cut off the key part that has been passed // over to the children. @@ -398,14 +296,7 @@ func (st *StackTrie) insert(key, value []byte, prefix []byte) { // - And the 'st.type' will be 'hashedNode' AGAIN // // This method also sets 'st.type' to hashedNode, and clears 'st.key'. -func (st *StackTrie) hash(path []byte) { - h := newHasher(false) - defer returnHasherToPool(h) - - st.hashRec(h, path) -} - -func (st *StackTrie) hashRec(hasher *hasher, path []byte) { +func (stack *StackTrie) hash(st *stNode, path []byte) { // The switch below sets this to the RLP-encoding of this node. var encodedNode []byte @@ -426,7 +317,7 @@ func (st *StackTrie) hashRec(hasher *hasher, path []byte) { nodes.Children[i] = nilValueNode continue } - child.hashRec(hasher, append(path, byte(i))) + stack.hash(child, append(path, byte(i))) if len(child.val) < 32 { nodes.Children[i] = rawNode(child.val) } else { @@ -435,14 +326,14 @@ func (st *StackTrie) hashRec(hasher *hasher, path []byte) { // Release child back to pool. st.children[i] = nil - returnToPool(child) + stPool.Put(child.Reset()) } - nodes.encode(hasher.encbuf) - encodedNode = hasher.encodedBytes() + nodes.encode(stack.h.encbuf) + encodedNode = stack.h.encodedBytes() case extNode: - st.children[0].hashRec(hasher, append(path, st.key...)) + stack.hash(st.children[0], append(path, st.key...)) n := shortNode{Key: hexToCompactInPlace(st.key)} if len(st.children[0].val) < 32 { @@ -451,19 +342,20 @@ func (st *StackTrie) hashRec(hasher *hasher, path []byte) { n.Val = hashNode(st.children[0].val) } - n.encode(hasher.encbuf) - encodedNode = hasher.encodedBytes() + n.encode(stack.h.encbuf) + encodedNode = stack.h.encodedBytes() // Release child back to pool. - returnToPool(st.children[0]) + stPool.Put(st.children[0].Reset()) + st.children[0] = nil case leafNode: st.key = append(st.key, byte(16)) n := shortNode{Key: hexToCompactInPlace(st.key), Val: valueNode(st.val)} - n.encode(hasher.encbuf) - encodedNode = hasher.encodedBytes() + n.encode(stack.h.encbuf) + encodedNode = stack.h.encodedBytes() default: panic("invalid node type") @@ -478,18 +370,16 @@ func (st *StackTrie) hashRec(hasher *hasher, path []byte) { // Write the hash to the 'val'. We allocate a new val here to not mutate // input values - st.val = hasher.hashData(encodedNode) - if st.writeFn != nil { - st.writeFn(st.owner, path, common.BytesToHash(st.val), encodedNode) + st.val = stack.h.hashData(encodedNode) + if stack.writeFn != nil { + stack.writeFn(stack.owner, path, common.BytesToHash(st.val), encodedNode) } } // Hash returns the hash of the current node. -func (st *StackTrie) Hash() (h common.Hash) { - hasher := newHasher(false) - defer returnHasherToPool(hasher) - - st.hashRec(hasher, nil) +func (stack *StackTrie) Hash() (h common.Hash) { + st := stack.root + stack.hash(st, nil) if len(st.val) == 32 { copy(h[:], st.val) return h @@ -497,9 +387,9 @@ func (st *StackTrie) Hash() (h common.Hash) { // If the node's RLP isn't 32 bytes long, the node will not // be hashed, and instead contain the rlp-encoding of the // node. For the top level node, we need to force the hashing. - hasher.sha.Reset() - hasher.sha.Write(st.val) - hasher.sha.Read(h[:]) + stack.h.sha.Reset() + stack.h.sha.Write(st.val) + stack.h.sha.Read(h[:]) return h } @@ -510,14 +400,12 @@ func (st *StackTrie) Hash() (h common.Hash) { // // The associated database is expected, otherwise the whole commit // functionality should be disabled. -func (st *StackTrie) Commit() (h common.Hash, err error) { - if st.writeFn == nil { +func (stack *StackTrie) Commit() (h common.Hash, err error) { + if stack.writeFn == nil { return common.Hash{}, ErrCommitDisabled } - hasher := newHasher(false) - defer returnHasherToPool(hasher) - - st.hashRec(hasher, nil) + st := stack.root + stack.hash(st, nil) if len(st.val) == 32 { copy(h[:], st.val) return h, nil @@ -525,10 +413,95 @@ func (st *StackTrie) Commit() (h common.Hash, err error) { // If the node's RLP isn't 32 bytes long, the node will not // be hashed (and committed), and instead contain the rlp-encoding of the // node. For the top level node, we need to force the hashing+commit. - hasher.sha.Reset() - hasher.sha.Write(st.val) - hasher.sha.Read(h[:]) + stack.h.sha.Reset() + stack.h.sha.Write(st.val) + stack.h.sha.Read(h[:]) - st.writeFn(st.owner, nil, h, st.val) + stack.writeFn(stack.owner, nil, h, st.val) return h, nil } + +//// NewFromBinary initialises a serialized stacktrie with the given db. +//func NewFromBinary(data []byte, writeFn NodeWriteFunc) (*StackTrie, error) { +// var st StackTrie +// if err := st.UnmarshalBinary(data); err != nil { +// return nil, err +// } +// // If a database is used, we need to recursively add it to every child +// if writeFn != nil { +// st.setWriter(writeFn) +// } +// return &st, nil +//} +// +//// MarshalBinary implements encoding.BinaryMarshaler +//func (st *StackTrie) MarshalBinary() (data []byte, err error) { +// var ( +// b bytes.Buffer +// w = bufio.NewWriter(&b) +// ) +// if err := gob.NewEncoder(w).Encode(struct { +// Owner common.Hash +// NodeType uint8 +// Val []byte +// Key []byte +// }{ +// st.owner, +// st.nodeType, +// st.val, +// st.key, +// }); err != nil { +// return nil, err +// } +// for _, child := range st.children { +// if child == nil { +// w.WriteByte(0) +// continue +// } +// w.WriteByte(1) +// if childData, err := child.MarshalBinary(); err != nil { +// return nil, err +// } else { +// w.Write(childData) +// } +// } +// w.Flush() +// return b.Bytes(), nil +//} +// +//// UnmarshalBinary implements encoding.BinaryUnmarshaler +//func (st *StackTrie) UnmarshalBinary(data []byte) error { +// r := bytes.NewReader(data) +// return st.unmarshalBinary(r) +//} +// +//func (st *StackTrie) unmarshalBinary(r io.Reader) error { +// var dec struct { +// Owner common.Hash +// NodeType uint8 +// Val []byte +// Key []byte +// } +// if err := gob.NewDecoder(r).Decode(&dec); err != nil { +// return err +// } +// st.owner = dec.Owner +// st.nodeType = dec.NodeType +// st.val = dec.Val +// st.key = dec.Key +// +// var hasChild = make([]byte, 1) +// for i := range st.children { +// if _, err := r.Read(hasChild); err != nil { +// return err +// } else if hasChild[0] == 0 { +// continue +// } +// var child StackTrie +// if err := child.unmarshalBinary(r); err != nil { +// return err +// } +// st.children[i] = &child +// } +// return nil +//} diff --git a/trie/stacktrie_test.go b/trie/stacktrie_test.go index 6bd0b83e39..cba8201308 100644 --- a/trie/stacktrie_test.go +++ b/trie/stacktrie_test.go @@ -198,12 +198,11 @@ func TestStackTrieInsertAndHash(t *testing.T) { {"000003", "XXXXXXXXXXXXXXXXXXXXXXXXXXXX", "962c0fffdeef7612a4f7bff1950d67e3e81c878e48b9ae45b3b374253b050bd8"}, }, } - st := NewStackTrie(nil) for i, test := range tests { // The StackTrie does not allow Insert(), Hash(), Insert(), ... // so we will create new trie for every sequence length of inserts. for l := 1; l <= len(test); l++ { - st.Reset() + st := NewStackTrie(nil) for j := 0; j < l; j++ { kv := &test[j] if err := st.Update(common.FromHex(kv.K), []byte(kv.V)); err != nil { @@ -380,45 +379,45 @@ func TestStacktrieNotModifyValues(t *testing.T) { // TestStacktrieSerialization tests that the stacktrie works well if we // serialize/unserialize it a lot -func TestStacktrieSerialization(t *testing.T) { - var ( - st = NewStackTrie(nil) - nt = NewEmpty(NewDatabase(rawdb.NewMemoryDatabase(), nil)) - keyB = big.NewInt(1) - keyDelta = big.NewInt(1) - vals [][]byte - keys [][]byte - ) - getValue := func(i int) []byte { - if i%2 == 0 { // large - return crypto.Keccak256(big.NewInt(int64(i)).Bytes()) - } else { //small - return big.NewInt(int64(i)).Bytes() - } - } - for i := 0; i < 10; i++ { - vals = append(vals, getValue(i)) - keys = append(keys, common.BigToHash(keyB).Bytes()) - keyB = keyB.Add(keyB, keyDelta) - keyDelta.Add(keyDelta, common.Big1) - } - for i, k := range keys { - nt.Update(k, common.CopyBytes(vals[i])) - } - - for i, k := range keys { - blob, err := st.MarshalBinary() - if err != nil { - t.Fatal(err) - } - newSt, err := NewFromBinary(blob, nil) - if err != nil { - t.Fatal(err) - } - st = newSt - st.Update(k, common.CopyBytes(vals[i])) - } - if have, want := st.Hash(), nt.Hash(); have != want { - t.Fatalf("have %#x want %#x", have, want) - } -} +//func TestStacktrieSerialization(t *testing.T) { +// var ( +// st = NewStackTrie(nil) +// nt = NewEmpty(NewDatabase(rawdb.NewMemoryDatabase(), nil)) +// keyB = big.NewInt(1) +// keyDelta = big.NewInt(1) +// vals [][]byte +// keys [][]byte +// ) +// getValue := func(i int) []byte { +// if i%2 == 0 { // large +// return crypto.Keccak256(big.NewInt(int64(i)).Bytes()) +// } else { //small +// return big.NewInt(int64(i)).Bytes() +// } +// } +// for i := 0; i < 10; i++ { +// vals = append(vals, getValue(i)) +// keys = append(keys, common.BigToHash(keyB).Bytes()) +// keyB = keyB.Add(keyB, keyDelta) +// keyDelta.Add(keyDelta, common.Big1) +// } +// for i, k := range keys { +// nt.Update(k, common.CopyBytes(vals[i])) +// } +// +// for i, k := range keys { +// blob, err := st.MarshalBinary() +// if err != nil { +// t.Fatal(err) +// } +// newSt, err := NewFromBinary(blob, nil) +// if err != nil { +// t.Fatal(err) +// } +// st = newSt +// st.Update(k, common.CopyBytes(vals[i])) +// } +// if have, want := st.Hash(), nt.Hash(); have != want { +// t.Fatalf("have %#x want %#x", have, want) +// } +//}