From bbcc6dd0ad4ac75241eb429ebe46e33ddb2b4878 Mon Sep 17 00:00:00 2001 From: Jerry Date: Sat, 16 Jul 2022 16:54:02 -0700 Subject: [PATCH] Parallel state processor --- core/blockstm/executor.go | 241 +++++++++++++++++++++++++++++++ core/blockstm/mvhashmap.go | 12 +- core/blockstm/status.go | 163 +++++++++++++++++++++ core/blockstm/txio.go | 22 ++- core/parallel_state_processor.go | 231 +++++++++++++++++++++++++++++ core/state/statedb.go | 62 +++++++- core/state/statedb_test.go | 25 ++++ 7 files changed, 742 insertions(+), 14 deletions(-) create mode 100644 core/blockstm/executor.go create mode 100644 core/blockstm/status.go create mode 100644 core/parallel_state_processor.go diff --git a/core/blockstm/executor.go b/core/blockstm/executor.go new file mode 100644 index 0000000000..cb5019bbb2 --- /dev/null +++ b/core/blockstm/executor.go @@ -0,0 +1,241 @@ +package blockstm + +import ( + "fmt" + + "github.com/ethereum/go-ethereum/log" +) + +type ExecResult struct { + err error + ver Version + txIn TxnInput + txOut TxnOutput + txAllOut TxnOutput +} + +type ExecTask interface { + Execute(mvh *MVHashMap, incarnation int) error + MVReadList() []ReadDescriptor + MVWriteList() []WriteDescriptor + MVFullWriteList() []WriteDescriptor +} + +type ExecVersionView struct { + ver Version + et ExecTask + mvh *MVHashMap +} + +func (ev *ExecVersionView) Execute() (er ExecResult) { + er.ver = ev.ver + if er.err = ev.et.Execute(ev.mvh, ev.ver.Incarnation); er.err != nil { + log.Debug("blockstm executed task failed", "Tx index", ev.ver.TxnIndex, "incarnation", ev.ver.Incarnation, "err", er.err) + return + } + + er.txIn = ev.et.MVReadList() + er.txOut = ev.et.MVWriteList() + er.txAllOut = ev.et.MVFullWriteList() + log.Debug("blockstm executed task", "Tx index", ev.ver.TxnIndex, "incarnation", ev.ver.Incarnation, "err", er.err) + + return +} + +var ErrExecAbort = fmt.Errorf("execution aborted with dependency") + +const numGoProcs = 4 + +// nolint: gocognit +func ExecuteParallel(tasks []ExecTask) (lastTxIO *TxnInputOutput, err error) { + if len(tasks) == 0 { + return MakeTxnInputOutput(len(tasks)), nil + } + + chTasks := make(chan ExecVersionView, len(tasks)) + chResults := make(chan ExecResult, len(tasks)) + chDone := make(chan bool) + + var cntExec, cntSuccess, cntAbort, cntTotalValidations, cntValidationFail int + + for i := 0; i < numGoProcs; i++ { + go func(procNum int, t chan ExecVersionView) { + Loop: + for { + select { + case task := <-t: + { + res := task.Execute() + chResults <- res + } + case <-chDone: + break Loop + } + } + log.Debug("blockstm", "proc done", procNum) // TODO: logging ... + }(i, chTasks) + } + + mvh := MakeMVHashMap() + + execTasks := makeStatusManager(len(tasks)) + validateTasks := makeStatusManager(0) + + // bootstrap execution + for x := 0; x < numGoProcs; x++ { + tx := execTasks.takeNextPending() + if tx != -1 { + cntExec++ + + log.Debug("blockstm", "bootstrap: proc", x, "executing task", tx) + chTasks <- ExecVersionView{ver: Version{tx, 0}, et: tasks[tx], mvh: mvh} + } + } + + lastTxIO = MakeTxnInputOutput(len(tasks)) + txIncarnations := make([]int, len(tasks)) + + diagExecSuccess := make([]int, len(tasks)) + diagExecAbort := make([]int, len(tasks)) + + for { + res := <-chResults + switch res.err { + case nil: + { + mvh.FlushMVWriteSet(res.txAllOut) + lastTxIO.recordRead(res.ver.TxnIndex, res.txIn) + if res.ver.Incarnation == 0 { + lastTxIO.recordWrite(res.ver.TxnIndex, res.txOut) + lastTxIO.recordAllWrite(res.ver.TxnIndex, res.txAllOut) + } else { + if res.txAllOut.hasNewWrite(lastTxIO.AllWriteSet(res.ver.TxnIndex)) { + log.Debug("blockstm", "Revalidate completed txs greater than current tx: ", res.ver.TxnIndex) + validateTasks.pushPendingSet(execTasks.getRevalidationRange(res.ver.TxnIndex)) + } + + prevWrite := lastTxIO.AllWriteSet(res.ver.TxnIndex) + + // Remove entries that were previously written but are no longer written + + cmpMap := make(map[string]bool) + + for _, w := range res.txAllOut { + cmpMap[string(w.Path)] = true + } + + for _, v := range prevWrite { + if _, ok := cmpMap[string(v.Path)]; !ok { + mvh.Delete(v.Path, res.ver.TxnIndex) + } + } + + lastTxIO.recordWrite(res.ver.TxnIndex, res.txOut) + lastTxIO.recordAllWrite(res.ver.TxnIndex, res.txAllOut) + } + validateTasks.pushPending(res.ver.TxnIndex) + execTasks.markComplete(res.ver.TxnIndex) + if diagExecSuccess[res.ver.TxnIndex] > 0 && diagExecAbort[res.ver.TxnIndex] == 0 { + log.Debug("blockstm", "got multiple successful execution w/o abort?", "Tx", res.ver.TxnIndex, "incarnation", res.ver.Incarnation) + } + diagExecSuccess[res.ver.TxnIndex]++ + cntSuccess++ + } + case ErrExecAbort: + { + // bit of a subtle / tricky bug here. this adds the tx back to pending ... + execTasks.revertInProgress(res.ver.TxnIndex) + // ... but the incarnation needs to be bumped + txIncarnations[res.ver.TxnIndex]++ + diagExecAbort[res.ver.TxnIndex]++ + cntAbort++ + } + default: + { + err = res.err + break + } + } + + // if we got more work, queue one up... + nextTx := execTasks.takeNextPending() + if nextTx != -1 { + cntExec++ + chTasks <- ExecVersionView{ver: Version{nextTx, txIncarnations[nextTx]}, et: tasks[nextTx], mvh: mvh} + } + + // do validations ... + maxComplete := execTasks.maxAllComplete() + + const validationIncrement = 2 + + cntValidate := validateTasks.countPending() + // if we're currently done with all execution tasks then let's validate everything; otherwise do one increment ... + if execTasks.countComplete() != len(tasks) && cntValidate > validationIncrement { + cntValidate = validationIncrement + } + + var toValidate []int + + for i := 0; i < cntValidate; i++ { + if validateTasks.minPending() <= maxComplete { + toValidate = append(toValidate, validateTasks.takeNextPending()) + } else { + break + } + } + + for i := 0; i < len(toValidate); i++ { + cntTotalValidations++ + + tx := toValidate[i] + log.Debug("blockstm", "validating task", tx) + + if ValidateVersion(tx, lastTxIO, mvh) { + log.Debug("blockstm", "* completed validation task", tx) + validateTasks.markComplete(tx) + } else { + log.Debug("blockstm", "* validation task FAILED", tx) + cntValidationFail++ + diagExecAbort[tx]++ + for _, v := range lastTxIO.AllWriteSet(tx) { + mvh.MarkEstimate(v.Path, tx) + } + // 'create validation tasks for all transactions > tx ...' + validateTasks.pushPendingSet(execTasks.getRevalidationRange(tx + 1)) + validateTasks.clearInProgress(tx) // clear in progress - pending will be added again once new incarnation executes + if execTasks.checkPending(tx) { + // println() // have to think about this ... + } else { + execTasks.pushPending(tx) + execTasks.clearComplete(tx) + txIncarnations[tx]++ + } + } + } + + // if we didn't queue work previously, do check again so we keep making progress ... + if nextTx == -1 { + nextTx = execTasks.takeNextPending() + if nextTx != -1 { + cntExec++ + + log.Debug("blockstm", "# tx queued up", nextTx) + chTasks <- ExecVersionView{ver: Version{nextTx, txIncarnations[nextTx]}, et: tasks[nextTx], mvh: mvh} + } + } + + if validateTasks.countComplete() == len(tasks) && execTasks.countComplete() == len(tasks) { + log.Debug("blockstm exec summary", "execs", cntExec, "success", cntSuccess, "aborts", cntAbort, "validations", cntTotalValidations, "failures", cntValidationFail) + break + } + } + + for i := 0; i < numGoProcs; i++ { + chDone <- true + } + close(chTasks) + close(chResults) + + return +} diff --git a/core/blockstm/mvhashmap.go b/core/blockstm/mvhashmap.go index ff9f1a6f9d..52a5487b5d 100644 --- a/core/blockstm/mvhashmap.go +++ b/core/blockstm/mvhashmap.go @@ -5,6 +5,8 @@ import ( "sync" "github.com/emirpasic/gods/maps/treemap" + + "github.com/ethereum/go-ethereum/log" ) const FlagDone = 0 @@ -81,7 +83,7 @@ func (mv *MVHashMap) Write(k []byte, v Version, data interface{}) { panic(fmt.Errorf("existing transaction value does not have lower incarnation: %v, %v", string(k), v.TxnIndex)) } else if ci.(*WriteCell).flag == FlagEstimate { - println("marking previous estimate as done tx", v.TxnIndex, v.Incarnation) + log.Debug("mvhashmap marking previous estimate as done", "tx index", v.TxnIndex, "incarnation", v.Incarnation) } ci.(*WriteCell).flag = FlagDone @@ -190,10 +192,16 @@ func (mv *MVHashMap) Read(k []byte, txIdx int) (res MVReadResult) { return } +func (mv *MVHashMap) FlushMVWriteSet(writes []WriteDescriptor) { + for _, v := range writes { + mv.Write(v.Path, v.V, v.Val) + } +} + func ValidateVersion(txIdx int, lastInputOutput *TxnInputOutput, versionedData *MVHashMap) (valid bool) { valid = true - for _, rd := range lastInputOutput.readSet(txIdx) { + for _, rd := range lastInputOutput.ReadSet(txIdx) { mvResult := versionedData.Read(rd.Path, txIdx) switch mvResult.Status() { case MVReadResultDone: diff --git a/core/blockstm/status.go b/core/blockstm/status.go new file mode 100644 index 0000000000..759abf63eb --- /dev/null +++ b/core/blockstm/status.go @@ -0,0 +1,163 @@ +package blockstm + +import ( + "fmt" + "sort" +) + +func makeStatusManager(numTasks int) (t taskStatusManager) { + t.pending = make([]int, numTasks) + for i := 0; i < numTasks; i++ { + t.pending[i] = i + } + + return +} + +type taskStatusManager struct { + pending []int + inProgress []int + complete []int +} + +func insertInList(l []int, v int) []int { + if len(l) == 0 || v > l[len(l)-1] { + return append(l, v) + } else { + x := sort.SearchInts(l, v) + if x < len(l) && l[x] == v { + // already in list + return l + } + a := append(l[:x+1], l[x:]...) + a[x] = v + return a + } +} + +func (m *taskStatusManager) takeNextPending() int { + if len(m.pending) == 0 { + return -1 + } + + x := m.pending[0] + m.pending = m.pending[1:] + m.inProgress = insertInList(m.inProgress, x) + + return x +} + +func hasNoGap(l []int) bool { + return l[0]+len(l) == l[len(l)-1]+1 +} + +func (m taskStatusManager) maxAllComplete() int { + if len(m.complete) == 0 || m.complete[0] != 0 { + return -1 + } else if m.complete[len(m.complete)-1] == len(m.complete)-1 { + return m.complete[len(m.complete)-1] + } else { + for i := len(m.complete) - 2; i >= 0; i-- { + if hasNoGap(m.complete[:i+1]) { + return m.complete[i] + } + } + } + + return -1 +} + +func (m *taskStatusManager) pushPending(tx int) { + m.pending = insertInList(m.pending, tx) +} + +func removeFromList(l []int, v int, expect bool) []int { + x := sort.SearchInts(l, v) + if x == -1 || l[x] != v { + if expect { + panic(fmt.Errorf("should not happen - element expected in list")) + } + + return l + } + + switch x { + case 0: + return l[1:] + case len(l) - 1: + return l[:len(l)-1] + default: + return append(l[:x], l[x+1:]...) + } +} + +func (m *taskStatusManager) markComplete(tx int) { + m.inProgress = removeFromList(m.inProgress, tx, true) + m.complete = insertInList(m.complete, tx) +} + +func (m *taskStatusManager) minPending() int { + if len(m.pending) == 0 { + return -1 + } else { + return m.pending[0] + } +} + +func (m *taskStatusManager) countComplete() int { + return len(m.complete) +} + +func (m *taskStatusManager) revertInProgress(tx int) { + m.inProgress = removeFromList(m.inProgress, tx, true) + m.pending = insertInList(m.pending, tx) +} + +func (m *taskStatusManager) clearInProgress(tx int) { + m.inProgress = removeFromList(m.inProgress, tx, true) +} + +func (m *taskStatusManager) countPending() int { + return len(m.pending) +} + +func (m *taskStatusManager) checkInProgress(tx int) bool { + x := sort.SearchInts(m.inProgress, tx) + if x < len(m.inProgress) && m.inProgress[x] == tx { + return true + } + + return false +} + +func (m *taskStatusManager) checkPending(tx int) bool { + x := sort.SearchInts(m.pending, tx) + if x < len(m.pending) && m.pending[x] == tx { + return true + } + + return false +} + +// getRevalidationRange: this range will be all tasks from tx (inclusive) that are not currently in progress up to the +// 'all complete' limit +func (m *taskStatusManager) getRevalidationRange(txFrom int) (ret []int) { + max := m.maxAllComplete() // haven't learned to trust compilers :) + for x := txFrom; x <= max; x++ { + if !m.checkInProgress(x) { + ret = append(ret, x) + } + } + + return +} + +func (m *taskStatusManager) pushPendingSet(set []int) { + for _, v := range set { + m.pushPending(v) + } +} + +func (m *taskStatusManager) clearComplete(tx int) { + m.complete = removeFromList(m.complete, tx, false) +} diff --git a/core/blockstm/txio.go b/core/blockstm/txio.go index 4325277c1d..7716197acd 100644 --- a/core/blockstm/txio.go +++ b/core/blockstm/txio.go @@ -47,22 +47,28 @@ func (txo TxnOutput) hasNewWrite(cmpSet []WriteDescriptor) bool { } type TxnInputOutput struct { - inputs []TxnInput - outputs []TxnOutput + inputs []TxnInput + outputs []TxnOutput // write sets that should be checked during validation + allOutputs []TxnOutput // entire write sets in MVHashMap. allOutputs should always be a parent set of outputs } -func (io *TxnInputOutput) readSet(txnIdx int) []ReadDescriptor { +func (io *TxnInputOutput) ReadSet(txnIdx int) []ReadDescriptor { return io.inputs[txnIdx] } -func (io *TxnInputOutput) writeSet(txnIdx int) []WriteDescriptor { +func (io *TxnInputOutput) WriteSet(txnIdx int) []WriteDescriptor { return io.outputs[txnIdx] } +func (io *TxnInputOutput) AllWriteSet(txnIdx int) []WriteDescriptor { + return io.allOutputs[txnIdx] +} + func MakeTxnInputOutput(numTx int) *TxnInputOutput { return &TxnInputOutput{ - inputs: make([]TxnInput, numTx), - outputs: make([]TxnOutput, numTx), + inputs: make([]TxnInput, numTx), + outputs: make([]TxnOutput, numTx), + allOutputs: make([]TxnOutput, numTx), } } @@ -73,3 +79,7 @@ func (io *TxnInputOutput) recordRead(txId int, input []ReadDescriptor) { func (io *TxnInputOutput) recordWrite(txId int, output []WriteDescriptor) { io.outputs[txId] = output } + +func (io *TxnInputOutput) recordAllWrite(txId int, output []WriteDescriptor) { + io.allOutputs[txId] = output +} diff --git a/core/parallel_state_processor.go b/core/parallel_state_processor.go new file mode 100644 index 0000000000..c2d6647f65 --- /dev/null +++ b/core/parallel_state_processor.go @@ -0,0 +1,231 @@ +// 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 ( + "fmt" + "math/big" + + "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/consensus" + "github.com/ethereum/go-ethereum/consensus/misc" + "github.com/ethereum/go-ethereum/core/blockstm" + "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/log" + "github.com/ethereum/go-ethereum/params" +) + +// StateProcessor is a basic Processor, which takes care of transitioning +// state from one point to another. +// +// StateProcessor implements Processor. +type ParallelStateProcessor struct { + config *params.ChainConfig // Chain configuration options + bc *BlockChain // Canonical block chain + engine consensus.Engine // Consensus engine used for block rewards +} + +// NewStateProcessor initialises a new StateProcessor. +func NewParallelStateProcessor(config *params.ChainConfig, bc *BlockChain, engine consensus.Engine) *ParallelStateProcessor { + return &ParallelStateProcessor{ + config: config, + bc: bc, + engine: engine, + } +} + +type ExecutionTask struct { + msg types.Message + config *params.ChainConfig + + gasLimit uint64 + blockNumber *big.Int + blockHash common.Hash + blockContext vm.BlockContext + tx *types.Transaction + index int + statedb *state.StateDB // State database that stores the modified values after tx execution. + cleanStateDB *state.StateDB // A clean copy of the initial statedb. It should not be modified. + evmConfig vm.Config + result *ExecutionResult +} + +func (task *ExecutionTask) Execute(mvh *blockstm.MVHashMap, incarnation int) (err error) { + task.statedb = task.cleanStateDB.Copy() + task.statedb.Prepare(task.tx.Hash(), task.index) + task.statedb.SetMVHashmap(mvh) + task.statedb.SetIncarnation(incarnation) + + evm := vm.NewEVM(task.blockContext, vm.TxContext{}, task.statedb, task.config, task.evmConfig) + + // Create a new context to be used in the EVM environment. + txContext := NewEVMTxContext(task.msg) + evm.Reset(txContext, task.statedb) + + defer func() { + if r := recover(); r != nil { + // In some pre-matured executions, EVM will panic. Recover from panic and retry the execution. + log.Debug("Recovered from EVM failure. Error:\n", r) + + err = blockstm.ErrExecAbort + + return + } + }() + + // Apply the transaction to the current state (included in the env). + result, err := ApplyMessage(evm, task.msg, new(GasPool).AddGas(task.gasLimit)) + + if task.statedb.HadInvalidRead() || err != nil { + err = blockstm.ErrExecAbort + return + } + + task.statedb.Finalise(false) + + task.result = result + + return +} + +func (task *ExecutionTask) MVReadList() []blockstm.ReadDescriptor { + return task.statedb.MVReadList() +} + +func (task *ExecutionTask) MVWriteList() []blockstm.WriteDescriptor { + return task.statedb.MVWriteList() +} + +func (task *ExecutionTask) MVFullWriteList() []blockstm.WriteDescriptor { + return task.statedb.MVFullWriteList() +} + +// Process processes the state changes according to the Ethereum rules by running +// the transaction messages using the statedb and applying any rewards to both +// the processor (coinbase) and any included uncles. +// +// Process returns the receipts and logs accumulated during the process and +// returns the amount of gas that was used in the process. If any of the +// transactions failed to execute due to insufficient gas it will return an error. +func (p *ParallelStateProcessor) Process(block *types.Block, statedb *state.StateDB, cfg vm.Config) (types.Receipts, []*types.Log, uint64, error) { + var ( + receipts types.Receipts + header = block.Header() + blockHash = block.Hash() + blockNumber = block.Number() + allLogs []*types.Log + usedGas = new(uint64) + ) + // Mutate the block and state according to any hard-fork specs + if p.config.DAOForkSupport && p.config.DAOForkBlock != nil && p.config.DAOForkBlock.Cmp(block.Number()) == 0 { + misc.ApplyDAOHardFork(statedb) + } + + tasks := make([]blockstm.ExecTask, 0, len(block.Transactions())) + + // Iterate over and process the individual transactions + for i, tx := range block.Transactions() { + msg, err := tx.AsMessage(types.MakeSigner(p.config, header.Number), header.BaseFee) + if err != nil { + log.Error("error creating message", "err", err) + return nil, nil, 0, fmt.Errorf("could not apply tx %d [%v]: %w", i, tx.Hash().Hex(), err) + } + + cleansdb := statedb.Copy() + + task := &ExecutionTask{ + msg: msg, + config: p.config, + gasLimit: block.GasLimit(), + blockNumber: blockNumber, + blockHash: blockHash, + tx: tx, + index: i, + cleanStateDB: cleansdb, + blockContext: NewEVMBlockContext(header, p.bc, nil), + } + + tasks = append(tasks, task) + } + + _, err := blockstm.ExecuteParallel(tasks) + + if err != nil { + log.Error("blockstm error executing block", "err", err) + return nil, nil, 0, err + } + + for _, task := range tasks { + task := task.(*ExecutionTask) + statedb.Prepare(task.tx.Hash(), task.index) + statedb.ApplyMVWriteSet(task.MVWriteList()) + + for _, l := range task.statedb.GetLogs(task.tx.Hash(), blockHash) { + statedb.AddLog(l) + } + + for k, v := range task.statedb.Preimages() { + statedb.AddPreimage(k, v) + } + + // Update the state with pending changes. + var root []byte + + if p.config.IsByzantium(blockNumber) { + statedb.Finalise(true) + } else { + root = statedb.IntermediateRoot(p.config.IsEIP158(blockNumber)).Bytes() + } + + *usedGas += task.result.UsedGas + + // Create a new receipt for the transaction, storing the intermediate root and gas used + // by the tx. + receipt := &types.Receipt{Type: task.tx.Type(), PostState: root, CumulativeGasUsed: *usedGas} + if task.result.Failed() { + receipt.Status = types.ReceiptStatusFailed + } else { + receipt.Status = types.ReceiptStatusSuccessful + } + + receipt.TxHash = task.tx.Hash() + receipt.GasUsed = task.result.UsedGas + + // If the transaction created a contract, store the creation address in the receipt. + if task.msg.To() == nil { + receipt.ContractAddress = crypto.CreateAddress(task.msg.From(), task.tx.Nonce()) + } + + // Set the receipt logs and create the bloom filter. + receipt.Logs = statedb.GetLogs(task.tx.Hash(), blockHash) + receipt.Bloom = types.CreateBloom(types.Receipts{receipt}) + receipt.BlockHash = blockHash + receipt.BlockNumber = blockNumber + receipt.TransactionIndex = uint(statedb.TxIndex()) + + receipts = append(receipts, receipt) + allLogs = append(allLogs, receipt.Logs...) + } + + // Finalize the block, applying any consensus engine specific extras (e.g. block rewards) + p.engine.Finalize(p.bc, header, statedb, block.Transactions(), block.Uncles()) + + return receipts, allLogs, *usedGas, nil +} diff --git a/core/state/statedb.go b/core/state/statedb.go index 2d6250994c..1b7af5dda3 100644 --- a/core/state/statedb.go +++ b/core/state/statedb.go @@ -81,10 +81,12 @@ type StateDB struct { stateObjectsDirty map[common.Address]struct{} // State objects modified in the current execution // Block-stm related fields - mvHashmap *blockstm.MVHashMap - incarnation int - readMap map[string]blockstm.ReadDescriptor - writeMap map[string]blockstm.WriteDescriptor + mvHashmap *blockstm.MVHashMap + incarnation int + readMap map[string]blockstm.ReadDescriptor + writeMap map[string]blockstm.WriteDescriptor + newStateObjects map[common.Address]struct{} + invalidRead bool // DB error. // State objects are used by the consensus core and VM which are @@ -145,6 +147,7 @@ func New(root common.Hash, db Database, snaps *snapshot.Tree) (*StateDB, error) stateObjects: make(map[common.Address]*stateObject), stateObjectsPending: make(map[common.Address]struct{}), stateObjectsDirty: make(map[common.Address]struct{}), + newStateObjects: make(map[common.Address]struct{}), logs: make(map[common.Hash][]*types.Log), preimages: make(map[common.Hash][]byte), journal: newJournal(), @@ -177,6 +180,20 @@ func (sdb *StateDB) SetMVHashmap(mvhm *blockstm.MVHashMap) { func (s *StateDB) MVWriteList() []blockstm.WriteDescriptor { writes := make([]blockstm.WriteDescriptor, 0, len(s.writeMap)) + for _, v := range s.writeMap { + if len(v.Path) != common.AddressLength { + writes = append(writes, v) + } else if _, ok := s.newStateObjects[common.BytesToAddress(v.Path)]; ok { + writes = append(writes, v) + } + } + + return writes +} + +func (s *StateDB) MVFullWriteList() []blockstm.WriteDescriptor { + writes := make([]blockstm.WriteDescriptor, 0, len(s.writeMap)) + for _, v := range s.writeMap { writes = append(writes, v) } @@ -206,6 +223,14 @@ func (s *StateDB) ensureWriteMap() { } } +func (s *StateDB) HadInvalidRead() bool { + return s.invalidRead +} + +func (s *StateDB) SetIncarnation(inc int) { + s.incarnation = inc +} + func MVRead[T any](s *StateDB, k []byte, defaultV T, readStorage func(s *StateDB) T) (v T) { if s.mvHashmap == nil { return readStorage(s) @@ -238,6 +263,7 @@ func MVRead[T any](s *StateDB, k []byte, defaultV T, readStorage func(s *StateDB } case blockstm.MVReadResultDependency: { + s.invalidRead = true return defaultV } case blockstm.MVReadResultNone: @@ -262,7 +288,6 @@ func MVRead[T any](s *StateDB, k []byte, defaultV T, readStorage func(s *StateDB func MVWrite(s *StateDB, k []byte) { if s.mvHashmap != nil { s.ensureWriteMap() - s.mvHashmap.Write(k, s.Version(), s) s.writeMap[string(k)] = blockstm.WriteDescriptor{ Path: k, V: s.Version(), @@ -281,6 +306,15 @@ func MVWritten(s *StateDB, k []byte) bool { return ok } +// Apply entries in the write set to MVHashMap. Note that this function does not clear the write set. +func (s *StateDB) FlushMVWriteSet() { + if s.mvHashmap != nil && s.writeMap != nil { + s.mvHashmap.FlushMVWriteSet(s.MVFullWriteList()) + } +} + +// Apply entries in a given write set to StateDB. Note that this function does not change MVHashMap nor write set +// of the current StateDB. func (sw *StateDB) ApplyMVWriteSet(writes []blockstm.WriteDescriptor) { for i := range writes { path := writes[i].Path @@ -304,7 +338,8 @@ func (sw *StateDB) ApplyMVWriteSet(writes []blockstm.WriteDescriptor) { case codePath: sw.SetCode(addr, sr.GetCode(addr)) case suicidePath: - if suicided := sr.HasSuicided(addr); suicided { + stateObject := sr.getDeletedStateObject(addr) + if stateObject != nil && stateObject.deleted { sw.Suicide(addr) } default: @@ -606,6 +641,12 @@ func (s *StateDB) HasSuicided(addr common.Address) bool { // AddBalance adds amount to the account associated with addr. func (s *StateDB) AddBalance(addr common.Address, amount *big.Int) { stateObject := s.GetOrNewStateObject(addr) + + if s.mvHashmap != nil { + // ensure a read balance operation is recorded in mvHashmap + s.GetBalance(addr) + } + if stateObject != nil { stateObject = s.mvRecordWritten(stateObject) stateObject.AddBalance(amount) @@ -616,6 +657,12 @@ func (s *StateDB) AddBalance(addr common.Address, amount *big.Int) { // SubBalance subtracts amount from the account associated with addr. func (s *StateDB) SubBalance(addr common.Address, amount *big.Int) { stateObject := s.GetOrNewStateObject(addr) + + if s.mvHashmap != nil { + // ensure a read balance operation is recorded in mvHashmap + s.GetBalance(addr) + } + if stateObject != nil { stateObject = s.mvRecordWritten(stateObject) stateObject.SubBalance(amount) @@ -859,6 +906,8 @@ func (s *StateDB) createObject(addr common.Address) (newobj, prev *stateObject) s.journal.append(resetObjectChange{prev: prev, prevdestruct: prevdestruct}) } s.setStateObject(newobj) + s.newStateObjects[addr] = struct{}{} + MVWrite(s, addr.Bytes()) if prev != nil && !prev.deleted { return newobj, prev @@ -923,6 +972,7 @@ func (s *StateDB) Copy() *StateDB { stateObjects: make(map[common.Address]*stateObject, len(s.journal.dirties)), stateObjectsPending: make(map[common.Address]struct{}, len(s.stateObjectsPending)), stateObjectsDirty: make(map[common.Address]struct{}, len(s.journal.dirties)), + newStateObjects: make(map[common.Address]struct{}, len(s.newStateObjects)), refund: s.refund, logs: make(map[common.Hash][]*types.Log, len(s.logs)), logSize: s.logSize, diff --git a/core/state/statedb_test.go b/core/state/statedb_test.go index c1f9f04812..1fd1f5477c 100644 --- a/core/state/statedb_test.go +++ b/core/state/statedb_test.go @@ -521,6 +521,7 @@ func TestMVHashMapReadWriteDelete(t *testing.T) { states[1].GetOrNewStateObject(addr) states[1].SetState(addr, key, val) states[1].SetBalance(addr, balance) + states[1].FlushMVWriteSet() // Tx1 read v = states[2].GetState(addr, key) @@ -547,6 +548,7 @@ func TestMVHashMapReadWriteDelete(t *testing.T) { states[3].Finalise(false) v = states[3].GetState(addr, key) assert.Equal(t, common.Hash{}, v) + states[3].FlushMVWriteSet() // Tx4 read v = states[4].GetState(addr, key) @@ -581,6 +583,7 @@ func TestMVHashMapRevert(t *testing.T) { states[0].GetOrNewStateObject(addr) states[0].SetState(addr, key, val) states[0].SetBalance(addr, balance) + states[0].FlushMVWriteSet() // Tx1 perform some ops and then revert snapshot := states[1].Snapshot() @@ -601,6 +604,7 @@ func TestMVHashMapRevert(t *testing.T) { assert.Equal(t, val, v) assert.Equal(t, balance, b) states[1].Finalise(false) + states[1].FlushMVWriteSet() // Tx2 check the state and balance v = states[2].GetState(addr, key) @@ -639,11 +643,13 @@ func TestMVHashMapMarkEstimate(t *testing.T) { states[0].SetState(addr, key, val) v = states[0].GetState(addr, key) assert.Equal(t, val, v) + states[0].FlushMVWriteSet() // Tx1 write states[1].GetOrNewStateObject(addr) states[1].SetState(addr, key, val) states[1].SetBalance(addr, balance) + states[1].FlushMVWriteSet() // Tx2 read v = states[2].GetState(addr, key) @@ -696,12 +702,14 @@ func TestMVHashMapOverwrite(t *testing.T) { states[0].GetOrNewStateObject(addr) states[0].SetState(addr, key, val1) states[0].SetBalance(addr, balance1) + states[0].FlushMVWriteSet() // Tx1 write states[1].SetState(addr, key, val2) states[1].SetBalance(addr, balance2) v := states[1].GetState(addr, key) b := states[1].GetBalance(addr) + states[1].FlushMVWriteSet() assert.Equal(t, val2, v) assert.Equal(t, balance2, b) @@ -774,14 +782,17 @@ func TestMVHashMapWriteNoConflict(t *testing.T) { // Tx0 write states[0].GetOrNewStateObject(addr) + states[0].FlushMVWriteSet() // Tx2 write states[2].SetState(addr, key2, val2) + states[2].FlushMVWriteSet() // Tx1 write tx1Snapshot := states[1].Snapshot() states[1].SetState(addr, key1, val1) states[1].SetBalance(addr, balance1) + states[1].FlushMVWriteSet() // Tx1 read assert.Equal(t, val1, states[1].GetState(addr, key1)) @@ -813,6 +824,7 @@ func TestMVHashMapWriteNoConflict(t *testing.T) { // Tx1 revert states[1].RevertToSnapshot(tx1Snapshot) + states[1].FlushMVWriteSet() assert.Equal(t, common.Hash{}, states[3].GetState(addr, key1)) assert.Equal(t, common.Hash{}, states[3].GetState(addr, key2)) @@ -853,6 +865,7 @@ func TestApplyMVWriteSet(t *testing.T) { addr1 := common.HexToAddress("0x01") addr2 := common.HexToAddress("0x02") + addr3 := common.HexToAddress("0x03") key1 := common.HexToHash("0x01") key2 := common.HexToHash("0x02") val1 := common.HexToHash("0x01") @@ -862,13 +875,19 @@ func TestApplyMVWriteSet(t *testing.T) { code := []byte{1, 2, 3} // Tx0 write + states[0].GetOrNewStateObject(addr1) states[0].SetState(addr1, key1, val1) states[0].SetBalance(addr1, balance1) states[0].SetState(addr2, key2, val2) + states[0].GetOrNewStateObject(addr3) + states[0].Finalise(false) + states[0].FlushMVWriteSet() + sSingleProcess.GetOrNewStateObject(addr1) sSingleProcess.SetState(addr1, key1, val1) sSingleProcess.SetBalance(addr1, balance1) sSingleProcess.SetState(addr2, key2, val2) + sSingleProcess.GetOrNewStateObject(addr3) sClean.ApplyMVWriteSet(states[0].MVWriteList()) @@ -878,6 +897,8 @@ func TestApplyMVWriteSet(t *testing.T) { states[1].SetState(addr1, key2, val2) states[1].SetBalance(addr1, balance2) states[1].SetNonce(addr1, 1) + states[1].Finalise(false) + states[1].FlushMVWriteSet() sSingleProcess.SetState(addr1, key2, val2) sSingleProcess.SetBalance(addr1, balance2) @@ -891,6 +912,8 @@ func TestApplyMVWriteSet(t *testing.T) { states[2].SetState(addr1, key1, val2) states[2].SetBalance(addr1, balance2) states[2].SetNonce(addr1, 2) + states[2].Finalise(false) + states[2].FlushMVWriteSet() sSingleProcess.SetState(addr1, key1, val2) sSingleProcess.SetBalance(addr1, balance2) @@ -903,6 +926,8 @@ func TestApplyMVWriteSet(t *testing.T) { // Tx3 write states[3].Suicide(addr2) states[3].SetCode(addr1, code) + states[3].Finalise(false) + states[3].FlushMVWriteSet() sSingleProcess.Suicide(addr2) sSingleProcess.SetCode(addr1, code)