les, light: simplify NodeSet, use readTraceDB

This commit is contained in:
Zsolt Felfoldi 2017-10-10 22:05:10 +02:00
parent d8bf82e354
commit abc5400b66
3 changed files with 57 additions and 98 deletions

View file

@ -239,16 +239,16 @@ func (r *TrieRequest) Validate(db ethdb.Database, msg *Msg) error {
case MsgProofsV2: case MsgProofsV2:
proofs := msg.Obj.(light.NodeList) proofs := msg.Obj.(light.NodeList)
// Verify the proof and store if checks out // Verify the proof and store if checks out
pdb := proofs.NodeSet() nodeSet := proofs.NodeSet()
cdb := pdb.ReadCache() reads := &readTraceDB{db: nodeSet}
if _, err, _ := trie.VerifyProof(r.Id.Root, r.Key, cdb); err != nil { if _, err, _ := trie.VerifyProof(r.Id.Root, r.Key, reads); err != nil {
return fmt.Errorf("merkle proof verification failed: %v", err) return fmt.Errorf("merkle proof verification failed: %v", err)
} }
// check if all nodes have been read by VerifyProof // check if all nodes have been read by VerifyProof
if pdb.KeyCount() != cdb.KeyCount() { if len(reads.reads) != nodeSet.KeyCount() {
return errUselessNodes return errUselessNodes
} }
r.Proof = pdb r.Proof = nodeSet
return nil return nil
default: default:
@ -417,7 +417,7 @@ func (r *ChtRequest) Validate(db ethdb.Database, msg *Msg) error {
if len(resp.AuxData) != 1 { if len(resp.AuxData) != 1 {
return errInvalidEntryCount return errInvalidEntryCount
} }
pdb := resp.Proofs.NodeSet() nodeSet := resp.Proofs.NodeSet()
headerEnc := resp.AuxData[0] headerEnc := resp.AuxData[0]
if len(headerEnc) == 0 { if len(headerEnc) == 0 {
return errHeaderUnavailable return errHeaderUnavailable
@ -431,12 +431,12 @@ func (r *ChtRequest) Validate(db ethdb.Database, msg *Msg) error {
var encNumber [8]byte var encNumber [8]byte
binary.BigEndian.PutUint64(encNumber[:], r.BlockNum) binary.BigEndian.PutUint64(encNumber[:], r.BlockNum)
cdb := pdb.ReadCache() reads := &readTraceDB{db: nodeSet}
value, err, _ := trie.VerifyProof(r.ChtRoot, encNumber[:], cdb) value, err, _ := trie.VerifyProof(r.ChtRoot, encNumber[:], reads)
if err != nil { if err != nil {
return fmt.Errorf("merkle proof verification failed: %v", err) return fmt.Errorf("merkle proof verification failed: %v", err)
} }
if pdb.KeyCount() != cdb.KeyCount() { if len(reads.reads) != nodeSet.KeyCount() {
return errUselessNodes return errUselessNodes
} }
@ -452,7 +452,7 @@ func (r *ChtRequest) Validate(db ethdb.Database, msg *Msg) error {
} }
// Verifications passed, store and return // Verifications passed, store and return
r.Header = header r.Header = header
r.Proof = pdb r.Proof = nodeSet
r.Td = node.Td r.Td = node.Td
default: default:
return errInvalidMessageType return errInvalidMessageType
@ -515,8 +515,8 @@ func (r *BloomRequest) Validate(db ethdb.Database, msg *Msg) error {
} }
resps := msg.Obj.(PPTResps) resps := msg.Obj.(PPTResps)
proofs := resps.Proofs proofs := resps.Proofs
pdb := proofs.NodeSet() nodeSet := proofs.NodeSet()
cdb := pdb.ReadCache() reads := &readTraceDB{db: nodeSet}
r.BloomBits = make([][]byte, len(r.SectionIdxList)) r.BloomBits = make([][]byte, len(r.SectionIdxList))
@ -526,16 +526,38 @@ func (r *BloomRequest) Validate(db ethdb.Database, msg *Msg) error {
for i, idx := range r.SectionIdxList { for i, idx := range r.SectionIdxList {
binary.BigEndian.PutUint64(encNumber[2:10], idx) binary.BigEndian.PutUint64(encNumber[2:10], idx)
value, err, _ := trie.VerifyProof(r.BltRoot, encNumber[:], cdb) value, err, _ := trie.VerifyProof(r.BltRoot, encNumber[:], reads)
if err != nil { if err != nil {
return err return err
} }
r.BloomBits[i] = value r.BloomBits[i] = value
} }
if pdb.KeyCount() != cdb.KeyCount() { if len(reads.reads) != nodeSet.KeyCount() {
return errUselessNodes return errUselessNodes
} }
r.Proofs = pdb r.Proofs = nodeSet
return nil return nil
} }
// readTraceDB stores the keys of database reads. We use this to check that received node
// sets contain only the trie nodes necessary to make proofs pass.
type readTraceDB struct {
db trie.DatabaseReader
reads map[string]struct{}
}
// Get returns a stored node
func (db *readTraceDB) Get(k []byte) ([]byte, error) {
if db.reads == nil {
db.reads = make(map[string]struct{})
}
db.reads[string(k)] = struct{}{}
return db.db.Get(k)
}
// Has returns true if the node set contains the given key
func (db *readTraceDB) Has(key []byte) (bool, error) {
_, err := db.Get(key)
return err == nil, nil
}

View file

@ -29,11 +29,9 @@ import (
// NodeSet stores a set of trie nodes. It implements trie.Database and can also // NodeSet stores a set of trie nodes. It implements trie.Database and can also
// act as a cache for another trie.Database. // act as a cache for another trie.Database.
type NodeSet struct { type NodeSet struct {
db map[string][]byte db map[string][]byte
dataSize int dataSize int
lock sync.RWMutex lock sync.RWMutex
fallback trie.Database
copyFromFallback, writeToFallback bool
} }
// NewNodeSet creates an empty node set // NewNodeSet creates an empty node set
@ -43,25 +41,6 @@ func NewNodeSet() *NodeSet {
} }
} }
// SetFallback will add a fallback database, making this node set a cache for the backing database.
// If copyFromFallback is true, it keeps any node it fetches from the fallback database.
// If writeToFallback is true, it writes stored nodes to the fallback database too.
func (db *NodeSet) SetFallback(fallback trie.Database, copyFromFallback, writeToFallback bool) {
db.lock.Lock()
defer db.lock.Unlock()
db.fallback = fallback
db.copyFromFallback = copyFromFallback
db.writeToFallback = writeToFallback
}
// ReadCache returns a new read cache (copyFromFallback=true) for this node set
func (db *NodeSet) ReadCache() *NodeSet {
cdb := NewNodeSet()
cdb.SetFallback(db, true, false)
return cdb
}
// Put stores a new node in the set // Put stores a new node in the set
func (db *NodeSet) Put(key []byte, value []byte) error { func (db *NodeSet) Put(key []byte, value []byte) error {
db.lock.Lock() db.lock.Lock()
@ -70,9 +49,6 @@ func (db *NodeSet) Put(key []byte, value []byte) error {
if _, ok := db.db[string(key)]; !ok { if _, ok := db.db[string(key)]; !ok {
db.db[string(key)] = common.CopyBytes(value) db.db[string(key)] = common.CopyBytes(value)
db.dataSize += len(value) db.dataSize += len(value)
if db.writeToFallback && db.fallback != nil {
db.fallback.Put(key, value)
}
} }
return nil return nil
} }
@ -85,14 +61,6 @@ func (db *NodeSet) Get(key []byte) ([]byte, error) {
if entry, ok := db.db[string(key)]; ok { if entry, ok := db.db[string(key)]; ok {
return entry, nil return entry, nil
} }
if db.fallback != nil {
value, err := db.fallback.Get(key)
if db.copyFromFallback && err == nil {
db.db[string(key)] = value
db.dataSize += len(value)
}
return value, err
}
return nil, errors.New("not found") return nil, errors.New("not found")
} }

View file

@ -19,13 +19,13 @@ package trie
import ( import (
"bytes" "bytes"
crand "crypto/rand" crand "crypto/rand"
"errors"
mrand "math/rand" mrand "math/rand"
"testing" "testing"
"time" "time"
"github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/crypto"
"github.com/ethereum/go-ethereum/ethdb"
) )
func init() { func init() {
@ -36,7 +36,7 @@ func TestProof(t *testing.T) {
trie, vals := randomTrie(500) trie, vals := randomTrie(500)
root := trie.Hash() root := trie.Hash()
for _, kv := range vals { for _, kv := range vals {
proofs := newTestProofDb() proofs, _ := ethdb.NewMemDatabase()
if trie.Prove(kv.k, 0, proofs) != nil { if trie.Prove(kv.k, 0, proofs) != nil {
t.Fatalf("missing key %x while constructing proof", kv.k) t.Fatalf("missing key %x while constructing proof", kv.k)
} }
@ -53,14 +53,14 @@ func TestProof(t *testing.T) {
func TestOneElementProof(t *testing.T) { func TestOneElementProof(t *testing.T) {
trie := new(Trie) trie := new(Trie)
updateString(trie, "k", "v") updateString(trie, "k", "v")
proofs := newTestProofDb() proofs, _ := ethdb.NewMemDatabase()
trie.Prove([]byte("k"), 0, proofs) trie.Prove([]byte("k"), 0, proofs)
if len(proofs.db) != 1 { if len(proofs.Keys()) != 1 {
t.Error("proof should have one element") t.Error("proof should have one element")
} }
val, err, _ := VerifyProof(trie.Hash(), []byte("k"), proofs) val, err, _ := VerifyProof(trie.Hash(), []byte("k"), proofs)
if err != nil { if err != nil {
t.Fatalf("VerifyProof error: %v\nraw proof: %v", err, proofs.db) t.Fatalf("VerifyProof error: %v\nproof hashes: %v", err, proofs.Keys())
} }
if !bytes.Equal(val, []byte("v")) { if !bytes.Equal(val, []byte("v")) {
t.Fatalf("VerifyProof returned wrong value: got %x, want 'k'", val) t.Fatalf("VerifyProof returned wrong value: got %x, want 'k'", val)
@ -71,21 +71,17 @@ func TestVerifyBadProof(t *testing.T) {
trie, vals := randomTrie(800) trie, vals := randomTrie(800)
root := trie.Hash() root := trie.Hash()
for _, kv := range vals { for _, kv := range vals {
proofs := newTestProofDb() proofs, _ := ethdb.NewMemDatabase()
trie.Prove(kv.k, 0, proofs) trie.Prove(kv.k, 0, proofs)
if len(proofs.db) == 0 { if len(proofs.Keys()) == 0 {
t.Fatal("zero length proof") t.Fatal("zero length proof")
} }
idx := mrand.Intn(len(proofs.db)) keys := proofs.Keys()
for key, node := range proofs.db { key := keys[mrand.Intn(len(keys))]
if idx == 0 { node, _ := proofs.Get(key)
delete(proofs.db, key) proofs.Delete(key)
mutateByte(node) mutateByte(node)
proofs.Put(crypto.Keccak256(node), node) proofs.Put(crypto.Keccak256(node), node)
break
}
idx--
}
if _, err, _ := VerifyProof(root, kv.k, proofs); err == nil { if _, err, _ := VerifyProof(root, kv.k, proofs); err == nil {
t.Fatalf("expected proof to fail for key %x", kv.k) t.Fatalf("expected proof to fail for key %x", kv.k)
} }
@ -113,8 +109,8 @@ func BenchmarkProve(b *testing.B) {
b.ResetTimer() b.ResetTimer()
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
kv := vals[keys[i%len(keys)]] kv := vals[keys[i%len(keys)]]
proofs := newTestProofDb() proofs, _ := ethdb.NewMemDatabase()
if trie.Prove(kv.k, 0, proofs); len(proofs.db) == 0 { if trie.Prove(kv.k, 0, proofs); len(proofs.Keys()) == 0 {
b.Fatalf("zero length proof for %x", kv.k) b.Fatalf("zero length proof for %x", kv.k)
} }
} }
@ -124,10 +120,10 @@ func BenchmarkVerifyProof(b *testing.B) {
trie, vals := randomTrie(100) trie, vals := randomTrie(100)
root := trie.Hash() root := trie.Hash()
var keys []string var keys []string
var proofs []*testProofDb var proofs []*ethdb.MemDatabase
for k := range vals { for k := range vals {
keys = append(keys, k) keys = append(keys, k)
proof := newTestProofDb() proof, _ := ethdb.NewMemDatabase()
trie.Prove([]byte(k), 0, proof) trie.Prove([]byte(k), 0, proof)
proofs = append(proofs, proof) proofs = append(proofs, proof)
} }
@ -165,30 +161,3 @@ func randBytes(n int) []byte {
crand.Read(r) crand.Read(r)
return r return r
} }
type testProofDb struct {
db map[string][]byte
}
func newTestProofDb() *testProofDb {
return &testProofDb{
db: make(map[string][]byte),
}
}
func (db *testProofDb) Put(key []byte, value []byte) error {
db.db[string(key)] = common.CopyBytes(value)
return nil
}
func (db *testProofDb) Get(key []byte) ([]byte, error) {
if entry, ok := db.db[string(key)]; ok {
return entry, nil
}
return nil, errors.New("not found")
}
func (db *testProofDb) Has(key []byte) (bool, error) {
_, err := db.Get(key)
return err == nil, nil
}