Fix a few bugs in the stack trie (#9)

* Fix a few bugs

* Fix the remaining bugs

* Move the stack code to its own file

* More PR grooming
This commit is contained in:
Guillaume Ballet 2020-03-04 13:25:11 +01:00 committed by Martin Holst Swende
parent 7d5fd3aab0
commit e64f57fe45
No known key found for this signature in database
GPG key ID: 683B438C05A5DDF0
3 changed files with 308 additions and 174 deletions

View file

@ -101,6 +101,48 @@ func TestTrieGenerationAppendonly(t *testing.T) {
}
}
func TestMultipleStackTrieInsertion(t *testing.T) {
// Get a fairly large trie
// Create a custom account factory to recreate the same addresses
makeAccounts := func(num int) map[common.Hash][]byte {
accounts := make(map[common.Hash][]byte)
for i := 0; i < num; i++ {
h := common.Hash{}
binary.BigEndian.PutUint64(h[:], uint64(i+1))
accounts[h] = randomAccountWithSmall()
}
return accounts
}
// Build up a large stack of snapshots
base := &diskLayer{
diskdb: rawdb.NewMemoryDatabase(),
root: common.HexToHash("0x01"),
cache: fastcache.New(1024 * 500),
}
snaps := &Tree{
layers: map[common.Hash]snapshot{
base.root: base,
},
}
// 4K accounts
snaps.Update(common.HexToHash("0x02"), common.HexToHash("0x01"), makeAccounts(4000), nil)
head := snaps.Snapshot(common.HexToHash("0x02"))
// Call it once to make it create the lists before test starts
head.(*diffLayer).AccountIterator(common.HexToHash("0x00"))
var got1 common.Hash
it := head.(*diffLayer).AccountIterator(common.HexToHash("0x00"))
got1 = generateTrie(it, PruneGenerate)
var got2 common.Hash
it = head.(*diffLayer).AccountIterator(common.HexToHash("0x00"))
got2 = generateTrie(it, StackGenerate)
if got2 != got1 {
t.Fatalf("Error: got %x exp %x", got2, got1)
}
}
// BenchmarkTrieGeneration/4K/standard-8 127 9429425 ns/op 6188077 B/op 58026 allocs/op
// BenchmarkTrieGeneration/4K/pruning-8 72 16544534 ns/op 6617322 B/op 55016 allocs/op
// BenchmarkTrieGeneration/4K/stack-8 159 6452936 ns/op 6308393 B/op 12022 allocs/op
@ -145,7 +187,7 @@ func BenchmarkTrieGeneration(b *testing.B) {
got = generateTrie(it, StdGenerate)
}
b.StopTimer()
if exp := common.HexToHash("fecc4e1fce05c888c8acc8baa2d7677a531714668b7a09b5ede6e3e110be266b"); got != exp{
if exp := common.HexToHash("fecc4e1fce05c888c8acc8baa2d7677a531714668b7a09b5ede6e3e110be266b"); got != exp {
b.Fatalf("Error: got %x exp %x", got, exp)
}
})
@ -158,7 +200,7 @@ func BenchmarkTrieGeneration(b *testing.B) {
got = generateTrie(it, PruneGenerate)
}
b.StopTimer()
if exp := common.HexToHash("fecc4e1fce05c888c8acc8baa2d7677a531714668b7a09b5ede6e3e110be266b"); got != exp{
if exp := common.HexToHash("fecc4e1fce05c888c8acc8baa2d7677a531714668b7a09b5ede6e3e110be266b"); got != exp {
b.Fatalf("Error: got %x exp %x", got, exp)
}
@ -172,7 +214,7 @@ func BenchmarkTrieGeneration(b *testing.B) {
got = generateTrie(it, StackGenerate)
}
b.StopTimer()
if exp := common.HexToHash("fecc4e1fce05c888c8acc8baa2d7677a531714668b7a09b5ede6e3e110be266b"); got != exp{
if exp := common.HexToHash("fecc4e1fce05c888c8acc8baa2d7677a531714668b7a09b5ede6e3e110be266b"); got != exp {
b.Fatalf("Error: got %x exp %x", got, exp)
}

View file

@ -17,7 +17,6 @@
package trie
import (
"bytes"
"fmt"
"github.com/ethereum/go-ethereum/common"
@ -124,173 +123,3 @@ func (t *HashTrie) Hash() common.Hash {
t.root = cached
return common.BytesToHash(hashed.(hashNode))
}
type StackTrieItem struct {
ext shortNode
branch fullNode
depth int
useBranch bool
keyUntilHere []byte
}
type StackTrie struct {
stack []StackTrieItem
top int
hasher *hasher
}
func NewStackTrie() *StackTrie {
return &StackTrie{
top: -1,
stack: []StackTrieItem{
StackTrieItem{},
},
hasher: newHasher(false),
}
}
func (st *StackTrie) TryUpdate(key, value []byte) error {
k := keybytesToHex(key)
if len(value) == 0 {
panic("deletion not supported")
}
st.insert(&st.stack[0].ext, nil, k, valueNode(value))
return nil
}
func (st *StackTrie) insert(n node, prefix, key []byte, value node) node {
// Special case: the trie is empty
if st.top == -1 {
st.top = 0
st.stack[st.top].depth = 0
st.stack[st.top].ext.Key = key
st.stack[st.top].ext.Val, _ = st.hasher.hash(value, false)
st.stack[st.top].keyUntilHere = []byte("")
return &st.stack[st.top].ext
}
// Use the prefix key to find the stack level in which the code needs to
// be inserted.
level := -1
for index := st.top; index >= 0; index-- {
level = index
if bytes.Equal(st.stack[level].keyUntilHere, key[:len(st.stack[level].keyUntilHere)]) {
// Found the common denominator, stop the search
break
}
}
// Already hash the value, which it will be anyway
hv, _ := st.hasher.hash(value, false)
// The difference happens at this level, find out where
// exactly. The extension part of the fullnode part?
extStart := len(st.stack[level].keyUntilHere)
extEnd := extStart + len(st.stack[level].ext.Key)
if bytes.Equal(st.stack[level].ext.Key, key[extStart:extEnd]) {
// The extension and the key are identical on the length of
// the extension, so st.stack[level].ext.Val should be a fullNode and
// the difference should be found there. Panic if this is
// not the case.
fn := st.stack[level].ext.Val.(*fullNode)
// The correct entry is the only one that isn't nil
for i := 15; i >= 0; i-- {
if fn.Children[i] != nil {
switch fn.Children[i].(type) {
// Only hash entries that are not already hashed
case *fullNode, *shortNode:
fn.Children[i], _ = st.hasher.hash(fn.Children[i], false)
st.top = level
default:
}
break
}
}
// That fullNode should have at most one non-hashNode child,
// hash it because no more nodes will be inserted in it.
if len(st.stack) == st.top+1 {
st.stack = append(st.stack, StackTrieItem{})
}
st.top++
keyUntilHere := len(st.stack[level].keyUntilHere) + len(st.stack[level].ext.Key) + 1
st.stack[level].branch.Children[key[keyUntilHere]] = &st.stack[st.top].ext
st.stack[st.top].keyUntilHere = key[:keyUntilHere]
st.stack[st.top].ext.Key = key[keyUntilHere:]
st.stack[st.top].ext.Val = hv
st.stack[st.top].ext.flags = nodeFlag{dirty: true}
st.stack[st.top].depth = st.stack[level].depth + 1
} else {
// extension keys differ, need to create a split and
// hash the former node.
whereitdiffers := 0
offset := len(st.stack[level].keyUntilHere)
for i := range st.stack[level].ext.Key {
if key[offset+i] != st.stack[level].ext.Key[i] {
whereitdiffers = i
break
}
}
// Start by hashing the node right after the extension,
// to free some space.
var hn node
switch st.stack[level].ext.Val.(type) {
case *fullNode:
h, _ := st.hasher.hash(st.stack[level].ext.Val, false)
hn = h.(hashNode)
case hashNode, valueNode:
hn = st.stack[level].ext.Val
default:
panic("Encountered unexpected node type")
}
// Allocate the next full node, it's going to be
// reused several times.
if len(st.stack) == st.top+1 {
st.stack = append(st.stack, StackTrieItem{})
}
st.top++
// Store the partially-hashed old node in the newly allocated
// slot, in order to finish the hashing.
slot := st.stack[level].ext.Key[whereitdiffers]
st.stack[st.top].ext.Key = st.stack[level].ext.Key[whereitdiffers+1:]
st.stack[st.top].ext.Val = hn
st.stack[st.top].ext.flags = nodeFlag{dirty: true}
// Hasher directement la branche si l'ext est vide
h, _ := st.hasher.hash(&st.stack[st.top].ext, false)
st.stack[level].branch.Children[slot] = h.(hashNode)
st.stack[level].ext.Val = &st.stack[level].branch
st.stack[level].ext.Key = st.stack[level].ext.Key[:whereitdiffers]
// Now use the newly allocated+hashed stack st.stack[level] to store
// the rest of the inserted (key, value) pair.
slot = key[whereitdiffers+len(st.stack[level].keyUntilHere)]
st.stack[level].branch.Children[slot] = &st.stack[st.top].ext
st.stack[st.top].ext.Key = key[whereitdiffers+len(st.stack[level].keyUntilHere)+1:]
st.stack[st.top].ext.Val = hv
st.stack[st.top].keyUntilHere = key[:whereitdiffers+len(st.stack[level].keyUntilHere)+1]
st.stack[st.top].depth = st.stack[level].depth + 1
}
// if ext.length == 0, directly return the full node.
if len(st.stack[0].ext.Key) == 0 {
return &st.stack[0].branch
}
return &st.stack[0].ext
}
func (st *StackTrie) Hash() common.Hash {
if st.top == -1 {
return emptyRoot
}
h, _ := st.hasher.hash(&st.stack[0].ext, false)
return common.BytesToHash(h.(hashNode))
}

263
trie/stacktrie.go Normal file
View file

@ -0,0 +1,263 @@
// Copyright 2020 The go-ethereum Authors
// This file is part of the go-ethereum library.
//
// The go-ethereum library is free software: you can redistribute it and/or modify
// it under the terms of the GNU Lesser General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// The go-ethereum library is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Lesser General Public License for more details.
//
// You should have received a copy of the GNU Lesser General Public License
// along with the go-ethereum library. If not, see <http://www.gnu.org/licenses/>.
package trie
import (
"bytes"
//"fmt"
"github.com/ethereum/go-ethereum/common"
)
// StackTrieItem represents an (extension, fullnode) tuple to be stored
// in a "stack" in order to be reused multiple times so as to save many
// allocations.
type StackTrieItem struct {
ext shortNode
branch fullNode
depth int
useBranch bool
keyUntilHere []byte
}
// StackTrie is a "stack" of (extension, fullnode) tuples that are
// used to calculate the hash of a trie. The core idea is that at
// any time, only one branch is expanded and the rest is hashed as
// soon as it is determined it is no longer needed.
type StackTrie struct {
stack []StackTrieItem
top int
hasher *hasher
}
// NewStackTrie builds a new stack trie. The whole stack space is
// pre-allocated so as to save reallocations down the road.
func NewStackTrie() *StackTrie {
return &StackTrie{
top: -1,
stack: make([]StackTrieItem, 65),
hasher: newHasher(false),
}
}
func (st *StackTrie) TryUpdate(key, value []byte) error {
k := keybytesToHex(key)
if len(value) == 0 {
panic("deletion not supported")
}
st.insert(&st.stack[0].ext, nil, k, valueNode(value))
//fmt.Println("trie=", &st.stack[0].ext)
return nil
}
// alloc prepares the next stage in the stack for reuse.
func (st *StackTrie) alloc() {
for i := 0; i < 16; i++ {
st.stack[st.top+1].branch.Children[i] = nil
}
st.top++
}
func (st *StackTrie) insert(n node, prefix, key []byte, value node) node {
// Special case: the trie is empty
if st.top == -1 {
st.top = 0
st.stack[st.top].depth = 0
st.stack[st.top].ext.Key = key
st.stack[st.top].ext.Val, _ = st.hasher.hash(value, false)
st.stack[st.top].keyUntilHere = []byte("")
return &st.stack[st.top].ext
}
// Use the prefix key to find the stack level in which the code needs to
// be inserted.
level := -1
for index := st.top; index >= 0; index-- {
level = index
if bytes.Equal(st.stack[level].keyUntilHere, key[:len(st.stack[level].keyUntilHere)]) {
// Found the common denominator, stop the search
break
}
}
// Already hash the value, which it will be anyway
hv, _ := st.hasher.hash(value, false)
// The difference happens at this level, find out where
// exactly. The extension part of the fullnode part?
extStart := len(st.stack[level].keyUntilHere)
extEnd := extStart + len(st.stack[level].ext.Key)
if bytes.Equal(st.stack[level].ext.Key, key[extStart:extEnd]) {
// The extension and the key are identical on the length of
// the extension, so st.stack[level].ext.Val should point to
// st.stack[level].branch, and the difference should be foud
// there.
var fn *fullNode
fn = &st.stack[level].branch
// The correct entry is the only one that isn't nil
for i := 15; i >= 0; i-- {
if fn.Children[i] != nil {
switch fn.Children[i].(type) {
// Only hash entries that are not already hashed
case *fullNode, *shortNode:
fn.Children[i], _ = st.hasher.hash(fn.Children[i], false)
st.top = level
default:
}
break
}
}
// That fullNode should have at most one non-hashNode child,
// hash it because no more nodes will be inserted in it.
st.alloc()
keyUntilHere := len(st.stack[level].keyUntilHere) + len(st.stack[level].ext.Key) + 1
st.stack[level].branch.Children[key[keyUntilHere-1]] = &st.stack[st.top].ext
st.stack[st.top].keyUntilHere = key[:keyUntilHere]
st.stack[st.top].ext.Key = key[keyUntilHere:]
st.stack[st.top].ext.Val = hv
st.stack[st.top].ext.flags = nodeFlag{dirty: true}
st.stack[st.top].depth = st.stack[level].depth + 1
} else {
// extension keys differ, need to create a split and
// hash the former node.
whereitdiffers := 0
offset := len(st.stack[level].keyUntilHere)
for i := range st.stack[level].ext.Key {
if key[offset+i] != st.stack[level].ext.Key[i] {
whereitdiffers = i
break
}
}
// Special case: the split is at the first byte, in this case
// the current ext needs to be skipped.
if whereitdiffers == 0 {
// Hash the existing node
saveSlot := st.stack[level].ext.Key[0]
st.stack[level].ext.Key = st.stack[level].ext.Key[1:]
var h node
if len(st.stack[level].ext.Key) == 0 {
h, _ = st.hasher.hash(&st.stack[level].branch, false)
} else {
h, _ = st.hasher.hash(&st.stack[level].ext, false)
}
for i := range st.stack[level].branch.Children {
st.stack[level].branch.Children[i] = nil
}
st.stack[level].branch.Children[saveSlot] = h
// Set the ext key to empty
st.stack[level].ext.Key = st.stack[level].ext.Key[:0]
st.top = level
// Insert the new leaf, starting with allocating more space
// if needed.
st.alloc()
st.stack[st.top].ext.Key = key[offset+1:]
st.stack[st.top].ext.Val = hv
st.stack[level].branch.Children[key[offset]] = &st.stack[st.top].ext
st.stack[st.top].keyUntilHere = key[:offset+1]
// Update parent reference if this isn't the root
if level > 0 {
parentslot := key[offset-1]
st.stack[level-1].branch.Children[parentslot] = &st.stack[level].branch
}
} else {
// Start by hashing the node right after the extension,
// to free some space.
var hashPrevBranch node
switch st.stack[level].ext.Val.(type) {
case *fullNode:
h, _ := st.hasher.hash(st.stack[level].ext.Val, false)
hashPrevBranch = h.(hashNode)
st.top = level
case hashNode, valueNode:
hashPrevBranch = st.stack[level].ext.Val
default:
panic("Encountered unexpected node type")
}
// Store the completed subtree in a fullNode at the slot
// where both keys differ.
slot := st.stack[level].ext.Key[whereitdiffers]
// Allocate the next full node, it's going to be
// reused several times.
st.alloc()
// Special case: the keys differ at the last element
if len(st.stack[level].ext.Key) == whereitdiffers+1 {
// Directly use the hashed value
for i := range st.stack[level].branch.Children {
st.stack[level].branch.Children[i] = nil
}
st.stack[level].branch.Children[slot] = hashPrevBranch
} else {
// Store the partially-hashed old node in the newly allocated
// slot, in order to finish the hashing.
st.stack[st.top].ext.Key = st.stack[level].ext.Key[whereitdiffers+1:]
st.stack[st.top].ext.Val = hashPrevBranch
st.stack[st.top].ext.flags = nodeFlag{dirty: true}
// Directly hash the branch if the extension is empty
var h node
if len(st.stack[st.top].ext.Key) == 0 {
h, _ = st.hasher.hash(&st.stack[st.top].branch, false)
} else {
h, _ = st.hasher.hash(&st.stack[st.top].ext, false)
}
st.stack[level].branch.Children[slot] = h
}
st.stack[level].ext.Val = &st.stack[level].branch
st.stack[level].ext.Key = st.stack[level].ext.Key[:whereitdiffers]
// Now use the newly allocated+hashed stack st.stack[level] to store
// the rest of the inserted (key, value) pair.
slot = key[whereitdiffers+len(st.stack[level].keyUntilHere)]
st.stack[st.top].ext.Key = key[whereitdiffers+len(st.stack[level].keyUntilHere)+1:]
if len(st.stack[st.top].ext.Key) == 0 {
st.stack[level].branch.Children[slot] = hv
} else {
st.stack[level].branch.Children[slot] = &st.stack[st.top].ext
st.stack[st.top].ext.Val = hv
}
st.stack[st.top].keyUntilHere = key[:whereitdiffers+len(st.stack[level].keyUntilHere)+1]
st.stack[st.top].depth = st.stack[level].depth + 1
}
}
// if ext.length == 0, directly return the full node.
if len(st.stack[0].ext.Key) == 0 {
return &st.stack[0].branch
}
return &st.stack[0].ext
}
// Hash hashes the stack trie by hashing the first entry in the stack
func (st *StackTrie) Hash() common.Hash {
if st.top == -1 {
return emptyRoot
}
h, _ := st.hasher.hash(&st.stack[0].ext, false)
return common.BytesToHash(h.(hashNode))
}