From a486c11fbe8e80b0d3ca94d2c1f59e1e70266c05 Mon Sep 17 00:00:00 2001 From: zsfelfoldi Date: Wed, 19 Aug 2015 03:27:16 +0200 Subject: [PATCH] light client stage 1 squashed commit --- cmd/evm/main.go | 5 +- cmd/geth/blocktestcmd.go | 3 +- cmd/geth/chaincmd.go | 7 +- cmd/geth/main.go | 7 +- cmd/utils/flags.go | 24 +- core/access/chain_access.go | 183 ++++++++++ core/access/context.go | 73 ++++ core/access/peer.go | 174 ++++++++++ core/bench_test.go | 6 +- core/block_processor.go | 24 +- core/block_processor_test.go | 10 +- core/blockchain.go | 237 +++++++++---- core/blockchain_test.go | 70 ++-- core/chain_makers.go | 8 +- core/chain_makers_test.go | 8 +- core/chain_util.go | 19 +- core/chain_util_test.go | 30 +- core/core_access.go | 153 +++++++++ core/genesis.go | 8 +- core/manager.go | 2 + core/state/dump.go | 3 +- core/state/managed_state.go | 3 +- core/state/managed_state_test.go | 3 +- core/state/state_access.go | 161 +++++++++ core/state/state_object.go | 39 +-- core/state/state_test.go | 11 +- core/state/statedb.go | 46 +-- core/state/sync_test.go | 5 +- core/transaction_pool_test.go | 7 +- core/transaction_util.go | 30 +- core/types/derive_sha.go | 3 +- core/vm_env.go | 5 +- eth/backend.go | 113 +++++-- eth/backend_test.go | 3 +- eth/downloader/downloader.go | 79 ++++- eth/downloader/downloader_test.go | 7 +- eth/downloader/modes.go | 1 + eth/filters/filter.go | 17 +- eth/filters/filter_test.go | 16 +- eth/gasprice.go | 6 +- eth/handler.go | 71 ++-- eth/handler_test.go | 115 ++++--- eth/helper_test.go | 16 +- eth/peer.go | 2 +- eth/protocol.go | 17 +- eth/sync.go | 2 +- les/handler.go | 543 ++++++++++++++++++++++++++++++ les/metrics.go | 157 +++++++++ les/peer.go | 349 +++++++++++++++++++ les/protocol.go | 197 +++++++++++ les/sync.go | 72 ++++ miner/miner.go | 7 + miner/worker.go | 31 +- rpc/api/admin.go | 2 +- rpc/api/debug.go | 3 +- rpc/api/eth.go | 58 ++-- rpc/comms/comms.go | 24 +- rpc/comms/http.go | 24 +- rpc/comms/inproc.go | 10 + rpc/shared/types.go | 10 + tests/block_test_util.go | 7 +- tests/state_test_util.go | 9 +- tests/util.go | 5 +- tests/vm_test_util.go | 7 +- trie/encoding.go | 17 + trie/iterator.go | 10 +- trie/iterator_test.go | 8 +- trie/proof.go | 38 ++- trie/proof_test.go | 7 +- trie/secure_trie.go | 17 +- trie/secure_trie_test.go | 11 +- trie/sync_test.go | 15 +- trie/trie.go | 94 ++++-- trie/trie_test.go | 39 +-- xeth/state.go | 2 +- xeth/types.go | 15 +- xeth/xeth.go | 54 ++- 77 files changed, 3124 insertions(+), 550 deletions(-) create mode 100644 core/access/chain_access.go create mode 100644 core/access/context.go create mode 100644 core/access/peer.go create mode 100644 core/core_access.go create mode 100644 core/state/state_access.go create mode 100644 les/handler.go create mode 100644 les/metrics.go create mode 100644 les/peer.go create mode 100644 les/protocol.go create mode 100644 les/sync.go diff --git a/cmd/evm/main.go b/cmd/evm/main.go index 64044c4214..e069411160 100644 --- a/cmd/evm/main.go +++ b/cmd/evm/main.go @@ -28,6 +28,7 @@ import ( "github.com/ethereum/go-ethereum/cmd/utils" "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" @@ -113,7 +114,7 @@ func run(ctx *cli.Context) { glog.SetV(ctx.GlobalInt(VerbosityFlag.Name)) db, _ := ethdb.NewMemDatabase() - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) sender := statedb.CreateAccount(common.StringToAddress("sender")) receiver := statedb.CreateAccount(common.StringToAddress("receiver")) receiver.SetCode(common.Hex2Bytes(ctx.GlobalString(CodeFlag.Name))) @@ -186,7 +187,7 @@ func NewEnv(state *state.StateDB, transactor common.Address, value *big.Int) *VM } func (self *VMEnv) Db() vm.Database { return self.state } -func (self *VMEnv) MakeSnapshot() vm.Database { return self.state.Copy() } +func (self *VMEnv) MakeSnapshot() vm.Database { return self.state.Copy(access.NullCtx) } func (self *VMEnv) SetSnapshot(db vm.Database) { self.state.Set(db.(*state.StateDB)) } func (self *VMEnv) Origin() common.Address { return *self.transactor } func (self *VMEnv) BlockNumber() *big.Int { return common.Big0 } diff --git a/cmd/geth/blocktestcmd.go b/cmd/geth/blocktestcmd.go index e4d97aa531..36986cdc7b 100644 --- a/cmd/geth/blocktestcmd.go +++ b/cmd/geth/blocktestcmd.go @@ -22,6 +22,7 @@ import ( "github.com/codegangsta/cli" "github.com/ethereum/go-ethereum/cmd/utils" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/eth" "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/tests" @@ -124,7 +125,7 @@ func runOneBlockTest(ctx *cli.Context, test *tests.BlockTest) (*eth.Ethereum, er if err != nil { return ethereum, fmt.Errorf("Block Test load error: %v", err) } - newDB, err := cm.State() + newDB, err := cm.State(access.NullCtx) if err != nil { return ethereum, fmt.Errorf("Block Test get state error: %v", err) } diff --git a/cmd/geth/chaincmd.go b/cmd/geth/chaincmd.go index 80f3777d6d..2de3002f90 100644 --- a/cmd/geth/chaincmd.go +++ b/cmd/geth/chaincmd.go @@ -27,6 +27,7 @@ import ( "github.com/ethereum/go-ethereum/cmd/utils" "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/ethdb" @@ -170,16 +171,16 @@ func dump(ctx *cli.Context) { for _, arg := range ctx.Args() { var block *types.Block if hashish(arg) { - block = chain.GetBlock(common.HexToHash(arg)) + block = chain.GetBlock(common.HexToHash(arg), access.NoOdr) } else { num, _ := strconv.Atoi(arg) - block = chain.GetBlockByNumber(uint64(num)) + block = chain.GetBlockByNumber(uint64(num), access.NoOdr) } if block == nil { fmt.Println("{}") utils.Fatalf("block not found") } else { - state, err := state.New(block.Root(), chainDb) + state, err := state.New(block.Root(), access.NewDbChainAccess(chainDb), access.NullCtx) if err != nil { utils.Fatalf("could not create new state: %v", err) return diff --git a/cmd/geth/main.go b/cmd/geth/main.go index e74ddd0d09..32200670fd 100644 --- a/cmd/geth/main.go +++ b/cmd/geth/main.go @@ -34,6 +34,7 @@ import ( "github.com/ethereum/go-ethereum/cmd/utils" "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/eth" "github.com/ethereum/go-ethereum/ethdb" @@ -304,6 +305,7 @@ JavaScript API. See https://github.com/ethereum/go-ethereum/wiki/Javascipt-Conso utils.BlockchainVersionFlag, utils.OlympicFlag, utils.FastSyncFlag, + utils.EthModeFlag, utils.CacheFlag, utils.JSpathFlag, utils.ListenPortFlag, @@ -554,12 +556,13 @@ func blockRecovery(ctx *cli.Context) { if err != nil { glog.Fatalln("could not open db:", err) } + ca := access.NewDbChainAccess(blockDb) var block *types.Block if arg[0] == '#' { - block = core.GetBlock(blockDb, core.GetCanonicalHash(blockDb, common.String2Big(arg[1:]).Uint64())) + block = core.GetBlock(ca, core.GetCanonicalHash(blockDb, common.String2Big(arg[1:]).Uint64()), access.NoOdr) } else { - block = core.GetBlock(blockDb, common.HexToHash(arg)) + block = core.GetBlock(ca, common.HexToHash(arg), access.NoOdr) } if block == nil { diff --git a/cmd/utils/flags.go b/cmd/utils/flags.go index d741d0544f..b8d62c7f91 100644 --- a/cmd/utils/flags.go +++ b/cmd/utils/flags.go @@ -28,12 +28,14 @@ import ( "path/filepath" "runtime" "strconv" + "strings" "github.com/codegangsta/cli" "github.com/ethereum/ethash" "github.com/ethereum/go-ethereum/accounts" "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/eth" @@ -156,6 +158,11 @@ var ( Name: "lightkdf", Usage: "Reduce KDF memory & CPU usage at some expense of KDF strength", } + EthModeFlag = cli.StringFlag{ + Name: "mode", + Value: "archive", + Usage: "Client mode of operation (archive, full, light)", + } // Miner settings // TODO: refactor CPU vs GPU mining flags MiningEnabledFlag = cli.BoolFlag{ @@ -424,12 +431,25 @@ func MakeEthConfig(clientID, version string, ctx *cli.Context) *eth.Config { if err != nil { glog.V(logger.Error).Infoln("WARNING: No etherbase set and no accounts found as default") } + // Resolve the mode of opeation from the string flag + var clientMode eth.Mode + switch strings.ToLower(ctx.GlobalString(EthModeFlag.Name)) { + case "archive": + clientMode = eth.ArchiveMode + case "full": + clientMode = eth.FullMode + case "light": + clientMode = eth.LightMode + default: + glog.Fatalf("Unknown node type requested: %s", ctx.GlobalString(EthModeFlag.Name)) + } // Assemble the entire eth configuration and return cfg := ð.Config{ Name: common.MakeName(clientID, version), DataDir: MustDataDir(ctx), GenesisFile: ctx.GlobalString(GenesisFileFlag.Name), FastSync: ctx.GlobalBool(FastSyncFlag.Name), + Mode: clientMode, BlockChainVersion: ctx.GlobalInt(BlockchainVersionFlag.Name), DatabaseCache: ctx.GlobalInt(CacheFlag.Name), SkipBcVersionCheck: false, @@ -552,12 +572,12 @@ func MakeChain(ctx *cli.Context) (chain *core.BlockChain, chainDb ethdb.Database eventMux := new(event.TypeMux) pow := ethash.New() //genesis := core.GenesisBlock(uint64(ctx.GlobalInt(GenesisNonceFlag.Name)), blockDB) - chain, err = core.NewBlockChain(chainDb, pow, eventMux) + chain, err = core.NewBlockChain(access.NewDbChainAccess(chainDb), pow, eventMux) if err != nil { Fatalf("Could not start chainmanager: %v", err) } - proc := core.NewBlockProcessor(chainDb, pow, chain, eventMux) + proc := core.NewBlockProcessor(access.NewDbChainAccess(chainDb), pow, chain, eventMux) chain.SetProcessor(proc) return chain, chainDb } diff --git a/core/access/chain_access.go b/core/access/chain_access.go new file mode 100644 index 0000000000..4439e2fbcb --- /dev/null +++ b/core/access/chain_access.go @@ -0,0 +1,183 @@ +// Copyright 2015 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 . + +// Package access provides a layer to handle local blockchain database and +// on-demand network retrieval +package access + +import ( + "errors" + "sync" + "time" + + "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/ethdb" + "github.com/ethereum/go-ethereum/logger" + "github.com/ethereum/go-ethereum/logger/glog" +) + +var ( + ErrCancel = errors.New("ODR cancelled") + errNotInDb = errors.New("object not found in database") +) + +const LogLevel = logger.Info + +var ( + requestTimeout = time.Millisecond * 300 + retryPeers = time.Second * 1 +) + +type ChainAccess struct { + db ethdb.Database + odr bool // light client mode, odr enabled + lock sync.Mutex + valFunc validatorFunc + deliverChan chan *Msg + peers *peerSet + + // p2p access objects + // parameters (light/full/archive) +} + +func NewDbChainAccess(db ethdb.Database) *ChainAccess { + return NewChainAccess(db, false) +} + +func NewChainAccess(db ethdb.Database, odr bool) *ChainAccess { + return &ChainAccess{db: db, peers: newPeerSet(), odr: odr} +} + +func (self *ChainAccess) Db() ethdb.Database { + return self.db +} + +func (self *ChainAccess) OdrEnabled() bool { + return self.odr +} + +func (self *ChainAccess) RegisterPeer(id string, version int, head common.Hash, getBlockBodies getBlockBodiesFn, getNodeData getNodeDataFn, getReceipts getReceiptsFn, getProofs getProofsFn) error { + glog.V(logger.Detail).Infoln("Registering peer", id) + if err := self.peers.Register(newPeer(id, version, head, getBlockBodies, getNodeData, getReceipts, getProofs)); err != nil { + glog.V(logger.Error).Infoln("Register failed:", err) + return err + } + return nil +} + +func (self *ChainAccess) UnregisterPeer(id string) { + self.peers.Unregister(id) +} + +const ( + MsgBlockBodies = iota + MsgNodeData + MsgReceipts + MsgProofs +) + +type Msg struct { + MsgType int + Obj interface{} +} + +type ObjectAccess interface { + // database storage + DbGet() bool + DbPut() + // network retrieval + Request(*Peer) error + Valid(*Msg) bool // if true, keeps the retrieved object +} + +type requestFunc func(*Peer) error +type validatorFunc func(*Msg) bool + +func (self *ChainAccess) Deliver(id string, msg *Msg) (processed bool) { + self.lock.Lock() + valFunc := self.valFunc + chn := self.deliverChan + self.lock.Unlock() + if (valFunc != nil) && (chn != nil) && valFunc(msg) { + chn <- msg + return true + } + return false +} + +func (self *ChainAccess) networkRequest(rqFunc requestFunc, valFunc validatorFunc, ctx *OdrContext) (*Msg, error) { + + self.lock.Lock() + self.deliverChan = make(chan *Msg) + self.valFunc = valFunc + self.lock.Unlock() + + defer func() { + self.lock.Lock() + self.deliverChan = nil + self.valFunc = nil + self.lock.Unlock() + }() + + //fmt.Println("networkRequest") + + var msg *Msg + + for { + peers := self.peers.BestPeers() + if len(peers) == 0 { + select { + case <-ctx.cancelOrTimeout: + return nil, ErrCancel + case <-time.After(retryPeers): + } + } + for _, peer := range peers { + rqFunc(peer) + select { + case <-ctx.cancelOrTimeout: + return nil, ErrCancel + case msg = <-self.deliverChan: + peer.Promote() + glog.V(LogLevel).Infof("networkRequest success") + return msg, nil + case <-time.After(requestTimeout): + peer.Demote() + glog.V(LogLevel).Infof("networkRequest timeout") + } + } + } +} + +func (self *ChainAccess) Retrieve(obj ObjectAccess, ctx *OdrContext) (err error) { + // look in db + if obj.DbGet() { + return nil + } + if ctx != nil { + // not found in db, trying the network + _, err = self.networkRequest(obj.Request, obj.Valid, ctx) + if err == nil { + // retrieved from network, store in db + obj.DbPut() + } else { + glog.V(LogLevel).Infof("networkRequest err = %v", err) + } + return + } else { + return errNotInDb + } +} diff --git a/core/access/context.go b/core/access/context.go new file mode 100644 index 0000000000..410d03e3e7 --- /dev/null +++ b/core/access/context.go @@ -0,0 +1,73 @@ +// Copyright 2015 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 . + +// Package access provides a layer to handle local blockchain database and +// on-demand network retrieval +package access + +import ( + "sync/atomic" + "time" +) + +type OdrContext struct { + cancel, cancelOrTimeout chan struct{} + id *OdrChannelID + cancelled int32 +} + +var NullCtx = (*OdrContext)(nil) // used when creating states +var NoOdr = (*OdrContext)(nil) // used for individual requests + +func NewContext(id *OdrChannelID) *OdrContext { + ctx := &OdrContext{ + cancel: make(chan struct{}), + cancelOrTimeout: make(chan struct{}), + id: id, + } + go func() { + select { + case <-ctx.cancel: + case <-time.After(id.timeout): + } + ctx.setCancelled() + close(ctx.cancelOrTimeout) + }() + return ctx +} + +func (self *OdrContext) Cancel() { + close(self.cancel) +} + +func (self *OdrContext) setCancelled() { + atomic.StoreInt32(&self.cancelled, 1) +} + +func (self *OdrContext) IsCancelled() bool { + if self == nil { + return false + } + return atomic.LoadInt32(&self.cancelled) == 1 +} + +type OdrChannelID struct { + timeout time.Duration +} + +func NewChannelID(timeout time.Duration) *OdrChannelID { + return &OdrChannelID{timeout: timeout} +} \ No newline at end of file diff --git a/core/access/peer.go b/core/access/peer.go new file mode 100644 index 0000000000..2b30a400b0 --- /dev/null +++ b/core/access/peer.go @@ -0,0 +1,174 @@ +// Copyright 2015 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 . + +// Package access provides a layer to handle local blockchain database and +// on-demand network retrieval +package access + +import ( + "errors" + "sync" + "sync/atomic" + + "github.com/ethereum/go-ethereum/common" +) + +var ( + errAlreadyRegistered = errors.New("peer is already registered") + errNotRegistered = errors.New("peer is not registered") + errNoOdr = errors.New("peer cannot serve on-demand requests") +) + +type ProofReq struct { + Root common.Hash + Key []byte +} + +type getBlockBodiesFn func([]common.Hash) error +type getNodeDataFn func([]common.Hash) error +type getReceiptsFn func([]common.Hash) error +type getProofsFn func([]*ProofReq) error + +type Peer struct { + id string // Unique identifier of the peer + head common.Hash // Hash of the peers latest known block + + rep int32 // Simple peer reputation + + GetBlockBodies getBlockBodiesFn + GetNodeData getNodeDataFn + GetReceipts getReceiptsFn + GetProofs getProofsFn + + version int // LES protocol version number +} + +func newPeer(id string, version int, head common.Hash, getBlockBodies getBlockBodiesFn, getNodeData getNodeDataFn, getReceipts getReceiptsFn, getProofs getProofsFn) *Peer { + return &Peer{ + id: id, + head: head, + GetBlockBodies: getBlockBodies, + GetNodeData: getNodeData, + GetReceipts: getReceipts, + GetProofs: getProofs, + version: version, + } +} + +func (p *Peer) Id() string { + return p.id +} + +func (p *Peer) Promote() { + atomic.AddInt32(&p.rep, 1) +} + +// Demote decreases the peer's reputation or leaves it at 0. +func (p *Peer) Demote() { + for { + // Calculate the new reputation value + prev := atomic.LoadInt32(&p.rep) + next := prev / 2 + + // Try to update the old value + if atomic.CompareAndSwapInt32(&p.rep, prev, next) { + return + } + } +} + +// peerSet represents the collection of active peer participating in the block +// download procedure. +type peerSet struct { + peers map[string]*Peer + lock sync.RWMutex +} + +// newPeerSet creates a new peer set top track the active download sources. +func newPeerSet() *peerSet { + return &peerSet{ + peers: make(map[string]*Peer), + } +} + +// Register injects a new peer into the working set, or returns an error if the +// peer is already known. +func (ps *peerSet) Register(p *Peer) error { + ps.lock.Lock() + defer ps.lock.Unlock() + + if _, ok := ps.peers[p.id]; ok { + return errAlreadyRegistered + } + ps.peers[p.id] = p + return nil +} + +// Unregister removes a remote peer from the active set, disabling any further +// actions to/from that particular entity. +func (ps *peerSet) Unregister(id string) error { + ps.lock.Lock() + defer ps.lock.Unlock() + + if _, ok := ps.peers[id]; !ok { + return errNotRegistered + } + delete(ps.peers, id) + return nil +} + +// Peer retrieves the registered peer with the given id. +func (ps *peerSet) Peer(id string) *Peer { + ps.lock.RLock() + defer ps.lock.RUnlock() + + return ps.peers[id] +} + +// Len returns if the current number of peers in the set. +func (ps *peerSet) Len() int { + ps.lock.RLock() + defer ps.lock.RUnlock() + + return len(ps.peers) +} + +// AllPeers retrieves a flat list of all the peers within the set. +func (ps *peerSet) AllPeers() []*Peer { + ps.lock.RLock() + defer ps.lock.RUnlock() + + list := make([]*Peer, 0, len(ps.peers)) + for _, p := range ps.peers { + list = append(list, p) + } + return list +} + +func (ps *peerSet) BestPeers() []*Peer { + list := ps.AllPeers() + ps.lock.RLock() + defer ps.lock.RUnlock() + + for i := 0; i < len(list); i++ { + for j := i + 1; j < len(list); j++ { + if atomic.LoadInt32(&list[i].rep) < atomic.LoadInt32(&list[j].rep) { + list[i], list[j] = list[j], list[i] + } + } + } + return list +} diff --git a/core/bench_test.go b/core/bench_test.go index b5eb518033..704e6a0ad4 100644 --- a/core/bench_test.go +++ b/core/bench_test.go @@ -24,6 +24,7 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/ethdb" @@ -168,8 +169,9 @@ func benchInsertChain(b *testing.B, disk bool, gen func(int, *BlockGen)) { // Time the insertion of the new chain. // State and blocks are stored in the same DB. evmux := new(event.TypeMux) - chainman, _ := NewBlockChain(db, FakePow{}, evmux) - chainman.SetProcessor(NewBlockProcessor(db, FakePow{}, chainman, evmux)) + ca := access.NewDbChainAccess(db) + chainman, _ := NewBlockChain(ca, FakePow{}, evmux) + chainman.SetProcessor(NewBlockProcessor(ca, FakePow{}, chainman, evmux)) defer chainman.Stop() b.ReportAllocs() b.ResetTimer() diff --git a/core/block_processor.go b/core/block_processor.go index e7b2f63e5a..d4a9ee3b52 100644 --- a/core/block_processor.go +++ b/core/block_processor.go @@ -23,11 +23,11 @@ import ( "time" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/crypto" - "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/event" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" @@ -43,7 +43,7 @@ const ( ) type BlockProcessor struct { - chainDb ethdb.Database + chainAccess *access.ChainAccess // Mutex for locking the block processor. Blocks can only be handled one at a time mutex sync.Mutex // Canonical block chain @@ -85,9 +85,9 @@ func (gp *GasPool) String() string { return (*big.Int)(gp).String() } -func NewBlockProcessor(db ethdb.Database, pow pow.PoW, blockchain *BlockChain, eventMux *event.TypeMux) *BlockProcessor { +func NewBlockProcessor(ca *access.ChainAccess, pow pow.PoW, blockchain *BlockChain, eventMux *event.TypeMux) *BlockProcessor { sm := &BlockProcessor{ - chainDb: db, + chainAccess: ca, mem: make(map[string]*big.Int), Pow: pow, bc: blockchain, @@ -189,7 +189,7 @@ func (sm *BlockProcessor) RetryProcess(block *types.Block) (logs vm.Logs, err er if !sm.bc.HasBlock(block.ParentHash()) { return nil, ParentError(block.ParentHash()) } - parent := sm.bc.GetBlock(block.ParentHash()) + parent := sm.bc.GetBlock(block.ParentHash(), access.NoOdr) // FIXME Change to full header validation. See #1225 errch := make(chan bool) @@ -212,12 +212,12 @@ func (sm *BlockProcessor) Process(block *types.Block) (logs vm.Logs, receipts ty defer sm.mutex.Unlock() if sm.bc.HasBlock(block.Hash()) { - if _, err := state.New(block.Root(), sm.chainDb); err == nil { + if _, err := state.New(block.Root(), sm.chainAccess, access.NullCtx); err == nil { return nil, nil, &KnownBlockError{block.Number(), block.Hash()} } } - if parent := sm.bc.GetBlock(block.ParentHash()); parent != nil { - if _, err := state.New(parent.Root(), sm.chainDb); err == nil { + if parent := sm.bc.GetBlock(block.ParentHash(), access.NoOdr); parent != nil { + if _, err := state.New(parent.Root(), sm.chainAccess, access.NullCtx); err == nil { return sm.processWithParent(block, parent) } } @@ -226,7 +226,7 @@ func (sm *BlockProcessor) Process(block *types.Block) (logs vm.Logs, receipts ty func (sm *BlockProcessor) processWithParent(block, parent *types.Block) (logs vm.Logs, receipts types.Receipts, err error) { // Create a new state based on the parent's root (e.g., create copy) - state, err := state.New(parent.Root(), sm.chainDb) + state, err := state.New(parent.Root(), sm.chainAccess, access.NullCtx) if err != nil { return nil, nil, err } @@ -369,8 +369,8 @@ func (sm *BlockProcessor) VerifyUncles(statedb *state.StateDB, block, parent *ty // GetBlockReceipts returns the receipts beloniging to the block hash func (sm *BlockProcessor) GetBlockReceipts(bhash common.Hash) types.Receipts { - if block := sm.BlockChain().GetBlock(bhash); block != nil { - return GetBlockReceipts(sm.chainDb, block.Hash()) + if block := sm.BlockChain().GetBlock(bhash, access.NoOdr); block != nil { + return GetBlockReceipts(sm.chainAccess, block.Hash(), access.NoOdr) } return nil @@ -380,7 +380,7 @@ func (sm *BlockProcessor) GetBlockReceipts(bhash common.Hash) types.Receipts { // where it tries to get it from the (updated) method which gets them from the receipts or // the depricated way by re-processing the block. func (sm *BlockProcessor) GetLogs(block *types.Block) (logs vm.Logs, err error) { - receipts := GetBlockReceipts(sm.chainDb, block.Hash()) + receipts := GetBlockReceipts(sm.chainAccess, block.Hash(), access.NoOdr) // coalesce logs for _, receipt := range receipts { logs = append(logs, receipt.Logs...) diff --git a/core/block_processor_test.go b/core/block_processor_test.go index 3050456b41..46c3cd6c99 100644 --- a/core/block_processor_test.go +++ b/core/block_processor_test.go @@ -22,6 +22,7 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" @@ -32,21 +33,22 @@ import ( func proc() (*BlockProcessor, *BlockChain) { db, _ := ethdb.NewMemDatabase() + ca := access.NewDbChainAccess(db) var mux event.TypeMux WriteTestNetGenesisBlock(db, 0) - blockchain, err := NewBlockChain(db, thePow(), &mux) + blockchain, err := NewBlockChain(ca, thePow(), &mux) if err != nil { fmt.Println(err) } - return NewBlockProcessor(db, ezp.New(), blockchain, &mux), blockchain + return NewBlockProcessor(ca, ezp.New(), blockchain, &mux), blockchain } func TestNumber(t *testing.T) { pow := ezp.New() _, chain := proc() - statedb, _ := state.New(chain.Genesis().Root(), chain.chainDb) + statedb, _ := state.New(chain.Genesis().Root(), access.NewDbChainAccess(chain.chainDb), access.NullCtx) header := makeHeader(chain.Genesis(), statedb) header.Number = big.NewInt(3) err := ValidateHeader(pow, header, chain.Genesis().Header(), false, false) @@ -82,7 +84,7 @@ func TestPutReceipt(t *testing.T) { }} PutReceipts(db, types.Receipts{receipt}) - receipt = GetReceipt(db, common.Hash{}) + receipt = GetReceipt(access.NewDbChainAccess(db), common.Hash{}, access.NoOdr) if receipt == nil { t.Error("expected to get 1 receipt, got none.") } diff --git a/core/blockchain.go b/core/blockchain.go index cea346e387..c036d9a472 100644 --- a/core/blockchain.go +++ b/core/blockchain.go @@ -19,6 +19,7 @@ package core import ( crand "crypto/rand" + "bytes" "errors" "fmt" "io" @@ -31,6 +32,7 @@ import ( "time" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" @@ -64,6 +66,7 @@ const ( ) type BlockChain struct { + chainAccess *access.ChainAccess chainDb ethdb.Database processor types.BlockProcessor eventMux *event.TypeMux @@ -73,9 +76,11 @@ type BlockChain struct { chainmu sync.RWMutex tsmu sync.RWMutex + lightMode bool // block head == header head, block data is only retrieved when necessary + checkpoint int // checkpoint counts towards the new checkpoint currentHeader *types.Header // Current head of the header chain (may be above the block chain!) - currentBlock *types.Block // Current head of the block chain + currentBlock *types.Block // Current head of the block chain (can be nil in light client mode) currentFastBlock *types.Block // Current head of the fast-sync chain (may be above the block chain!) headerCache *lru.Cache // Cache for the most recent block headers @@ -95,7 +100,7 @@ type BlockChain struct { rand *mrand.Rand } -func NewBlockChain(chainDb ethdb.Database, pow pow.PoW, mux *event.TypeMux) (*BlockChain, error) { +func NewBlockChain(chainAccess *access.ChainAccess, pow pow.PoW, mux *event.TypeMux) (*BlockChain, error) { headerCache, _ := lru.New(headerCacheLimit) bodyCache, _ := lru.New(bodyCacheLimit) bodyRLPCache, _ := lru.New(bodyCacheLimit) @@ -104,7 +109,8 @@ func NewBlockChain(chainDb ethdb.Database, pow pow.PoW, mux *event.TypeMux) (*Bl futureBlocks, _ := lru.New(maxFutureBlocks) bc := &BlockChain{ - chainDb: chainDb, + chainDb: chainAccess.Db(), + chainAccess: chainAccess, eventMux: mux, quit: make(chan struct{}), headerCache: headerCache, @@ -114,6 +120,7 @@ func NewBlockChain(chainDb ethdb.Database, pow pow.PoW, mux *event.TypeMux) (*Bl blockCache: blockCache, futureBlocks: futureBlocks, pow: pow, + lightMode: chainAccess.OdrEnabled(), } // Seed a fast but crypto originating random generator seed, err := crand.Int(crand.Reader, big.NewInt(math.MaxInt64)) @@ -122,13 +129,13 @@ func NewBlockChain(chainDb ethdb.Database, pow pow.PoW, mux *event.TypeMux) (*Bl } bc.rand = mrand.New(mrand.NewSource(seed.Int64())) - bc.genesisBlock = bc.GetBlockByNumber(0) + bc.genesisBlock = bc.GetBlockByNumber(0, access.NoOdr) if bc.genesisBlock == nil { reader, err := NewDefaultGenesisReader() if err != nil { return nil, err } - bc.genesisBlock, err = WriteGenesisBlock(chainDb, reader) + bc.genesisBlock, err = WriteGenesisBlock(bc.chainDb, reader) if err != nil { return nil, err } @@ -153,42 +160,50 @@ func NewBlockChain(chainDb ethdb.Database, pow pow.PoW, mux *event.TypeMux) (*Bl // loadLastState loads the last known chain state from the database. This method // assumes that the chain manager mutex is held. func (self *BlockChain) loadLastState() error { - // Restore the last known head block - head := GetHeadBlockHash(self.chainDb) - if head == (common.Hash{}) { - // Corrupt or empty database, init from scratch - self.Reset() - } else { - if block := self.GetBlock(head); block != nil { - // Block found, set as the current head - self.currentBlock = block - } else { + if !self.lightMode { + // Restore the last known head block + head := GetHeadBlockHash(self.chainDb) + if head == (common.Hash{}) { // Corrupt or empty database, init from scratch self.Reset() + } else { + if block := self.GetBlock(head, access.NoOdr); block != nil { + // Block found, set as the current head + self.currentBlock = block + } else { + // Corrupt or empty database, init from scratch + self.Reset() + } } + self.currentHeader = self.currentBlock.Header() } // Restore the last known head header - self.currentHeader = self.currentBlock.Header() if head := GetHeadHeaderHash(self.chainDb); head != (common.Hash{}) { if header := self.GetHeader(head); header != nil { self.currentHeader = header } } - // Restore the last known head fast block - self.currentFastBlock = self.currentBlock - if head := GetHeadFastBlockHash(self.chainDb); head != (common.Hash{}) { - if block := self.GetBlock(head); block != nil { - self.currentFastBlock = block + if !self.lightMode { + // Restore the last known head fast block + self.currentFastBlock = self.currentBlock + if head := GetHeadFastBlockHash(self.chainDb); head != (common.Hash{}) { + if block := self.GetBlock(head, access.NoOdr); block != nil { + self.currentFastBlock = block + } } } // Issue a status log and return - headerTd := self.GetTd(self.currentHeader.Hash()) - blockTd := self.GetTd(self.currentBlock.Hash()) - fastTd := self.GetTd(self.currentFastBlock.Hash()) + if self.currentHeader != nil { + headerTd := self.GetTd(self.currentHeader.Hash()) + glog.V(logger.Info).Infof("Last header: #%d [%x…] TD=%v", self.currentHeader.Number, self.currentHeader.Hash().Bytes()[:4], headerTd) + } - glog.V(logger.Info).Infof("Last header: #%d [%x…] TD=%v", self.currentHeader.Number, self.currentHeader.Hash().Bytes()[:4], headerTd) - glog.V(logger.Info).Infof("Last block: #%d [%x…] TD=%v", self.currentBlock.Number(), self.currentBlock.Hash().Bytes()[:4], blockTd) - glog.V(logger.Info).Infof("Fast block: #%d [%x…] TD=%v", self.currentFastBlock.Number(), self.currentFastBlock.Hash().Bytes()[:4], fastTd) + if !self.lightMode { + blockTd := self.GetTd(self.currentBlock.Hash()) + fastTd := self.GetTd(self.currentFastBlock.Hash()) + glog.V(logger.Info).Infof("Last block: #%d [%x…] TD=%v", self.currentBlock.Number(), self.currentBlock.Hash().Bytes()[:4], blockTd) + glog.V(logger.Info).Infof("Fast block: #%d [%x…] TD=%v", self.currentFastBlock.Number(), self.currentFastBlock.Hash().Bytes()[:4], fastTd) + } return nil } @@ -208,14 +223,16 @@ func (bc *BlockChain) SetHead(head uint64) { height = hh } } - if bc.currentBlock != nil { - if bh := bc.currentBlock.NumberU64(); bh > height { - height = bh + if !bc.lightMode { + if bc.currentBlock != nil { + if bh := bc.currentBlock.NumberU64(); bh > height { + height = bh + } } - } - if bc.currentFastBlock != nil { - if fbh := bc.currentFastBlock.NumberU64(); fbh > height { - height = fbh + if bc.currentFastBlock != nil { + if fbh := bc.currentFastBlock.NumberU64(); fbh > height { + height = fbh + } } } // Gather all the hashes that need deletion @@ -225,13 +242,15 @@ func (bc *BlockChain) SetHead(head uint64) { drop[bc.currentHeader.Hash()] = struct{}{} bc.currentHeader = bc.GetHeader(bc.currentHeader.ParentHash) } - for bc.currentBlock != nil && bc.currentBlock.NumberU64() > head { - drop[bc.currentBlock.Hash()] = struct{}{} - bc.currentBlock = bc.GetBlock(bc.currentBlock.ParentHash()) - } - for bc.currentFastBlock != nil && bc.currentFastBlock.NumberU64() > head { - drop[bc.currentFastBlock.Hash()] = struct{}{} - bc.currentFastBlock = bc.GetBlock(bc.currentFastBlock.ParentHash()) + if !bc.lightMode { + for bc.currentBlock != nil && bc.currentBlock.NumberU64() > head { + drop[bc.currentBlock.Hash()] = struct{}{} + bc.currentBlock = bc.GetBlock(bc.currentBlock.ParentHash(), access.NoOdr) + } + for bc.currentFastBlock != nil && bc.currentFastBlock.NumberU64() > head { + drop[bc.currentFastBlock.Hash()] = struct{}{} + bc.currentFastBlock = bc.GetBlock(bc.currentFastBlock.ParentHash(), access.NoOdr) + } } // Roll back the canonical chain numbering for i := height; i > head; i-- { @@ -250,9 +269,14 @@ func (bc *BlockChain) SetHead(head uint64) { bc.blockCache.Purge() bc.futureBlocks.Purge() - // Update all computed fields to the new head - if bc.currentBlock == nil { - bc.currentBlock = bc.genesisBlock + if !bc.lightMode { + // Update all computed fields to the new head + if bc.currentBlock == nil { + bc.currentBlock = bc.genesisBlock + } + bc.insert(bc.currentBlock) + } else { + bc.writeHeader(bc.currentHeader) } if bc.currentHeader == nil { bc.currentHeader = bc.genesisBlock.Header() @@ -276,11 +300,11 @@ func (bc *BlockChain) SetHead(head uint64) { // irrelevant what the chain contents were prior. func (self *BlockChain) FastSyncCommitHead(hash common.Hash) error { // Make sure that both the block as well at its state trie exists - block := self.GetBlock(hash) + block := self.GetBlock(hash, access.NoOdr) if block == nil { return fmt.Errorf("non existent block [%x…]", hash[:4]) } - if _, err := trie.NewSecure(block.Root(), self.chainDb); err != nil { + if _, err := trie.NewSecure(block.Root(), self.chainDb, nil); err != nil { return err } // If all checks out, manually set the head block @@ -296,14 +320,46 @@ func (self *BlockChain) GasLimit() *big.Int { self.mu.RLock() defer self.mu.RUnlock() - return self.currentBlock.GasLimit() + if self.lightMode { + return self.currentHeader.GasLimit + } else { + return self.currentBlock.GasLimit() + } } func (self *BlockChain) LastBlockHash() common.Hash { self.mu.RLock() defer self.mu.RUnlock() - return self.currentBlock.Hash() + if self.lightMode { + return self.currentHeader.Hash() + } else { + return self.currentBlock.Hash() + } +} + +func (self *BlockChain) LastBlockNumberU64() uint64 { + self.mu.RLock() + defer self.mu.RUnlock() + + if self.lightMode { + return self.currentHeader.Number.Uint64() + } else { + return self.currentBlock.NumberU64() + } +} + +// CurrentBlockHeader returns the header of the current block, which is +// currentBlock.Header in full client mode and currentHeader in light mode +func (self *BlockChain) CurrentBlockHeader() *types.Header { + self.mu.RLock() + defer self.mu.RUnlock() + + if self.lightMode { + return self.currentHeader + } else { + return self.currentBlock.Header() + } } // CurrentHeader retrieves the current head header of the canonical chain. The @@ -318,9 +374,24 @@ func (self *BlockChain) CurrentHeader() *types.Header { // CurrentBlock retrieves the current head block of the canonical chain. The // block is retrieved from the blockchain's internal cache. func (self *BlockChain) CurrentBlock() *types.Block { + return self.CurrentBlockOdr(access.NoOdr) +} + +func (self *BlockChain) CurrentBlockOdr(ctx *access.OdrContext) *types.Block { self.mu.RLock() defer self.mu.RUnlock() + if self.lightMode { + hash := self.LastBlockHash() + var bhash common.Hash + if self.currentBlock != nil { + bhash = self.currentBlock.Hash() + } + if self.currentBlock == nil || bytes.Compare(bhash[:], hash[:]) != 0 { + self.currentBlock = self.GetBlock(hash, ctx) + } + } + return self.currentBlock } @@ -337,15 +408,29 @@ func (self *BlockChain) Status() (td *big.Int, currentBlock common.Hash, genesis self.mu.RLock() defer self.mu.RUnlock() - return self.GetTd(self.currentBlock.Hash()), self.currentBlock.Hash(), self.genesisBlock.Hash() + return self.GetTd(self.LastBlockHash()), self.LastBlockHash(), self.genesisBlock.Hash() } func (self *BlockChain) SetProcessor(proc types.BlockProcessor) { self.processor = proc } -func (self *BlockChain) State() (*state.StateDB, error) { - return state.New(self.CurrentBlock().Root(), self.chainDb) +func (self *BlockChain) StateNoOdr() (*state.StateDB, error) { + return self.State(access.NullCtx) +} + +func (self *BlockChain) State(ctx *access.OdrContext) (*state.StateDB, error) { + var root common.Hash + if self.lightMode { + if self.currentHeader != nil { + root = self.currentHeader.Root + } else { + root = common.Hash{} + } + } else { + root = self.CurrentBlock().Root() + } + return state.New(root, self.chainAccess, ctx) } // Reset purges the entire blockchain, restoring it to its genesis state. @@ -378,7 +463,7 @@ func (bc *BlockChain) ResetWithGenesisBlock(genesis *types.Block) { // Export writes the active chain to the given writer. func (self *BlockChain) Export(w io.Writer) error { - if err := self.ExportN(w, uint64(0), self.currentBlock.NumberU64()); err != nil { + if err := self.ExportN(w, uint64(0), self.CurrentBlock().NumberU64()); err != nil { return err } return nil @@ -396,7 +481,7 @@ func (self *BlockChain) ExportN(w io.Writer, first uint64, last uint64) error { glog.V(logger.Info).Infof("exporting %d blocks...\n", last-first+1) for nr := first; nr <= last; nr++ { - block := self.GetBlockByNumber(nr) + block := self.GetBlockByNumber(nr, access.NoOdr) if block == nil { return fmt.Errorf("export failed on #%d: not found", nr) } @@ -481,13 +566,13 @@ func (self *BlockChain) GetHeaderByNumber(number uint64) *types.Header { // GetBody retrieves a block body (transactions and uncles) from the database by // hash, caching it if found. -func (self *BlockChain) GetBody(hash common.Hash) *types.Body { +func (self *BlockChain) GetBody(hash common.Hash, ctx *access.OdrContext) *types.Body { // Short circuit if the body's already in the cache, retrieve otherwise if cached, ok := self.bodyCache.Get(hash); ok { body := cached.(*types.Body) return body } - body := GetBody(self.chainDb, hash) + body := GetBody(self.chainAccess, hash, ctx) if body == nil { return nil } @@ -498,12 +583,12 @@ func (self *BlockChain) GetBody(hash common.Hash) *types.Body { // GetBodyRLP retrieves a block body in RLP encoding from the database by hash, // caching it if found. -func (self *BlockChain) GetBodyRLP(hash common.Hash) rlp.RawValue { +func (self *BlockChain) GetBodyRLP(hash common.Hash, ctx *access.OdrContext) rlp.RawValue { // Short circuit if the body's already in the cache, retrieve otherwise if cached, ok := self.bodyRLPCache.Get(hash); ok { return cached.(rlp.RawValue) } - body := GetBodyRLP(self.chainDb, hash) + body := GetBodyRLP(self.chainAccess, hash, ctx) if len(body) == 0 { return nil } @@ -521,7 +606,7 @@ func (self *BlockChain) GetTd(hash common.Hash) *big.Int { } td := GetTd(self.chainDb, hash) if td == nil { - return nil + return big.NewInt(0) } // Cache the found body for next time and return self.tdCache.Add(hash, td) @@ -531,16 +616,20 @@ func (self *BlockChain) GetTd(hash common.Hash) *big.Int { // HasBlock checks if a block is fully present in the database or not, caching // it if present. func (bc *BlockChain) HasBlock(hash common.Hash) bool { - return bc.GetBlock(hash) != nil + return bc.GetBlock(hash, access.NoOdr) != nil +} + +func (self *BlockChain) GetBlockNoOdr(hash common.Hash) *types.Block { + return self.GetBlock(hash, access.NoOdr) } // GetBlock retrieves a block from the database by hash, caching it if found. -func (self *BlockChain) GetBlock(hash common.Hash) *types.Block { +func (self *BlockChain) GetBlock(hash common.Hash, ctx *access.OdrContext) *types.Block { // Short circuit if the block's already in the cache, retrieve otherwise if block, ok := self.blockCache.Get(hash); ok { return block.(*types.Block) } - block := GetBlock(self.chainDb, hash) + block := GetBlock(self.chainAccess, hash, ctx) if block == nil { return nil } @@ -551,12 +640,12 @@ func (self *BlockChain) GetBlock(hash common.Hash) *types.Block { // GetBlockByNumber retrieves a block from the database by number, caching it // (associated with its hash) if found. -func (self *BlockChain) GetBlockByNumber(number uint64) *types.Block { +func (self *BlockChain) GetBlockByNumber(number uint64, ctx *access.OdrContext) *types.Block { hash := GetCanonicalHash(self.chainDb, number) if hash == (common.Hash{}) { return nil } - return self.GetBlock(hash) + return self.GetBlock(hash, ctx) } // GetBlockHashesFromHash retrieves a number of block hashes starting at a given @@ -585,7 +674,7 @@ func (self *BlockChain) GetBlockHashesFromHash(hash common.Hash, max uint64) []c // GetBlocksFromHash returns the block corresponding to hash and up to n-1 ancestors. func (self *BlockChain) GetBlocksFromHash(hash common.Hash, n int) (blocks []*types.Block) { for i := 0; i < n; i++ { - block := self.GetBlock(hash) + block := self.GetBlock(hash, access.NoOdr) if block == nil { break } @@ -601,7 +690,7 @@ func (self *BlockChain) GetUnclesInChain(block *types.Block, length int) []*type uncles := []*types.Header{} for i := 0; block != nil && i < length; i++ { uncles = append(uncles, block.Uncles()...) - block = self.GetBlock(block.ParentHash()) + block = self.GetBlock(block.ParentHash(), access.NoOdr) } return uncles } @@ -824,16 +913,16 @@ func (self *BlockChain) Rollback(chain []common.Hash) { for i := len(chain) - 1; i >= 0; i-- { hash := chain[i] - if self.currentHeader.Hash() == hash { + if self.currentHeader != nil && self.currentHeader.Hash() == hash { self.currentHeader = self.GetHeader(self.currentHeader.ParentHash) WriteHeadHeaderHash(self.chainDb, self.currentHeader.Hash()) } - if self.currentFastBlock.Hash() == hash { - self.currentFastBlock = self.GetBlock(self.currentFastBlock.ParentHash()) + if self.currentFastBlock != nil && self.currentFastBlock.Hash() == hash { + self.currentFastBlock = self.GetBlock(self.currentFastBlock.ParentHash(), access.NoOdr) WriteHeadFastBlockHash(self.chainDb, self.currentFastBlock.Hash()) } - if self.currentBlock.Hash() == hash { - self.currentBlock = self.GetBlock(self.currentBlock.ParentHash()) + if self.currentBlock != nil && self.currentBlock.Hash() == hash { + self.currentBlock = self.GetBlock(self.currentBlock.ParentHash(), access.NoOdr) WriteHeadBlockHash(self.chainDb, self.currentBlock.Hash()) } } @@ -1161,12 +1250,12 @@ func (self *BlockChain) reorg(oldBlock, newBlock *types.Block) error { // first reduce whoever is higher bound if oldBlock.NumberU64() > newBlock.NumberU64() { // reduce old chain - for oldBlock = oldBlock; oldBlock != nil && oldBlock.NumberU64() != newBlock.NumberU64(); oldBlock = self.GetBlock(oldBlock.ParentHash()) { + for oldBlock = oldBlock; oldBlock != nil && oldBlock.NumberU64() != newBlock.NumberU64(); oldBlock = self.GetBlock(oldBlock.ParentHash(), access.NoOdr) { deletedTxs = append(deletedTxs, oldBlock.Transactions()...) } } else { // reduce new chain and append new chain blocks for inserting later on - for newBlock = newBlock; newBlock != nil && newBlock.NumberU64() != oldBlock.NumberU64(); newBlock = self.GetBlock(newBlock.ParentHash()) { + for newBlock = newBlock; newBlock != nil && newBlock.NumberU64() != oldBlock.NumberU64(); newBlock = self.GetBlock(newBlock.ParentHash(), access.NoOdr) { newChain = append(newChain, newBlock) } } @@ -1186,7 +1275,7 @@ func (self *BlockChain) reorg(oldBlock, newBlock *types.Block) error { newChain = append(newChain, newBlock) deletedTxs = append(deletedTxs, oldBlock.Transactions()...) - oldBlock, newBlock = self.GetBlock(oldBlock.ParentHash()), self.GetBlock(newBlock.ParentHash()) + oldBlock, newBlock = self.GetBlock(oldBlock.ParentHash(), access.NoOdr), self.GetBlock(newBlock.ParentHash(), access.NoOdr) if oldBlock == nil { return fmt.Errorf("Invalid old chain") } @@ -1209,7 +1298,7 @@ func (self *BlockChain) reorg(oldBlock, newBlock *types.Block) error { if err := PutTransactions(self.chainDb, block, block.Transactions()); err != nil { return err } - receipts := GetBlockReceipts(self.chainDb, block.Hash()) + receipts := GetBlockReceipts(self.chainAccess, block.Hash(), access.NoOdr) // write receipts if err := PutReceipts(self.chainDb, receipts); err != nil { return err diff --git a/core/blockchain_test.go b/core/blockchain_test.go index 8ddc5032b1..a82379446c 100644 --- a/core/blockchain_test.go +++ b/core/blockchain_test.go @@ -28,6 +28,7 @@ import ( "github.com/ethereum/ethash" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/crypto" @@ -51,13 +52,14 @@ func thePow() pow.PoW { func theBlockChain(db ethdb.Database, t *testing.T) *BlockChain { var eventMux event.TypeMux WriteTestNetGenesisBlock(db, 0) - blockchain, err := NewBlockChain(db, thePow(), &eventMux) + ca := access.NewDbChainAccess(db) + blockchain, err := NewBlockChain(ca, thePow(), &eventMux) if err != nil { t.Error("failed creating chainmanager:", err) t.FailNow() return nil } - blockMan := NewBlockProcessor(db, nil, blockchain, &eventMux) + blockMan := NewBlockProcessor(ca, nil, blockchain, &eventMux) blockchain.SetProcessor(blockMan) return blockchain @@ -73,8 +75,8 @@ func testFork(t *testing.T, processor *BlockProcessor, i, n int, full bool, comp // Assert the chains have the same header/block at #i var hash1, hash2 common.Hash if full { - hash1 = processor.bc.GetBlockByNumber(uint64(i)).Hash() - hash2 = processor2.bc.GetBlockByNumber(uint64(i)).Hash() + hash1 = processor.bc.GetBlockByNumber(uint64(i), access.NoOdr).Hash() + hash2 = processor2.bc.GetBlockByNumber(uint64(i), access.NoOdr).Hash() } else { hash1 = processor.bc.GetHeaderByNumber(uint64(i)).Hash() hash2 = processor2.bc.GetHeaderByNumber(uint64(i)).Hash() @@ -120,7 +122,7 @@ func testFork(t *testing.T, processor *BlockProcessor, i, n int, full bool, comp func printChain(bc *BlockChain) { for i := bc.CurrentBlock().Number().Uint64(); i > 0; i-- { - b := bc.GetBlockByNumber(uint64(i)) + b := bc.GetBlockByNumber(uint64(i), access.NoOdr) fmt.Printf("\t%x %v\n", b.Hash(), b.Difficulty()) } } @@ -138,8 +140,8 @@ func testBlockChainImport(chain []*types.Block, processor *BlockProcessor) error } // Manually insert the block into the database, but don't reorganize (allows subsequent testing) processor.bc.mu.Lock() - WriteTd(processor.chainDb, block.Hash(), new(big.Int).Add(block.Difficulty(), processor.bc.GetTd(block.ParentHash()))) - WriteBlock(processor.chainDb, block) + WriteTd(processor.chainAccess.Db(), block.Hash(), new(big.Int).Add(block.Difficulty(), processor.bc.GetTd(block.ParentHash()))) + WriteBlock(processor.chainAccess.Db(), block) processor.bc.mu.Unlock() } return nil @@ -155,8 +157,8 @@ func testHeaderChainImport(chain []*types.Header, processor *BlockProcessor) err } // Manually insert the header into the database, but don't reorganize (allows subsequent testing) processor.bc.mu.Lock() - WriteTd(processor.chainDb, header.Hash(), new(big.Int).Add(header.Difficulty, processor.bc.GetTd(header.ParentHash))) - WriteHeader(processor.chainDb, header) + WriteTd(processor.chainAccess.Db(), header.Hash(), new(big.Int).Add(header.Difficulty, processor.bc.GetTd(header.ParentHash))) + WriteHeader(processor.chainAccess.Db(), header) processor.bc.mu.Unlock() } return nil @@ -452,7 +454,8 @@ func makeBlockChainWithDiff(genesis *types.Block, d []int, seed byte) []*types.B func chm(genesis *types.Block, db ethdb.Database) *BlockChain { var eventMux event.TypeMux - bc := &BlockChain{chainDb: db, genesisBlock: genesis, eventMux: &eventMux, pow: FakePow{}, rand: rand.New(rand.NewSource(0))} + ca := access.NewDbChainAccess(db) + bc := &BlockChain{chainAccess: ca, chainDb: db, genesisBlock: genesis, eventMux: &eventMux, pow: FakePow{}, rand: rand.New(rand.NewSource(0))} bc.headerCache, _ = lru.New(100) bc.bodyCache, _ = lru.New(100) bc.bodyRLPCache, _ = lru.New(100) @@ -500,7 +503,7 @@ func testReorg(t *testing.T, first, second []int, td int64, full bool) { // Check that the chain is valid number and link wise if full { prev := bc.CurrentBlock() - for block := bc.GetBlockByNumber(bc.CurrentBlock().NumberU64() - 1); block.NumberU64() != 0; prev, block = block, bc.GetBlockByNumber(block.NumberU64()-1) { + for block := bc.GetBlockByNumber(bc.CurrentBlock().NumberU64()-1, access.NoOdr); block.NumberU64() != 0; prev, block = block, bc.GetBlockByNumber(block.NumberU64()-1, access.NoOdr) { if prev.ParentHash() != block.Hash() { t.Errorf("parent block hash mismatch: have %x, want %x", prev.ParentHash(), block.Hash()) } @@ -587,7 +590,7 @@ func testReorgBadHashes(t *testing.T, full bool) { defer func() { delete(BadHashes, headers[3].Hash()) }() } // Create a new chain manager and check it rolled back the state - ncm, err := NewBlockChain(db, FakePow{}, new(event.TypeMux)) + ncm, err := NewBlockChain(access.NewDbChainAccess(db), FakePow{}, new(event.TypeMux)) if err != nil { t.Fatalf("failed to create new chain manager: %v", err) } @@ -665,7 +668,7 @@ func testInsertNonceError(t *testing.T, full bool) { // Check that all no blocks after the failing block have been inserted. for j := 0; j < i-failAt; j++ { if full { - if block := bc.GetBlockByNumber(failNum + uint64(j)); block != nil { + if block := bc.GetBlockByNumber(failNum+uint64(j), access.NoOdr); block != nil { t.Errorf("test %d: invalid block in chain: %v", i, block) } } else { @@ -708,19 +711,21 @@ func TestFastVsFullChains(t *testing.T) { }) // Import the chain as an archive node for the comparison baseline archiveDb, _ := ethdb.NewMemDatabase() + archiveCa := access.NewDbChainAccess(archiveDb) WriteGenesisBlockForTesting(archiveDb, GenesisAccount{address, funds}) - archive, _ := NewBlockChain(archiveDb, FakePow{}, new(event.TypeMux)) - archive.SetProcessor(NewBlockProcessor(archiveDb, FakePow{}, archive, new(event.TypeMux))) + archive, _ := NewBlockChain(archiveCa, FakePow{}, new(event.TypeMux)) + archive.SetProcessor(NewBlockProcessor(archiveCa, FakePow{}, archive, new(event.TypeMux))) if n, err := archive.InsertChain(blocks); err != nil { t.Fatalf("failed to process block %d: %v", n, err) } // Fast import the chain as a non-archive node to test fastDb, _ := ethdb.NewMemDatabase() + fastCa := access.NewDbChainAccess(fastDb) WriteGenesisBlockForTesting(fastDb, GenesisAccount{address, funds}) - fast, _ := NewBlockChain(fastDb, FakePow{}, new(event.TypeMux)) - fast.SetProcessor(NewBlockProcessor(fastDb, FakePow{}, fast, new(event.TypeMux))) + fast, _ := NewBlockChain(fastCa, FakePow{}, new(event.TypeMux)) + fast.SetProcessor(NewBlockProcessor(fastCa, FakePow{}, fast, new(event.TypeMux))) headers := make([]*types.Header, len(blocks)) for i, block := range blocks { @@ -742,14 +747,14 @@ func TestFastVsFullChains(t *testing.T) { if fheader, aheader := fast.GetHeader(hash), archive.GetHeader(hash); fheader.Hash() != aheader.Hash() { t.Errorf("block #%d [%x]: header mismatch: have %v, want %v", num, hash, fheader, aheader) } - if fblock, ablock := fast.GetBlock(hash), archive.GetBlock(hash); fblock.Hash() != ablock.Hash() { + if fblock, ablock := fast.GetBlock(hash, access.NoOdr), archive.GetBlock(hash, access.NoOdr); fblock.Hash() != ablock.Hash() { t.Errorf("block #%d [%x]: block mismatch: have %v, want %v", num, hash, fblock, ablock) } else if types.DeriveSha(fblock.Transactions()) != types.DeriveSha(ablock.Transactions()) { t.Errorf("block #%d [%x]: transactions mismatch: have %v, want %v", num, hash, fblock.Transactions(), ablock.Transactions()) } else if types.CalcUncleHash(fblock.Uncles()) != types.CalcUncleHash(ablock.Uncles()) { t.Errorf("block #%d [%x]: uncles mismatch: have %v, want %v", num, hash, fblock.Uncles(), ablock.Uncles()) } - if freceipts, areceipts := GetBlockReceipts(fastDb, hash), GetBlockReceipts(archiveDb, hash); types.DeriveSha(freceipts) != types.DeriveSha(areceipts) { + if freceipts, areceipts := GetBlockReceipts(fastCa, hash, access.NoOdr), GetBlockReceipts(archiveCa, hash, access.NoOdr); types.DeriveSha(freceipts) != types.DeriveSha(areceipts) { t.Errorf("block #%d [%x]: receipts mismatch: have %v, want %v", num, hash, freceipts, areceipts) } } @@ -794,10 +799,10 @@ func TestLightVsFastVsFullChainHeads(t *testing.T) { } // Import the chain as an archive node and ensure all pointers are updated archiveDb, _ := ethdb.NewMemDatabase() - WriteGenesisBlockForTesting(archiveDb, GenesisAccount{address, funds}) + archiveCa := access.NewDbChainAccess(archiveDb) - archive, _ := NewBlockChain(archiveDb, FakePow{}, new(event.TypeMux)) - archive.SetProcessor(NewBlockProcessor(archiveDb, FakePow{}, archive, new(event.TypeMux))) + archive, _ := NewBlockChain(archiveCa, FakePow{}, new(event.TypeMux)) + archive.SetProcessor(NewBlockProcessor(archiveCa, FakePow{}, archive, new(event.TypeMux))) if n, err := archive.InsertChain(blocks); err != nil { t.Fatalf("failed to process block %d: %v", n, err) @@ -808,9 +813,10 @@ func TestLightVsFastVsFullChainHeads(t *testing.T) { // Import the chain as a non-archive node and ensure all pointers are updated fastDb, _ := ethdb.NewMemDatabase() + fastCa := access.NewDbChainAccess(fastDb) WriteGenesisBlockForTesting(fastDb, GenesisAccount{address, funds}) - fast, _ := NewBlockChain(fastDb, FakePow{}, new(event.TypeMux)) - fast.SetProcessor(NewBlockProcessor(fastDb, FakePow{}, fast, new(event.TypeMux))) + fast, _ := NewBlockChain(fastCa, FakePow{}, new(event.TypeMux)) + fast.SetProcessor(NewBlockProcessor(fastCa, FakePow{}, fast, new(event.TypeMux))) headers := make([]*types.Header, len(blocks)) for i, block := range blocks { @@ -828,9 +834,10 @@ func TestLightVsFastVsFullChainHeads(t *testing.T) { // Import the chain as a light node and ensure all pointers are updated lightDb, _ := ethdb.NewMemDatabase() + lightCa := access.NewDbChainAccess(lightDb) WriteGenesisBlockForTesting(lightDb, GenesisAccount{address, funds}) - light, _ := NewBlockChain(lightDb, FakePow{}, new(event.TypeMux)) - light.SetProcessor(NewBlockProcessor(lightDb, FakePow{}, light, new(event.TypeMux))) + light, _ := NewBlockChain(lightCa, FakePow{}, new(event.TypeMux)) + light.SetProcessor(NewBlockProcessor(lightCa, FakePow{}, light, new(event.TypeMux))) if n, err := light.InsertHeaderChain(headers, 1); err != nil { t.Fatalf("failed to insert header %d: %v", n, err) @@ -895,8 +902,9 @@ func TestChainTxReorgs(t *testing.T) { }) // Import the chain. This runs all block validation rules. evmux := &event.TypeMux{} - chainman, _ := NewBlockChain(db, FakePow{}, evmux) - chainman.SetProcessor(NewBlockProcessor(db, FakePow{}, chainman, evmux)) + ca := access.NewDbChainAccess(db) + chainman, _ := NewBlockChain(ca, FakePow{}, evmux) + chainman.SetProcessor(NewBlockProcessor(ca, FakePow{}, chainman, evmux)) if i, err := chainman.InsertChain(chain); err != nil { t.Fatalf("failed to insert original chain[%d]: %v", i, err) } @@ -929,7 +937,7 @@ func TestChainTxReorgs(t *testing.T) { if GetTransaction(db, tx.Hash()) != nil { t.Errorf("drop %d: tx found while shouldn't have been", i) } - if GetReceipt(db, tx.Hash()) != nil { + if GetReceipt(ca, tx.Hash(), access.NoOdr) != nil { t.Errorf("drop %d: receipt found while shouldn't have been", i) } } @@ -938,7 +946,7 @@ func TestChainTxReorgs(t *testing.T) { if GetTransaction(db, tx.Hash()) == nil { t.Errorf("add %d: expected tx to be found", i) } - if GetReceipt(db, tx.Hash()) == nil { + if GetReceipt(ca, tx.Hash(), access.NoOdr) == nil { t.Errorf("add %d: expected receipt to be found", i) } } @@ -947,7 +955,7 @@ func TestChainTxReorgs(t *testing.T) { if GetTransaction(db, tx.Hash()) == nil { t.Errorf("share %d: expected tx to be found", i) } - if GetReceipt(db, tx.Hash()) == nil { + if GetReceipt(ca, tx.Hash(), access.NoOdr) == nil { t.Errorf("share %d: expected receipt to be found", i) } } diff --git a/core/chain_makers.go b/core/chain_makers.go index 56e37a0fc2..2fd107324a 100644 --- a/core/chain_makers.go +++ b/core/chain_makers.go @@ -21,6 +21,7 @@ import ( "math/big" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/ethdb" @@ -164,7 +165,7 @@ func (b *BlockGen) OffsetTime(seconds int64) { // values. Inserting them into BlockChain requires use of FakePow or // a similar non-validating proof of work implementation. func GenerateChain(parent *types.Block, db ethdb.Database, n int, gen func(int, *BlockGen)) ([]*types.Block, []types.Receipts) { - statedb, err := state.New(parent.Root(), db) + statedb, err := state.New(parent.Root(), access.NewDbChainAccess(db), access.NullCtx) if err != nil { panic(err) } @@ -222,8 +223,9 @@ func newCanonical(n int, full bool) (ethdb.Database, *BlockProcessor, error) { // Initialize a fresh chain with only a genesis block genesis, _ := WriteTestNetGenesisBlock(db, 0) - blockchain, _ := NewBlockChain(db, FakePow{}, evmux) - processor := NewBlockProcessor(db, FakePow{}, blockchain, evmux) + ca := access.NewDbChainAccess(db) + blockchain, _ := NewBlockChain(ca, FakePow{}, evmux) + processor := NewBlockProcessor(ca, FakePow{}, blockchain, evmux) processor.bc.SetProcessor(processor) // Create and inject the requested chain diff --git a/core/chain_makers_test.go b/core/chain_makers_test.go index 7f47cf2888..b391118f70 100644 --- a/core/chain_makers_test.go +++ b/core/chain_makers_test.go @@ -20,6 +20,7 @@ import ( "fmt" "math/big" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/ethdb" @@ -77,14 +78,15 @@ func ExampleGenerateChain() { // Import the chain. This runs all block validation rules. evmux := &event.TypeMux{} - chainman, _ := NewBlockChain(db, FakePow{}, evmux) - chainman.SetProcessor(NewBlockProcessor(db, FakePow{}, chainman, evmux)) + ca := access.NewDbChainAccess(db) + chainman, _ := NewBlockChain(ca, FakePow{}, evmux) + chainman.SetProcessor(NewBlockProcessor(ca, FakePow{}, chainman, evmux)) if i, err := chainman.InsertChain(chain); err != nil { fmt.Printf("insert error (block %d): %v\n", i, err) return } - state, _ := chainman.State() + state, _ := chainman.State(access.NullCtx) fmt.Printf("last block: #%d\n", chainman.CurrentBlock().Number()) fmt.Println("balance of addr1:", state.GetBalance(addr1)) fmt.Println("balance of addr2:", state.GetBalance(addr2)) diff --git a/core/chain_util.go b/core/chain_util.go index ddff381a1d..1deef6b1c4 100644 --- a/core/chain_util.go +++ b/core/chain_util.go @@ -23,6 +23,7 @@ import ( "math/big" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/logger" @@ -183,15 +184,17 @@ func GetHeader(db ethdb.Database, hash common.Hash) *types.Header { } // GetBodyRLP retrieves the block body (transactions and uncles) in RLP encoding. -func GetBodyRLP(db ethdb.Database, hash common.Hash) rlp.RawValue { - data, _ := db.Get(append(append(blockPrefix, hash[:]...), bodySuffix...)) - return data +func GetBodyRLP(ca *access.ChainAccess, hash common.Hash, ctx *access.OdrContext) rlp.RawValue { + //fmt.Println("request block %v", hash) + r := &BlockAccess{db: ca.Db(), blockHash: hash} + ca.Retrieve(r, ctx) + return r.rlp } // GetBody retrieves the block body (transactons, uncles) corresponding to the // hash, nil if none found. -func GetBody(db ethdb.Database, hash common.Hash) *types.Body { - data := GetBodyRLP(db, hash) +func GetBody(ca *access.ChainAccess, hash common.Hash, ctx *access.OdrContext) *types.Body { + data := GetBodyRLP(ca, hash, ctx) if len(data) == 0 { return nil } @@ -220,13 +223,13 @@ func GetTd(db ethdb.Database, hash common.Hash) *big.Int { // GetBlock retrieves an entire block corresponding to the hash, assembling it // back from the stored header and body. -func GetBlock(db ethdb.Database, hash common.Hash) *types.Block { +func GetBlock(ca *access.ChainAccess, hash common.Hash, ctx *access.OdrContext) *types.Block { // Retrieve the block header and body contents - header := GetHeader(db, hash) + header := GetHeader(ca.Db(), hash) if header == nil { return nil } - body := GetBody(db, hash) + body := GetBody(ca, hash, ctx) if body == nil { return nil } diff --git a/core/chain_util_test.go b/core/chain_util_test.go index 0bbcbbe53d..522d4b528a 100644 --- a/core/chain_util_test.go +++ b/core/chain_util_test.go @@ -24,6 +24,7 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/crypto" @@ -121,6 +122,7 @@ func TestHeaderStorage(t *testing.T) { // Tests block body storage and retrieval operations. func TestBodyStorage(t *testing.T) { db, _ := ethdb.NewMemDatabase() + ca := access.NewDbChainAccess(db) // Create a test body to move around the database and make sure it's really new body := &types.Body{Uncles: []*types.Header{{Extra: []byte("test header")}}} @@ -129,19 +131,19 @@ func TestBodyStorage(t *testing.T) { rlp.Encode(hasher, body) hash := common.BytesToHash(hasher.Sum(nil)) - if entry := GetBody(db, hash); entry != nil { + if entry := GetBody(ca, hash, access.NoOdr); entry != nil { t.Fatalf("Non existent body returned: %v", entry) } // Write and verify the body in the database if err := WriteBody(db, hash, body); err != nil { t.Fatalf("Failed to write body into database: %v", err) } - if entry := GetBody(db, hash); entry == nil { + if entry := GetBody(ca, hash, access.NoOdr); entry == nil { t.Fatalf("Stored body not found") } else if types.DeriveSha(types.Transactions(entry.Transactions)) != types.DeriveSha(types.Transactions(body.Transactions)) || types.CalcUncleHash(entry.Uncles) != types.CalcUncleHash(body.Uncles) { t.Fatalf("Retrieved body mismatch: have %v, want %v", entry, body) } - if entry := GetBodyRLP(db, hash); entry == nil { + if entry := GetBodyRLP(ca, hash, access.NoOdr); entry == nil { t.Fatalf("Stored body RLP not found") } else { hasher := sha3.NewKeccak256() @@ -153,7 +155,7 @@ func TestBodyStorage(t *testing.T) { } // Delete the body and verify the execution DeleteBody(db, hash) - if entry := GetBody(db, hash); entry != nil { + if entry := GetBody(ca, hash, access.NoOdr); entry != nil { t.Fatalf("Deleted body returned: %v", entry) } } @@ -161,6 +163,7 @@ func TestBodyStorage(t *testing.T) { // Tests block storage and retrieval operations. func TestBlockStorage(t *testing.T) { db, _ := ethdb.NewMemDatabase() + ca := access.NewDbChainAccess(db) // Create a test block to move around the database and make sure it's really new block := types.NewBlockWithHeader(&types.Header{ @@ -169,20 +172,20 @@ func TestBlockStorage(t *testing.T) { TxHash: types.EmptyRootHash, ReceiptHash: types.EmptyRootHash, }) - if entry := GetBlock(db, block.Hash()); entry != nil { + if entry := GetBlock(ca, block.Hash(), access.NoOdr); entry != nil { t.Fatalf("Non existent block returned: %v", entry) } if entry := GetHeader(db, block.Hash()); entry != nil { t.Fatalf("Non existent header returned: %v", entry) } - if entry := GetBody(db, block.Hash()); entry != nil { + if entry := GetBody(ca, block.Hash(), access.NoOdr); entry != nil { t.Fatalf("Non existent body returned: %v", entry) } // Write and verify the block in the database if err := WriteBlock(db, block); err != nil { t.Fatalf("Failed to write block into database: %v", err) } - if entry := GetBlock(db, block.Hash()); entry == nil { + if entry := GetBlock(ca, block.Hash(), access.NoOdr); entry == nil { t.Fatalf("Stored block not found") } else if entry.Hash() != block.Hash() { t.Fatalf("Retrieved block mismatch: have %v, want %v", entry, block) @@ -192,20 +195,20 @@ func TestBlockStorage(t *testing.T) { } else if entry.Hash() != block.Header().Hash() { t.Fatalf("Retrieved header mismatch: have %v, want %v", entry, block.Header()) } - if entry := GetBody(db, block.Hash()); entry == nil { + if entry := GetBody(ca, block.Hash(), access.NoOdr); entry == nil { t.Fatalf("Stored body not found") } else if types.DeriveSha(types.Transactions(entry.Transactions)) != types.DeriveSha(block.Transactions()) || types.CalcUncleHash(entry.Uncles) != types.CalcUncleHash(block.Uncles()) { t.Fatalf("Retrieved body mismatch: have %v, want %v", entry, &types.Body{block.Transactions(), block.Uncles()}) } // Delete the block and verify the execution DeleteBlock(db, block.Hash()) - if entry := GetBlock(db, block.Hash()); entry != nil { + if entry := GetBlock(ca, block.Hash(), access.NoOdr); entry != nil { t.Fatalf("Deleted block returned: %v", entry) } if entry := GetHeader(db, block.Hash()); entry != nil { t.Fatalf("Deleted header returned: %v", entry) } - if entry := GetBody(db, block.Hash()); entry != nil { + if entry := GetBody(ca, block.Hash(), access.NoOdr); entry != nil { t.Fatalf("Deleted body returned: %v", entry) } } @@ -213,6 +216,7 @@ func TestBlockStorage(t *testing.T) { // Tests that partial block contents don't get reassembled into full blocks. func TestPartialBlockStorage(t *testing.T) { db, _ := ethdb.NewMemDatabase() + ca := access.NewDbChainAccess(db) block := types.NewBlockWithHeader(&types.Header{ Extra: []byte("test block"), UncleHash: types.EmptyUncleHash, @@ -223,7 +227,7 @@ func TestPartialBlockStorage(t *testing.T) { if err := WriteHeader(db, block.Header()); err != nil { t.Fatalf("Failed to write header into database: %v", err) } - if entry := GetBlock(db, block.Hash()); entry != nil { + if entry := GetBlock(ca, block.Hash(), access.NoOdr); entry != nil { t.Fatalf("Non existent block returned: %v", entry) } DeleteHeader(db, block.Hash()) @@ -232,7 +236,7 @@ func TestPartialBlockStorage(t *testing.T) { if err := WriteBody(db, block.Hash(), &types.Body{block.Transactions(), block.Uncles()}); err != nil { t.Fatalf("Failed to write body into database: %v", err) } - if entry := GetBlock(db, block.Hash()); entry != nil { + if entry := GetBlock(ca, block.Hash(), access.NoOdr); entry != nil { t.Fatalf("Non existent block returned: %v", entry) } DeleteBody(db, block.Hash()) @@ -244,7 +248,7 @@ func TestPartialBlockStorage(t *testing.T) { if err := WriteBody(db, block.Hash(), &types.Body{block.Transactions(), block.Uncles()}); err != nil { t.Fatalf("Failed to write body into database: %v", err) } - if entry := GetBlock(db, block.Hash()); entry == nil { + if entry := GetBlock(ca, block.Hash(), access.NoOdr); entry == nil { t.Fatalf("Stored block not found") } else if entry.Hash() != block.Hash() { t.Fatalf("Retrieved block mismatch: have %v, want %v", entry, block) diff --git a/core/core_access.go b/core/core_access.go new file mode 100644 index 0000000000..4c8d7f81fe --- /dev/null +++ b/core/core_access.go @@ -0,0 +1,153 @@ +// Copyright 2015 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 . + +package core + +import ( + "bytes" + + "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" + "github.com/ethereum/go-ethereum/core/types" + "github.com/ethereum/go-ethereum/crypto" + "github.com/ethereum/go-ethereum/ethdb" + "github.com/ethereum/go-ethereum/logger" + "github.com/ethereum/go-ethereum/logger/glog" + "github.com/ethereum/go-ethereum/rlp" +) + +type BlockAccess struct { + db ethdb.Database + blockHash common.Hash + rlp []byte +} + +func (self *BlockAccess) Request(peer *access.Peer) error { + glog.V(access.LogLevel).Infof("ODR: requesting body of block %08x from peer %v", self.blockHash[:4], peer.Id()) + return peer.GetBlockBodies([]common.Hash{self.blockHash}) +} + +func (self *BlockAccess) Valid(msg *access.Msg) bool { + glog.V(access.LogLevel).Infof("ODR: validating body of block %08x", self.blockHash[:4]) + if msg.MsgType != access.MsgBlockBodies { + glog.V(access.LogLevel).Infof("ODR: invalid message type") + return false + } + bodies := msg.Obj.([]*types.Body) + if len(bodies) != 1 { + glog.V(access.LogLevel).Infof("ODR: invalid number of entries: %d", len(bodies)) + return false + } + body := bodies[0] + header := GetHeader(self.db, self.blockHash) + if header == nil { + glog.V(access.LogLevel).Infof("ODR: header not found for block %08x", self.blockHash[:4]) + return false + } + txHash := types.DeriveSha(types.Transactions(body.Transactions)) + if header.TxHash != txHash { + glog.V(access.LogLevel).Infof("ODR: header.TxHash %08x does not match received txHash %08x", header.TxHash[:4], txHash[:4]) + return false + } + uncleHash := types.CalcUncleHash(body.Uncles) + if header.UncleHash != uncleHash { + glog.V(access.LogLevel).Infof("ODR: header.UncleHash %08x does not match received uncleHash %08x", header.UncleHash[:4], uncleHash[:4]) + return false + } + data, err := rlp.EncodeToBytes(body) + if err != nil { + glog.V(access.LogLevel).Infof("ODR: body RLP encode error: %v", err) + return false + } + self.rlp = data + glog.V(access.LogLevel).Infof("ODR: validation successful") + return true +} + +func (self *BlockAccess) DbGet() bool { + self.rlp, _ = self.db.Get(append(append(blockPrefix, self.blockHash[:]...), bodySuffix...)) + glog.V(access.LogLevel).Infof("ODR: get body %08x len = %d", self.blockHash[:4], len(self.rlp)) + return len(self.rlp) != 0 +} + +func (self *BlockAccess) DbPut() { + self.db.Put(append(append(blockPrefix, self.blockHash[:]...), bodySuffix...), self.rlp) + glog.V(access.LogLevel).Infof("ODR: put body %08x len = %d", self.blockHash[:4], len(self.rlp)) +} + +type ReceiptsAccess struct { + db ethdb.Database + blockHash common.Hash + receipts types.Receipts +} + +func (self *ReceiptsAccess) Request(peer *access.Peer) error { + glog.V(access.LogLevel).Infof("ODR: requesting receipts for block %08x from peer %v", self.blockHash[:4], peer.Id()) + return peer.GetReceipts([]common.Hash{self.blockHash}) +} + +func (self *ReceiptsAccess) Valid(msg *access.Msg) bool { + glog.V(access.LogLevel).Infof("ODR: validating receipts for block %08x", self.blockHash[:4]) + if msg.MsgType != access.MsgReceipts { + glog.V(access.LogLevel).Infof("ODR: invalid message type") + return false + } + receipts := msg.Obj.([]types.Receipts) + if len(receipts) != 1 { + glog.V(access.LogLevel).Infof("ODR: invalid number of entries: %d", len(receipts)) + return false + } + data, err := rlp.EncodeToBytes(receipts[0]) + if err != nil { + glog.V(access.LogLevel).Infof("ODR: RLP encode error: %v", err) + return false + } + hash := crypto.Sha3Hash(data) + header := GetHeader(self.db, self.blockHash) + if header == nil { + glog.V(access.LogLevel).Infof("ODR: header not found for block %08x", self.blockHash[:4]) + return false + } + if bytes.Compare(header.ReceiptHash[:], hash[:]) != 0 { + glog.V(access.LogLevel).Infof("ODR: header receipts hash %08x does not match calculated RLP hash %08x", header.ReceiptHash[:4], hash[:4]) + return false + } + self.receipts = receipts[0] + glog.V(access.LogLevel).Infof("ODR: validation successful") + return true +} + +func (self *ReceiptsAccess) DbGet() bool { + data, _ := self.db.Get(append(blockReceiptsPre, self.blockHash[:]...)) + if len(data) == 0 { + return false + } + rs := []*types.ReceiptForStorage{} + if err := rlp.DecodeBytes(data, &rs); err != nil { + glog.V(logger.Error).Infof("invalid receipt array RLP for hash %x: %v", self.blockHash, err) + return false + } + self.receipts = make(types.Receipts, len(rs)) + for i, receipt := range rs { + self.receipts[i] = (*types.Receipt)(receipt) + } + return true +} + +func (self *ReceiptsAccess) DbPut() { + PutBlockReceipts(self.db, self.blockHash, self.receipts) + PutReceipts(self.db, self.receipts) +} diff --git a/core/genesis.go b/core/genesis.go index dac5de92ff..ff5ec73159 100644 --- a/core/genesis.go +++ b/core/genesis.go @@ -25,6 +25,7 @@ import ( "strings" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/ethdb" @@ -61,7 +62,8 @@ func WriteGenesisBlock(chainDb ethdb.Database, reader io.Reader) (*types.Block, } // creating with empty hash always works - statedb, _ := state.New(common.Hash{}, chainDb) + ca := access.NewDbChainAccess(chainDb) + statedb, _ := state.New(common.Hash{}, ca, access.NullCtx) for addr, account := range genesis.Alloc { address := common.HexToAddress(addr) statedb.AddBalance(address, common.String2Big(account.Balance)) @@ -85,7 +87,7 @@ func WriteGenesisBlock(chainDb ethdb.Database, reader io.Reader) (*types.Block, Root: root, }, nil, nil, nil) - if block := GetBlock(chainDb, block.Hash()); block != nil { + if block := GetBlock(ca, block.Hash(), access.NoOdr); block != nil { glog.V(logger.Info).Infoln("Genesis block already in chain. Writing canonical number") err := WriteCanonicalHash(chainDb, block.Hash(), block.NumberU64()) if err != nil { @@ -118,7 +120,7 @@ func WriteGenesisBlock(chainDb ethdb.Database, reader io.Reader) (*types.Block, // GenesisBlockForTesting creates a block in which addr has the given wei balance. // The state trie of the block is written to db. the passed db needs to contain a state root func GenesisBlockForTesting(db ethdb.Database, addr common.Address, balance *big.Int) *types.Block { - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) obj := statedb.GetOrNewStateObject(addr) obj.SetBalance(balance) root, err := statedb.Commit() diff --git a/core/manager.go b/core/manager.go index 289c87c112..c2dbe8bfef 100644 --- a/core/manager.go +++ b/core/manager.go @@ -18,6 +18,7 @@ package core import ( "github.com/ethereum/go-ethereum/accounts" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/event" ) @@ -29,6 +30,7 @@ type Backend interface { BlockChain() *BlockChain TxPool() *TxPool ChainDb() ethdb.Database + ChainAccess() *access.ChainAccess DappDb() ethdb.Database EventMux() *event.TypeMux } diff --git a/core/state/dump.go b/core/state/dump.go index 9acb8a0244..173edc847b 100644 --- a/core/state/dump.go +++ b/core/state/dump.go @@ -21,6 +21,7 @@ import ( "fmt" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" ) type Account struct { @@ -45,7 +46,7 @@ func (self *StateDB) RawDump() World { it := self.trie.Iterator() for it.Next() { addr := self.trie.GetKey(it.Key) - stateObject := NewStateObjectFromBytes(common.BytesToAddress(addr), it.Value, self.db) + stateObject := NewStateObjectFromBytes(common.BytesToAddress(addr), it.Value, self.ca, access.NullCtx) account := Account{Balance: stateObject.balance.String(), Nonce: stateObject.nonce, Root: common.Bytes2Hex(stateObject.Root()), CodeHash: common.Bytes2Hex(stateObject.codeHash)} account.Storage = make(map[string]string) diff --git a/core/state/managed_state.go b/core/state/managed_state.go index 4df0479791..1a1e7e4909 100644 --- a/core/state/managed_state.go +++ b/core/state/managed_state.go @@ -20,6 +20,7 @@ import ( "sync" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" ) type account struct { @@ -39,7 +40,7 @@ type ManagedState struct { // ManagedState returns a new managed state with the statedb as it's backing layer func ManageState(statedb *StateDB) *ManagedState { return &ManagedState{ - StateDB: statedb.Copy(), + StateDB: statedb.Copy(access.NullCtx), accounts: make(map[string]*account), } } diff --git a/core/state/managed_state_test.go b/core/state/managed_state_test.go index 0b53a42c54..8910bea4c6 100644 --- a/core/state/managed_state_test.go +++ b/core/state/managed_state_test.go @@ -20,6 +20,7 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/ethdb" ) @@ -27,7 +28,7 @@ var addr = common.BytesToAddress([]byte("test")) func create() (*ManagedState, *account) { db, _ := ethdb.NewMemDatabase() - statedb, _ := New(common.Hash{}, db) + statedb, _ := New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) ms := ManageState(statedb) so := &StateObject{address: addr, nonce: 100} ms.StateDB.stateObjects[addr.Str()] = so diff --git a/core/state/state_access.go b/core/state/state_access.go new file mode 100644 index 0000000000..d112b28188 --- /dev/null +++ b/core/state/state_access.go @@ -0,0 +1,161 @@ +// Copyright 2014 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 . + +// Package state provides a caching layer atop the Ethereum state trie. +package state + +import ( + "bytes" + //"fmt" + + "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" + "github.com/ethereum/go-ethereum/crypto" + "github.com/ethereum/go-ethereum/crypto/sha3" + "github.com/ethereum/go-ethereum/ethdb" + "github.com/ethereum/go-ethereum/logger/glog" + "github.com/ethereum/go-ethereum/trie" +) + +type TrieAccess struct { + ca *access.ChainAccess + root common.Hash + trieDb trie.Database +} + +func NewTrieAccess(ca *access.ChainAccess, root common.Hash, trieDb trie.Database) *TrieAccess { + return &TrieAccess{ + ca: ca, + root: root, + trieDb: trieDb, + } +} + +func (self *TrieAccess) RetrieveKey(key []byte, ctx *access.OdrContext) bool { + //fmt.Println("request trie %v key %v", self.root, key) + r := &TrieEntryAccess{root: self.root, trieDb: self.trieDb, key: key} + return self.ca.Retrieve(r, ctx) == nil +} + +func (self *TrieAccess) OdrEnabled() bool { + return self.ca.OdrEnabled() +} + +type TrieEntryAccess struct { + root common.Hash + trieDb trie.Database + key, value []byte + proof trie.MerkleProof + skipLevels int // set by DbGet() if unsuccessful +} + +func (self *TrieEntryAccess) Request(peer *access.Peer) error { + glog.V(access.LogLevel).Infof("ODR: requesting trie root %08x key %08x from peer %v", self.root[:4], self.key[:4], peer.Id()) + req := &access.ProofReq{ + Root: self.root, + Key: self.key, + } + return peer.GetProofs([]*access.ProofReq{req}) +} + +func (self *TrieEntryAccess) Valid(msg *access.Msg) bool { + glog.V(access.LogLevel).Infof("ODR: validating trie root %08x key %08x", self.root[:4], self.key[:4]) + + if msg.MsgType != access.MsgProofs { + glog.V(access.LogLevel).Infof("ODR: invalid message type") + return false + } + proofs := msg.Obj.([]trie.MerkleProof) + if len(proofs) != 1 { + glog.V(access.LogLevel).Infof("ODR: invalid number of entries: %d", len(proofs)) + return false + } + value, err := trie.VerifyProof(self.root, self.key, proofs[0]) + if err != nil { + glog.V(access.LogLevel).Infof("ODR: merkle proof verification error: %v", err) + return false + } + self.proof = proofs[0] + self.value = value + glog.V(access.LogLevel).Infof("ODR: validation successful") + return true +} + +func (self *TrieEntryAccess) DbGet() bool { + return false // not used +} + +func (self *TrieEntryAccess) DbPut() { + trie.StoreProof(self.trieDb, self.proof) +} + +type NodeDataAccess struct { + db ethdb.Database + hash common.Hash + data []byte +} + +func (self *NodeDataAccess) Request(peer *access.Peer) error { + glog.V(access.LogLevel).Infof("ODR: requesting node data for hash %08x from peer %v", self.hash[:4], peer.Id()) + return peer.GetNodeData([]common.Hash{self.hash}) +} + +func (self *NodeDataAccess) Valid(msg *access.Msg) bool { + glog.V(access.LogLevel).Infof("ODR: validating node data for hash %08x", self.hash[:4]) + if msg.MsgType != access.MsgNodeData { + glog.V(access.LogLevel).Infof("ODR: invalid message type") + return false + } + reply := msg.Obj.([][]byte) + if len(reply) != 1 { + glog.V(access.LogLevel).Infof("ODR: invalid number of entries: %d", len(reply)) + return false + } + data := reply[0] + hash := crypto.Sha3Hash(data) + if bytes.Compare(self.hash[:], hash[:]) != 0 { + glog.V(access.LogLevel).Infof("ODR: requested hash %08x does not match received data hash %08x", self.hash[:4], hash[:4]) + return false + } + self.data = data + glog.V(access.LogLevel).Infof("ODR: validation successful") + return true +} + +func (self *NodeDataAccess) DbGet() bool { + data, _ := self.db.Get(self.hash[:]) + if len(data) == 0 { + return false + } + self.data = data + return true +} + +func (self *NodeDataAccess) DbPut() { + self.db.Put(self.hash[:], self.data) +} + +var sha3_nil = sha3.NewKeccak256().Sum(nil) + +func RetrieveNodeData(ca *access.ChainAccess, hash common.Hash, ctx *access.OdrContext) []byte { + //fmt.Println("request node data %v", hash) + if bytes.Compare(hash[:], sha3_nil) == 0 { + return nil + } + r := &NodeDataAccess{db: ca.Db(), hash: hash} + ca.Retrieve(r, ctx) + return r.data +} diff --git a/core/state/state_object.go b/core/state/state_object.go index c06e3d2276..27897b1dc8 100644 --- a/core/state/state_object.go +++ b/core/state/state_object.go @@ -22,8 +22,8 @@ import ( "math/big" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/crypto" - "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" "github.com/ethereum/go-ethereum/rlp" @@ -57,7 +57,7 @@ func (self Storage) Copy() Storage { type StateObject struct { // State database for storing state changes - db ethdb.Database + ca *access.ChainAccess trie *trie.SecureTrie // Address belonging to this account @@ -83,14 +83,14 @@ type StateObject struct { dirty bool } -func NewStateObject(address common.Address, db ethdb.Database) *StateObject { - object := &StateObject{db: db, address: address, balance: new(big.Int), dirty: true} - object.trie, _ = trie.NewSecure(common.Hash{}, db) +func NewStateObject(address common.Address, ca *access.ChainAccess) *StateObject { + object := &StateObject{ca: ca, address: address, balance: new(big.Int), dirty: true} + object.trie, _ = trie.NewSecure(common.Hash{}, ca.Db(), nil) object.storage = make(Storage) return object } -func NewStateObjectFromBytes(address common.Address, data []byte, db ethdb.Database) *StateObject { +func NewStateObjectFromBytes(address common.Address, data []byte, ca *access.ChainAccess, ctx *access.OdrContext) *StateObject { var extobject struct { Nonce uint64 Balance *big.Int @@ -102,20 +102,20 @@ func NewStateObjectFromBytes(address common.Address, data []byte, db ethdb.Datab glog.Errorf("can't decode state object %x: %v", address, err) return nil } - trie, err := trie.NewSecure(extobject.Root, db) + trie, err := trie.NewSecure(extobject.Root, ca.Db(), NewTrieAccess(ca, extobject.Root, ca.Db())) if err != nil { // TODO: bubble this up or panic glog.Errorf("can't create account trie with root %x: %v", extobject.Root[:], err) return nil } - object := &StateObject{address: address, db: db} + object := &StateObject{address: address, ca: ca} object.nonce = extobject.Nonce object.balance = extobject.Balance object.codeHash = extobject.CodeHash object.trie = trie object.storage = make(map[string]common.Hash) - object.code, _ = db.Get(extobject.CodeHash) + object.code = RetrieveNodeData(ca, common.BytesToHash(extobject.CodeHash), ctx) return object } @@ -128,30 +128,31 @@ func (self *StateObject) MarkForDeletion() { } } -func (c *StateObject) getAddr(addr common.Hash) common.Hash { +func (c *StateObject) getAddr(addr common.Hash, ctx *access.OdrContext) common.Hash { var ret []byte - rlp.DecodeBytes(c.trie.Get(addr[:]), &ret) + value := c.trie.Get(addr[:], ctx) + rlp.DecodeBytes(value, &ret) return common.BytesToHash(ret) } -func (c *StateObject) setAddr(addr []byte, value common.Hash) { +func (c *StateObject) setAddr(addr []byte, value common.Hash, ctx *access.OdrContext) { v, err := rlp.EncodeToBytes(bytes.TrimLeft(value[:], "\x00")) if err != nil { // if RLPing failed we better panic and not fail silently. This would be considered a consensus issue panic(err) } - c.trie.Update(addr, v) + c.trie.Update(addr, v, ctx) } func (self *StateObject) Storage() Storage { return self.storage } -func (self *StateObject) GetState(key common.Hash) common.Hash { +func (self *StateObject) GetState(key common.Hash, ctx *access.OdrContext) common.Hash { strkey := key.Str() value, exists := self.storage[strkey] if !exists { - value = self.getAddr(key) + value = self.getAddr(key, ctx) if (value != common.Hash{}) { self.storage[strkey] = value } @@ -166,14 +167,14 @@ func (self *StateObject) SetState(k, value common.Hash) { } // Update updates the current cached storage to the trie -func (self *StateObject) Update() { +func (self *StateObject) Update(ctx *access.OdrContext) { for key, value := range self.storage { if (value == common.Hash{}) { - self.trie.Delete([]byte(key)) + self.trie.Delete([]byte(key), ctx) continue } - self.setAddr([]byte(key), value) + self.setAddr([]byte(key), value, ctx) } } @@ -206,7 +207,7 @@ func (c *StateObject) St() Storage { func (c *StateObject) ReturnGas(gas, price *big.Int) {} func (self *StateObject) Copy() *StateObject { - stateObject := NewStateObject(self.Address(), self.db) + stateObject := NewStateObject(self.Address(), self.ca) stateObject.balance.Set(self.balance) stateObject.codeHash = common.CopyBytes(self.codeHash) stateObject.nonce = self.nonce diff --git a/core/state/state_test.go b/core/state/state_test.go index 7ddbe11a13..ba4b085bab 100644 --- a/core/state/state_test.go +++ b/core/state/state_test.go @@ -24,6 +24,7 @@ import ( checker "gopkg.in/check.v1" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/ethdb" ) @@ -77,12 +78,12 @@ func (s *StateSuite) TestDump(c *checker.C) { func (s *StateSuite) SetUpTest(c *checker.C) { db, _ := ethdb.NewMemDatabase() - s.state, _ = New(common.Hash{}, db) + s.state, _ = New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) } func TestNull(t *testing.T) { db, _ := ethdb.NewMemDatabase() - state, _ := New(common.Hash{}, db) + state, _ := New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) address := common.HexToAddress("0x823140710bf13990e4500136726d8b55") state.CreateAccount(address) @@ -105,7 +106,7 @@ func (s *StateSuite) TestSnapshot(c *checker.C) { // set inital state object value s.state.SetState(stateobjaddr, storageaddr, data1) // get snapshot of current state - snapshot := s.state.Copy() + snapshot := s.state.Copy(access.NullCtx) // set new state object value s.state.SetState(stateobjaddr, storageaddr, data2) @@ -122,7 +123,7 @@ func (s *StateSuite) TestSnapshot(c *checker.C) { // printing/logging in tests (-check.vv does not work) func TestSnapshot2(t *testing.T) { db, _ := ethdb.NewMemDatabase() - state, _ := New(common.Hash{}, db) + state, _ := New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) stateobjaddr0 := toAddr([]byte("so0")) stateobjaddr1 := toAddr([]byte("so1")) @@ -161,7 +162,7 @@ func TestSnapshot2(t *testing.T) { t.Fatalf("deleted object not nil when getting") } - snapshot := state.Copy() + snapshot := state.Copy(access.NullCtx) state.Set(snapshot) so0Restored := state.GetStateObject(stateobjaddr0) diff --git a/core/state/statedb.go b/core/state/statedb.go index a9de71409c..39632407ac 100644 --- a/core/state/statedb.go +++ b/core/state/statedb.go @@ -22,6 +22,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core/vm" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" @@ -38,8 +39,9 @@ var StartingNonce uint64 // * Contracts // * Accounts type StateDB struct { - db ethdb.Database + ca *access.ChainAccess trie *trie.SecureTrie + ctx *access.OdrContext stateObjects map[string]*StateObject @@ -52,15 +54,16 @@ type StateDB struct { } // Create a new state from a given trie -func New(root common.Hash, db ethdb.Database) (*StateDB, error) { - tr, err := trie.NewSecure(root, db) +func New(root common.Hash, ca *access.ChainAccess, ctx *access.OdrContext) (*StateDB, error) { + tr, err := trie.NewSecure(root, ca.Db(), NewTrieAccess(ca, root, ca.Db())) if err != nil { glog.Errorf("can't create state trie with root %x: %v", root[:], err) return nil, err } return &StateDB{ - db: db, + ca: ca, trie: tr, + ctx: ctx, stateObjects: make(map[string]*StateObject), refund: new(big.Int), logs: make(map[common.Hash]vm.Logs), @@ -141,7 +144,7 @@ func (self *StateDB) GetCode(addr common.Address) []byte { func (self *StateDB) GetState(a common.Address, b common.Hash) common.Hash { stateObject := self.GetStateObject(a) if stateObject != nil { - return stateObject.GetState(b) + return stateObject.GetState(b, self.ctx) } return common.Hash{} @@ -208,10 +211,10 @@ func (self *StateDB) UpdateStateObject(stateObject *StateObject) { //addr := stateObject.Address() if len(stateObject.CodeHash()) > 0 { - self.db.Put(stateObject.CodeHash(), stateObject.code) + self.ca.Db().Put(stateObject.CodeHash(), stateObject.code) } addr := stateObject.Address() - self.trie.Update(addr[:], stateObject.RlpEncode()) + self.trie.Update(addr[:], stateObject.RlpEncode(), self.ctx) } // Delete the given state object and delete it from the state trie @@ -219,8 +222,7 @@ func (self *StateDB) DeleteStateObject(stateObject *StateObject) { stateObject.deleted = true addr := stateObject.Address() - self.trie.Delete(addr[:]) - //delete(self.stateObjects, addr.Str()) + self.trie.Delete(addr[:], self.ctx) } // Retrieve a state object given my the address. Nil if not found @@ -234,12 +236,12 @@ func (self *StateDB) GetStateObject(addr common.Address) (stateObject *StateObje return stateObject } - data := self.trie.Get(addr[:]) + data := self.trie.Get(addr[:], self.ctx) if len(data) == 0 { return nil } - stateObject = NewStateObjectFromBytes(addr, []byte(data), self.db) + stateObject = NewStateObjectFromBytes(addr, []byte(data), self.ca, self.ctx) self.SetStateObject(stateObject) return stateObject @@ -265,7 +267,7 @@ func (self *StateDB) newStateObject(addr common.Address) *StateObject { glog.Infof("(+) %x\n", addr) } - stateObject := NewStateObject(addr, self.db) + stateObject := NewStateObject(addr, self.ca) stateObject.SetNonce(StartingNonce) self.stateObjects[addr.Str()] = stateObject @@ -295,9 +297,13 @@ func (self *StateDB) CreateAccount(addr common.Address) vm.Account { // Setting, copying of the state methods // -func (self *StateDB) Copy() *StateDB { +func (self *StateDB) CopyWithCtx() *StateDB { + return self.Copy(self.ctx) +} + +func (self *StateDB) Copy(ctx *access.OdrContext) *StateDB { // ignore error - we assume state-to-be-copied always exists - state, _ := New(common.Hash{}, self.db) + state, _ := New(common.Hash{}, self.ca, ctx) state.trie = self.trie for k, stateObject := range self.stateObjects { state.stateObjects[k] = stateObject.Copy() @@ -337,7 +343,7 @@ func (s *StateDB) IntermediateRoot() common.Hash { if stateObject.remove { s.DeleteStateObject(stateObject) } else { - stateObject.Update() + stateObject.Update(s.ctx) s.UpdateStateObject(stateObject) } stateObject.dirty = false @@ -348,14 +354,14 @@ func (s *StateDB) IntermediateRoot() common.Hash { // Commit commits all state changes to the database. func (s *StateDB) Commit() (root common.Hash, err error) { - return s.commit(s.db) + return s.commit(s.ca.Db()) } // CommitBatch commits all state changes to a write batch but does not // execute the batch. It is used to validate state changes against // the root hash stored in a block. func (s *StateDB) CommitBatch() (root common.Hash, batch ethdb.Batch) { - batch = s.db.NewBatch() + batch = s.ca.Db().NewBatch() root, _ = s.commit(batch) return root, batch } @@ -370,12 +376,12 @@ func (s *StateDB) commit(db trie.DatabaseWriter) (common.Hash, error) { s.DeleteStateObject(stateObject) } else { // Write any storage changes in the state object to its trie. - stateObject.Update() + stateObject.Update(s.ctx) // Commit the trie of the object to the batch. // This updates the trie root internally, so // getting the root hash of the storage trie // through UpdateStateObject is fast. - if _, err := stateObject.trie.CommitTo(db); err != nil { + if _, err := stateObject.trie.CommitTo(db, s.ctx); err != nil { return common.Hash{}, err } // Update the object in the account trie. @@ -383,7 +389,7 @@ func (s *StateDB) commit(db trie.DatabaseWriter) (common.Hash, error) { } stateObject.dirty = false } - return s.trie.CommitTo(db) + return s.trie.CommitTo(db, s.ctx) } func (self *StateDB) Refunds() *big.Int { diff --git a/core/state/sync_test.go b/core/state/sync_test.go index 0dab372ba0..8f6196556c 100644 --- a/core/state/sync_test.go +++ b/core/state/sync_test.go @@ -22,6 +22,7 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/trie" ) @@ -38,7 +39,7 @@ type testAccount struct { func makeTestState() (ethdb.Database, common.Hash, []*testAccount) { // Create an empty state db, _ := ethdb.NewMemDatabase() - state, _ := New(common.Hash{}, db) + state, _ := New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) // Fill it with some arbitrary data accounts := []*testAccount{} @@ -68,7 +69,7 @@ func makeTestState() (ethdb.Database, common.Hash, []*testAccount) { // checkStateAccounts cross references a reconstructed state with an expected // account array. func checkStateAccounts(t *testing.T, db ethdb.Database, root common.Hash, accounts []*testAccount) { - state, _ := New(root, db) + state, _ := New(root, access.NewDbChainAccess(db), access.NullCtx) for i, acc := range accounts { if balance := state.GetBalance(acc.address); balance.Cmp(acc.balance) != 0 { diff --git a/core/transaction_pool_test.go b/core/transaction_pool_test.go index 229dcacf37..c419192c30 100644 --- a/core/transaction_pool_test.go +++ b/core/transaction_pool_test.go @@ -22,6 +22,7 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" @@ -36,7 +37,7 @@ func transaction(nonce uint64, gaslimit *big.Int, key *ecdsa.PrivateKey) *types. func setupTxPool() (*TxPool, *ecdsa.PrivateKey) { db, _ := ethdb.NewMemDatabase() - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) var m event.TypeMux key, _ := crypto.GenerateKey() @@ -163,7 +164,7 @@ func TestTransactionChainFork(t *testing.T) { addr := crypto.PubkeyToAddress(key.PublicKey) resetState := func() { db, _ := ethdb.NewMemDatabase() - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) pool.currentState = func() (*state.StateDB, error) { return statedb, nil } currentState, _ := pool.currentState() currentState.AddBalance(addr, big.NewInt(100000000000000)) @@ -189,7 +190,7 @@ func TestTransactionDoubleNonce(t *testing.T) { addr := crypto.PubkeyToAddress(key.PublicKey) resetState := func() { db, _ := ethdb.NewMemDatabase() - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) pool.currentState = func() (*state.StateDB, error) { return statedb, nil } currentState, _ := pool.currentState() currentState.AddBalance(addr, big.NewInt(100000000000000)) diff --git a/core/transaction_util.go b/core/transaction_util.go index e2e5b9aeea..6008765bf9 100644 --- a/core/transaction_util.go +++ b/core/transaction_util.go @@ -20,6 +20,7 @@ import ( "fmt" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/logger" @@ -29,8 +30,8 @@ import ( ) var ( - receiptsPre = []byte("receipts-") blockReceiptsPre = []byte("receipts-block-") + receiptsPre = []byte("receipts-") ) // PutTransactions stores the transactions in the given database @@ -119,9 +120,13 @@ func DeleteReceipt(db ethdb.Database, txHash common.Hash) { } // GetReceipt returns a receipt by hash -func GetReceipt(db ethdb.Database, txHash common.Hash) *types.Receipt { - data, _ := db.Get(append(receiptsPre, txHash[:]...)) +func GetReceipt(ca *access.ChainAccess, txHash common.Hash, ctx *access.OdrContext) *types.Receipt { + data, _ := ca.Db().Get(append(receiptsPre, txHash[:]...)) if len(data) == 0 { + if ctx != nil { + // we need a txHash -> blockHash mapping or and ODR request by txHash to do this + glog.V(logger.Error).Infoln("GetReceipt ODR not yet supported") + } return nil } var receipt types.ReceiptForStorage @@ -134,21 +139,10 @@ func GetReceipt(db ethdb.Database, txHash common.Hash) *types.Receipt { // GetBlockReceipts returns the receipts generated by the transactions // included in block's given hash. -func GetBlockReceipts(db ethdb.Database, hash common.Hash) types.Receipts { - data, _ := db.Get(append(blockReceiptsPre, hash[:]...)) - if len(data) == 0 { - return nil - } - rs := []*types.ReceiptForStorage{} - if err := rlp.DecodeBytes(data, &rs); err != nil { - glog.V(logger.Error).Infof("invalid receipt array RLP for hash %x: %v", hash, err) - return nil - } - receipts := make(types.Receipts, len(rs)) - for i, receipt := range rs { - receipts[i] = (*types.Receipt)(receipt) - } - return receipts +func GetBlockReceipts(ca *access.ChainAccess, hash common.Hash, ctx *access.OdrContext) types.Receipts { + r := &ReceiptsAccess{db: ca.Db(), blockHash: hash} + ca.Retrieve(r, ctx) + return r.receipts } // PutBlockReceipts stores the block's transactions associated receipts diff --git a/core/types/derive_sha.go b/core/types/derive_sha.go index 00c42c5bc6..cfee56599f 100644 --- a/core/types/derive_sha.go +++ b/core/types/derive_sha.go @@ -20,6 +20,7 @@ import ( "bytes" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/rlp" "github.com/ethereum/go-ethereum/trie" ) @@ -35,7 +36,7 @@ func DeriveSha(list DerivableList) common.Hash { for i := 0; i < list.Len(); i++ { keybuf.Reset() rlp.Encode(keybuf, uint(i)) - trie.Update(keybuf.Bytes(), list.GetRlp(i)) + trie.Update(keybuf.Bytes(), list.GetRlp(i), access.NoOdr) } return trie.Hash() } diff --git a/core/vm_env.go b/core/vm_env.go index c8b50debc6..be9d993bcb 100644 --- a/core/vm_env.go +++ b/core/vm_env.go @@ -20,6 +20,7 @@ import ( "math/big" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" @@ -59,7 +60,7 @@ func (self *VMEnv) SetDepth(i int) { self.depth = i } func (self *VMEnv) VmType() vm.Type { return self.typ } func (self *VMEnv) SetVmType(t vm.Type) { self.typ = t } func (self *VMEnv) GetHash(n uint64) common.Hash { - for block := self.chain.GetBlock(self.header.ParentHash); block != nil; block = self.chain.GetBlock(block.ParentHash()) { + for block := self.chain.GetBlock(self.header.ParentHash, access.NoOdr); block != nil; block = self.chain.GetBlock(block.ParentHash(), access.NoOdr) { if block.NumberU64() == n { return block.Hash() } @@ -76,7 +77,7 @@ func (self *VMEnv) CanTransfer(from common.Address, balance *big.Int) bool { } func (self *VMEnv) MakeSnapshot() vm.Database { - return self.state.Copy() + return self.state.CopyWithCtx() } func (self *VMEnv) SetSnapshot(copy vm.Database) { diff --git a/eth/backend.go b/eth/backend.go index 9eb211e31f..fb574b12f4 100644 --- a/eth/backend.go +++ b/eth/backend.go @@ -38,12 +38,14 @@ import ( "github.com/ethereum/go-ethereum/common/httpclient" "github.com/ethereum/go-ethereum/core" "github.com/ethereum/go-ethereum/core/state" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/eth/downloader" "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/event" + "github.com/ethereum/go-ethereum/les" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" "github.com/ethereum/go-ethereum/miner" @@ -60,6 +62,8 @@ const ( autoDAGcheckInterval = 10 * time.Hour autoDAGepochHeight = epochLength / 2 + + useLightProtocol = true ) var ( @@ -96,6 +100,7 @@ type Config struct { GenesisBlock *types.Block // used by block tests FastSync bool Olympic bool + Mode Mode BlockChainVersion int SkipBcVersionCheck bool // e.g. blockchain export @@ -231,17 +236,20 @@ type Ethereum struct { chainDb ethdb.Database // Block chain database dappDb ethdb.Database // Dapp database + chainAccess *access.ChainAccess // blockchain access layer + //*** SERVICES *** // State manager for processing new blocks and managing the over all states - blockProcessor *core.BlockProcessor - txPool *core.TxPool - blockchain *core.BlockChain - accountManager *accounts.Manager - whisper *whisper.Whisper - pow *ethash.Ethash - protocolManager *ProtocolManager - SolcPath string - solc *compiler.Solidity + blockProcessor *core.BlockProcessor + txPool *core.TxPool + blockchain *core.BlockChain + accountManager *accounts.Manager + whisper *whisper.Whisper + pow *ethash.Ethash + protocolManager *ProtocolManager + lightProtocolManager *les.ProtocolManager + SolcPath string + solc *compiler.Solidity GpoMinGasPrice *big.Int GpoMaxGasPrice *big.Int @@ -251,6 +259,7 @@ type Ethereum struct { GpobaseCorrectionFactor int httpclient *httpclient.HTTPClient + Mode Mode net *p2p.Server eventMux *event.TypeMux @@ -305,10 +314,10 @@ func New(config *Config) (*Ethereum, error) { if err := upgradeChainDatabase(chainDb); err != nil { return nil, err } - if err := addMipmapBloomBins(chainDb); err != nil { + chainAccess := access.NewChainAccess(chainDb, config.Mode == LightMode) + if err := addMipmapBloomBins(chainAccess); err != nil { return nil, err } - dappDb, err := newdb(filepath.Join(config.DataDir, "dapp")) if err != nil { var ok bool @@ -381,6 +390,7 @@ func New(config *Config) (*Ethereum, error) { eth := &Ethereum{ shutdownChan: make(chan bool), chainDb: chainDb, + chainAccess: chainAccess, dappDb: dappDb, eventMux: &event.TypeMux{}, accountManager: config.AccountManager, @@ -400,6 +410,7 @@ func New(config *Config) (*Ethereum, error) { GpobaseStepUp: config.GpobaseStepUp, GpobaseCorrectionFactor: config.GpobaseCorrectionFactor, httpclient: httpclient.New(config.DocRoot), + Mode: config.Mode, } if config.PowTest { @@ -412,22 +423,36 @@ func New(config *Config) (*Ethereum, error) { eth.pow = ethash.New() } //genesis := core.GenesisBlock(uint64(config.GenesisNonce), stateDb) - eth.blockchain, err = core.NewBlockChain(chainDb, eth.pow, eth.EventMux()) + eth.blockchain, err = core.NewBlockChain(chainAccess, eth.pow, eth.EventMux()) if err != nil { if err == core.ErrNoGenesis { return nil, fmt.Errorf(`Genesis block not found. Please supply a genesis block with the "--genesis /path/to/file" argument`) } return nil, err } - newPool := core.NewTxPool(eth.EventMux(), eth.blockchain.State, eth.blockchain.GasLimit) + newPool := core.NewTxPool(eth.EventMux(), eth.blockchain.StateNoOdr, eth.blockchain.GasLimit) eth.txPool = newPool - eth.blockProcessor = core.NewBlockProcessor(chainDb, eth.pow, eth.blockchain, eth.EventMux()) + eth.blockProcessor = core.NewBlockProcessor(chainAccess, eth.pow, eth.blockchain, eth.EventMux()) eth.blockchain.SetProcessor(eth.blockProcessor) - if eth.protocolManager, err = NewProtocolManager(config.FastSync, config.NetworkId, eth.eventMux, eth.txPool, eth.pow, eth.blockchain, chainDb); err != nil { - return nil, err + + if config.Mode != LightMode { + if eth.protocolManager, err = NewProtocolManager(config.Mode, config.NetworkId, eth.eventMux, eth.txPool, eth.pow, eth.blockchain, chainAccess); err != nil { + return nil, err + } + } + + if useLightProtocol { + if eth.lightProtocolManager, err = les.NewProtocolManager(les.Mode(config.Mode), config.NetworkId, eth.eventMux, eth.pow, eth.blockchain, chainAccess); err != nil { + return nil, err + } + } + + if config.Mode == LightMode { + eth.miner = miner.NewDummy(eth, eth.EventMux(), eth.pow) + } else { + eth.miner = miner.New(eth, eth.EventMux(), eth.pow) } - eth.miner = miner.New(eth, eth.EventMux(), eth.pow) eth.miner.SetGasPrice(config.GasPrice) eth.miner.SetExtra(config.ExtraData) @@ -440,7 +465,14 @@ func New(config *Config) (*Ethereum, error) { if err != nil { return nil, err } - protocols := append([]p2p.Protocol{}, eth.protocolManager.SubProtocols...) + protocols := []p2p.Protocol{} + + if eth.protocolManager != nil { + protocols = append(protocols, eth.protocolManager.SubProtocols...) + } + if eth.lightProtocolManager != nil { + protocols = append(protocols, eth.lightProtocolManager.SubProtocols...) + } if config.Shh { protocols = append(protocols, eth.whisper.Protocol()) } @@ -508,16 +540,29 @@ func (s *Ethereum) TxPool() *core.TxPool { return s.txPool } func (s *Ethereum) Whisper() *whisper.Whisper { return s.whisper } func (s *Ethereum) EventMux() *event.TypeMux { return s.eventMux } func (s *Ethereum) ChainDb() ethdb.Database { return s.chainDb } +func (s *Ethereum) ChainAccess() *access.ChainAccess { return s.chainAccess } func (s *Ethereum) DappDb() ethdb.Database { return s.dappDb } func (s *Ethereum) IsListening() bool { return true } // Always listening func (s *Ethereum) PeerCount() int { return s.net.PeerCount() } func (s *Ethereum) Peers() []*p2p.Peer { return s.net.Peers() } func (s *Ethereum) MaxPeers() int { return s.net.MaxPeers } func (s *Ethereum) ClientVersion() string { return s.clientVersion } -func (s *Ethereum) EthVersion() int { return int(s.protocolManager.SubProtocols[0].Version) } -func (s *Ethereum) NetVersion() int { return s.netVersionId } -func (s *Ethereum) ShhVersion() int { return s.shhVersionId } -func (s *Ethereum) Downloader() *downloader.Downloader { return s.protocolManager.downloader } +func (s *Ethereum) EthVersion() int { + if s.protocolManager == nil { + return 0 + } else { + return int(s.protocolManager.SubProtocols[0].Version) + } +} +func (s *Ethereum) NetVersion() int { return s.netVersionId } +func (s *Ethereum) ShhVersion() int { return s.shhVersionId } +func (s *Ethereum) Downloader() *downloader.Downloader { + if s.protocolManager == nil { + return nil + } else { + return s.protocolManager.downloader + } +} // Start the ethereum func (s *Ethereum) Start() error { @@ -537,7 +582,13 @@ func (s *Ethereum) Start() error { s.StartAutoDAG() } - s.protocolManager.Start() + if s.protocolManager != nil { + s.protocolManager.Start() + } + + if s.lightProtocolManager != nil { + s.lightProtocolManager.Start() + } if s.whisper != nil { s.whisper.Start() @@ -569,9 +620,14 @@ func (self *Ethereum) AddPeer(nodeURL string) error { func (s *Ethereum) Stop() { s.net.Stop() s.blockchain.Stop() - s.protocolManager.Stop() + if s.protocolManager != nil { + s.protocolManager.Stop() + } s.txPool.Stop() s.eventMux.Stop() + if s.lightProtocolManager != nil { + s.lightProtocolManager.Stop() + } if s.whisper != nil { s.whisper.Stop() } @@ -609,7 +665,7 @@ func (self *Ethereum) StartAutoDAG() { select { case <-timer: glog.V(logger.Info).Infof("checking DAG (ethash dir: %s)", ethash.DefaultDir) - currentBlock := self.BlockChain().CurrentBlock().NumberU64() + currentBlock := self.BlockChain().LastBlockNumberU64() thisEpoch := currentBlock / epochLength if nextEpoch <= thisEpoch { if currentBlock%epochLength > autoDAGepochHeight { @@ -747,9 +803,10 @@ func upgradeChainDatabase(db ethdb.Database) error { return nil } -func addMipmapBloomBins(db ethdb.Database) (err error) { +func addMipmapBloomBins(ca *access.ChainAccess) (err error) { const mipmapVersion uint = 2 + db := ca.Db() // check if the version is set. We ignore data for now since there's // only one version so we can easily ignore it for now var data []byte @@ -771,7 +828,7 @@ func addMipmapBloomBins(db ethdb.Database) (err error) { return } }() - latestBlock := core.GetBlock(db, core.GetHeadBlockHash(db)) + latestBlock := core.GetBlock(ca, core.GetHeadBlockHash(db), access.NoOdr) if latestBlock == nil { // clean database return } @@ -783,7 +840,7 @@ func addMipmapBloomBins(db ethdb.Database) (err error) { if (hash == common.Hash{}) { return fmt.Errorf("chain db corrupted. Could not find block %d.", i) } - core.WriteMipmapBloom(db, i, core.GetBlockReceipts(db, hash)) + core.WriteMipmapBloom(db, i, core.GetBlockReceipts(ca, hash, access.NoOdr)) } glog.V(logger.Info).Infoln("upgrade completed in", time.Since(tstart)) return nil diff --git a/eth/backend_test.go b/eth/backend_test.go index 0379fc843a..fdabf4d64e 100644 --- a/eth/backend_test.go +++ b/eth/backend_test.go @@ -6,6 +6,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/ethdb" @@ -50,7 +51,7 @@ func TestMipmapUpgrade(t *testing.T) { } } - err := addMipmapBloomBins(db) + err := addMipmapBloomBins(access.NewDbChainAccess(db)) if err != nil { t.Fatal(err) } diff --git a/eth/downloader/downloader.go b/eth/downloader/downloader.go index 153427ee48..33c4a8ef92 100644 --- a/eth/downloader/downloader.go +++ b/eth/downloader/downloader.go @@ -44,6 +44,7 @@ var ( MaxBodyFetch = 128 // Amount of block bodies to be fetched per retrieval request MaxReceiptFetch = 256 // Amount of transaction receipts to allow fetching per request MaxStateFetch = 384 // Amount of node state values to allow fetching per request + MaxProofsFetch = 64 // Amount of merkle proofs to allow fetching per request hashTTL = 5 * time.Second // [eth/61] Time it takes for a hash request to time out blockSoftTTL = 3 * time.Second // [eth/61] Request completion threshold for increasing or decreasing a peer's bandwidth @@ -156,13 +157,13 @@ type Downloader struct { } // New creates a new downloader to fetch hashes and blocks from remote peers. -func New(stateDb ethdb.Database, mux *event.TypeMux, hasHeader headerCheckFn, hasBlock blockCheckFn, getHeader headerRetrievalFn, +func New(mode SyncMode, stateDb ethdb.Database, mux *event.TypeMux, hasHeader headerCheckFn, hasBlock blockCheckFn, getHeader headerRetrievalFn, getBlock blockRetrievalFn, headHeader headHeaderRetrievalFn, headBlock headBlockRetrievalFn, headFastBlock headFastBlockRetrievalFn, commitHeadBlock headBlockCommitterFn, getTd tdRetrievalFn, insertHeaders headerChainInsertFn, insertBlocks blockChainInsertFn, insertReceipts receiptChainInsertFn, rollback chainRollbackFn, dropPeer peerDropFn) *Downloader { return &Downloader{ - mode: FullSync, + mode: mode, mux: mux, queue: newQueue(stateDb), peers: newPeerSet(), @@ -225,6 +226,10 @@ func (d *Downloader) RegisterPeer(id string, version int, head common.Hash, getRelHeaders relativeHeaderFetcherFn, getAbsHeaders absoluteHeaderFetcherFn, getBlockBodies blockBodyFetcherFn, getReceipts receiptFetcherFn, getNodeData stateFetcherFn) error { + if d.mode == NoSync { + return nil + } + glog.V(logger.Detail).Infoln("Registering peer", id) if err := d.peers.Register(newPeer(id, version, head, getRelHashes, getAbsHashes, getBlocks, getRelHeaders, getAbsHeaders, getBlockBodies, getReceipts, getNodeData)); err != nil { glog.V(logger.Error).Infoln("Register failed:", err) @@ -237,6 +242,10 @@ func (d *Downloader) RegisterPeer(id string, version int, head common.Hash, // the specified peer. An effort is also made to return any pending fetches into // the queue. func (d *Downloader) UnregisterPeer(id string) error { + if d.mode == NoSync { + return nil + } + glog.V(logger.Detail).Infoln("Unregistering peer", id) if err := d.peers.Unregister(id); err != nil { glog.V(logger.Error).Infoln("Unregister failed:", err) @@ -249,6 +258,10 @@ func (d *Downloader) UnregisterPeer(id string) error { // Synchronise tries to sync up our local block chain with a remote peer, both // adding various sanity checks as well as wrapping it with various log entries. func (d *Downloader) Synchronise(id string, head common.Hash, td *big.Int, mode SyncMode) error { + if d.mode == NoSync { + return nil + } + glog.V(logger.Detail).Infof("Attempting synchronisation: %v, head [%x…], TD %v", id, head[:4], td) err := d.synchronise(id, head, td, mode) @@ -276,6 +289,10 @@ func (d *Downloader) Synchronise(id string, head common.Hash, td *big.Int, mode // it will use the best peer possible and synchronize if it's TD is higher than our own. If any of the // checks fail an error will be returned. This method is synchronous func (d *Downloader) synchronise(id string, hash common.Hash, td *big.Int, mode SyncMode) error { + if d.mode == NoSync { + return nil + } + // Mock out the synchonisation if testing if d.synchroniseMock != nil { return d.synchroniseMock(id, hash) @@ -331,6 +348,10 @@ func (d *Downloader) synchronise(id string, hash common.Hash, td *big.Int, mode // syncWithPeer starts a block synchronization based on the hash chain from the // specified peer and head hash. func (d *Downloader) syncWithPeer(p *peer, hash common.Hash, td *big.Int) (err error) { + if d.mode == NoSync { + return nil + } + d.mux.Post(StartEvent{}) defer func() { // reset on error @@ -459,6 +480,10 @@ func (d *Downloader) syncWithPeer(p *peer, hash common.Hash, td *big.Int) (err e // cancel cancels all of the operations and resets the queue. It returns true // if the cancel operation was completed. func (d *Downloader) cancel() { + if d.mode == NoSync { + return + } + // Close the current cancel channel d.cancelLock.Lock() if d.cancelCh != nil { @@ -484,6 +509,10 @@ func (d *Downloader) Terminate() { // fetchHeight61 retrieves the head block of the remote peer to aid in estimating // the total time a pending synchronisation would take. func (d *Downloader) fetchHeight61(p *peer) (uint64, error) { + if d.mode == NoSync { + return 0, nil + } + glog.V(logger.Debug).Infof("%v: retrieving remote chain height", p) // Request the advertised remote head block and wait for the response @@ -531,6 +560,10 @@ func (d *Downloader) fetchHeight61(p *peer) (uint64, error) { // In the rare scenario when we ended up on a long reorganization (i.e. none of // the head blocks match), we do a binary search to find the common ancestor. func (d *Downloader) findAncestor61(p *peer) (uint64, error) { + if d.mode == NoSync { + return 0, nil + } + glog.V(logger.Debug).Infof("%v: looking for common ancestor", p) // Request out head blocks to short circuit ancestor location @@ -652,6 +685,10 @@ func (d *Downloader) findAncestor61(p *peer) (uint64, error) { // fetchHashes61 keeps retrieving hashes from the requested number, until no more // are returned, potentially throttling on the way. func (d *Downloader) fetchHashes61(p *peer, td *big.Int, from uint64) error { + if d.mode == NoSync { + return nil + } + glog.V(logger.Debug).Infof("%v: downloading hashes from #%d", p, from) // Create a timeout timer, and the associated hash fetcher @@ -758,6 +795,10 @@ func (d *Downloader) fetchHashes61(p *peer, td *big.Int, from uint64) error { // peers, reserving a chunk of blocks for each, waiting for delivery and also // periodically checking for timeouts. func (d *Downloader) fetchBlocks61(from uint64) error { + if d.mode == NoSync { + return nil + } + glog.V(logger.Debug).Infof("Downloading blocks from #%d", from) defer glog.V(logger.Debug).Infof("Block download terminated") @@ -916,6 +957,10 @@ func (d *Downloader) fetchBlocks61(from uint64) error { // fetchHeight retrieves the head header of the remote peer to aid in estimating // the total time a pending synchronisation would take. func (d *Downloader) fetchHeight(p *peer) (uint64, error) { + if d.mode == NoSync { + return 0, nil + } + glog.V(logger.Debug).Infof("%v: retrieving remote chain height", p) // Request the advertised remote head block and wait for the response @@ -963,10 +1008,17 @@ func (d *Downloader) fetchHeight(p *peer) (uint64, error) { // In the rare scenario when we ended up on a long reorganization (i.e. none of // the head links match), we do a binary search to find the common ancestor. func (d *Downloader) findAncestor(p *peer) (uint64, error) { + if d.mode == NoSync { + return 0, nil + } + glog.V(logger.Debug).Infof("%v: looking for common ancestor", p) // Request our head headers to short circuit ancestor location - head := d.headHeader().Number.Uint64() + var head uint64 + if header := d.headHeader(); header != nil { + head = header.Number.Uint64() + } if d.mode == FullSync { head = d.headBlock().NumberU64() } else if d.mode == FastSync { @@ -1092,6 +1144,10 @@ func (d *Downloader) findAncestor(p *peer) (uint64, error) { // The queue parameter can be used to switch between queuing headers for block // body download too, or directly import as pure header chains. func (d *Downloader) fetchHeaders(p *peer, td *big.Int, from uint64) error { + if d.mode == NoSync { + return nil + } + glog.V(logger.Debug).Infof("%v: downloading headers from #%d", p, from) defer glog.V(logger.Debug).Infof("%v: header download terminated", p) @@ -1276,6 +1332,10 @@ func (d *Downloader) fetchHeaders(p *peer, td *big.Int, from uint64) error { // available peers, reserving a chunk of blocks for each, waiting for delivery // and also periodically checking for timeouts. func (d *Downloader) fetchBodies(from uint64) error { + if d.mode == NoSync { + return nil + } + glog.V(logger.Debug).Infof("Downloading block bodies from #%d", from) var ( @@ -1324,6 +1384,10 @@ func (d *Downloader) fetchReceipts(from uint64) error { // available peers, reserving a chunk of nodes for each, waiting for delivery and // also periodically checking for timeouts. func (d *Downloader) fetchNodeData() error { + if d.mode == NoSync { + return nil + } + glog.V(logger.Debug).Infof("Downloading node state data") var ( @@ -1378,7 +1442,11 @@ func (d *Downloader) fetchParts(errCancel error, deliveryCh chan dataPack, deliv fetchHook func([]*types.Header), fetch func(*peer, *fetchRequest) error, cancel func(*fetchRequest), capacity func(*peer) int, idle func() ([]*peer, int), setIdle func(*peer), kind string) error { - // Create a ticker to detect expired retrieval tasks + if d.mode == NoSync { + return nil + } + + // Create a ticker to detect expired retreival tasks ticker := time.NewTicker(100 * time.Millisecond) defer ticker.Stop() @@ -1678,6 +1746,9 @@ func (d *Downloader) DeliverNodeData(id string, data [][]byte) (err error) { // deliver injects a new batch of data received from a remote node. func (d *Downloader) deliver(id string, destCh chan dataPack, packet dataPack, inMeter, dropMeter metrics.Meter) (err error) { + if d.mode == NoSync { + return nil + } // Update the delivery metrics for both good and failed deliveries inMeter.Mark(int64(packet.Items())) defer func() { diff --git a/eth/downloader/downloader_test.go b/eth/downloader/downloader_test.go index ef6f74a6ba..13e2d8d53f 100644 --- a/eth/downloader/downloader_test.go +++ b/eth/downloader/downloader_test.go @@ -27,6 +27,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" @@ -150,7 +151,7 @@ func newTester() *downloadTester { peerChainTds: make(map[string]map[common.Hash]*big.Int), } tester.stateDb, _ = ethdb.NewMemDatabase() - tester.downloader = New(tester.stateDb, new(event.TypeMux), tester.hasHeader, tester.hasBlock, tester.getHeader, + tester.downloader = New(FullSync, tester.stateDb, new(event.TypeMux), tester.hasHeader, tester.hasBlock, tester.getHeader, tester.getBlock, tester.headHeader, tester.headBlock, tester.headFastBlock, tester.commitHeadBlock, tester.getTd, tester.insertHeaders, tester.insertBlocks, tester.insertReceipts, tester.rollback, tester.dropPeer) @@ -253,7 +254,7 @@ func (dl *downloadTester) headFastBlock() *types.Block { func (dl *downloadTester) commitHeadBlock(hash common.Hash) error { // For now only check that the state trie is correct if block := dl.getBlock(hash); block != nil { - _, err := trie.NewSecure(block.Root(), dl.stateDb) + _, err := trie.NewSecure(block.Root(), dl.stateDb, nil) return err } return fmt.Errorf("non existent block: %x", hash[:4]) @@ -682,7 +683,7 @@ func assertOwnForkedChain(t *testing.T, tester *downloadTester, common int, leng index = len(tester.ownHashes) - lengths[len(lengths)-1] + int(tester.downloader.queue.fastSyncPivot) } if index > 0 { - if statedb, err := state.New(tester.ownHeaders[tester.ownHashes[index]].Root, tester.stateDb); statedb == nil || err != nil { + if statedb, err := state.New(tester.ownHeaders[tester.ownHashes[index]].Root, access.NewDbChainAccess(tester.stateDb), access.NullCtx); statedb == nil || err != nil { t.Fatalf("state reconstruction failed: %v", err) } } diff --git a/eth/downloader/modes.go b/eth/downloader/modes.go index ec339c0740..7bc3047ee2 100644 --- a/eth/downloader/modes.go +++ b/eth/downloader/modes.go @@ -23,4 +23,5 @@ const ( FullSync SyncMode = iota // Synchronise the entire blockchain history from full blocks FastSync // Quickly download the headers, full sync only at the chain head LightSync // Download only the headers and terminate afterwards + NoSync // Dummy downloader, only one of the eth/les downloaders is actually working ) diff --git a/eth/filters/filter.go b/eth/filters/filter.go index ff192cdf6b..609ea36eae 100644 --- a/eth/filters/filter.go +++ b/eth/filters/filter.go @@ -21,6 +21,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/ethdb" @@ -32,6 +33,7 @@ type AccountChange struct { // Filtering interface type Filter struct { + ca *access.ChainAccess db ethdb.Database begin, end int64 addresses []common.Address @@ -44,8 +46,12 @@ type Filter struct { // Create a new filter which uses a bloom filter on blocks to figure out whether a particular block // is interesting or not. -func New(db ethdb.Database) *Filter { - return &Filter{db: db} +func New(ca *access.ChainAccess) *Filter { + return &Filter{ca: ca, db: ca.Db()} +} + +func NewWithDb(db ethdb.Database) *Filter { + return &Filter{ca: access.NewDbChainAccess(db), db: db} } // Set the earliest and latest block for filtering. @@ -69,7 +75,7 @@ func (self *Filter) SetTopics(topics [][]common.Hash) { // Run filters logs with the current parameters set func (self *Filter) Find() vm.Logs { - latestBlock := core.GetBlock(self.db, core.GetHeadBlockHash(self.db)) + latestBlock := core.GetBlock(self.ca, core.GetHeadBlockHash(self.db), access.NoOdr) var beginBlockNo uint64 = uint64(self.begin) if self.begin == -1 { beginBlockNo = latestBlock.NumberU64() @@ -124,7 +130,7 @@ func (self *Filter) getLogs(start, end uint64) (logs vm.Logs) { for i := start; i <= end; i++ { hash := core.GetCanonicalHash(self.db, i) if hash != (common.Hash{}) { - block = core.GetBlock(self.db, hash) + block = core.GetBlock(self.ca, hash, access.NoOdr) } else { // block not found return logs } @@ -134,7 +140,7 @@ func (self *Filter) getLogs(start, end uint64) (logs vm.Logs) { if self.bloomFilter(block) { // Get the logs of the block var ( - receipts = core.GetBlockReceipts(self.db, block.Hash()) + receipts = core.GetBlockReceipts(self.ca, block.Hash(), access.NoOdr) unfiltered vm.Logs ) for _, receipt := range receipts { @@ -142,6 +148,7 @@ func (self *Filter) getLogs(start, end uint64) (logs vm.Logs) { } logs = append(logs, self.FilterLogs(unfiltered)...) } + block = core.GetBlock(self.ca, block.ParentHash(), access.NoOdr) } return logs diff --git a/eth/filters/filter_test.go b/eth/filters/filter_test.go index a5418e2e74..22e515aa87 100644 --- a/eth/filters/filter_test.go +++ b/eth/filters/filter_test.go @@ -84,7 +84,7 @@ func BenchmarkMipmaps(b *testing.B) { } b.ResetTimer() - filter := New(db) + filter := NewWithDb(db) filter.SetAddresses([]common.Address{addr1, addr2, addr3, addr4}) filter.SetBeginBlock(0) filter.SetEndBlock(-1) @@ -185,7 +185,7 @@ func TestFilters(t *testing.T) { } } - filter := New(db) + filter := NewWithDb(db) filter.SetAddresses([]common.Address{addr}) filter.SetTopics([][]common.Hash{[]common.Hash{hash1, hash2, hash3, hash4}}) filter.SetBeginBlock(0) @@ -196,7 +196,7 @@ func TestFilters(t *testing.T) { t.Error("expected 4 log, got", len(logs)) } - filter = New(db) + filter = NewWithDb(db) filter.SetAddresses([]common.Address{addr}) filter.SetTopics([][]common.Hash{[]common.Hash{hash3}}) filter.SetBeginBlock(900) @@ -209,7 +209,7 @@ func TestFilters(t *testing.T) { t.Errorf("expected log[0].Topics[0] to be %x, got %x", hash3, logs[0].Topics[0]) } - filter = New(db) + filter = NewWithDb(db) filter.SetAddresses([]common.Address{addr}) filter.SetTopics([][]common.Hash{[]common.Hash{hash3}}) filter.SetBeginBlock(990) @@ -222,7 +222,7 @@ func TestFilters(t *testing.T) { t.Errorf("expected log[0].Topics[0] to be %x, got %x", hash3, logs[0].Topics[0]) } - filter = New(db) + filter = NewWithDb(db) filter.SetTopics([][]common.Hash{[]common.Hash{hash1, hash2}}) filter.SetBeginBlock(1) filter.SetEndBlock(10) @@ -233,7 +233,7 @@ func TestFilters(t *testing.T) { } failHash := common.BytesToHash([]byte("fail")) - filter = New(db) + filter = NewWithDb(db) filter.SetTopics([][]common.Hash{[]common.Hash{failHash}}) filter.SetBeginBlock(0) filter.SetEndBlock(-1) @@ -244,7 +244,7 @@ func TestFilters(t *testing.T) { } failAddr := common.BytesToAddress([]byte("failmenow")) - filter = New(db) + filter = NewWithDb(db) filter.SetAddresses([]common.Address{failAddr}) filter.SetBeginBlock(0) filter.SetEndBlock(-1) @@ -254,7 +254,7 @@ func TestFilters(t *testing.T) { t.Error("expected 0 log, got", len(logs)) } - filter = New(db) + filter = NewWithDb(db) filter.SetTopics([][]common.Hash{[]common.Hash{failHash}, []common.Hash{hash1}}) filter.SetBeginBlock(0) filter.SetEndBlock(-1) diff --git a/eth/gasprice.go b/eth/gasprice.go index b752c22dd1..3bf6adeef4 100644 --- a/eth/gasprice.go +++ b/eth/gasprice.go @@ -22,6 +22,7 @@ import ( "sync" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" @@ -74,6 +75,9 @@ func NewGasPriceOracle(eth *Ethereum) *GasPriceOracle { } func (gpo *GasPriceOracle) init() { + if gpo.eth.Mode == LightMode { + return + } gpo.initOnce.Do(func() { gpo.processPastBlocks(gpo.eth.BlockChain()) go gpo.listenLoop() @@ -92,7 +96,7 @@ func (self *GasPriceOracle) processPastBlocks(chain *core.BlockChain) { } self.firstProcessed = uint64(first) for i := first; i <= last; i++ { - block := chain.GetBlockByNumber(uint64(i)) + block := chain.GetBlockByNumber(uint64(i), access.NoOdr) if block != nil { self.processBlock(block) } diff --git a/eth/handler.go b/eth/handler.go index d8c5b4b648..09496db974 100644 --- a/eth/handler.go +++ b/eth/handler.go @@ -26,10 +26,10 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/eth/downloader" "github.com/ethereum/go-ethereum/eth/fetcher" - "github.com/ethereum/go-ethereum/ethdb" "github.com/ethereum/go-ethereum/event" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" @@ -58,10 +58,11 @@ type blockFetcherFn func([]common.Hash) error type ProtocolManager struct { networkId int - fastSync bool - txpool txPool - blockchain *core.BlockChain - chaindb ethdb.Database + fastSync bool + mode Mode + txpool txPool + blockchain *core.BlockChain + chainAccess *access.ChainAccess downloader *downloader.Downloader fetcher *fetcher.Fetcher @@ -86,30 +87,31 @@ type ProtocolManager struct { // NewProtocolManager returns a new ethereum sub protocol manager. The Ethereum sub protocol manages peers capable // with the ethereum network. -func NewProtocolManager(fastSync bool, networkId int, mux *event.TypeMux, txpool txPool, pow pow.PoW, blockchain *core.BlockChain, chaindb ethdb.Database) (*ProtocolManager, error) { +func NewProtocolManager(mode Mode, networkId int, mux *event.TypeMux, txpool txPool, pow pow.PoW, blockchain *core.BlockChain, ca *access.ChainAccess) (*ProtocolManager, error) { // Figure out whether to allow fast sync or not - if fastSync && blockchain.CurrentBlock().NumberU64() > 0 { + if mode == FullMode && blockchain.CurrentBlock().NumberU64() > 0 { glog.V(logger.Info).Infof("blockchain not empty, fast sync disabled") - fastSync = false + mode = ArchiveMode } // Create the protocol manager with the base fields manager := &ProtocolManager{ networkId: networkId, - fastSync: fastSync, - eventMux: mux, - txpool: txpool, - blockchain: blockchain, - chaindb: chaindb, - peers: newPeerSet(), - newPeerCh: make(chan *peer, 1), - txsyncCh: make(chan *txsync), - quitSync: make(chan struct{}), + fastSync: mode == FullMode, + mode: mode, + eventMux: mux, + txpool: txpool, + blockchain: blockchain, + chainAccess: ca, + peers: newPeerSet(), + newPeerCh: make(chan *peer, 1), + txsyncCh: make(chan *txsync), + quitSync: make(chan struct{}), } // Initiate a sub-protocol for every implemented version we can handle manager.SubProtocols = make([]p2p.Protocol, 0, len(ProtocolVersions)) for i, version := range ProtocolVersions { // Skip protocol version if incompatible with the mode of operation - if fastSync && version < eth63 { + if manager.fastSync && version < eth63 { continue } // Compatible; initialise the sub-protocol @@ -138,17 +140,27 @@ func NewProtocolManager(fastSync bool, networkId int, mux *event.TypeMux, txpool return nil, errIncompatibleConfig } // Construct the different synchronisation mechanisms - manager.downloader = downloader.New(chaindb, manager.eventMux, blockchain.HasHeader, blockchain.HasBlock, blockchain.GetHeader, blockchain.GetBlock, - blockchain.CurrentHeader, blockchain.CurrentBlock, blockchain.CurrentFastBlock, blockchain.FastSyncCommitHead, blockchain.GetTd, + var syncMode downloader.SyncMode + switch mode { + case ArchiveMode: + syncMode = downloader.FullSync + case FullMode: + syncMode = downloader.FastSync + case LightMode: + syncMode = downloader.NoSync + } + + manager.downloader = downloader.New(syncMode, ca.Db(), manager.eventMux, blockchain.HasHeader, blockchain.HasBlock, blockchain.GetHeader, + blockchain.GetBlockNoOdr, blockchain.CurrentHeader, blockchain.CurrentBlock, blockchain.CurrentFastBlock, blockchain.FastSyncCommitHead, blockchain.GetTd, blockchain.InsertHeaderChain, blockchain.InsertChain, blockchain.InsertReceiptChain, blockchain.Rollback, manager.removePeer) validator := func(block *types.Block, parent *types.Block) error { return core.ValidateHeader(pow, block.Header(), parent.Header(), true, false) } heighter := func() uint64 { - return blockchain.CurrentBlock().NumberU64() + return blockchain.LastBlockNumberU64() } - manager.fetcher = fetcher.New(blockchain.GetBlock, validator, manager.BroadcastBlock, heighter, blockchain.InsertChain, manager.removePeer) + manager.fetcher = fetcher.New(blockchain.GetBlockNoOdr, validator, manager.BroadcastBlock, heighter, blockchain.InsertChain, manager.removePeer) return manager, nil } @@ -233,6 +245,7 @@ func (pm *ProtocolManager) handle(p *peer) error { p.RequestHeadersByNumber, p.RequestBodies, p.RequestReceipts, p.RequestNodeData); err != nil { return err } + // Propagate existing transactions. new transactions appearing // after this will be sent via broadcasts. pm.syncTransactions(p) @@ -292,7 +305,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { request.Amount = uint64(downloader.MaxHashFetch) } // Calculate the last block that should be retrieved, and short circuit if unavailable - last := pm.blockchain.GetBlockByNumber(request.Number + request.Amount - 1) + last := pm.blockchain.GetBlockByNumber(request.Number+request.Amount-1, access.NoOdr) if last == nil { last = pm.blockchain.CurrentBlock() request.Amount = last.NumberU64() - request.Number + 1 @@ -342,7 +355,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrDecode, "msg %v: %v", msg, err) } // Retrieve the requested block, stopping if enough was found - if block := pm.blockchain.GetBlock(hash); block != nil { + if block := pm.blockchain.GetBlock(hash, access.NoOdr); block != nil { blocks = append(blocks, block) bytes += block.Size() } @@ -468,7 +481,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrDecode, "msg %v: %v", msg, err) } // Retrieve the requested block body, stopping if enough was found - if data := pm.blockchain.GetBodyRLP(hash); len(data) != 0 { + if data := pm.blockchain.GetBodyRLP(hash, access.NoOdr); len(data) != 0 { bodies = append(bodies, data) bytes += len(data) } @@ -517,7 +530,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrDecode, "msg %v: %v", msg, err) } // Retrieve the requested state entry, stopping if enough was found - if entry, err := pm.chaindb.Get(hash.Bytes()); err == nil { + if entry, err := pm.chainAccess.Db().Get(hash.Bytes()); err == nil { data = append(data, entry) bytes += len(entry) } @@ -555,7 +568,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrDecode, "msg %v: %v", msg, err) } // Retrieve the requested block's receipts, skipping if unknown to us - results := core.GetBlockReceipts(pm.chaindb, hash) + results := core.GetBlockReceipts(pm.chainAccess, hash, access.NoOdr) if results == nil { if header := pm.blockchain.GetHeader(hash); header == nil || header.ReceiptHash != types.EmptyRootHash { continue @@ -649,7 +662,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { // Update the peers total difficulty if needed, schedule a download if gapped if request.TD.Cmp(p.Td()) > 0 { p.SetTd(request.TD) - td := pm.blockchain.GetTd(pm.blockchain.CurrentBlock().Hash()) + td := pm.blockchain.GetTd(pm.blockchain.LastBlockHash()) if request.TD.Cmp(new(big.Int).Add(td, request.Block.Difficulty())) > 0 { go pm.synchronise(p) } @@ -686,7 +699,7 @@ func (pm *ProtocolManager) BroadcastBlock(block *types.Block, propagate bool) { if propagate { // Calculate the TD of the block (it's not imported yet, so block.Td is not valid) var td *big.Int - if parent := pm.blockchain.GetBlock(block.ParentHash()); parent != nil { + if parent := pm.blockchain.GetBlock(block.ParentHash(), access.NoOdr); parent != nil { td = new(big.Int).Add(block.Difficulty(), pm.blockchain.GetTd(block.ParentHash())) } else { glog.V(logger.Error).Infof("propagating dangling block #%d [%x]", block.NumberU64(), hash[:4]) diff --git a/eth/handler_test.go b/eth/handler_test.go index ab2ce54b13..c976a9b793 100644 --- a/eth/handler_test.go +++ b/eth/handler_test.go @@ -8,6 +8,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" @@ -36,7 +37,11 @@ func TestProtocolCompatibility(t *testing.T) { for i, tt := range tests { ProtocolVersions = []uint{tt.version} - pm, err := newTestProtocolManager(tt.fastSync, 0, nil, nil) + mode := ArchiveMode + if tt.fastSync { + mode = FullMode + } + pm, err := newTestProtocolManager(mode, 0, nil, nil) if pm != nil { defer pm.Stop() } @@ -62,23 +67,23 @@ func testGetBlockHashes(t *testing.T, protocol int) { number int result int }{ - {common.Hash{}, 1, 0}, // Make sure non existent hashes don't return results - {pm.blockchain.Genesis().Hash(), 1, 0}, // There are no hashes to retrieve up from the genesis - {pm.blockchain.GetBlockByNumber(5).Hash(), 5, 5}, // All the hashes including the genesis requested - {pm.blockchain.GetBlockByNumber(5).Hash(), 10, 5}, // More hashes than available till the genesis requested - {pm.blockchain.GetBlockByNumber(100).Hash(), 10, 10}, // All hashes available from the middle of the chain - {pm.blockchain.CurrentBlock().Hash(), 10, 10}, // All hashes available from the head of the chain - {pm.blockchain.CurrentBlock().Hash(), limit, limit}, // Request the maximum allowed hash count - {pm.blockchain.CurrentBlock().Hash(), limit + 1, limit}, // Request more than the maximum allowed hash count + {common.Hash{}, 1, 0}, // Make sure non existent hashes don't return results + {pm.blockchain.Genesis().Hash(), 1, 0}, // There are no hashes to retrieve up from the genesis + {pm.blockchain.GetBlockByNumber(5, access.NoOdr).Hash(), 5, 5}, // All the hashes including the genesis requested + {pm.blockchain.GetBlockByNumber(5, access.NoOdr).Hash(), 10, 5}, // More hashes than available till the genesis requested + {pm.blockchain.GetBlockByNumber(100, access.NoOdr).Hash(), 10, 10}, // All hashes available from the middle of the chain + {pm.blockchain.CurrentBlock().Hash(), 10, 10}, // All hashes available from the head of the chain + {pm.blockchain.CurrentBlock().Hash(), limit, limit}, // Request the maximum allowed hash count + {pm.blockchain.CurrentBlock().Hash(), limit + 1, limit}, // Request more than the maximum allowed hash count } // Run each of the tests and verify the results against the chain for i, tt := range tests { // Assemble the hash response we would like to receive resp := make([]common.Hash, tt.result) if len(resp) > 0 { - from := pm.blockchain.GetBlock(tt.origin).NumberU64() - 1 + from := pm.blockchain.GetBlock(tt.origin, access.NoOdr).NumberU64() - 1 for j := 0; j < len(resp); j++ { - resp[j] = pm.blockchain.GetBlockByNumber(uint64(int(from) - j)).Hash() + resp[j] = pm.blockchain.GetBlockByNumber(uint64(int(from)-j), access.NoOdr).Hash() } } // Send the hash request and verify the response @@ -120,7 +125,7 @@ func testGetBlockHashesFromNumber(t *testing.T, protocol int) { // Assemble the hash response we would like to receive resp := make([]common.Hash, tt.result) for j := 0; j < len(resp); j++ { - resp[j] = pm.blockchain.GetBlockByNumber(tt.origin + uint64(j)).Hash() + resp[j] = pm.blockchain.GetBlockByNumber(tt.origin+uint64(j), access.NoOdr).Hash() } // Send the hash request and verify the response p2p.Send(peer.app, 0x08, getBlockHashesFromNumberData{tt.origin, uint64(tt.number)}) @@ -157,11 +162,11 @@ func testGetBlocks(t *testing.T, protocol int) { // Existing and non-existing blocks interleaved should not cause problems {0, []common.Hash{ common.Hash{}, - pm.blockchain.GetBlockByNumber(1).Hash(), + pm.blockchain.GetBlockByNumber(1, access.NoOdr).Hash(), common.Hash{}, - pm.blockchain.GetBlockByNumber(10).Hash(), + pm.blockchain.GetBlockByNumber(10, access.NoOdr).Hash(), common.Hash{}, - pm.blockchain.GetBlockByNumber(100).Hash(), + pm.blockchain.GetBlockByNumber(100, access.NoOdr).Hash(), common.Hash{}, }, []bool{false, true, false, true, false, true, false}, 3}, } @@ -177,7 +182,7 @@ func testGetBlocks(t *testing.T, protocol int) { if !seen[num] { seen[num] = true - block := pm.blockchain.GetBlockByNumber(uint64(num)) + block := pm.blockchain.GetBlockByNumber(uint64(num), access.NoOdr) hashes = append(hashes, block.Hash()) if len(blocks) < tt.expected { blocks = append(blocks, block) @@ -189,7 +194,7 @@ func testGetBlocks(t *testing.T, protocol int) { for j, hash := range tt.explicit { hashes = append(hashes, hash) if tt.available[j] && len(blocks) < tt.expected { - blocks = append(blocks, pm.blockchain.GetBlock(hash)) + blocks = append(blocks, pm.blockchain.GetBlock(hash, access.NoOdr)) } } // Send the hash request and verify the response @@ -222,48 +227,48 @@ func testGetBlockHeaders(t *testing.T, protocol int) { }{ // A single random block should be retrievable by hash and number too { - &getBlockHeadersData{Origin: hashOrNumber{Hash: pm.blockchain.GetBlockByNumber(limit / 2).Hash()}, Amount: 1}, - []common.Hash{pm.blockchain.GetBlockByNumber(limit / 2).Hash()}, + &getBlockHeadersData{Origin: hashOrNumber{Hash: pm.blockchain.GetBlockByNumber(limit/2, access.NoOdr).Hash()}, Amount: 1}, + []common.Hash{pm.blockchain.GetBlockByNumber(limit/2, access.NoOdr).Hash()}, }, { &getBlockHeadersData{Origin: hashOrNumber{Number: limit / 2}, Amount: 1}, - []common.Hash{pm.blockchain.GetBlockByNumber(limit / 2).Hash()}, + []common.Hash{pm.blockchain.GetBlockByNumber(limit/2, access.NoOdr).Hash()}, }, // Multiple headers should be retrievable in both directions { &getBlockHeadersData{Origin: hashOrNumber{Number: limit / 2}, Amount: 3}, []common.Hash{ - pm.blockchain.GetBlockByNumber(limit / 2).Hash(), - pm.blockchain.GetBlockByNumber(limit/2 + 1).Hash(), - pm.blockchain.GetBlockByNumber(limit/2 + 2).Hash(), + pm.blockchain.GetBlockByNumber(limit/2, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(limit/2+1, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(limit/2+2, access.NoOdr).Hash(), }, }, { &getBlockHeadersData{Origin: hashOrNumber{Number: limit / 2}, Amount: 3, Reverse: true}, []common.Hash{ - pm.blockchain.GetBlockByNumber(limit / 2).Hash(), - pm.blockchain.GetBlockByNumber(limit/2 - 1).Hash(), - pm.blockchain.GetBlockByNumber(limit/2 - 2).Hash(), + pm.blockchain.GetBlockByNumber(limit/2, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(limit/2-1, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(limit/2-2, access.NoOdr).Hash(), }, }, // Multiple headers with skip lists should be retrievable { &getBlockHeadersData{Origin: hashOrNumber{Number: limit / 2}, Skip: 3, Amount: 3}, []common.Hash{ - pm.blockchain.GetBlockByNumber(limit / 2).Hash(), - pm.blockchain.GetBlockByNumber(limit/2 + 4).Hash(), - pm.blockchain.GetBlockByNumber(limit/2 + 8).Hash(), + pm.blockchain.GetBlockByNumber(limit/2, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(limit/2+4, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(limit/2+8, access.NoOdr).Hash(), }, }, { &getBlockHeadersData{Origin: hashOrNumber{Number: limit / 2}, Skip: 3, Amount: 3, Reverse: true}, []common.Hash{ - pm.blockchain.GetBlockByNumber(limit / 2).Hash(), - pm.blockchain.GetBlockByNumber(limit/2 - 4).Hash(), - pm.blockchain.GetBlockByNumber(limit/2 - 8).Hash(), + pm.blockchain.GetBlockByNumber(limit/2, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(limit/2-4, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(limit/2-8, access.NoOdr).Hash(), }, }, // The chain endpoints should be retrievable { &getBlockHeadersData{Origin: hashOrNumber{Number: 0}, Amount: 1}, - []common.Hash{pm.blockchain.GetBlockByNumber(0).Hash()}, + []common.Hash{pm.blockchain.GetBlockByNumber(0, access.NoOdr).Hash()}, }, { &getBlockHeadersData{Origin: hashOrNumber{Number: pm.blockchain.CurrentBlock().NumberU64()}, Amount: 1}, []common.Hash{pm.blockchain.CurrentBlock().Hash()}, @@ -277,28 +282,28 @@ func testGetBlockHeaders(t *testing.T, protocol int) { { &getBlockHeadersData{Origin: hashOrNumber{Number: pm.blockchain.CurrentBlock().NumberU64() - 4}, Skip: 3, Amount: 3}, []common.Hash{ - pm.blockchain.GetBlockByNumber(pm.blockchain.CurrentBlock().NumberU64() - 4).Hash(), - pm.blockchain.GetBlockByNumber(pm.blockchain.CurrentBlock().NumberU64()).Hash(), + pm.blockchain.GetBlockByNumber(pm.blockchain.CurrentBlock().NumberU64()-4, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(pm.blockchain.CurrentBlock().NumberU64(), access.NoOdr).Hash(), }, }, { &getBlockHeadersData{Origin: hashOrNumber{Number: 4}, Skip: 3, Amount: 3, Reverse: true}, []common.Hash{ - pm.blockchain.GetBlockByNumber(4).Hash(), - pm.blockchain.GetBlockByNumber(0).Hash(), + pm.blockchain.GetBlockByNumber(4, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(0, access.NoOdr).Hash(), }, }, // Check that requesting more than available is handled gracefully, even if mid skip { &getBlockHeadersData{Origin: hashOrNumber{Number: pm.blockchain.CurrentBlock().NumberU64() - 4}, Skip: 2, Amount: 3}, []common.Hash{ - pm.blockchain.GetBlockByNumber(pm.blockchain.CurrentBlock().NumberU64() - 4).Hash(), - pm.blockchain.GetBlockByNumber(pm.blockchain.CurrentBlock().NumberU64() - 1).Hash(), + pm.blockchain.GetBlockByNumber(pm.blockchain.CurrentBlock().NumberU64()-4, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(pm.blockchain.CurrentBlock().NumberU64()-1, access.NoOdr).Hash(), }, }, { &getBlockHeadersData{Origin: hashOrNumber{Number: 4}, Skip: 2, Amount: 3, Reverse: true}, []common.Hash{ - pm.blockchain.GetBlockByNumber(4).Hash(), - pm.blockchain.GetBlockByNumber(1).Hash(), + pm.blockchain.GetBlockByNumber(4, access.NoOdr).Hash(), + pm.blockchain.GetBlockByNumber(1, access.NoOdr).Hash(), }, }, // Check that non existing headers aren't returned @@ -315,7 +320,7 @@ func testGetBlockHeaders(t *testing.T, protocol int) { // Collect the headers to expect in the response headers := []*types.Header{} for _, hash := range tt.expect { - headers = append(headers, pm.blockchain.GetBlock(hash).Header()) + headers = append(headers, pm.blockchain.GetBlock(hash, access.NoOdr).Header()) } // Send the hash request and verify the response p2p.Send(peer.app, 0x03, tt.query) @@ -353,11 +358,11 @@ func testGetBlockBodies(t *testing.T, protocol int) { // Existing and non-existing blocks interleaved should not cause problems {0, []common.Hash{ common.Hash{}, - pm.blockchain.GetBlockByNumber(1).Hash(), + pm.blockchain.GetBlockByNumber(1, access.NoOdr).Hash(), common.Hash{}, - pm.blockchain.GetBlockByNumber(10).Hash(), + pm.blockchain.GetBlockByNumber(10, access.NoOdr).Hash(), common.Hash{}, - pm.blockchain.GetBlockByNumber(100).Hash(), + pm.blockchain.GetBlockByNumber(100, access.NoOdr).Hash(), common.Hash{}, }, []bool{false, true, false, true, false, true, false}, 3}, } @@ -365,7 +370,7 @@ func testGetBlockBodies(t *testing.T, protocol int) { for i, tt := range tests { // Collect the hashes to request, and the response to expect hashes, seen := []common.Hash{}, make(map[int64]bool) - bodies := []*blockBody{} + bodies := []*types.Body{} for j := 0; j < tt.random; j++ { for { @@ -373,10 +378,10 @@ func testGetBlockBodies(t *testing.T, protocol int) { if !seen[num] { seen[num] = true - block := pm.blockchain.GetBlockByNumber(uint64(num)) + block := pm.blockchain.GetBlockByNumber(uint64(num), access.NoOdr) hashes = append(hashes, block.Hash()) if len(bodies) < tt.expected { - bodies = append(bodies, &blockBody{Transactions: block.Transactions(), Uncles: block.Uncles()}) + bodies = append(bodies, &types.Body{Transactions: block.Transactions(), Uncles: block.Uncles()}) } break } @@ -385,8 +390,8 @@ func testGetBlockBodies(t *testing.T, protocol int) { for j, hash := range tt.explicit { hashes = append(hashes, hash) if tt.available[j] && len(bodies) < tt.expected { - block := pm.blockchain.GetBlock(hash) - bodies = append(bodies, &blockBody{Transactions: block.Transactions(), Uncles: block.Uncles()}) + block := pm.blockchain.GetBlock(hash, access.NoOdr) + bodies = append(bodies, &types.Body{Transactions: block.Transactions(), Uncles: block.Uncles()}) } } // Send the hash request and verify the response @@ -442,7 +447,7 @@ func testGetNodeData(t *testing.T, protocol int) { // Fetch for now the entire chain db hashes := []common.Hash{} - for _, key := range pm.chaindb.(*ethdb.MemDatabase).Keys() { + for _, key := range pm.chainAccess.Db().(*ethdb.MemDatabase).Keys() { if len(key) == len(common.Hash{}) { hashes = append(hashes, common.BytesToHash(key)) } @@ -471,10 +476,10 @@ func testGetNodeData(t *testing.T, protocol int) { } accounts := []common.Address{testBankAddress, acc1Addr, acc2Addr} for i := uint64(0); i <= pm.blockchain.CurrentBlock().NumberU64(); i++ { - trie, _ := state.New(pm.blockchain.GetBlockByNumber(i).Root(), statedb) + trie, _ := state.New(pm.blockchain.GetBlockByNumber(i, access.NoOdr).Root(), access.NewDbChainAccess(statedb), access.NullCtx) for j, acc := range accounts { - state, _ := pm.blockchain.State() + state, _ := pm.blockchain.State(access.NullCtx) bw := state.GetBalance(acc) bh := trie.GetBalance(acc) @@ -534,10 +539,10 @@ func testGetReceipt(t *testing.T, protocol int) { // Collect the hashes to request, and the response to expect hashes, receipts := []common.Hash{}, []types.Receipts{} for i := uint64(0); i <= pm.blockchain.CurrentBlock().NumberU64(); i++ { - block := pm.blockchain.GetBlockByNumber(i) + block := pm.blockchain.GetBlockByNumber(i, access.NoOdr) hashes = append(hashes, block.Hash()) - receipts = append(receipts, core.GetBlockReceipts(pm.chaindb, block.Hash())) + receipts = append(receipts, core.GetBlockReceipts(pm.chainAccess, block.Hash(), access.NoOdr)) } // Send the hash request and verify the response p2p.Send(peer.app, 0x0f, hashes) diff --git a/eth/helper_test.go b/eth/helper_test.go index 65fccf7b4d..1057016fbc 100644 --- a/eth/helper_test.go +++ b/eth/helper_test.go @@ -11,6 +11,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/ethdb" @@ -28,21 +29,22 @@ var ( // newTestProtocolManager creates a new protocol manager for testing purposes, // with the given number of blocks already known, and potential notification // channels for different events. -func newTestProtocolManager(fastSync bool, blocks int, generator func(int, *core.BlockGen), newtx chan<- []*types.Transaction) (*ProtocolManager, error) { +func newTestProtocolManager(mode Mode, blocks int, generator func(int, *core.BlockGen), newtx chan<- []*types.Transaction) (*ProtocolManager, error) { var ( evmux = new(event.TypeMux) pow = new(core.FakePow) db, _ = ethdb.NewMemDatabase() + ca = access.NewDbChainAccess(db) genesis = core.WriteGenesisBlockForTesting(db, core.GenesisAccount{testBankAddress, testBankFunds}) - blockchain, _ = core.NewBlockChain(db, pow, evmux) - blockproc = core.NewBlockProcessor(db, pow, blockchain, evmux) + blockchain, _ = core.NewBlockChain(ca, pow, evmux) + blockproc = core.NewBlockProcessor(ca, pow, blockchain, evmux) ) blockchain.SetProcessor(blockproc) chain, _ := core.GenerateChain(genesis, db, blocks, generator) if _, err := blockchain.InsertChain(chain); err != nil { panic(err) } - pm, err := NewProtocolManager(fastSync, NetworkId, evmux, &testTxPool{added: newtx}, pow, blockchain, db) + pm, err := NewProtocolManager(mode, NetworkId, evmux, &testTxPool{added: newtx}, pow, blockchain, ca) if err != nil { return nil, err } @@ -55,7 +57,11 @@ func newTestProtocolManager(fastSync bool, blocks int, generator func(int, *core // channels for different events. In case of an error, the constructor force- // fails the test. func newTestProtocolManagerMust(t *testing.T, fastSync bool, blocks int, generator func(int, *core.BlockGen), newtx chan<- []*types.Transaction) *ProtocolManager { - pm, err := newTestProtocolManager(fastSync, blocks, generator, newtx) + mode := ArchiveMode + if fastSync { + mode = FullMode + } + pm, err := newTestProtocolManager(mode, blocks, generator, newtx) if err != nil { t.Fatalf("Failed to create protocol manager: %v", err) } diff --git a/eth/peer.go b/eth/peer.go index 15ba22ff53..0e20d084c7 100644 --- a/eth/peer.go +++ b/eth/peer.go @@ -196,7 +196,7 @@ func (p *peer) SendBlockHeaders(headers []*types.Header) error { } // SendBlockBodies sends a batch of block contents to the remote peer. -func (p *peer) SendBlockBodies(bodies []*blockBody) error { +func (p *peer) SendBlockBodies(bodies []*types.Body) error { return p2p.Send(p.rw, BlockBodiesMsg, blockBodiesData(bodies)) } diff --git a/eth/protocol.go b/eth/protocol.go index 808ac06011..78404bdf5a 100644 --- a/eth/protocol.go +++ b/eth/protocol.go @@ -26,6 +26,15 @@ import ( "github.com/ethereum/go-ethereum/rlp" ) +// Mode represents the mode of operation of the eth client. +type Mode int + +const ( + ArchiveMode Mode = iota // Maintain the entire blockchain history + FullMode // Maintain only a recent view of the blockchain + LightMode // Don't maintain any history, rather fetch on demand +) + // Constants to match up protocol versions and messages const ( eth61 = 61 @@ -201,14 +210,8 @@ type newBlockData struct { TD *big.Int } -// blockBody represents the data content of a single block. -type blockBody struct { - Transactions []*types.Transaction // Transactions contained within a block - Uncles []*types.Header // Uncles contained within a block -} - // blockBodiesData is the network packet for block content distribution. -type blockBodiesData []*blockBody +type blockBodiesData []*types.Body // nodeDataData is the network response packet for a node data retrieval. type nodeDataData []struct { diff --git a/eth/sync.go b/eth/sync.go index bbf2abc048..4d7d25a65d 100644 --- a/eth/sync.go +++ b/eth/sync.go @@ -161,7 +161,7 @@ func (pm *ProtocolManager) synchronise(peer *peer) { return } // Make sure the peer's TD is higher than our own. If not drop. - td := pm.blockchain.GetTd(pm.blockchain.CurrentBlock().Hash()) + td := pm.blockchain.GetTd(pm.blockchain.LastBlockHash()) if peer.Td().Cmp(td) <= 0 { return } diff --git a/les/handler.go b/les/handler.go new file mode 100644 index 0000000000..8855588e3c --- /dev/null +++ b/les/handler.go @@ -0,0 +1,543 @@ +// Copyright 2015 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 . + +// Package les implements the Light Ethereum Subprotocol. +package les + +import ( + "errors" + "fmt" + "sync" + //"time" + + "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" + "github.com/ethereum/go-ethereum/core/types" + "github.com/ethereum/go-ethereum/eth/downloader" + //"github.com/ethereum/go-ethereum/eth/fetcher" + "github.com/ethereum/go-ethereum/event" + "github.com/ethereum/go-ethereum/logger" + "github.com/ethereum/go-ethereum/logger/glog" + "github.com/ethereum/go-ethereum/p2p" + "github.com/ethereum/go-ethereum/pow" + "github.com/ethereum/go-ethereum/rlp" + "github.com/ethereum/go-ethereum/trie" +) + +const ( + softResponseLimit = 2 * 1024 * 1024 // Target maximum size of returned blocks, headers or node data. + estHeaderRlpSize = 500 // Approximate size of an RLP encoded block header + + ethVersion = 63 // equivalent eth version for the downloader +) + +// errIncompatibleConfig is returned if the requested protocols and configs are +// not compatible (low protocol version restrictions and high requirements). +var errIncompatibleConfig = errors.New("incompatible configuration") + +func errResp(code errCode, format string, v ...interface{}) error { + return fmt.Errorf("%v - %v", code, fmt.Sprintf(format, v...)) +} + +type hashFetcherFn func(common.Hash) error + +type ProtocolManager struct { + mode Mode + blockchain *core.BlockChain + chainAccess *access.ChainAccess + + downloader *downloader.Downloader + //fetcher *fetcher.Fetcher + peers *peerSet + + SubProtocols []p2p.Protocol + + eventMux *event.TypeMux + + // channels for fetcher, syncer, txsyncLoop + newPeerCh chan *peer + quitSync chan struct{} + + // wait group is used for graceful shutdowns during downloading + // and processing + wg sync.WaitGroup + quit bool +} + +// NewProtocolManager returns a new ethereum sub protocol manager. The Ethereum sub protocol manages peers capable +// with the ethereum network. +func NewProtocolManager(mode Mode, networkId int, mux *event.TypeMux, pow pow.PoW, blockchain *core.BlockChain, ca *access.ChainAccess) (*ProtocolManager, error) { + // Create the protocol manager with the base fields + manager := &ProtocolManager{ + mode: mode, + eventMux: mux, + blockchain: blockchain, + chainAccess: ca, + peers: newPeerSet(), + newPeerCh: make(chan *peer, 1), + quitSync: make(chan struct{}), + } + // Initiate a sub-protocol for every implemented version we can handle + manager.SubProtocols = make([]p2p.Protocol, 0, len(ProtocolVersions)) + for i, version := range ProtocolVersions { + // Skip protocol version if incompatible with the mode of operation + if minimumProtocolVersion[mode] > version { + continue + } + // Compatible, initialize the sub-protocol + version := version // Closure for the run + manager.SubProtocols = append(manager.SubProtocols, p2p.Protocol{ + Name: "les", + Version: version, + Length: ProtocolLengths[i], + Run: func(p *p2p.Peer, rw p2p.MsgReadWriter) error { + peer := manager.newPeer(int(version), networkId, p, rw) + manager.newPeerCh <- peer + return manager.handle(peer) + }, + }) + } + if len(manager.SubProtocols) == 0 { + return nil, errIncompatibleConfig + } + // Construct the different synchronisation mechanisms + var syncMode downloader.SyncMode + switch mode { + case ArchiveMode: + syncMode = downloader.NoSync + case FullMode: + syncMode = downloader.NoSync + case LightMode: + syncMode = downloader.LightSync + } + glog.V(access.LogLevel).Infof("LES: create downloader") + + manager.downloader = downloader.New(syncMode, ca.Db(), manager.eventMux, blockchain.HasHeader, blockchain.HasBlock, blockchain.GetHeader, + blockchain.GetBlockNoOdr, blockchain.CurrentHeader, blockchain.CurrentBlock, blockchain.CurrentFastBlock, blockchain.FastSyncCommitHead, blockchain.GetTd, + blockchain.InsertHeaderChain, blockchain.InsertChain, blockchain.InsertReceiptChain, blockchain.Rollback, manager.removePeer) + + /*validator := func(block *types.Block, parent *types.Block) error { + return core.ValidateHeader(pow, block.Header(), parent.Header(), true, false) + } + heighter := func() uint64 { + return chainman.LastBlockNumberU64() + } + manager.fetcher = fetcher.New(chainman.GetBlockNoOdr, validator, nil, heighter, chainman.InsertChain, manager.removePeer) + */ + return manager, nil +} + +func (pm *ProtocolManager) removePeer(id string) { + // Short circuit if the peer was already removed + peer := pm.peers.Peer(id) + if peer == nil { + return + } + glog.V(logger.Debug).Infoln("Removing peer", id) + + // Unregister the peer from the downloader and Ethereum peer set + glog.V(access.LogLevel).Infof("LES: unregister peer %v", id) + pm.downloader.UnregisterPeer(id) + if err := pm.peers.Unregister(id); err != nil { + glog.V(logger.Error).Infoln("Removal failed:", err) + } + pm.chainAccess.UnregisterPeer(id) + // Hard disconnect at the networking layer + if peer != nil { + peer.Peer.Disconnect(p2p.DiscUselessPeer) + } +} + +func (pm *ProtocolManager) Start() { + // start sync handler + go pm.syncer() +} + +func (pm *ProtocolManager) Stop() { + // Showing a log message. During download / process this could actually + // take between 5 to 10 seconds and therefor feedback is required. + glog.V(logger.Info).Infoln("Stopping light ethereum protocol handler...") + + pm.quit = true + close(pm.quitSync) // quits syncer, fetcher + + // Wait for any process action + pm.wg.Wait() + + glog.V(logger.Info).Infoln("Ethereum protocol handler stopped") +} + +func (pm *ProtocolManager) newPeer(pv, nv int, p *p2p.Peer, rw p2p.MsgReadWriter) *peer { + return newPeer(pv, nv, p, newMeteredMsgWriter(rw)) +} + +// handle is the callback invoked to manage the life cycle of a les peer. When +// this function terminates, the peer is disconnected. +func (pm *ProtocolManager) handle(p *peer) error { + glog.V(logger.Debug).Infof("%v: peer connected [%s]", p, p.Name()) + + // Execute the LES handshake + td, head, genesis := pm.blockchain.Status() + if err := p.Handshake(td, head, genesis); err != nil { + glog.V(logger.Debug).Infof("%v: handshake failed: %v", p, err) + return err + } + if rw, ok := p.rw.(*meteredMsgReadWriter); ok { + rw.Init(p.version) + } + // Register the peer locally + glog.V(logger.Detail).Infof("%v: adding peer", p) + if err := pm.peers.Register(p); err != nil { + glog.V(logger.Error).Infof("%v: addition failed: %v", p, err) + return err + } + defer pm.removePeer(p.id) + + // Register the peer in the downloader. If the downloader considers it banned, we disconnect + glog.V(access.LogLevel).Infof("LES: register peer %v", p.id) + if err := pm.downloader.RegisterPeer(p.id, ethVersion, p.Head(), + nil, nil, nil, p.RequestHeadersByHash, p.RequestHeadersByNumber, + p.RequestBodies, p.RequestReceipts, p.RequestNodeData); err != nil { + return err + } + + pm.chainAccess.RegisterPeer(p.id, p.version, p.Head(), p.RequestBodies, p.RequestNodeData, p.RequestReceipts, p.RequestProofs) + + // main loop. handle incoming messages. + for { + if err := pm.handleMsg(p); err != nil { + glog.V(logger.Debug).Infof("%v: message handling failed: %v", p, err) + return err + } + } + return nil +} + +// handleMsg is invoked whenever an inbound message is received from a remote +// peer. The remote connection is torn down upon returning any error. +func (pm *ProtocolManager) handleMsg(p *peer) error { + // Read the next message from the remote peer, and ensure it's fully consumed + msg, err := p.rw.ReadMsg() + if err != nil { + return err + } + if msg.Size > ProtocolMaxMsgSize { + return errResp(ErrMsgTooLarge, "%v > %v", msg.Size, ProtocolMaxMsgSize) + } + defer msg.Discard() + + // Handle the message depending on its contents + switch { + case msg.Code == StatusMsg: + glog.V(access.LogLevel).Infof("LES: received StatusMsg from peer %v", p.id) + // Status messages should never arrive after the handshake + return errResp(ErrExtraStatusMsg, "uncontrolled status message") + + // Block header query, collect the requested headers and reply + case msg.Code == GetBlockHeadersMsg: + glog.V(access.LogLevel).Infof("LES: received GetBlockHeadersMsg from peer %v", p.id) + // Decode the complex header query + var query getBlockHeadersData + if err := msg.Decode(&query); err != nil { + return errResp(ErrDecode, "%v: %v", msg, err) + } + // Gather headers until the fetch or network limits is reached + var ( + bytes common.StorageSize + headers []*types.Header + unknown bool + ) + for !unknown && len(headers) < int(query.Amount) && bytes < softResponseLimit && len(headers) < downloader.MaxHeaderFetch { + // Retrieve the next header satisfying the query + var origin *types.Header + if query.Origin.Hash != (common.Hash{}) { + origin = pm.blockchain.GetHeader(query.Origin.Hash) + } else { + origin = pm.blockchain.GetHeaderByNumber(query.Origin.Number) + } + if origin == nil { + break + } + headers = append(headers, origin) + bytes += estHeaderRlpSize + + // Advance to the next header of the query + switch { + case query.Origin.Hash != (common.Hash{}) && query.Reverse: + // Hash based traversal towards the genesis block + for i := 0; i < int(query.Skip)+1; i++ { + if header := pm.blockchain.GetHeader(query.Origin.Hash); header != nil { + query.Origin.Hash = header.ParentHash + } else { + unknown = true + break + } + } + case query.Origin.Hash != (common.Hash{}) && !query.Reverse: + // Hash based traversal towards the leaf block + if header := pm.blockchain.GetHeaderByNumber(origin.Number.Uint64() + query.Skip + 1); header != nil { + if pm.blockchain.GetBlockHashesFromHash(header.Hash(), query.Skip+1)[query.Skip] == query.Origin.Hash { + query.Origin.Hash = header.Hash() + } else { + unknown = true + } + } else { + unknown = true + } + case query.Reverse: + // Number based traversal towards the genesis block + if query.Origin.Number >= query.Skip+1 { + query.Origin.Number -= (query.Skip + 1) + } else { + unknown = true + } + + case !query.Reverse: + // Number based traversal towards the leaf block + query.Origin.Number += (query.Skip + 1) + } + } + return p.SendBlockHeaders(headers) + + case msg.Code == BlockHeadersMsg: + glog.V(access.LogLevel).Infof("LES: received BlockHeadersMsg from peer %v", p.id) + // A batch of headers arrived to one of our previous requests + var headers []*types.Header + if err := msg.Decode(&headers); err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + // Filter out any explicitly requested headers, deliver the rest to the downloader + /*filter := len(headers) == 1 + if filter { + headers = pm.fetcher.FilterHeaders(headers, time.Now()) + } + if len(headers) > 0 || !filter {*/ + err := pm.downloader.DeliverHeaders(p.id, headers) + if err != nil { + glog.V(logger.Debug).Infoln(err) + } + //} + + case msg.Code == GetBlockBodiesMsg: + glog.V(access.LogLevel).Infof("LES: received GetBlockBodiesMsg from peer %v", p.id) + // Decode the retrieval message + msgStream := rlp.NewStream(msg.Payload, uint64(msg.Size)) + if _, err := msgStream.List(); err != nil { + return err + } + // Gather blocks until the fetch or network limits is reached + var ( + hash common.Hash + bytes int + bodies []rlp.RawValue + ) + for bytes < softResponseLimit && len(bodies) < downloader.MaxBlockFetch { + // Retrieve the hash of the next block + if err := msgStream.Decode(&hash); err == rlp.EOL { + break + } else if err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + // Retrieve the requested block body, stopping if enough was found + if data := pm.blockchain.GetBodyRLP(hash, access.NoOdr); len(data) != 0 { + bodies = append(bodies, data) + bytes += len(data) + } + } + return p.SendBlockBodiesRLP(bodies) + + case msg.Code == BlockBodiesMsg: + glog.V(access.LogLevel).Infof("LES: received BlockBodiesMsg from peer %v", p.id) + // A batch of block bodies arrived to one of our previous requests + var data []*types.Body + if err := msg.Decode(&data); err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + pm.chainAccess.Deliver(p.id, &access.Msg{ + MsgType: access.MsgBlockBodies, + Obj: data, + }) + + case msg.Code == GetNodeDataMsg: + glog.V(access.LogLevel).Infof("LES: received GetNodeDataMsg from peer %v", p.id) + // Decode the retrieval message + msgStream := rlp.NewStream(msg.Payload, uint64(msg.Size)) + if _, err := msgStream.List(); err != nil { + return err + } + // Gather state data until the fetch or network limits is reached + var ( + hash common.Hash + bytes int + data [][]byte + ) + for bytes < softResponseLimit && len(data) < downloader.MaxStateFetch { + // Retrieve the hash of the next state entry + if err := msgStream.Decode(&hash); err == rlp.EOL { + break + } else if err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + // Retrieve the requested state entry, stopping if enough was found + if entry, err := pm.chainAccess.Db().Get(hash.Bytes()); err == nil { + data = append(data, entry) + bytes += len(entry) + } + } + return p.SendNodeData(data) + + case msg.Code == NodeDataMsg: + glog.V(access.LogLevel).Infof("LES: received NodeDataMsg from peer %v", p.id) + // A batch of node state data arrived to one of our previous requests + var data [][]byte + if err := msg.Decode(&data); err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + pm.chainAccess.Deliver(p.id, &access.Msg{ + MsgType: access.MsgNodeData, + Obj: data, + }) + + case msg.Code == GetReceiptsMsg: + glog.V(access.LogLevel).Infof("LES: received GetReceiptsMsg from peer %v", p.id) + // Decode the retrieval message + msgStream := rlp.NewStream(msg.Payload, uint64(msg.Size)) + if _, err := msgStream.List(); err != nil { + return err + } + // Gather state data until the fetch or network limits is reached + var ( + hash common.Hash + bytes int + receipts []rlp.RawValue + ) + for bytes < softResponseLimit && len(receipts) < downloader.MaxReceiptFetch { + // Retrieve the hash of the next block + if err := msgStream.Decode(&hash); err == rlp.EOL { + break + } else if err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + // Retrieve the requested block's receipts, skipping if unknown to us + results := core.GetBlockReceipts(pm.chainAccess, hash, access.NoOdr) + if results == nil { + if header := pm.blockchain.GetHeader(hash); header == nil || header.ReceiptHash != types.EmptyRootHash { + continue + } + } + // If known, encode and queue for response packet + if encoded, err := rlp.EncodeToBytes(results); err != nil { + glog.V(logger.Error).Infof("failed to encode receipt: %v", err) + } else { + receipts = append(receipts, encoded) + bytes += len(encoded) + } + } + return p.SendReceiptsRLP(receipts) + + case msg.Code == ReceiptsMsg: + glog.V(access.LogLevel).Infof("LES: received ReceiptsMsg from peer %v", p.id) + // A batch of receipts arrived to one of our previous requests + var receipts [][]*types.Receipt + if err := msg.Decode(&receipts); err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + pm.chainAccess.Deliver(p.id, &access.Msg{ + MsgType: access.MsgReceipts, + Obj: receipts[0], + }) + + case msg.Code == GetProofsMsg: + glog.V(access.LogLevel).Infof("LES: received GetProofsMsg from peer %v", p.id) + // Decode the retrieval message + msgStream := rlp.NewStream(msg.Payload, uint64(msg.Size)) + if _, err := msgStream.List(); err != nil { + return err + } + // Gather state data until the fetch or network limits is reached + var ( + req access.ProofReq + bytes int + proofs proofsData + ) + for bytes < softResponseLimit && len(proofs) < downloader.MaxProofsFetch { + // Retrieve the hash of the next state entry + if err := msgStream.Decode(&req); err == rlp.EOL { + break + } else if err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + // Retrieve the requested state entry, stopping if enough was found + if trie, _ := trie.NewSecure(req.Root, pm.chainAccess.Db(), nil); trie != nil { + proof := trie.Prove(req.Key) + proofs = append(proofs, proof) + bytes += len(proof) + } + } + return p.SendProofs(proofs) + + case msg.Code == ProofsMsg: + glog.V(access.LogLevel).Infof("LES: received ProofsMsg from peer %v", p.id) + // A batch of merkle proofs arrived to one of our previous requests + var data []trie.MerkleProof + if err := msg.Decode(&data); err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + pm.chainAccess.Deliver(p.id, &access.Msg{ + MsgType: access.MsgProofs, + Obj: data, + }) + + case msg.Code == NewBlockHashesMsg: + glog.V(access.LogLevel).Infof("LES: received NewBlockHashesMsg from peer %v", p.id) + // Retrieve and deseralize the remote new block hashes notification + type announce struct { + Hash common.Hash + Number uint64 + } + var announces = []announce{} + + // We're running the old protocol, make block number unknown (0) + var hashes []common.Hash + if err := msg.Decode(&hashes); err != nil { + return errResp(ErrDecode, "%v: %v", msg, err) + } + for _, hash := range hashes { + announces = append(announces, announce{hash, 0}) + } + // Mark the hashes as present at the remote node + for _, block := range announces { + p.MarkBlock(block.Hash) + p.SetHead(block.Hash) + } + // Schedule all the unknown hashes for retrieval + unknown := make([]announce, 0, len(announces)) + for _, block := range announces { + if !pm.blockchain.HasBlock(block.Hash) { + unknown = append(unknown, block) + } + } + /*for _, block := range unknown { + pm.fetcher.Notify(p.id, block.Hash, block.Number, time.Now(), nil, p.RequestOneHeader, p.RequestBodies) + }*/ + + default: + glog.V(access.LogLevel).Infof("LES: received unknown message with code %d from peer %v", msg.Code, p.id) + return errResp(ErrInvalidMsgCode, "%v", msg.Code) + } + return nil +} diff --git a/les/metrics.go b/les/metrics.go new file mode 100644 index 0000000000..d533625a8a --- /dev/null +++ b/les/metrics.go @@ -0,0 +1,157 @@ +// Copyright 2015 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 . + +package les + +import ( + "github.com/ethereum/go-ethereum/metrics" + "github.com/ethereum/go-ethereum/p2p" +) + +var ( + /* propTxnInPacketsMeter = metrics.NewMeter("eth/prop/txns/in/packets") + propTxnInTrafficMeter = metrics.NewMeter("eth/prop/txns/in/traffic") + propTxnOutPacketsMeter = metrics.NewMeter("eth/prop/txns/out/packets") + propTxnOutTrafficMeter = metrics.NewMeter("eth/prop/txns/out/traffic") + propHashInPacketsMeter = metrics.NewMeter("eth/prop/hashes/in/packets") + propHashInTrafficMeter = metrics.NewMeter("eth/prop/hashes/in/traffic") + propHashOutPacketsMeter = metrics.NewMeter("eth/prop/hashes/out/packets") + propHashOutTrafficMeter = metrics.NewMeter("eth/prop/hashes/out/traffic") + propBlockInPacketsMeter = metrics.NewMeter("eth/prop/blocks/in/packets") + propBlockInTrafficMeter = metrics.NewMeter("eth/prop/blocks/in/traffic") + propBlockOutPacketsMeter = metrics.NewMeter("eth/prop/blocks/out/packets") + propBlockOutTrafficMeter = metrics.NewMeter("eth/prop/blocks/out/traffic") + reqHashInPacketsMeter = metrics.NewMeter("eth/req/hashes/in/packets") + reqHashInTrafficMeter = metrics.NewMeter("eth/req/hashes/in/traffic") + reqHashOutPacketsMeter = metrics.NewMeter("eth/req/hashes/out/packets") + reqHashOutTrafficMeter = metrics.NewMeter("eth/req/hashes/out/traffic") + reqBlockInPacketsMeter = metrics.NewMeter("eth/req/blocks/in/packets") + reqBlockInTrafficMeter = metrics.NewMeter("eth/req/blocks/in/traffic") + reqBlockOutPacketsMeter = metrics.NewMeter("eth/req/blocks/out/packets") + reqBlockOutTrafficMeter = metrics.NewMeter("eth/req/blocks/out/traffic") + reqHeaderInPacketsMeter = metrics.NewMeter("eth/req/headers/in/packets") + reqHeaderInTrafficMeter = metrics.NewMeter("eth/req/headers/in/traffic") + reqHeaderOutPacketsMeter = metrics.NewMeter("eth/req/headers/out/packets") + reqHeaderOutTrafficMeter = metrics.NewMeter("eth/req/headers/out/traffic") + reqBodyInPacketsMeter = metrics.NewMeter("eth/req/bodies/in/packets") + reqBodyInTrafficMeter = metrics.NewMeter("eth/req/bodies/in/traffic") + reqBodyOutPacketsMeter = metrics.NewMeter("eth/req/bodies/out/packets") + reqBodyOutTrafficMeter = metrics.NewMeter("eth/req/bodies/out/traffic") + reqStateInPacketsMeter = metrics.NewMeter("eth/req/states/in/packets") + reqStateInTrafficMeter = metrics.NewMeter("eth/req/states/in/traffic") + reqStateOutPacketsMeter = metrics.NewMeter("eth/req/states/out/packets") + reqStateOutTrafficMeter = metrics.NewMeter("eth/req/states/out/traffic") + reqReceiptInPacketsMeter = metrics.NewMeter("eth/req/receipts/in/packets") + reqReceiptInTrafficMeter = metrics.NewMeter("eth/req/receipts/in/traffic") + reqReceiptOutPacketsMeter = metrics.NewMeter("eth/req/receipts/out/packets") + reqReceiptOutTrafficMeter = metrics.NewMeter("eth/req/receipts/out/traffic")*/ + miscInPacketsMeter = metrics.NewMeter("les/misc/in/packets") + miscInTrafficMeter = metrics.NewMeter("les/misc/in/traffic") + miscOutPacketsMeter = metrics.NewMeter("les/misc/out/packets") + miscOutTrafficMeter = metrics.NewMeter("les/misc/out/traffic") +) + +// meteredMsgReadWriter is a wrapper around a p2p.MsgReadWriter, capable of +// accumulating the above defined metrics based on the data stream contents. +type meteredMsgReadWriter struct { + p2p.MsgReadWriter // Wrapped message stream to meter + version int // Protocol version to select correct meters +} + +// newMeteredMsgWriter wraps a p2p MsgReadWriter with metering support. If the +// metrics system is disabled, this fucntion returns the original object. +func newMeteredMsgWriter(rw p2p.MsgReadWriter) p2p.MsgReadWriter { + if !metrics.Enabled { + return rw + } + return &meteredMsgReadWriter{MsgReadWriter: rw} +} + +// Init sets the protocol version used by the stream to know which meters to +// increment in case of overlapping message ids between protocol versions. +func (rw *meteredMsgReadWriter) Init(version int) { + rw.version = version +} + +func (rw *meteredMsgReadWriter) ReadMsg() (p2p.Msg, error) { + // Read the message and short circuit in case of an error + msg, err := rw.MsgReadWriter.ReadMsg() + if err != nil { + return msg, err + } + // Account for the data traffic + packets, traffic := miscInPacketsMeter, miscInTrafficMeter + /* switch { + case rw.version < eth62 && msg.Code == BlockHashesMsg: + packets, traffic = reqHashInPacketsMeter, reqHashInTrafficMeter + case rw.version < eth62 && msg.Code == BlocksMsg: + packets, traffic = reqBlockInPacketsMeter, reqBlockInTrafficMeter + + case rw.version >= eth62 && msg.Code == BlockHeadersMsg: + packets, traffic = reqBlockInPacketsMeter, reqBlockInTrafficMeter + case rw.version >= eth62 && msg.Code == BlockBodiesMsg: + packets, traffic = reqBodyInPacketsMeter, reqBodyInTrafficMeter + + case rw.version >= eth63 && msg.Code == NodeDataMsg: + packets, traffic = reqStateInPacketsMeter, reqStateInTrafficMeter + case rw.version >= eth63 && msg.Code == ReceiptsMsg: + packets, traffic = reqReceiptInPacketsMeter, reqReceiptInTrafficMeter + + case msg.Code == NewBlockHashesMsg: + packets, traffic = propHashInPacketsMeter, propHashInTrafficMeter + case msg.Code == NewBlockMsg: + packets, traffic = propBlockInPacketsMeter, propBlockInTrafficMeter + case msg.Code == TxMsg: + packets, traffic = propTxnInPacketsMeter, propTxnInTrafficMeter + }*/ + packets.Mark(1) + traffic.Mark(int64(msg.Size)) + + return msg, err +} + +func (rw *meteredMsgReadWriter) WriteMsg(msg p2p.Msg) error { + // Account for the data traffic + packets, traffic := miscOutPacketsMeter, miscOutTrafficMeter + /* switch { + case rw.version < eth62 && msg.Code == BlockHashesMsg: + packets, traffic = reqHashOutPacketsMeter, reqHashOutTrafficMeter + case rw.version < eth62 && msg.Code == BlocksMsg: + packets, traffic = reqBlockOutPacketsMeter, reqBlockOutTrafficMeter + + case rw.version >= eth62 && msg.Code == BlockHeadersMsg: + packets, traffic = reqHeaderOutPacketsMeter, reqHeaderOutTrafficMeter + case rw.version >= eth62 && msg.Code == BlockBodiesMsg: + packets, traffic = reqBodyOutPacketsMeter, reqBodyOutTrafficMeter + + case rw.version >= eth63 && msg.Code == NodeDataMsg: + packets, traffic = reqStateOutPacketsMeter, reqStateOutTrafficMeter + case rw.version >= eth63 && msg.Code == ReceiptsMsg: + packets, traffic = reqReceiptOutPacketsMeter, reqReceiptOutTrafficMeter + + case msg.Code == NewBlockHashesMsg: + packets, traffic = propHashOutPacketsMeter, propHashOutTrafficMeter + case msg.Code == NewBlockMsg: + packets, traffic = propBlockOutPacketsMeter, propBlockOutTrafficMeter + case msg.Code == TxMsg: + packets, traffic = propTxnOutPacketsMeter, propTxnOutTrafficMeter + }*/ + packets.Mark(1) + traffic.Mark(int64(msg.Size)) + + // Send the packet to the p2p layer + return rw.MsgReadWriter.WriteMsg(msg) +} diff --git a/les/peer.go b/les/peer.go new file mode 100644 index 0000000000..f997330529 --- /dev/null +++ b/les/peer.go @@ -0,0 +1,349 @@ +// Copyright 2015 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 . + +// Package les implements the Light Ethereum Subprotocol. +package les + +import ( + "errors" + "fmt" + "math/big" + "sync" + + "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" + "github.com/ethereum/go-ethereum/core/types" + "github.com/ethereum/go-ethereum/logger" + "github.com/ethereum/go-ethereum/logger/glog" + "github.com/ethereum/go-ethereum/p2p" + "github.com/ethereum/go-ethereum/rlp" + "gopkg.in/fatih/set.v0" +) + +var ( + errAlreadyRegistered = errors.New("peer is already registered") + errNotRegistered = errors.New("peer is not registered") +) + +const ( + maxKnownBlocks = 1024 // Maximum block hashes to keep in the known list (prevent DOS) +) + +type peer struct { + *p2p.Peer + + rw p2p.MsgReadWriter + + version int // Protocol version negotiated + network int // Network ID being on + + id string + + head common.Hash + td *big.Int + lock sync.RWMutex + + knownBlocks *set.Set // Set of block hashes known to be known by this peer +} + +func newPeer(version, network int, p *p2p.Peer, rw p2p.MsgReadWriter) *peer { + id := p.ID() + + return &peer{ + Peer: p, + rw: rw, + version: version, + network: network, + id: fmt.Sprintf("%x", id[:8]), + knownBlocks: set.New(), + } +} + +// Head retrieves a copy of the current head (most recent) hash of the peer. +func (p *peer) Head() (hash common.Hash) { + p.lock.RLock() + defer p.lock.RUnlock() + + copy(hash[:], p.head[:]) + return hash +} + +// SetHead updates the head (most recent) hash of the peer. +func (p *peer) SetHead(hash common.Hash) { + p.lock.Lock() + defer p.lock.Unlock() + + copy(p.head[:], hash[:]) +} + +// Td retrieves the current total difficulty of a peer. +func (p *peer) Td() *big.Int { + p.lock.RLock() + defer p.lock.RUnlock() + + return new(big.Int).Set(p.td) +} + +// SetTd updates the current total difficulty of a peer. +func (p *peer) SetTd(td *big.Int) { + p.lock.Lock() + defer p.lock.Unlock() + + p.td.Set(td) +} + +// MarkBlock marks a block as known for the peer, ensuring that the block will +// never be propagated to this particular peer. +func (p *peer) MarkBlock(hash common.Hash) { + // If we reached the memory allowance, drop a previously known block hash + for p.knownBlocks.Size() >= maxKnownBlocks { + p.knownBlocks.Pop() + } + p.knownBlocks.Add(hash) +} + +// SendNewBlockHashes announces the availability of a number of blocks through +// a hash notification. +func (p *peer) SendNewBlockHashes(hashes []common.Hash, numbers []uint64) error { + for _, hash := range hashes { + p.knownBlocks.Add(hash) + } + request := make(newBlockHashesData, len(hashes)) + for i := 0; i < len(hashes); i++ { + request[i].Hash = hashes[i] + request[i].Number = numbers[i] + } + return p2p.Send(p.rw, NewBlockHashesMsg, request) +} + +// SendBlockHeaders sends a batch of block headers to the remote peer. +func (p *peer) SendBlockHeaders(headers []*types.Header) error { + return p2p.Send(p.rw, BlockHeadersMsg, headers) +} + +// SendBlockBodies sends a batch of block contents to the remote peer. +func (p *peer) SendBlockBodies(bodies []*types.Body) error { + return p2p.Send(p.rw, BlockBodiesMsg, blockBodiesData(bodies)) +} + +// SendBlockBodiesRLP sends a batch of block contents to the remote peer from +// an already RLP encoded format. +func (p *peer) SendBlockBodiesRLP(bodies []rlp.RawValue) error { + return p2p.Send(p.rw, BlockBodiesMsg, bodies) +} + +// SendNodeDataRLP sends a batch of arbitrary internal data, corresponding to the +// hashes requested. +func (p *peer) SendNodeData(data [][]byte) error { + return p2p.Send(p.rw, NodeDataMsg, data) +} + +// SendReceiptsRLP sends a batch of transaction receipts, corresponding to the +// ones requested from an already RLP encoded format. +func (p *peer) SendReceiptsRLP(receipts []rlp.RawValue) error { + return p2p.Send(p.rw, ReceiptsMsg, receipts) +} + +// SendProofs sends a batch of merkle proofs, corresponding to the ones requested. +func (p *peer) SendProofs(proofs proofsData) error { + return p2p.Send(p.rw, ProofsMsg, proofs) +} + +// RequestHeaders is a wrapper around the header query functions to fetch a +// single header. It is used solely by the fetcher. +func (p *peer) RequestOneHeader(hash common.Hash) error { + glog.V(logger.Debug).Infof("%v fetching a single header: %x", p, hash) + return p2p.Send(p.rw, GetBlockHeadersMsg, &getBlockHeadersData{Origin: hashOrNumber{Hash: hash}, Amount: uint64(1), Skip: uint64(0), Reverse: false}) +} + +// RequestHeadersByHash fetches a batch of blocks' headers corresponding to the +// specified header query, based on the hash of an origin block. +func (p *peer) RequestHeadersByHash(origin common.Hash, amount int, skip int, reverse bool) error { + glog.V(logger.Debug).Infof("%v fetching %d headers from %x, skipping %d (reverse = %v)", p, amount, origin[:4], skip, reverse) + return p2p.Send(p.rw, GetBlockHeadersMsg, &getBlockHeadersData{Origin: hashOrNumber{Hash: origin}, Amount: uint64(amount), Skip: uint64(skip), Reverse: reverse}) +} + +// RequestHeadersByNumber fetches a batch of blocks' headers corresponding to the +// specified header query, based on the number of an origin block. +func (p *peer) RequestHeadersByNumber(origin uint64, amount int, skip int, reverse bool) error { + glog.V(logger.Debug).Infof("%v fetching %d headers from #%d, skipping %d (reverse = %v)", p, amount, origin, skip, reverse) + return p2p.Send(p.rw, GetBlockHeadersMsg, &getBlockHeadersData{Origin: hashOrNumber{Number: origin}, Amount: uint64(amount), Skip: uint64(skip), Reverse: reverse}) +} + +// RequestBodies fetches a batch of blocks' bodies corresponding to the hashes +// specified. +func (p *peer) RequestBodies(hashes []common.Hash) error { + glog.V(logger.Debug).Infof("%v fetching %d block bodies", p, len(hashes)) + return p2p.Send(p.rw, GetBlockBodiesMsg, hashes) +} + +// RequestNodeData fetches a batch of arbitrary data from a node's known state +// data, corresponding to the specified hashes. +func (p *peer) RequestNodeData(hashes []common.Hash) error { + glog.V(logger.Debug).Infof("%v fetching %v state data", p, len(hashes)) + return p2p.Send(p.rw, GetNodeDataMsg, hashes) +} + +// RequestReceipts fetches a batch of transaction receipts from a remote node. +func (p *peer) RequestReceipts(hashes []common.Hash) error { + glog.V(logger.Debug).Infof("%v fetching %v receipts", p, len(hashes)) + return p2p.Send(p.rw, GetReceiptsMsg, hashes) +} + +// RequestProofs fetches a batch of merkle proofs from a remote node. +func (p *peer) RequestProofs(reqs []*access.ProofReq) error { + glog.V(logger.Debug).Infof("%v fetching %v proofs", p, len(reqs)) + return p2p.Send(p.rw, GetProofsMsg, reqs) +} + +// Handshake executes the les protocol handshake, negotiating version number, +// network IDs, difficulties, head and genesis blocks. +func (p *peer) Handshake(td *big.Int, head common.Hash, genesis common.Hash) error { + // Send out own handshake in a new thread + errc := make(chan error, 1) + go func() { + errc <- p2p.Send(p.rw, StatusMsg, &statusData{ + ProtocolVersion: uint32(p.version), + NetworkId: uint32(p.network), + TD: td, + CurrentBlock: head, + GenesisBlock: genesis, + }) + }() + // In the mean time retrieve the remote status message + msg, err := p.rw.ReadMsg() + if err != nil { + return err + } + if msg.Code != StatusMsg { + return errResp(ErrNoStatusMsg, "first msg has code %x (!= %x)", msg.Code, StatusMsg) + } + if msg.Size > ProtocolMaxMsgSize { + return errResp(ErrMsgTooLarge, "%v > %v", msg.Size, ProtocolMaxMsgSize) + } + // Decode the handshake and make sure everything matches + var status statusData + if err := msg.Decode(&status); err != nil { + return errResp(ErrDecode, "msg %v: %v", msg, err) + } + if status.GenesisBlock != genesis { + return errResp(ErrGenesisBlockMismatch, "%x (!= %x)", status.GenesisBlock, genesis) + } + if int(status.NetworkId) != p.network { + return errResp(ErrNetworkIdMismatch, "%d (!= %d)", status.NetworkId, p.network) + } + if int(status.ProtocolVersion) != p.version { + return errResp(ErrProtocolVersionMismatch, "%d (!= %d)", status.ProtocolVersion, p.version) + } + // Configure the remote peer, and sanity check out handshake too + p.td, p.head = status.TD, status.CurrentBlock + return <-errc +} + +// String implements fmt.Stringer. +func (p *peer) String() string { + return fmt.Sprintf("Peer %s [%s]", p.id, + fmt.Sprintf("les/%d", p.version), + ) +} + +// peerSet represents the collection of active peers currently participating in +// the Light Ethereum sub-protocol. +type peerSet struct { + peers map[string]*peer + lock sync.RWMutex +} + +// newPeerSet creates a new peer set to track the active participants. +func newPeerSet() *peerSet { + return &peerSet{ + peers: make(map[string]*peer), + } +} + +// Register injects a new peer into the working set, or returns an error if the +// peer is already known. +func (ps *peerSet) Register(p *peer) error { + ps.lock.Lock() + defer ps.lock.Unlock() + + if _, ok := ps.peers[p.id]; ok { + return errAlreadyRegistered + } + ps.peers[p.id] = p + return nil +} + +// Unregister removes a remote peer from the active set, disabling any further +// actions to/from that particular entity. +func (ps *peerSet) Unregister(id string) error { + ps.lock.Lock() + defer ps.lock.Unlock() + + if _, ok := ps.peers[id]; !ok { + return errNotRegistered + } + delete(ps.peers, id) + return nil +} + +// Peer retrieves the registered peer with the given id. +func (ps *peerSet) Peer(id string) *peer { + ps.lock.RLock() + defer ps.lock.RUnlock() + + return ps.peers[id] +} + +// Len returns if the current number of peers in the set. +func (ps *peerSet) Len() int { + ps.lock.RLock() + defer ps.lock.RUnlock() + + return len(ps.peers) +} + +// PeersWithoutBlock retrieves a list of peers that do not have a given block in +// their set of known hashes. +func (ps *peerSet) PeersWithoutBlock(hash common.Hash) []*peer { + ps.lock.RLock() + defer ps.lock.RUnlock() + + list := make([]*peer, 0, len(ps.peers)) + for _, p := range ps.peers { + if !p.knownBlocks.Has(hash) { + list = append(list, p) + } + } + return list +} + +// BestPeer retrieves the known peer with the currently highest total difficulty. +func (ps *peerSet) BestPeer() *peer { + ps.lock.RLock() + defer ps.lock.RUnlock() + + var ( + bestPeer *peer + bestTd *big.Int + ) + for _, p := range ps.peers { + if td := p.Td(); bestPeer == nil || td.Cmp(bestTd) > 0 { + bestPeer, bestTd = p, td + } + } + return bestPeer +} diff --git a/les/protocol.go b/les/protocol.go new file mode 100644 index 0000000000..41461580c2 --- /dev/null +++ b/les/protocol.go @@ -0,0 +1,197 @@ +// Copyright 2015 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 . + +// Package les implements the Light Ethereum Subprotocol. +package les + +import ( + "fmt" + "io" + "math/big" + + "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/types" + "github.com/ethereum/go-ethereum/rlp" + "github.com/ethereum/go-ethereum/trie" +) + +// Mode represents the mode of operation of the light client. +type Mode int + +const ( + ArchiveMode Mode = iota // Maintain the entire blockchain history + FullMode // Maintain only a recent view of the blockchain + LightMode // Don't maintain any history, rather fetch on demand +) + +// Constants to match up protocol versions and messages +const ( + lpv1 = 1 +) + +// minimumProtocolVersion is the minimum version of the protocol les must run to +// support the desired mode of operation. +var minimumProtocolVersion = map[Mode]uint{ + ArchiveMode: lpv1, + FullMode: lpv1, + LightMode: lpv1, +} + +// Supported versions of the les protocol (first is primary). +var ProtocolVersions = []uint{lpv1} + +// Number of implemented message corresponding to different protocol versions. +var ProtocolLengths = []uint64{19} + +const ( + NetworkId = 1 + ProtocolMaxMsgSize = 10 * 1024 * 1024 // Maximum cap on the size of a protocol message +) + +// les protocol message codes +const ( + // Protocol messages belonging to LPV1 + StatusMsg = 0x00 + NewBlockHashesMsg = 0x01 + GetBlockHeadersMsg = 0x03 + BlockHeadersMsg = 0x04 + GetBlockBodiesMsg = 0x05 + BlockBodiesMsg = 0x06 + GetNodeDataMsg = 0x0d + NodeDataMsg = 0x0e + GetReceiptsMsg = 0x0f + ReceiptsMsg = 0x10 + GetProofsMsg = 0x11 + ProofsMsg = 0x12 +) + +type errCode int + +const ( + ErrMsgTooLarge = iota + ErrDecode + ErrInvalidMsgCode + ErrProtocolVersionMismatch + ErrNetworkIdMismatch + ErrGenesisBlockMismatch + ErrNoStatusMsg + ErrExtraStatusMsg + ErrSuspendedPeer +) + +func (e errCode) String() string { + return errorToString[int(e)] +} + +// XXX change once legacy code is out +var errorToString = map[int]string{ + ErrMsgTooLarge: "Message too long", + ErrDecode: "Invalid message", + ErrInvalidMsgCode: "Invalid message code", + ErrProtocolVersionMismatch: "Protocol version mismatch", + ErrNetworkIdMismatch: "NetworkId mismatch", + ErrGenesisBlockMismatch: "Genesis block mismatch", + ErrNoStatusMsg: "No status message", + ErrExtraStatusMsg: "Extra status message", + ErrSuspendedPeer: "Suspended peer", +} + +type chainManager interface { + GetBlockHashesFromHash(hash common.Hash, amount uint64) (hashes []common.Hash) + GetBlock(hash common.Hash) (block *types.Block) + Status() (td *big.Int, currentBlock common.Hash, genesisBlock common.Hash) +} + +// statusData is the network packet for the status message. +type statusData struct { + ProtocolVersion uint32 + NetworkId uint32 + TD *big.Int + CurrentBlock common.Hash + GenesisBlock common.Hash +} + +// newBlockHashesData is the network packet for the block announcements. +type newBlockHashesData []struct { + Hash common.Hash // Hash of one particular block being announced + Number uint64 // Number of one particular block being announced +} + +// getBlockHashesData is the network packet for the hash based hash retrieval. +type getBlockHashesData struct { + Hash common.Hash + Amount uint64 +} + +// getBlockHeadersData represents a block header query. +type getBlockHeadersData struct { + Origin hashOrNumber // Block from which to retrieve headers + Amount uint64 // Maximum number of headers to retrieve + Skip uint64 // Blocks to skip between consecutive headers + Reverse bool // Query direction (false = rising towards latest, true = falling towards genesis) +} + +// hashOrNumber is a combined field for specifying an origin block. +type hashOrNumber struct { + Hash common.Hash // Block hash from which to retrieve headers (excludes Number) + Number uint64 // Block hash from which to retrieve headers (excludes Hash) +} + +// EncodeRLP is a specialized encoder for hashOrNumber to encode only one of the +// two contained union fields. +func (hn *hashOrNumber) EncodeRLP(w io.Writer) error { + if hn.Hash == (common.Hash{}) { + return rlp.Encode(w, hn.Number) + } + if hn.Number != 0 { + return fmt.Errorf("both origin hash (%x) and number (%d) provided", hn.Hash, hn.Number) + } + return rlp.Encode(w, hn.Hash) +} + +// DecodeRLP is a specialized decoder for hashOrNumber to decode the contents +// into either a block hash or a block number. +func (hn *hashOrNumber) DecodeRLP(s *rlp.Stream) error { + _, size, _ := s.Kind() + origin, err := s.Raw() + if err == nil { + switch { + case size == 32: + err = rlp.DecodeBytes(origin, &hn.Hash) + case size <= 8: + err = rlp.DecodeBytes(origin, &hn.Number) + default: + err = fmt.Errorf("invalid input size %d for origin", size) + } + } + return err +} + +// newBlockData is the network packet for the block propagation message. +type newBlockData struct { + Block *types.Block + TD *big.Int +} + +// blockBodiesData is the network packet for block content distribution. +type blockBodiesData []*types.Body + +// nodeDataData is the network response packet for a node data retrieval. +type nodeDataData []struct { + Value []byte +} + +type proofsData []trie.MerkleProof diff --git a/les/sync.go b/les/sync.go new file mode 100644 index 0000000000..b1b397b13f --- /dev/null +++ b/les/sync.go @@ -0,0 +1,72 @@ +// Copyright 2015 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 . + +package les + +import ( + "time" + + "github.com/ethereum/go-ethereum/eth/downloader" +) + +const ( + forceSyncCycle = 10 * time.Second // Time interval to force syncs, even if few peers are available + minDesiredPeerCount = 5 // Amount of peers desired to start syncing +) + +// syncer is responsible for periodically synchronising with the network, both +// downloading hashes and blocks as well as handling the announcement handler. +func (pm *ProtocolManager) syncer() { + // Start and ensure cleanup of sync mechanisms + //pm.fetcher.Start() + //defer pm.fetcher.Stop() + defer pm.downloader.Terminate() + + // Wait for different events to fire synchronisation operations + forceSync := time.Tick(forceSyncCycle) + for { + select { + case <-pm.newPeerCh: + // Make sure we have peers to select from, then sync + if pm.peers.Len() < minDesiredPeerCount { + break + } + go pm.synchronise(pm.peers.BestPeer()) + + case <-forceSync: + // Force a sync even if not enough peers are present + go pm.synchronise(pm.peers.BestPeer()) + + case <-pm.quitSync: + return + } + } +} + +// synchronise tries to sync up our local block chain with a remote peer. +func (pm *ProtocolManager) synchronise(peer *peer) { + // Short circuit if no peers are available + if peer == nil { + return + } + // Make sure the peer's TD is higher than our own. If not drop. + td := pm.blockchain.GetTd(pm.blockchain.LastBlockHash()) + if peer.Td().Cmp(td) <= 0 { + return + } + // Otherwise try to sync with the downloader + pm.downloader.Synchronise(peer.id, peer.Head(), peer.Td(), downloader.LightSync) +} diff --git a/miner/miner.go b/miner/miner.go index 6d4a84f1a4..772c7f5ff8 100644 --- a/miner/miner.go +++ b/miner/miner.go @@ -58,6 +58,13 @@ func New(eth core.Backend, mux *event.TypeMux, pow pow.PoW) *Miner { return miner } +// dummy miner for light client mode +func NewDummy(eth core.Backend, mux *event.TypeMux, pow pow.PoW) *Miner { + miner := &Miner{eth: eth, mux: mux, pow: pow, worker: newDummyWorker(common.Address{}, eth), canStart: 1} + + return miner +} + // update keeps track of the downloader events. Please be aware that this is a one shot type of update loop. // It's entered once and as soon as `Done` or `Failed` has been broadcasted the events are unregistered and // the loop is exited. This to prevent a major security vuln where external parties can DOS you with blocks diff --git a/miner/worker.go b/miner/worker.go index 2d072ef60b..af1ad5b8c0 100644 --- a/miner/worker.go +++ b/miner/worker.go @@ -27,6 +27,7 @@ import ( "github.com/ethereum/go-ethereum/accounts" "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" @@ -147,6 +148,26 @@ func newWorker(coinbase common.Address, eth core.Backend) *worker { return worker } +// dummy worker for light client mode +func newDummyWorker(coinbase common.Address, eth core.Backend) *worker { + worker := &worker{ + eth: eth, + mux: eth.EventMux(), + chainDb: eth.ChainDb(), + recv: make(chan *Result, resultQueueSize), + gasPrice: new(big.Int), + chain: eth.BlockChain(), + proc: eth.BlockProcessor(), + possibleUncles: make(map[common.Hash]*types.Block), + coinbase: coinbase, + txQueue: make(map[common.Hash]*types.Transaction), + quit: make(chan struct{}), + fullValidation: false, + } + + return worker +} + func (self *worker) setEtherbase(addr common.Address) { self.mu.Lock() defer self.mu.Unlock() @@ -285,7 +306,7 @@ func (self *worker) wait() { go self.mux.Post(core.NewMinedBlockEvent{block}) } else { work.state.Commit() - parent := self.chain.GetBlock(block.ParentHash()) + parent := self.chain.GetBlock(block.ParentHash(), access.NoOdr) if parent == nil { glog.V(logger.Error).Infoln("Invalid block found during mining") continue @@ -326,7 +347,7 @@ func (self *worker) wait() { // check staleness and display confirmation var stale, confirm string - canonBlock := self.chain.GetBlockByNumber(block.NumberU64()) + canonBlock := self.chain.GetBlockByNumber(block.NumberU64(), access.NoOdr) if canonBlock != nil && canonBlock.Hash() != block.Hash() { stale = "stale " } else { @@ -359,7 +380,7 @@ func (self *worker) push(work *Work) { // makeCurrent creates a new environment for the current cycle. func (self *worker) makeCurrent(parent *types.Block, header *types.Header) error { - state, err := state.New(parent.Root(), self.eth.ChainDb()) + state, err := state.New(parent.Root(), self.eth.ChainAccess(), access.NullCtx) if err != nil { return err } @@ -422,7 +443,7 @@ func (self *worker) isBlockLocallyMined(current *Work, deepBlockNum uint64) bool } //Does the block at {deepBlockNum} send earnings to my coinbase? - var block = self.chain.GetBlockByNumber(deepBlockNum) + var block = self.chain.GetBlockByNumber(deepBlockNum, access.NoOdr) return block != nil && block.Coinbase() == self.coinbase } @@ -636,7 +657,7 @@ func (env *Work) commitTransactions(transactions types.Transactions, gasPrice *b } func (env *Work) commitTransaction(tx *types.Transaction, proc *core.BlockProcessor, gp *core.GasPool) error { - snap := env.state.Copy() + snap := env.state.Copy(access.NullCtx) receipt, _, err := proc.ApplyTransaction(gp, env.state, env.header, tx, env.header.GasUsed, true) if err != nil { env.state.Set(snap) diff --git a/rpc/api/admin.go b/rpc/api/admin.go index b359d52a10..0220f78e0f 100644 --- a/rpc/api/admin.go +++ b/rpc/api/admin.go @@ -285,7 +285,7 @@ func (self *adminApi) SleepBlocks(req *shared.Request) (interface{}, error) { timer = time.NewTimer(time.Duration(args.Timeout) * time.Second).C } - height = new(big.Int).Add(self.xeth.CurrentBlock().Number(), big.NewInt(args.N)) + height = new(big.Int).Add(self.xeth.CurrentHeader().Number, big.NewInt(args.N)) height, err = sleepBlocks(self.xeth.UpdateState(), height, timer) if err != nil { return nil, err diff --git a/rpc/api/debug.go b/rpc/api/debug.go index d2cbc7f19d..e28dee70d0 100644 --- a/rpc/api/debug.go +++ b/rpc/api/debug.go @@ -22,6 +22,7 @@ import ( "time" "github.com/ethereum/ethash" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/eth" @@ -119,7 +120,7 @@ func (self *debugApi) DumpBlock(req *shared.Request) (interface{}, error) { return nil, fmt.Errorf("block #%d not found", args.BlockNumber) } - stateDb, err := state.New(block.Root(), self.ethereum.ChainDb()) + stateDb, err := state.New(block.Root(), self.ethereum.ChainAccess(), access.NullCtx) if err != nil { return nil, err } diff --git a/rpc/api/eth.go b/rpc/api/eth.go index b84ae31da2..bc620d78a3 100644 --- a/rpc/api/eth.go +++ b/rpc/api/eth.go @@ -19,12 +19,12 @@ package api import ( "bytes" "encoding/json" - "math/big" - "fmt" + "math/big" "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common/natspec" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/eth" "github.com/ethereum/go-ethereum/rpc/codec" "github.com/ethereum/go-ethereum/rpc/shared" @@ -144,7 +144,12 @@ func (self *ethApi) Hashrate(req *shared.Request) (interface{}, error) { } func (self *ethApi) BlockNumber(req *shared.Request) (interface{}, error) { - num := self.xeth.CurrentBlock().Number() + var num *big.Int + if self.xeth.CurrentHeader() != nil { + num = self.xeth.CurrentHeader().Number + } else { + num = big.NewInt(0) + } return newHexNum(num.Bytes()), nil } @@ -153,8 +158,8 @@ func (self *ethApi) GetBalance(req *shared.Request) (interface{}, error) { if err := self.codec.Decode(req.Params, &args); err != nil { return nil, shared.NewDecodeParamError(err.Error()) } - - return self.xeth.AtStateNum(args.BlockNumber).BalanceAt(args.Address), nil + res := self.xeth.WithCtx(req.GetCtx()).AtStateNum(args.BlockNumber).BalanceAt(args.Address) + return res, nil } func (self *ethApi) ProtocolVersion(req *shared.Request) (interface{}, error) { @@ -170,6 +175,9 @@ func (self *ethApi) IsMining(req *shared.Request) (interface{}, error) { } func (self *ethApi) IsSyncing(req *shared.Request) (interface{}, error) { + if self.ethereum.Downloader() == nil { + return false, nil + } origin, current, height := self.ethereum.Downloader().Progress() if current < height { return map[string]interface{}{ @@ -191,7 +199,7 @@ func (self *ethApi) GetStorage(req *shared.Request) (interface{}, error) { return nil, shared.NewDecodeParamError(err.Error()) } - return self.xeth.AtStateNum(args.BlockNumber).State().SafeGet(args.Address).Storage(), nil + return self.xeth.WithCtx(req.GetCtx()).AtStateNum(args.BlockNumber).State().SafeGet(args.Address).Storage(), nil } func (self *ethApi) GetStorageAt(req *shared.Request) (interface{}, error) { @@ -200,7 +208,7 @@ func (self *ethApi) GetStorageAt(req *shared.Request) (interface{}, error) { return nil, shared.NewDecodeParamError(err.Error()) } - return self.xeth.AtStateNum(args.BlockNumber).StorageAt(args.Address, args.Key), nil + return self.xeth.WithCtx(req.GetCtx()).AtStateNum(args.BlockNumber).StorageAt(args.Address, args.Key), nil } func (self *ethApi) GetTransactionCount(req *shared.Request) (interface{}, error) { @@ -209,7 +217,7 @@ func (self *ethApi) GetTransactionCount(req *shared.Request) (interface{}, error return nil, shared.NewDecodeParamError(err.Error()) } - count := self.xeth.AtStateNum(args.BlockNumber).TxCountAt(args.Address) + count := self.xeth.WithCtx(req.GetCtx()).AtStateNum(args.BlockNumber).TxCountAt(args.Address) return fmt.Sprintf("%#x", count), nil } @@ -218,7 +226,7 @@ func (self *ethApi) GetBlockTransactionCountByHash(req *shared.Request) (interfa if err := self.codec.Decode(req.Params, &args); err != nil { return nil, shared.NewDecodeParamError(err.Error()) } - block := self.xeth.EthBlockByHash(args.Hash) + block := self.xeth.WithCtx(req.GetCtx()).EthBlockByHash(args.Hash) if block == nil { return nil, nil } @@ -231,7 +239,7 @@ func (self *ethApi) GetBlockTransactionCountByNumber(req *shared.Request) (inter return nil, shared.NewDecodeParamError(err.Error()) } - block := self.xeth.EthBlockByNumber(args.BlockNumber) + block := self.xeth.WithCtx(req.GetCtx()).EthBlockByNumber(args.BlockNumber) if block == nil { return nil, nil } @@ -244,7 +252,7 @@ func (self *ethApi) GetUncleCountByBlockHash(req *shared.Request) (interface{}, return nil, shared.NewDecodeParamError(err.Error()) } - block := self.xeth.EthBlockByHash(args.Hash) + block := self.xeth.WithCtx(req.GetCtx()).EthBlockByHash(args.Hash) if block == nil { return nil, nil } @@ -257,7 +265,7 @@ func (self *ethApi) GetUncleCountByBlockNumber(req *shared.Request) (interface{} return nil, shared.NewDecodeParamError(err.Error()) } - block := self.xeth.EthBlockByNumber(args.BlockNumber) + block := self.xeth.WithCtx(req.GetCtx()).EthBlockByNumber(args.BlockNumber) if block == nil { return nil, nil } @@ -269,7 +277,7 @@ func (self *ethApi) GetData(req *shared.Request) (interface{}, error) { if err := self.codec.Decode(req.Params, &args); err != nil { return nil, shared.NewDecodeParamError(err.Error()) } - v := self.xeth.AtStateNum(args.BlockNumber).CodeAtBytes(args.Address) + v := self.xeth.WithCtx(req.GetCtx()).AtStateNum(args.BlockNumber).CodeAtBytes(args.Address) return newHexData(v), nil } @@ -337,7 +345,7 @@ func (self *ethApi) GetNatSpec(req *shared.Request) (interface{}, error) { } func (self *ethApi) EstimateGas(req *shared.Request) (interface{}, error) { - _, gas, err := self.doCall(req.Params) + _, gas, err := self.doCall(req.Params, req.GetCtx()) if err != nil { return nil, err } @@ -351,7 +359,7 @@ func (self *ethApi) EstimateGas(req *shared.Request) (interface{}, error) { } func (self *ethApi) Call(req *shared.Request) (interface{}, error) { - v, _, err := self.doCall(req.Params) + v, _, err := self.doCall(req.Params, req.GetCtx()) if err != nil { return nil, err } @@ -368,13 +376,13 @@ func (self *ethApi) Flush(req *shared.Request) (interface{}, error) { return nil, shared.NewNotImplementedError(req.Method) } -func (self *ethApi) doCall(params json.RawMessage) (string, string, error) { +func (self *ethApi) doCall(params json.RawMessage, ctx *access.OdrContext) (string, string, error) { args := new(CallArgs) if err := self.codec.Decode(params, &args); err != nil { return "", "", err } - return self.xeth.AtStateNum(args.BlockNumber).Call(args.From, args.To, args.Value.String(), args.Gas.String(), args.GasPrice.String(), args.Data) + return self.xeth.WithCtx(ctx).AtStateNum(args.BlockNumber).Call(args.From, args.To, args.Value.String(), args.Gas.String(), args.GasPrice.String(), args.Data) } func (self *ethApi) GetBlockByHash(req *shared.Request) (interface{}, error) { @@ -382,7 +390,7 @@ func (self *ethApi) GetBlockByHash(req *shared.Request) (interface{}, error) { if err := self.codec.Decode(req.Params, &args); err != nil { return nil, shared.NewDecodeParamError(err.Error()) } - block := self.xeth.EthBlockByHash(args.BlockHash) + block := self.xeth.WithCtx(req.GetCtx()).EthBlockByHash(args.BlockHash) if block == nil { return nil, nil } @@ -395,7 +403,7 @@ func (self *ethApi) GetBlockByNumber(req *shared.Request) (interface{}, error) { return nil, shared.NewDecodeParamError(err.Error()) } - block := self.xeth.EthBlockByNumber(args.BlockNumber) + block := self.xeth.WithCtx(req.GetCtx()).EthBlockByNumber(args.BlockNumber) if block == nil { return nil, nil } @@ -408,7 +416,7 @@ func (self *ethApi) GetTransactionByHash(req *shared.Request) (interface{}, erro return nil, shared.NewDecodeParamError(err.Error()) } - tx, bhash, bnum, txi := self.xeth.EthTransactionByHash(args.Hash) + tx, bhash, bnum, txi := self.xeth.EthTransactionByHash(args.Hash) //ODR if tx != nil { v := NewTransactionRes(tx) // if the blockhash is 0, assume this is a pending transaction @@ -428,7 +436,7 @@ func (self *ethApi) GetTransactionByBlockHashAndIndex(req *shared.Request) (inte return nil, shared.NewDecodeParamError(err.Error()) } - raw := self.xeth.EthBlockByHash(args.Hash) + raw := self.xeth.WithCtx(req.GetCtx()).EthBlockByHash(args.Hash) if raw == nil { return nil, nil } @@ -446,7 +454,7 @@ func (self *ethApi) GetTransactionByBlockNumberAndIndex(req *shared.Request) (in return nil, shared.NewDecodeParamError(err.Error()) } - raw := self.xeth.EthBlockByNumber(args.BlockNumber) + raw := self.xeth.WithCtx(req.GetCtx()).EthBlockByNumber(args.BlockNumber) if raw == nil { return nil, nil } @@ -464,7 +472,7 @@ func (self *ethApi) GetUncleByBlockHashAndIndex(req *shared.Request) (interface{ return nil, shared.NewDecodeParamError(err.Error()) } - raw := self.xeth.EthBlockByHash(args.Hash) + raw := self.xeth.WithCtx(req.GetCtx()).EthBlockByHash(args.Hash) if raw == nil { return nil, nil } @@ -482,7 +490,7 @@ func (self *ethApi) GetUncleByBlockNumberAndIndex(req *shared.Request) (interfac return nil, shared.NewDecodeParamError(err.Error()) } - raw := self.xeth.EthBlockByNumber(args.BlockNumber) + raw := self.xeth.WithCtx(req.GetCtx()).EthBlockByNumber(args.BlockNumber) if raw == nil { return nil, nil } @@ -662,7 +670,7 @@ func (self *ethApi) GetTransactionReceipt(req *shared.Request) (interface{}, err txhash := common.BytesToHash(common.FromHex(args.Hash)) tx, bhash, bnum, txi := self.xeth.EthTransactionByHash(args.Hash) - rec := self.xeth.GetTxReceipt(txhash) + rec := self.xeth.GetTxReceipt(txhash) //ODR // We could have an error of "not found". Should disambiguate // if err != nil { // return err, nil diff --git a/rpc/comms/comms.go b/rpc/comms/comms.go index 61fba5722c..794810efa2 100644 --- a/rpc/comms/comms.go +++ b/rpc/comms/comms.go @@ -19,12 +19,12 @@ package comms import ( "io" "net" - "fmt" "strings" - "strconv" + "time" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" "github.com/ethereum/go-ethereum/rpc/codec" @@ -69,6 +69,8 @@ func handle(id int, conn net.Conn, api shared.EthereumApi, c codec.Codec) { codec.Close() }() + channelID := access.NewChannelID(time.Second) + for { requests, isBatch, err := codec.ReadRequest() if err == io.EOF { @@ -78,11 +80,22 @@ func handle(id int, conn net.Conn, api shared.EthereumApi, c codec.Codec) { return } + ctx := access.NewContext(channelID) + if isBatch { responses := make([]*interface{}, len(requests)) responseCount := 0 + var res interface{} + var err error for _, req := range requests { - res, err := api.Execute(req) + if !ctx.IsCancelled() { + req.SetCtx(ctx) + res, err = api.Execute(req) + } + if ctx.IsCancelled() { + res = nil + err = access.ErrCancel + } if req.Id != nil { rpcResponse := shared.NewRpcResponse(req.Id, req.Jsonrpc, res, err) responses[responseCount] = rpcResponse @@ -97,7 +110,12 @@ func handle(id int, conn net.Conn, api shared.EthereumApi, c codec.Codec) { } } else { var rpcResponse interface{} + requests[0].SetCtx(ctx) res, err := api.Execute(requests[0]) + if ctx.IsCancelled() { + res = nil + err = access.ErrCancel + } rpcResponse = shared.NewRpcResponse(requests[0].Id, requests[0].Jsonrpc, res, err) err = codec.WriteResponse(rpcResponse) diff --git a/rpc/comms/http.go b/rpc/comms/http.go index f4a930d0ef..85c8edaafd 100644 --- a/rpc/comms/http.go +++ b/rpc/comms/http.go @@ -29,6 +29,7 @@ import ( "io" "io/ioutil" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" "github.com/ethereum/go-ethereum/rpc/codec" @@ -67,13 +68,14 @@ type stopServer struct { type handler struct { codec codec.Codec api shared.EthereumApi + channelID *access.OdrChannelID } // StartHTTP starts listening for RPC requests sent via HTTP. func StartHttp(cfg HttpConfig, codec codec.Codec, api shared.EthereumApi) error { httpServerMu.Lock() defer httpServerMu.Unlock() - + addr := fmt.Sprintf("%s:%d", cfg.ListenAddress, cfg.ListenPort) if httpServer != nil { if addr != httpServer.Addr { @@ -82,7 +84,8 @@ func StartHttp(cfg HttpConfig, codec codec.Codec, api shared.EthereumApi) error return nil // RPC service already running on given host/port } // Set up the request handler, wrapping it with CORS headers if configured. - handler := http.Handler(&handler{codec, api}) + channelID := access.NewChannelID(time.Second) + handler := http.Handler(&handler{codec, api, channelID}) if len(cfg.CorsDomain) > 0 { opts := cors.Options{ AllowedMethods: []string{"POST"}, @@ -121,9 +124,15 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, req *http.Request) { } c := h.codec.New(nil) + ctx := access.NewContext(h.channelID) var rpcReq shared.Request if err = c.Decode(payload, &rpcReq); err == nil { + rpcReq.SetCtx(ctx) reply, err := h.api.Execute(&rpcReq) + if ctx.IsCancelled() { + reply = nil + err = access.ErrCancel + } res := shared.NewRpcResponse(rpcReq.Id, rpcReq.Jsonrpc, reply, err) sendJSON(w, &res) return @@ -133,8 +142,17 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, req *http.Request) { if err = c.Decode(payload, &reqBatch); err == nil { resBatch := make([]*interface{}, len(reqBatch)) resCount := 0 + var reply interface{} + var err error for i, rpcReq := range reqBatch { - reply, err := h.api.Execute(&rpcReq) + if !ctx.IsCancelled() { + rpcReq.SetCtx(ctx) + reply, err = h.api.Execute(&rpcReq) + } + if ctx.IsCancelled() { + reply = nil + err = access.ErrCancel + } if rpcReq.Id != nil { // this leaves nil entries in the response batch for later removal resBatch[i] = shared.NewRpcResponse(rpcReq.Id, rpcReq.Jsonrpc, reply, err) resCount += 1 diff --git a/rpc/comms/inproc.go b/rpc/comms/inproc.go index e8058e32bf..4e4e9e9b84 100644 --- a/rpc/comms/inproc.go +++ b/rpc/comms/inproc.go @@ -18,7 +18,9 @@ package comms import ( "fmt" + "time" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/rpc/codec" "github.com/ethereum/go-ethereum/rpc/shared" ) @@ -30,12 +32,14 @@ type InProcClient struct { lastJsonrpc string lastErr error lastRes interface{} + channelID *access.OdrChannelID } // Create a new in process client func NewInProcClient(codec codec.Codec) *InProcClient { return &InProcClient{ codec: codec, + channelID: access.NewChannelID(time.Second), } } @@ -52,7 +56,13 @@ func (self *InProcClient) Send(req interface{}) error { if r, ok := req.(*shared.Request); ok { self.lastId = r.Id self.lastJsonrpc = r.Jsonrpc + ctx := access.NewContext(self.channelID) + r.SetCtx(ctx) self.lastRes, self.lastErr = self.api.Execute(r) + if ctx.IsCancelled() { + self.lastRes = nil + self.lastErr = access.ErrCancel + } return self.lastErr } diff --git a/rpc/shared/types.go b/rpc/shared/types.go index db328234d2..7dd6fb2d9a 100644 --- a/rpc/shared/types.go +++ b/rpc/shared/types.go @@ -19,6 +19,7 @@ package shared import ( "encoding/json" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/logger" "github.com/ethereum/go-ethereum/logger/glog" ) @@ -44,6 +45,7 @@ type Request struct { Jsonrpc string `json:"jsonrpc"` Method string `json:"method"` Params json.RawMessage `json:"params"` + ctx *access.OdrContext } // RPC response @@ -73,6 +75,14 @@ type ErrorObject struct { // Data interface{} `json:"data"` } +func (req *Request) GetCtx() *access.OdrContext { + return req.ctx +} + +func (req *Request) SetCtx(ctx *access.OdrContext) { + req.ctx = ctx +} + // Create RPC error response, this allows for custom error codes func NewRpcErrorResponse(id interface{}, jsonrpcver string, errCode int, err error) *ErrorResponse { jsonerr := &ErrorObject{errCode, err.Error()} diff --git a/tests/block_test_util.go b/tests/block_test_util.go index 6a2eb96a42..70ac8024d8 100644 --- a/tests/block_test_util.go +++ b/tests/block_test_util.go @@ -31,6 +31,7 @@ import ( "github.com/ethereum/go-ethereum/accounts" "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" @@ -203,7 +204,7 @@ func runBlockTest(test *BlockTest) error { return fmt.Errorf("lastblockhash validation mismatch: want: %x, have: %x", lastblockhash, cmlast) } - newDB, err := cm.State() + newDB, err := cm.State(access.NullCtx) if err != nil { return err } @@ -217,7 +218,7 @@ func runBlockTest(test *BlockTest) error { // InsertPreState populates the given database with the genesis // accounts defined by the test. func (t *BlockTest) InsertPreState(db ethdb.Database, am *accounts.Manager) (*state.StateDB, error) { - statedb, err := state.New(common.Hash{}, db) + statedb, err := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) if err != nil { return nil, err } @@ -441,7 +442,7 @@ func (test *BlockTest) ValidateImportedHeaders(cm *core.BlockChain, validBlocks // block-by-block, so we can only validate imported headers after // all blocks have been processed by ChainManager, as they may not // be part of the longest chain until last block is imported. - for b := cm.CurrentBlock(); b != nil && b.NumberU64() != 0; b = cm.GetBlock(b.Header().ParentHash) { + for b := cm.CurrentBlock(); b != nil && b.NumberU64() != 0; b = cm.GetBlock(b.Header().ParentHash, access.NoOdr) { bHash := common.Bytes2Hex(b.Hash().Bytes()) // hex without 0x prefix if err := validateHeader(bmap[bHash].BlockHeader, b.Header()); err != nil { return fmt.Errorf("Imported block header validation failed: %v", err) diff --git a/tests/state_test_util.go b/tests/state_test_util.go index 352fe3570b..2499872ac7 100644 --- a/tests/state_test_util.go +++ b/tests/state_test_util.go @@ -27,6 +27,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/crypto" @@ -103,7 +104,7 @@ func BenchStateTest(p string, conf bconf, b *testing.B) error { func benchStateTest(test VmTest, env map[string]string, b *testing.B) { b.StopTimer() db, _ := ethdb.NewMemDatabase() - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) for addr, account := range test.Pre { obj := StateObjectFromAccount(db, addr, account) statedb.SetStateObject(obj) @@ -142,7 +143,7 @@ func runStateTests(tests map[string]VmTest, skipTests []string) error { func runStateTest(test VmTest) error { db, _ := ethdb.NewMemDatabase() - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) for addr, account := range test.Pre { obj := StateObjectFromAccount(db, addr, account) statedb.SetStateObject(obj) @@ -192,7 +193,7 @@ func runStateTest(test VmTest) error { } for addr, value := range account.Storage { - v := obj.GetState(common.HexToHash(addr)) + v := obj.GetState(common.HexToHash(addr), access.NullCtx) vexp := common.HexToHash(value) if v != vexp { @@ -233,7 +234,7 @@ func RunState(statedb *state.StateDB, env, tx map[string]string) ([]byte, vm.Log // Set pre compiled contracts vm.Precompiled = vm.PrecompiledContracts() - snapshot := statedb.Copy() + snapshot := statedb.Copy(access.NullCtx) gaspool := new(core.GasPool).AddGas(common.Big(env["currentGasLimit"])) key, _ := hex.DecodeString(tx["secretKey"]) diff --git a/tests/util.go b/tests/util.go index 571161683f..3ba6a36112 100644 --- a/tests/util.go +++ b/tests/util.go @@ -23,6 +23,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" @@ -89,7 +90,7 @@ func (self Log) Topics() [][]byte { } func StateObjectFromAccount(db ethdb.Database, addr string, account Account) *state.StateObject { - obj := state.NewStateObject(common.HexToAddress(addr), db) + obj := state.NewStateObject(common.HexToAddress(addr), access.NewDbChainAccess(db)) obj.SetBalance(common.Big(account.Balance)) if common.IsHex(account.Code) { @@ -201,7 +202,7 @@ func (self *Env) CanTransfer(from common.Address, balance *big.Int) bool { return self.state.GetBalance(from).Cmp(balance) >= 0 } func (self *Env) MakeSnapshot() vm.Database { - return self.state.Copy() + return self.state.Copy(access.NullCtx) } func (self *Env) SetSnapshot(copy vm.Database) { self.state.Set(copy.(*state.StateDB)) diff --git a/tests/vm_test_util.go b/tests/vm_test_util.go index ddd14b1a3e..858acfde8e 100644 --- a/tests/vm_test_util.go +++ b/tests/vm_test_util.go @@ -25,6 +25,7 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/ethdb" @@ -108,7 +109,7 @@ func BenchVmTest(p string, conf bconf, b *testing.B) error { func benchVmTest(test VmTest, env map[string]string, b *testing.B) { b.StopTimer() db, _ := ethdb.NewMemDatabase() - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) for addr, account := range test.Pre { obj := StateObjectFromAccount(db, addr, account) statedb.SetStateObject(obj) @@ -159,7 +160,7 @@ func runVmTests(tests map[string]VmTest, skipTests []string) error { func runVmTest(test VmTest) error { db, _ := ethdb.NewMemDatabase() - statedb, _ := state.New(common.Hash{}, db) + statedb, _ := state.New(common.Hash{}, access.NewDbChainAccess(db), access.NullCtx) for addr, account := range test.Pre { obj := StateObjectFromAccount(db, addr, account) statedb.SetStateObject(obj) @@ -214,7 +215,7 @@ func runVmTest(test VmTest) error { } for addr, value := range account.Storage { - v := obj.GetState(common.HexToHash(addr)) + v := obj.GetState(common.HexToHash(addr), access.NoOdr) vexp := common.HexToHash(value) if v != vexp { diff --git a/trie/encoding.go b/trie/encoding.go index 3c172b8438..0067af6a20 100644 --- a/trie/encoding.go +++ b/trie/encoding.go @@ -69,6 +69,23 @@ func compactHexDecode(str []byte) []byte { return nibbles } +func compactHexEncode(nibbles []byte) []byte { + nl := len(nibbles) + if nibbles[nl-1] == 16 { + nl-- + } + l := (nl + 1) / 2 + var str = make([]byte, l) + for i, _ := range str { + b := nibbles[i*2] * 16 + if nl > i*2 { + b += nibbles[i*2+1] + } + str[i] = b + } + return str +} + func decodeCompact(key []byte) []byte { l := len(key) / 2 var res = make([]byte, l) diff --git a/trie/iterator.go b/trie/iterator.go index 38555fe08d..324cc655ce 100644 --- a/trie/iterator.go +++ b/trie/iterator.go @@ -16,7 +16,11 @@ package trie -import "bytes" +import ( + "bytes" + + "github.com/ethereum/go-ethereum/core/access" +) type Iterator struct { trie *Trie @@ -100,7 +104,7 @@ func (self *Iterator) next(node interface{}, key []byte, isIterStart bool) []byt } case hashNode: - return self.next(self.trie.resolveHash(node), key, isIterStart) + return self.next(self.trie.resolveHash(node, nil, nil, access.NoOdr), key, isIterStart) } return nil } @@ -127,7 +131,7 @@ func (self *Iterator) key(node interface{}) []byte { } } case hashNode: - return self.key(self.trie.resolveHash(node)) + return self.key(self.trie.resolveHash(node, nil, nil, access.NoOdr)) } return nil diff --git a/trie/iterator_test.go b/trie/iterator_test.go index fdc60b4122..ab057eb37b 100644 --- a/trie/iterator_test.go +++ b/trie/iterator_test.go @@ -16,7 +16,11 @@ package trie -import "testing" +import ( + "testing" + + "github.com/ethereum/go-ethereum/core/access" +) func TestIterator(t *testing.T) { trie := newEmpty() @@ -32,7 +36,7 @@ func TestIterator(t *testing.T) { v := make(map[string]bool) for _, val := range vals { v[val.k] = false - trie.Update([]byte(val.k), []byte(val.v)) + trie.Update([]byte(val.k), []byte(val.v), access.NoOdr) } trie.Commit() diff --git a/trie/proof.go b/trie/proof.go index a705c49db0..bfadceb5ab 100644 --- a/trie/proof.go +++ b/trie/proof.go @@ -6,10 +6,13 @@ import ( "fmt" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/crypto/sha3" "github.com/ethereum/go-ethereum/rlp" ) +type MerkleProof []rlp.RawValue + // Prove constructs a merkle proof for key. The result contains all // encoded nodes on the path to the value at key. The value itself is // also included in the last node and can be retrieved by verifying @@ -17,29 +20,28 @@ import ( // // The returned proof is nil if the trie does not contain a value for key. // For existing keys, the proof will have at least one element. -func (t *Trie) Prove(key []byte) []rlp.RawValue { +func (t *Trie) Prove(key []byte) MerkleProof { // Collect all nodes on the path to key. key = compactHexDecode(key) nodes := []node{} tn := t.root - for len(key) > 0 { + for len(key) > 0 && tn != nil { switch n := tn.(type) { case shortNode: if len(key) < len(n.Key) || !bytes.Equal(n.Key, key[:len(n.Key)]) { // The trie doesn't contain the key. - return nil + tn = nil + } else { + tn = n.Val } - tn = n.Val key = key[len(n.Key):] nodes = append(nodes, n) case fullNode: tn = n[key[0]] key = key[1:] nodes = append(nodes, n) - case nil: - return nil case hashNode: - tn = t.resolveHash(n) + tn = t.resolveHash(n, nil, nil, access.NoOdr) default: panic(fmt.Sprintf("%T: invalid node: %v", tn, tn)) } @@ -67,7 +69,7 @@ func (t *Trie) Prove(key []byte) []rlp.RawValue { // value for key in a trie with the given root hash. VerifyProof // returns an error if the proof contains invalid trie nodes or the // wrong value. -func VerifyProof(rootHash common.Hash, key []byte, proof []rlp.RawValue) (value []byte, err error) { +func VerifyProof(rootHash common.Hash, key []byte, proof MerkleProof) (value []byte, err error) { key = compactHexDecode(key) sha := sha3.NewKeccak256() wantHash := rootHash.Bytes() @@ -84,7 +86,12 @@ func VerifyProof(rootHash common.Hash, key []byte, proof []rlp.RawValue) (value keyrest, cld := get(n, key) switch cld := cld.(type) { case nil: - return nil, fmt.Errorf("key mismatch at proof node %d", i) + if i != len(proof)-1 { + return nil, fmt.Errorf("key mismatch at proof node %d", i) + } else { + // The trie doesn't contain the key. + return nil, nil + } case hashNode: key = keyrest wantHash = cld @@ -98,6 +105,19 @@ func VerifyProof(rootHash common.Hash, key []byte, proof []rlp.RawValue) (value return nil, errors.New("unexpected end of proof") } +func StoreProof(db Database, proof MerkleProof) { + sha := sha3.NewKeccak256() + for _, buf := range proof { + sha.Reset() + sha.Write(buf) + hash := sha.Sum(nil) + val, _ := db.Get(hash) + if val == nil { + db.Put(hash, buf) + } + } +} + func get(tn node, key []byte) ([]byte, node) { for len(key) > 0 { switch n := tn.(type) { diff --git a/trie/proof_test.go b/trie/proof_test.go index 6b5bef05c4..909936a7a0 100644 --- a/trie/proof_test.go +++ b/trie/proof_test.go @@ -8,6 +8,7 @@ import ( "time" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/rlp" ) @@ -119,14 +120,14 @@ func randomTrie(n int) (*Trie, map[string]*kv) { for i := byte(0); i < 100; i++ { value := &kv{common.LeftPadBytes([]byte{i}, 32), []byte{i}, false} value2 := &kv{common.LeftPadBytes([]byte{i + 10}, 32), []byte{i}, false} - trie.Update(value.k, value.v) - trie.Update(value2.k, value2.v) + trie.Update(value.k, value.v, access.NoOdr) + trie.Update(value2.k, value2.v, access.NoOdr) vals[string(value.k)] = value vals[string(value2.k)] = value2 } for i := 0; i < n; i++ { value := &kv{randBytes(32), randBytes(20), false} - trie.Update(value.k, value.v) + trie.Update(value.k, value.v, access.NoOdr) vals[string(value.k)] = value } return trie, vals diff --git a/trie/secure_trie.go b/trie/secure_trie.go index 47d1934d05..0623fab6bb 100644 --- a/trie/secure_trie.go +++ b/trie/secure_trie.go @@ -20,6 +20,7 @@ import ( "hash" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/crypto/sha3" ) @@ -49,11 +50,11 @@ type SecureTrie struct { // trie is initially empty. Otherwise, New will panics if db is nil // and returns ErrMissingRoot if the root node cannpt be found. // Accessing the trie loads nodes from db on demand. -func NewSecure(root common.Hash, db Database) (*SecureTrie, error) { +func NewSecure(root common.Hash, db Database, access OdrAccess) (*SecureTrie, error) { if db == nil { panic("NewSecure called with nil database") } - trie, err := New(root, db) + trie, err := New(root, db, access) if err != nil { return nil, err } @@ -62,8 +63,8 @@ func NewSecure(root common.Hash, db Database) (*SecureTrie, error) { // Get returns the value for key stored in the trie. // The value bytes must not be modified by the caller. -func (t *SecureTrie) Get(key []byte) []byte { - return t.Trie.Get(t.hashKey(key)) +func (t *SecureTrie) Get(key []byte, ctx *access.OdrContext) []byte { + return t.Trie.Get(t.hashKey(key), ctx) } // Update associates key with value in the trie. Subsequent calls to @@ -72,15 +73,15 @@ func (t *SecureTrie) Get(key []byte) []byte { // // The value bytes must not be modified by the caller while they are // stored in the trie. -func (t *SecureTrie) Update(key, value []byte) { +func (t *SecureTrie) Update(key, value []byte, ctx *access.OdrContext) { hk := t.hashKey(key) - t.Trie.Update(hk, value) + t.Trie.Update(hk, value, ctx) t.Trie.db.Put(t.secKey(hk), key) } // Delete removes any existing value for key from the trie. -func (t *SecureTrie) Delete(key []byte) { - t.Trie.Delete(t.hashKey(key)) +func (t *SecureTrie) Delete(key []byte, ctx *access.OdrContext) { + t.Trie.Delete(t.hashKey(key), ctx) } // GetKey returns the sha3 preimage of a hashed key that was diff --git a/trie/secure_trie_test.go b/trie/secure_trie_test.go index 13c6cd02e9..f22566fcfa 100644 --- a/trie/secure_trie_test.go +++ b/trie/secure_trie_test.go @@ -21,13 +21,14 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/ethdb" ) func newEmptySecure() *SecureTrie { db, _ := ethdb.NewMemDatabase() - trie, _ := NewSecure(common.Hash{}, db) + trie, _ := NewSecure(common.Hash{}, db, nil) return trie } @@ -45,9 +46,9 @@ func TestSecureDelete(t *testing.T) { } for _, val := range vals { if val.v != "" { - trie.Update([]byte(val.k), []byte(val.v)) + trie.Update([]byte(val.k), []byte(val.v), access.NoOdr) } else { - trie.Delete([]byte(val.k)) + trie.Delete([]byte(val.k), access.NoOdr) } } hash := trie.Hash() @@ -59,13 +60,13 @@ func TestSecureDelete(t *testing.T) { func TestSecureGetKey(t *testing.T) { trie := newEmptySecure() - trie.Update([]byte("foo"), []byte("bar")) + trie.Update([]byte("foo"), []byte("bar"), access.NoOdr) key := []byte("foo") value := []byte("bar") seckey := crypto.Sha3(key) - if !bytes.Equal(trie.Get(key), value) { + if !bytes.Equal(trie.Get(key, access.NoOdr), value) { t.Errorf("Get did not return bar") } if k := trie.GetKey(seckey); !bytes.Equal(k, key) { diff --git a/trie/sync_test.go b/trie/sync_test.go index 9c036a3a92..fcabdf39d8 100644 --- a/trie/sync_test.go +++ b/trie/sync_test.go @@ -21,6 +21,7 @@ import ( "testing" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/ethdb" ) @@ -28,18 +29,18 @@ import ( func makeTestTrie() (ethdb.Database, *Trie, map[string][]byte) { // Create an empty trie db, _ := ethdb.NewMemDatabase() - trie, _ := New(common.Hash{}, db) + trie, _ := New(common.Hash{}, db, nil) // Fill it with some arbitrary data content := make(map[string][]byte) for i := byte(0); i < 255; i++ { key, val := common.LeftPadBytes([]byte{1, i}, 32), []byte{i} content[string(key)] = val - trie.Update(key, val) + trie.Update(key, val, access.NoOdr) key, val = common.LeftPadBytes([]byte{2, i}, 32), []byte{i} content[string(key)] = val - trie.Update(key, val) + trie.Update(key, val, access.NoOdr) } trie.Commit() @@ -50,12 +51,12 @@ func makeTestTrie() (ethdb.Database, *Trie, map[string][]byte) { // checkTrieContents cross references a reconstructed trie with an expected data // content map. func checkTrieContents(t *testing.T, db Database, root []byte, content map[string][]byte) { - trie, err := New(common.BytesToHash(root), db) + trie, err := New(common.BytesToHash(root), db, nil) if err != nil { t.Fatalf("failed to create trie at %x: %v", root, err) } for key, val := range content { - if have := trie.Get([]byte(key)); bytes.Compare(have, val) != 0 { + if have := trie.Get([]byte(key), access.NoOdr); bytes.Compare(have, val) != 0 { t.Errorf("entry %x: content mismatch: have %x, want %x", key, have, val) } } @@ -63,8 +64,8 @@ func checkTrieContents(t *testing.T, db Database, root []byte, content map[strin // Tests that an empty trie is not scheduled for syncing. func TestEmptyTrieSync(t *testing.T) { - emptyA, _ := New(common.Hash{}, nil) - emptyB, _ := New(emptyRoot, nil) + emptyA, _ := New(common.Hash{}, nil, nil) + emptyB, _ := New(emptyRoot, nil, nil) for i, trie := range []*Trie{emptyA, emptyB} { db, _ := ethdb.NewMemDatabase() diff --git a/trie/trie.go b/trie/trie.go index a3a383fb58..d1a5d6b854 100644 --- a/trie/trie.go +++ b/trie/trie.go @@ -24,6 +24,7 @@ import ( "hash" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/crypto" "github.com/ethereum/go-ethereum/crypto/sha3" "github.com/ethereum/go-ethereum/logger" @@ -45,6 +46,13 @@ var ( ) var ErrMissingRoot = errors.New("missing root node") +var ErrTrieOdrFailure = errors.New("can't commit to db because of ODR failure") + +// OdrAccess must be implemented by on-demand network access layer +type OdrAccess interface { + OdrEnabled() bool + RetrieveKey(key []byte, ctx *access.OdrContext) bool +} // Database must be implemented by backing stores for the trie. type Database interface { @@ -67,8 +75,9 @@ type DatabaseWriter interface { // // Trie is not safe for concurrent use. type Trie struct { - root node - db Database + root node + db Database + access OdrAccess *hasher } @@ -78,14 +87,16 @@ type Trie struct { // trie is initially empty and does not require a database. Otherwise, // New will panics if db is nil or root does not exist in the // database. Accessing the trie loads nodes from db on demand. -func New(root common.Hash, db Database) (*Trie, error) { - trie := &Trie{db: db} +func New(root common.Hash, db Database, access OdrAccess) (*Trie, error) { + trie := &Trie{db: db, access: access} if (root != common.Hash{}) && root != emptyRoot { if db == nil { panic("trie.New: cannot use existing root without a database") } - if v, _ := trie.db.Get(root[:]); len(v) == 0 { - return nil, ErrMissingRoot + if access == nil || !access.OdrEnabled() { + if v, _ := trie.db.Get(root[:]); len(v) == 0 { + return nil, ErrMissingRoot + } } trie.root = hashNode(root.Bytes()) } @@ -99,24 +110,25 @@ func (t *Trie) Iterator() *Iterator { // Get returns the value for key stored in the trie. // The value bytes must not be modified by the caller. -func (t *Trie) Get(key []byte) []byte { +func (t *Trie) Get(key []byte, ctx *access.OdrContext) []byte { key = compactHexDecode(key) + pos := 0 tn := t.root - for len(key) > 0 { + for pos < len(key) { switch n := tn.(type) { case shortNode: - if len(key) < len(n.Key) || !bytes.Equal(n.Key, key[:len(n.Key)]) { + if len(key)-pos < len(n.Key) || !bytes.Equal(n.Key, key[pos:pos+len(n.Key)]) { return nil } tn = n.Val - key = key[len(n.Key):] + pos += len(n.Key) case fullNode: - tn = n[key[0]] - key = key[1:] + tn = n[key[pos]] + pos++ case nil: return nil case hashNode: - tn = t.resolveHash(n) + tn = t.resolveHash(n, key[:pos], key[pos:], ctx) default: panic(fmt.Sprintf("%T: invalid node: %v", tn, tn)) } @@ -130,16 +142,16 @@ func (t *Trie) Get(key []byte) []byte { // // The value bytes must not be modified by the caller while they are // stored in the trie. -func (t *Trie) Update(key, value []byte) { +func (t *Trie) Update(key, value []byte, ctx *access.OdrContext) { k := compactHexDecode(key) if len(value) != 0 { - t.root = t.insert(t.root, k, valueNode(value)) + t.root = t.insert(t.root, nil, k, valueNode(value), ctx) } else { - t.root = t.delete(t.root, k) + t.root = t.delete(t.root, nil, k, ctx) } } -func (t *Trie) insert(n node, key []byte, value node) node { +func (t *Trie) insert(n node, prefix, key []byte, value node, ctx *access.OdrContext) node { if len(key) == 0 { return value } @@ -149,12 +161,12 @@ func (t *Trie) insert(n node, key []byte, value node) node { // If the whole key matches, keep this short node as is // and only update the value. if matchlen == len(n.Key) { - return shortNode{n.Key, t.insert(n.Val, key[matchlen:], value)} + return shortNode{n.Key, t.insert(n.Val, append(prefix, key[:matchlen]...), key[matchlen:], value, ctx)} } // Otherwise branch out at the index where they differ. var branch fullNode - branch[n.Key[matchlen]] = t.insert(nil, n.Key[matchlen+1:], n.Val) - branch[key[matchlen]] = t.insert(nil, key[matchlen+1:], value) + branch[n.Key[matchlen]] = t.insert(nil, append(prefix, n.Key[:matchlen+1]...), n.Key[matchlen+1:], n.Val, ctx) + branch[key[matchlen]] = t.insert(nil, append(prefix, key[:matchlen+1]...), key[matchlen+1:], value, ctx) // Replace this shortNode with the branch if it occurs at index 0. if matchlen == 0 { return branch @@ -163,7 +175,7 @@ func (t *Trie) insert(n node, key []byte, value node) node { return shortNode{key[:matchlen], branch} case fullNode: - n[key[0]] = t.insert(n[key[0]], key[1:], value) + n[key[0]] = t.insert(n[key[0]], append(prefix, key[0]), key[1:], value, ctx) return n case nil: @@ -176,7 +188,7 @@ func (t *Trie) insert(n node, key []byte, value node) node { // // TODO: track whether insertion changed the value and keep // n as a hash node if it didn't. - return t.insert(t.resolveHash(n), key, value) + return t.insert(t.resolveHash(n, prefix, key, ctx), prefix, key, value, ctx) default: panic(fmt.Sprintf("%T: invalid node: %v", n, n)) @@ -184,15 +196,15 @@ func (t *Trie) insert(n node, key []byte, value node) node { } // Delete removes any existing value for key from the trie. -func (t *Trie) Delete(key []byte) { +func (t *Trie) Delete(key []byte, ctx *access.OdrContext) { k := compactHexDecode(key) - t.root = t.delete(t.root, k) + t.root = t.delete(t.root, nil, k, ctx) } // delete returns the new root of the trie with key deleted. // It reduces the trie to minimal form by simplifying // nodes on the way up after deleting recursively. -func (t *Trie) delete(n node, key []byte) node { +func (t *Trie) delete(n node, prefix, key []byte, ctx *access.OdrContext) node { switch n := n.(type) { case shortNode: matchlen := prefixLen(key, n.Key) @@ -206,7 +218,7 @@ func (t *Trie) delete(n node, key []byte) node { // from the subtrie. Child can never be nil here since the // subtrie must contain at least two other values with keys // longer than n.Key. - child := t.delete(n.Val, key[len(n.Key):]) + child := t.delete(n.Val, append(prefix, key[:len(n.Key)]...), key[len(n.Key):], ctx) switch child := child.(type) { case shortNode: // Deleting from the subtrie reduced it to another @@ -221,7 +233,7 @@ func (t *Trie) delete(n node, key []byte) node { } case fullNode: - n[key[0]] = t.delete(n[key[0]], key[1:]) + n[key[0]] = t.delete(n[key[0]], append(prefix, key[0]), key[1:], ctx) // Check how many non-nil entries are left after deleting and // reduce the full node to a short node if only one entry is // left. Since n must've contained at least two children @@ -250,7 +262,7 @@ func (t *Trie) delete(n node, key []byte) node { // shortNode{..., shortNode{...}}. Since the entry // might not be loaded yet, resolve it just for this // check. - cnode := t.resolve(n[pos]) + cnode := t.resolve(n[pos], prefix, []byte{byte(pos)}, ctx) if cnode, ok := cnode.(shortNode); ok { k := append([]byte{byte(pos)}, cnode.Key...) return shortNode{k, cnode.Val} @@ -273,7 +285,7 @@ func (t *Trie) delete(n node, key []byte) node { // // TODO: track whether deletion actually hit a key and keep // n as a hash node if it didn't. - return t.delete(t.resolveHash(n), key) + return t.delete(t.resolveHash(n, prefix, key, ctx), prefix, key, ctx) default: panic(fmt.Sprintf("%T: invalid node: %v (%v)", n, n, key)) @@ -287,19 +299,28 @@ func concat(s1 []byte, s2 ...byte) []byte { return r } -func (t *Trie) resolve(n node) node { +func (t *Trie) resolve(n node, prefix, suffix []byte, ctx *access.OdrContext) node { if n, ok := n.(hashNode); ok { - return t.resolveHash(n) + return t.resolveHash(n, prefix, suffix, ctx) } return n } -func (t *Trie) resolveHash(n hashNode) node { +func (t *Trie) resolveHash(n hashNode, prefix, suffix []byte, ctx *access.OdrContext) node { if v, ok := globalCache.Get(n); ok { return v } enc, err := t.db.Get(n) if err != nil || enc == nil { + if prefix != nil || suffix != nil { + key := compactHexEncode(append(prefix, suffix...)) + if t.access != nil && t.access.RetrieveKey(key, ctx) { + enc, err = t.db.Get(n) + } + } + } + if err != nil || enc == nil { + // TODO: This needs to be improved to properly distinguish errors. // Disk I/O errors shouldn't produce nil (and cause a // consensus failure or weird crash), but it is unclear how @@ -337,7 +358,7 @@ func (t *Trie) Commit() (root common.Hash, err error) { if t.db == nil { panic("Commit called on trie with nil database") } - return t.CommitTo(t.db) + return t.CommitTo(t.db, access.NullCtx) } // CommitTo writes all nodes to the given database. @@ -347,7 +368,12 @@ func (t *Trie) Commit() (root common.Hash, err error) { // load nodes from the trie's database. Calling code must ensure that // the changes made to db are written back to the trie's attached // database before using the trie. -func (t *Trie) CommitTo(db DatabaseWriter) (root common.Hash, err error) { +func (t *Trie) CommitTo(db DatabaseWriter, ctx *access.OdrContext) (root common.Hash, err error) { + if ctx != nil && ctx.IsCancelled() { + // we should never commit potentially corrupt trie nodes to the db + return common.Hash{}, ErrTrieOdrFailure + } + n, err := t.hashRoot(db) if err != nil { return (common.Hash{}), err diff --git a/trie/trie_test.go b/trie/trie_test.go index c96861bedb..4c100a20b5 100644 --- a/trie/trie_test.go +++ b/trie/trie_test.go @@ -26,6 +26,7 @@ import ( "github.com/davecgh/go-spew/spew" "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/ethdb" ) @@ -37,7 +38,7 @@ func init() { // Used for testing func newEmpty() *Trie { db, _ := ethdb.NewMemDatabase() - trie, _ := New(common.Hash{}, db) + trie, _ := New(common.Hash{}, db, nil) return trie } @@ -54,13 +55,13 @@ func TestNull(t *testing.T) { var trie Trie key := make([]byte, 32) value := common.FromHex("0x823140710bf13990e4500136726d8b55") - trie.Update(key, value) - value = trie.Get(key) + trie.Update(key, value, access.NoOdr) + value = trie.Get(key, access.NoOdr) } func TestMissingRoot(t *testing.T) { db, _ := ethdb.NewMemDatabase() - trie, err := New(common.HexToHash("0beec7b5ea3f0fdbc95d0dd47f3c5bc275da8a33"), db) + trie, err := New(common.HexToHash("0beec7b5ea3f0fdbc95d0dd47f3c5bc275da8a33"), db, nil) if trie != nil { t.Error("New returned non-nil trie for invalid root") } @@ -190,7 +191,7 @@ func TestReplication(t *testing.T) { } // create a new trie on top of the database and check that lookups work. - trie2, err := New(exp, trie.db) + trie2, err := New(exp, trie.db, nil) if err != nil { t.Fatalf("can't recreate trie at %x: %v", exp, err) } @@ -231,7 +232,7 @@ func paranoiaCheck(t1 *Trie) (bool, *Trie) { t2 := new(Trie) it := NewIterator(t1) for it.Next() { - t2.Update(it.Key, it.Value) + t2.Update(it.Key, it.Value, access.NoOdr) } return t2.Hash() == t1.Hash(), t2 } @@ -276,15 +277,15 @@ func TestOutput(t *testing.T) { trie.Commit() fmt.Println("############################## SMALL ################################") - trie2, _ := New(trie.Hash(), trie.db) + trie2, _ := New(trie.Hash(), trie.db, nil) getString(trie2, base+"20") fmt.Println(trie2.root) } func TestLargeValue(t *testing.T) { trie := newEmpty() - trie.Update([]byte("key1"), []byte{99, 99, 99, 99}) - trie.Update([]byte("key2"), bytes.Repeat([]byte{1}, 32)) + trie.Update([]byte("key1"), []byte{99, 99, 99, 99}, access.NoOdr) + trie.Update([]byte("key2"), bytes.Repeat([]byte{1}, 32), access.NoOdr) trie.Hash() } @@ -301,8 +302,8 @@ func TestLargeData(t *testing.T) { for i := byte(0); i < 255; i++ { value := &kv{common.LeftPadBytes([]byte{i}, 32), []byte{i}, false} value2 := &kv{common.LeftPadBytes([]byte{10, i}, 32), []byte{i}, false} - trie.Update(value.k, value.v) - trie.Update(value2.k, value2.v) + trie.Update(value.k, value.v, access.NoOdr) + trie.Update(value2.k, value2.v, access.NoOdr) vals[string(value.k)] = value vals[string(value2.k)] = value2 } @@ -341,12 +342,12 @@ func benchGet(b *testing.B, commit bool) { if commit { dir, tmpdb := tempDB() defer os.RemoveAll(dir) - trie, _ = New(common.Hash{}, tmpdb) + trie, _ = New(common.Hash{}, tmpdb, nil) } k := make([]byte, 32) for i := 0; i < benchElemCount; i++ { binary.LittleEndian.PutUint64(k, uint64(i)) - trie.Update(k, k) + trie.Update(k, k, access.NoOdr) } binary.LittleEndian.PutUint64(k, benchElemCount/2) if commit { @@ -355,7 +356,7 @@ func benchGet(b *testing.B, commit bool) { b.ResetTimer() for i := 0; i < b.N; i++ { - trie.Get(k) + trie.Get(k, access.NoOdr) } } @@ -364,7 +365,7 @@ func benchUpdate(b *testing.B, e binary.ByteOrder) *Trie { k := make([]byte, 32) for i := 0; i < b.N; i++ { e.PutUint64(k, uint64(i)) - trie.Update(k, k) + trie.Update(k, k, access.NoOdr) } return trie } @@ -374,7 +375,7 @@ func benchHash(b *testing.B, e binary.ByteOrder) { k := make([]byte, 32) for i := 0; i < benchElemCount; i++ { e.PutUint64(k, uint64(i)) - trie.Update(k, k) + trie.Update(k, k, access.NoOdr) } b.ResetTimer() @@ -396,13 +397,13 @@ func tempDB() (string, Database) { } func getString(trie *Trie, k string) []byte { - return trie.Get([]byte(k)) + return trie.Get([]byte(k), access.NoOdr) } func updateString(trie *Trie, k, v string) { - trie.Update([]byte(k), []byte(v)) + trie.Update([]byte(k), []byte(v), access.NoOdr) } func deleteString(trie *Trie, k string) { - trie.Delete([]byte(k)) + trie.Delete([]byte(k), access.NoOdr) } diff --git a/xeth/state.go b/xeth/state.go index 981fe63b7e..d67f0b9de0 100644 --- a/xeth/state.go +++ b/xeth/state.go @@ -45,7 +45,7 @@ func (self *State) SafeGet(addr string) *Object { func (self *State) safeGet(addr string) *state.StateObject { object := self.state.GetStateObject(common.HexToAddress(addr)) if object == nil { - object = state.NewStateObject(common.HexToAddress(addr), self.xeth.backend.ChainDb()) + object = state.NewStateObject(common.HexToAddress(addr), self.xeth.backend.ChainAccess()) } return object diff --git a/xeth/types.go b/xeth/types.go index 218c8dc7cd..6207d423b1 100644 --- a/xeth/types.go +++ b/xeth/types.go @@ -24,6 +24,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/crypto" @@ -39,20 +40,20 @@ func NewObject(state *state.StateObject) *Object { return &Object{state} } -func (self *Object) StorageString(str string) []byte { +func (self *Object) StorageString(str string, ctx *access.OdrContext) []byte { if common.IsHex(str) { - return self.storage(common.Hex2Bytes(str[2:])) + return self.storage(common.Hex2Bytes(str[2:]), ctx) } else { - return self.storage(common.RightPadBytes([]byte(str), 32)) + return self.storage(common.RightPadBytes([]byte(str), 32), ctx) } } -func (self *Object) StorageValue(addr *common.Value) []byte { - return self.storage(addr.Bytes()) +func (self *Object) StorageValue(addr *common.Value, ctx *access.OdrContext) []byte { + return self.storage(addr.Bytes(), ctx) } -func (self *Object) storage(addr []byte) []byte { - return self.StateObject.GetState(common.BytesToHash(addr)).Bytes() +func (self *Object) storage(addr []byte, ctx *access.OdrContext) []byte { + return self.StateObject.GetState(common.BytesToHash(addr), ctx).Bytes() } func (self *Object) Storage() (storage map[string]string) { diff --git a/xeth/xeth.go b/xeth/xeth.go index 35e6dd52d8..3181ee6d62 100644 --- a/xeth/xeth.go +++ b/xeth/xeth.go @@ -31,6 +31,7 @@ import ( "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common/compiler" "github.com/ethereum/go-ethereum/core" + "github.com/ethereum/go-ethereum/core/access" "github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/vm" @@ -84,6 +85,7 @@ type XEth struct { state *State whisper *Whisper filterManager *filters.FilterSystem + ctx *access.OdrContext } func NewTest(eth *eth.Ethereum, frontend Frontend) *XEth { @@ -113,7 +115,7 @@ func New(ethereum *eth.Ethereum, frontend Frontend) *XEth { if frontend == nil { xeth.frontend = dummyFrontend{} } - state, _ := xeth.backend.BlockChain().State() + state, _ := xeth.backend.BlockChain().State(access.NullCtx) xeth.state = NewState(xeth, state) go xeth.start() return xeth @@ -207,15 +209,15 @@ func (self *XEth) AtStateNum(num int64) *XEth { var err error switch num { case -2: - st = self.backend.Miner().PendingState().Copy() + st = self.backend.Miner().PendingState().Copy(self.ctx) default: if block := self.getBlockByHeight(num); block != nil { - st, err = state.New(block.Root(), self.backend.ChainDb()) + st, err = state.New(block.Root(), self.backend.ChainAccess(), self.ctx) if err != nil { return nil } } else { - st, err = state.New(self.backend.BlockChain().GetBlockByNumber(0).Root(), self.backend.ChainDb()) + st, err = state.New(self.backend.BlockChain().GetBlockByNumber(0, self.ctx).Root(), self.backend.ChainAccess(), self.ctx) if err != nil { return nil } @@ -238,6 +240,18 @@ func (self *XEth) WithState(statedb *state.StateDB) *XEth { func (self *XEth) State() *State { return self.state } +func (self *XEth) WithCtx(ctx *access.OdrContext) *XEth { + xeth := &XEth{ + backend: self.backend, + frontend: self.frontend, + gpo: self.gpo, + ctx: ctx, + } + + xeth.state = NewState(xeth, self.state.State().Copy(ctx)) + return xeth +} + // subscribes to new head block events and // waits until blockchain height is greater n at any time // given the current head, waits for the next chain event @@ -270,7 +284,7 @@ func (self *XEth) UpdateState() (wait chan *big.Int) { wait <- n n = nil } - statedb, err := state.New(event.Block.Root(), self.backend.ChainDb()) + statedb, err := state.New(event.Block.Root(), self.backend.ChainAccess(), access.NullCtx) if err != nil { glog.V(logger.Error).Infoln("Could not create new state: %v", err) return @@ -305,19 +319,19 @@ func (self *XEth) getBlockByHeight(height int64) *types.Block { num = uint64(height) } - return self.backend.BlockChain().GetBlockByNumber(num) + return self.backend.BlockChain().GetBlockByNumber(num, self.ctx) } func (self *XEth) BlockByHash(strHash string) *Block { hash := common.HexToHash(strHash) - block := self.backend.BlockChain().GetBlock(hash) + block := self.backend.BlockChain().GetBlock(hash, self.ctx) return NewBlock(block) } func (self *XEth) EthBlockByHash(strHash string) *types.Block { hash := common.HexToHash(strHash) - block := self.backend.BlockChain().GetBlock(hash) + block := self.backend.BlockChain().GetBlock(hash, self.ctx) return block } @@ -375,15 +389,19 @@ func (self *XEth) Td(hash common.Hash) *big.Int { } func (self *XEth) CurrentBlock() *types.Block { - return self.backend.BlockChain().CurrentBlock() + return self.backend.BlockChain().CurrentBlockOdr(self.ctx) +} + +func (self *XEth) CurrentHeader() *types.Header { + return self.backend.BlockChain().CurrentHeader() } func (self *XEth) GetBlockReceipts(bhash common.Hash) types.Receipts { - return self.backend.BlockProcessor().GetBlockReceipts(bhash) + return self.backend.BlockProcessor().GetBlockReceipts(bhash) //ODR } -func (self *XEth) GetTxReceipt(txhash common.Hash) *types.Receipt { - return core.GetReceipt(self.backend.ChainDb(), txhash) +func (self *XEth) GetTxReceipt(txhash common.Hash) *types.Receipt { //ODR + return core.GetReceipt(self.backend.ChainAccess(), txhash, access.NoOdr) //TODO } func (self *XEth) GasLimit() *big.Int { @@ -547,7 +565,7 @@ func (self *XEth) NewLogFilter(earliest, latest int64, skip, max int, address [] self.logMu.Lock() defer self.logMu.Unlock() - filter := filters.New(self.backend.ChainDb()) + filter := filters.New(self.backend.ChainAccess()) id := self.filterManager.Add(filter) self.logQueue[id] = &logQueue{timeout: time.Now()} @@ -571,7 +589,7 @@ func (self *XEth) NewTransactionFilter() int { self.transactionMu.Lock() defer self.transactionMu.Unlock() - filter := filters.New(self.backend.ChainDb()) + filter := filters.New(self.backend.ChainAccess()) id := self.filterManager.Add(filter) self.transactionQueue[id] = &hashQueue{timeout: time.Now()} @@ -590,7 +608,7 @@ func (self *XEth) NewBlockFilter() int { self.blockMu.Lock() defer self.blockMu.Unlock() - filter := filters.New(self.backend.ChainDb()) + filter := filters.New(self.backend.ChainAccess()) id := self.filterManager.Add(filter) self.blockQueue[id] = &hashQueue{timeout: time.Now()} @@ -657,7 +675,7 @@ func (self *XEth) Logs(id int) vm.Logs { } func (self *XEth) AllLogs(earliest, latest int64, skip, max int, address []string, topics [][]string) vm.Logs { - filter := filters.New(self.backend.ChainDb()) + filter := filters.New(self.backend.ChainAccess()) filter.SetBeginBlock(earliest) filter.SetEndBlock(latest) filter.SetAddresses(cAddress(address)) @@ -829,7 +847,7 @@ func (self *XEth) PushTx(encodedTx string) (string, error) { } func (self *XEth) Call(fromStr, toStr, valueStr, gasStr, gasPriceStr, dataStr string) (string, string, error) { - statedb := self.State().State().Copy() + statedb := self.State().State().CopyWithCtx() var from *state.StateObject if len(fromStr) == 0 { accounts, err := self.backend.AccountManager().Accounts() @@ -864,7 +882,7 @@ func (self *XEth) Call(fromStr, toStr, valueStr, gasStr, gasPriceStr, dataStr st msg.gasPrice = self.DefaultGasPrice() } - header := self.CurrentBlock().Header() + header := self.CurrentHeader() vmenv := core.NewEnv(statedb, self.backend.BlockChain(), msg, header) gp := new(core.GasPool).AddGas(common.MaxBig) res, gas, err := core.ApplyMessage(vmenv, msg, gp)