trie: avoid panic in stacktrie, return errors instead

This commit is contained in:
Martin Holst Swende 2023-10-17 09:14:11 +02:00
parent a8617c6d4d
commit 566818c629
No known key found for this signature in database
GPG key ID: 683B438C05A5DDF0
3 changed files with 40 additions and 11 deletions

View file

@ -230,7 +230,9 @@ func (dl *diskLayer) proveRange(ctx *generatorContext, trieId *trie.ID, prefix [
if origin == nil && !diskMore { if origin == nil && !diskMore {
stackTr := trie.NewStackTrie(nil) stackTr := trie.NewStackTrie(nil)
for i, key := range keys { for i, key := range keys {
stackTr.Update(key, vals[i]) if err := stackTr.Update(key, vals[i]); err != nil {
return nil, err
}
} }
if gotRoot := stackTr.Hash(); gotRoot != root { if gotRoot := stackTr.Hash(); gotRoot != root {
return &proofResult{ return &proofResult{

View file

@ -18,6 +18,7 @@ package trie
import ( import (
"bytes" "bytes"
"errors"
"sync" "sync"
"github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common"
@ -92,22 +93,26 @@ func NewStackTrie(options *StackTrieOptions) *StackTrie {
// Update inserts a (key, value) pair into the stack trie. // Update inserts a (key, value) pair into the stack trie.
func (t *StackTrie) Update(key, value []byte) error { func (t *StackTrie) Update(key, value []byte) error {
k := keybytesToHex(key)
if len(value) == 0 { if len(value) == 0 {
panic("deletion not supported") return errors.New("trying to insert empty (deletion)")
} }
k := keybytesToHex(key)
k = k[:len(k)-1] // chop the termination flag k = k[:len(k)-1] // chop the termination flag
// track the first and last inserted entries. // track the first and last inserted entries.
if t.first == nil { if t.first == nil {
t.first = append([]byte{}, k...) t.first = append([]byte{}, k...)
} }
if bytes.Compare(t.last, k) >= 0 {
return errors.New("non-ascending key order")
}
if t.last == nil { if t.last == nil {
t.last = append([]byte{}, k...) // allocate key slice t.last = append([]byte{}, k...) // allocate key slice
} else { } else {
t.last = append(t.last[:0], k...) // reuse key slice t.last = append(t.last[:0], k...) // reuse key slice
} }
t.insert(t.root, k, value, nil) if err := t.insert(t.root, k, value, nil); err != nil {
return err
}
return nil return nil
} }
@ -189,7 +194,7 @@ func (n *stNode) getDiffIndex(key []byte) int {
// Helper function to that inserts a (key, value) pair into // Helper function to that inserts a (key, value) pair into
// the trie. // the trie.
func (t *StackTrie) insert(st *stNode, key, value []byte, path []byte) { func (t *StackTrie) insert(st *stNode, key, value []byte, path []byte) error {
switch st.typ { switch st.typ {
case branchNode: /* Branch */ case branchNode: /* Branch */
idx := int(key[0]) idx := int(key[0])
@ -208,7 +213,7 @@ func (t *StackTrie) insert(st *stNode, key, value []byte, path []byte) {
if st.children[idx] == nil { if st.children[idx] == nil {
st.children[idx] = newLeaf(key[1:], value) st.children[idx] = newLeaf(key[1:], value)
} else { } else {
t.insert(st.children[idx], key[1:], value, append(path, key[0])) return t.insert(st.children[idx], key[1:], value, append(path, key[0]))
} }
case extNode: /* Ext */ case extNode: /* Ext */
@ -223,8 +228,7 @@ func (t *StackTrie) insert(st *stNode, key, value []byte, path []byte) {
if diffidx == len(st.key) { if diffidx == len(st.key) {
// Ext key and key segment are identical, recurse into // Ext key and key segment are identical, recurse into
// the child node. // the child node.
t.insert(st.children[0], key[diffidx:], value, append(path, key[:diffidx]...)) return t.insert(st.children[0], key[diffidx:], value, append(path, key[:diffidx]...))
return
} }
// Save the original part. Depending if the break is // Save the original part. Depending if the break is
// at the extension's last byte or not, create an // at the extension's last byte or not, create an
@ -282,7 +286,7 @@ func (t *StackTrie) insert(st *stNode, key, value []byte, path []byte) {
// keys differ, and 3) one leaf for the differentiated // keys differ, and 3) one leaf for the differentiated
// component of each key. // component of each key.
if diffidx >= len(st.key) { if diffidx >= len(st.key) {
panic("Trying to insert into existing key") return errors.New("trying to insert into existing key")
} }
// Check if the split occurs at the first nibble of the // Check if the split occurs at the first nibble of the
@ -324,11 +328,12 @@ func (t *StackTrie) insert(st *stNode, key, value []byte, path []byte) {
st.val = value st.val = value
case hashedNode: case hashedNode:
panic("trying to insert into hash") return errors.New("trying to insert into hash")
default: default:
panic("invalid type") panic("invalid type")
} }
return nil
} }
// hash converts st into a 'hashedNode', if possible. Possible outcomes: // hash converts st into a 'hashedNode', if possible. Possible outcomes:

View file

@ -26,6 +26,7 @@ import (
"github.com/ethereum/go-ethereum/core/rawdb" "github.com/ethereum/go-ethereum/core/rawdb"
"github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/crypto"
"github.com/ethereum/go-ethereum/trie/testutil" "github.com/ethereum/go-ethereum/trie/testutil"
"github.com/stretchr/testify/assert"
"golang.org/x/exp/slices" "golang.org/x/exp/slices"
) )
@ -463,3 +464,24 @@ func TestPartialStackTrie(t *testing.T) {
} }
} }
} }
func TestStacstaturieErrors(t *testing.T) {
s := NewStackTrie(nil)
// Deletion
if err := s.Update(nil, nil); err == nil {
t.Fatal("expected error")
}
if err := s.Update(nil, []byte{}); err == nil {
t.Fatal("expected error")
}
if err := s.Update([]byte{0xa}, []byte{}); err == nil {
t.Fatal("expected error")
}
// Non-ascending keys (going backwards or repeating)
assert.Nil(t, s.Update([]byte{0xaa}, []byte{0xa}))
assert.NotNil(t, s.Update([]byte{0xaa}, []byte{0xa}), "repeat insert same key")
assert.NotNil(t, s.Update([]byte{0xaa}, []byte{0xb}), "repeat insert same key")
assert.Nil(t, s.Update([]byte{0xab}, []byte{0xa}))
assert.NotNil(t, s.Update([]byte{0x10}, []byte{0xb}), "out of order insert")
assert.NotNil(t, s.Update([]byte{0xaa}, []byte{0xb}), "repeat insert same key")
}