Block-stm optimization

Added tests for executor and some improvements:

1. Add a dependency map during execution. This will prevent aborted tasks from being sent for execution immedaitely after failure.
2. Change the key of MVHashMap from string to a byte array. This will reduce time to convert byte slices to strings.
3. Use sync.Map to reduce the time spent in global mutex.
4. Skip applying intermediate states.
5. Estimate dependency when an execution fails without dependency information.
6. Divide execution task queue into two separate queues. One for relatively certain transactions, and the other for speculative future transactions.
7. Setting dependencies of Txs coming from the same sender before starting parallel execution.
8. Process results in their semantic order (transaction index) instead of the order when they arrive. Replace result channel with a priority queue.
This commit is contained in:
Jerry 2022-08-08 12:38:39 -07:00
parent 4507b2e057
commit c36ad88aec
12 changed files with 1473 additions and 383 deletions

119
core/blockstm/dag.go Normal file
View file

@ -0,0 +1,119 @@
package blockstm
import (
"fmt"
"sort"
"strings"
"github.com/heimdalr/dag"
"github.com/ethereum/go-ethereum/log"
)
type DAG struct {
*dag.DAG
}
func HasReadDep(txFrom TxnOutput, txTo TxnInput) bool {
reads := make(map[Key]bool)
for _, v := range txTo {
reads[v.Path] = true
}
for _, rd := range txFrom {
if _, ok := reads[rd.Path]; ok {
return true
}
}
return false
}
func BuildDAG(deps TxnInputOutput) (d DAG) {
d = DAG{dag.NewDAG()}
ids := make(map[int]string)
for i := len(deps.inputs) - 1; i > 0; i-- {
txTo := deps.inputs[i]
var txToId string
if _, ok := ids[i]; ok {
txToId = ids[i]
} else {
txToId, _ = d.AddVertex(i)
ids[i] = txToId
}
for j := i - 1; j >= 0; j-- {
txFrom := deps.allOutputs[j]
if HasReadDep(txFrom, txTo) {
var txFromId string
if _, ok := ids[j]; ok {
txFromId = ids[j]
} else {
txFromId, _ = d.AddVertex(j)
ids[j] = txFromId
}
err := d.AddEdge(txFromId, txToId)
if err != nil {
log.Warn("Failed to add edge", "from", txFromId, "to", txToId, "err", err)
}
break // once we add a 'backward' dep we can't execute before that transaction so no need to proceed
}
}
}
return
}
func (d DAG) Report(out func(string)) {
roots := make([]int, 0)
rootIds := make([]string, 0)
for k, i := range d.GetRoots() {
roots = append(roots, i.(int))
rootIds = append(rootIds, k)
}
sort.Ints(roots)
fmt.Println(roots)
makeStrs := func(ints []int) (ret []string) {
for _, v := range ints {
ret = append(ret, fmt.Sprint(v))
}
return
}
maxDesc := 0
maxDeps := 0
totalDeps := 0
for k, v := range roots {
ids := []int{v}
desc, _ := d.GetDescendants(rootIds[k])
for _, i := range desc {
ids = append(ids, i.(int))
}
sort.Ints(ids)
out(fmt.Sprintf("(%v) %v", len(ids), strings.Join(makeStrs(ids), "->")))
if len(desc) > maxDesc {
maxDesc = len(desc)
}
}
numTx := len(d.DAG.GetVertices())
out(fmt.Sprintf("max chain length: %v of %v (%v%%)", maxDesc+1, numTx,
fmt.Sprintf("%.1f", float64(maxDesc+1)*100.0/float64(numTx))))
out(fmt.Sprintf("max dep count: %v of %v (%v%%)", maxDeps, totalDeps,
fmt.Sprintf("%.1f", float64(maxDeps)*100.0/float64(totalDeps))))
}

View file

@ -1,8 +1,11 @@
package blockstm
import (
"container/heap"
"fmt"
"sort"
"sync"
"time"
"github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/log"
@ -22,6 +25,7 @@ type ExecTask interface {
MVWriteList() []WriteDescriptor
MVFullWriteList() []WriteDescriptor
Sender() common.Address
Settle()
}
type ExecVersionView struct {
@ -34,186 +38,362 @@ type ExecVersionView struct {
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")
type ErrExecAbortError struct {
Dependency int
}
const numGoProcs = 4
func (e ErrExecAbortError) Error() string {
if e.Dependency >= 0 {
return fmt.Sprintf("Execution aborted due to dependency %d", e.Dependency)
} else {
return "Execution aborted"
}
}
type IntHeap []int
func (h IntHeap) Len() int { return len(h) }
func (h IntHeap) Less(i, j int) bool { return h[i] < h[j] }
func (h IntHeap) Swap(i, j int) { h[i], h[j] = h[j], h[i] }
func (h *IntHeap) Push(x any) {
// Push and Pop use pointer receivers because they modify the slice's length,
// not just its contents.
*h = append(*h, x.(int))
}
func (h *IntHeap) Pop() any {
old := *h
n := len(old)
x := old[n-1]
*h = old[0 : n-1]
return x
}
// A thread safe priority queue
type SafePriorityQueue struct {
m sync.Mutex
queue *IntHeap
data map[int]interface{}
}
func NewSafePriorityQueue(capacity int) *SafePriorityQueue {
q := make(IntHeap, 0, capacity)
return &SafePriorityQueue{
m: sync.Mutex{},
queue: &q,
data: make(map[int]interface{}, capacity),
}
}
func (pq *SafePriorityQueue) Push(v int, d interface{}) {
pq.m.Lock()
heap.Push(pq.queue, v)
pq.data[v] = d
pq.m.Unlock()
}
func (pq *SafePriorityQueue) Pop() interface{} {
pq.m.Lock()
defer pq.m.Unlock()
v := heap.Pop(pq.queue).(int)
return pq.data[v]
}
func (pq *SafePriorityQueue) Len() int {
return pq.queue.Len()
}
type ParallelExecutionResult struct {
TxIO *TxnInputOutput
Stats *[][]uint64
Deps *DAG
}
const numGoProcs = 2
const numSpeculativeProcs = 16
// Max number of pre-validation to run per loop
const preValidateLimit = 5
// Max number of times a transaction (t) can be executed before its dependency is resolved to its previous tx (t-1)
const maxIncarnation = 2
// nolint: gocognit
func ExecuteParallel(tasks []ExecTask) (lastTxIO *TxnInputOutput, err error) {
// A stateless executor that executes transactions in parallel
func ExecuteParallel(tasks []ExecTask, profile bool) (ParallelExecutionResult, error) {
if len(tasks) == 0 {
return MakeTxnInputOutput(len(tasks)), nil
return ParallelExecutionResult{MakeTxnInputOutput(len(tasks)), nil, nil}, nil
}
// Stores the execution statistics for each task
stats := make([][]uint64, 0, len(tasks))
statsMutex := sync.Mutex{}
// Channel for tasks that should be prioritized
chTasks := make(chan ExecVersionView, len(tasks))
chResults := make(chan ExecResult, len(tasks))
chDone := make(chan bool)
mutMap := map[common.Address]*sync.RWMutex{}
for _, t := range tasks {
if _, ok := mutMap[t.Sender()]; !ok {
mutMap[t.Sender()] = &sync.RWMutex{}
}
// Channel for speculative tasks
chSpeculativeTasks := make(chan struct{}, len(tasks))
// A priority queue that stores speculative tasks
specTaskQueue := NewSafePriorityQueue(len(tasks))
// Channel to signal that the result of a transaction could be written to storage
chSettle := make(chan int, len(tasks))
// Channel to signal that a transaction has finished executing
chResults := make(chan struct{}, len(tasks))
// A priority queue that stores the transaction index of results, so we can validate the results in order
resultQueue := NewSafePriorityQueue(len(tasks))
// A wait group to wait for all settling tasks to finish
var settleWg sync.WaitGroup
// An integer that tracks the index of last settled transaction
lastSettled := -1
// For a task that runs only after all of its preceding tasks have finished and passed validation,
// its result will be absolutely valid and therefore its validation could be skipped.
// This map stores the boolean value indicating whether a task satisfy this condition ( absolutely valid).
skipCheck := make(map[int]bool)
for i := 0; i < len(tasks); i++ {
skipCheck[i] = false
}
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:
{
m := mutMap[task.sender]
if !m.TryLock() {
// why not this? -> chTasks <- task
t <- task
} else {
res := task.Execute()
chResults <- res
m.Unlock()
}
}
case <-chDone:
break Loop
}
}
log.Debug("blockstm", "proc done", procNum) // TODO: logging ...
}(i, chTasks)
}
mvh := MakeMVHashMap()
// Execution tasks stores the state of each execution task
execTasks := makeStatusManager(len(tasks))
// Validate tasks stores the state of each validation task
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, sender: tasks[tx].Sender()}
}
}
lastTxIO = MakeTxnInputOutput(len(tasks))
txIncarnations := make([]int, len(tasks))
// Stats for debugging purposes
var cntExec, cntSuccess, cntAbort, cntTotalValidations, cntValidationFail int
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))
}
// Initialize MVHashMap
mvh := MakeMVHashMap()
prevWrite := lastTxIO.AllWriteSet(res.ver.TxnIndex)
// Stores the inputs and outputs of the last incardanotion of all transactions
lastTxIO := MakeTxnInputOutput(len(tasks))
// Remove entries that were previously written but are no longer written
// Tracks the incarnation number of each transaction
txIncarnations := make([]int, len(tasks))
cmpMap := make(map[string]bool)
// A map that stores the estimated dependency of a transaction if it is aborted without any known dependency
estimateDeps := make(map[int][]int, len(tasks))
for _, w := range res.txAllOut {
cmpMap[string(w.Path)] = true
}
for i := 0; i < len(tasks); i++ {
estimateDeps[i] = make([]int, 0)
}
for _, v := range prevWrite {
if _, ok := cmpMap[string(v.Path)]; !ok {
mvh.Delete(v.Path, res.ver.TxnIndex)
}
}
// A map that records whether a transaction result has been speculatively validated
preValidated := make(map[int]bool, len(tasks))
lastTxIO.recordWrite(res.ver.TxnIndex, res.txOut)
lastTxIO.recordAllWrite(res.ver.TxnIndex, res.txAllOut)
begin := time.Now()
workerWg := sync.WaitGroup{}
workerWg.Add(numSpeculativeProcs + numGoProcs)
// Launch workers that execute transactions
for i := 0; i < numSpeculativeProcs+numGoProcs; i++ {
go func(procNum int) {
defer workerWg.Done()
doWork := func(task ExecVersionView) {
start := time.Duration(0)
if profile {
start = time.Since(begin)
}
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)
res := task.Execute()
if res.err == nil {
mvh.FlushMVWriteSet(res.txAllOut)
}
resultQueue.Push(res.ver.TxnIndex, res)
chResults <- struct{}{}
if profile {
end := time.Since(begin)
stat := []uint64{uint64(res.ver.TxnIndex), uint64(res.ver.Incarnation), uint64(start), uint64(end), uint64(procNum)}
statsMutex.Lock()
stats = append(stats, stat)
statsMutex.Unlock()
}
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 procNum < numSpeculativeProcs {
for range chSpeculativeTasks {
doWork(specTaskQueue.Pop().(ExecVersionView))
}
} else {
for task := range chTasks {
doWork(task)
}
}
}(i)
}
// Launch a worker that settles valid transactions
settleWg.Add(len(tasks))
go func() {
for t := range chSettle {
tasks[t].Settle()
settleWg.Done()
}
}()
// bootstrap first execution
tx := execTasks.takeNextPending()
if tx != -1 {
cntExec++
chTasks <- ExecVersionView{ver: Version{tx, 0}, et: tasks[tx], mvh: mvh, sender: tasks[tx].Sender()}
}
// Before starting execution, going through each task to check their explicit dependencies (whether they are coming from the same account)
prevSenderTx := make(map[common.Address]int)
for i, t := range tasks {
if tx, ok := prevSenderTx[t.Sender()]; ok {
execTasks.addDependencies(tx, i)
execTasks.clearPending(i)
}
// 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, sender: tasks[nextTx].Sender()}
prevSenderTx[t.Sender()] = i
}
var res ExecResult
var err error
// Start main validation loop
// nolint:nestif
for range chResults {
res = resultQueue.Pop().(ExecResult)
tx := res.ver.TxnIndex
if res.err == nil {
lastTxIO.recordRead(tx, res.txIn)
if res.ver.Incarnation == 0 {
lastTxIO.recordWrite(tx, res.txOut)
lastTxIO.recordAllWrite(tx, res.txAllOut)
} else {
if res.txAllOut.hasNewWrite(lastTxIO.AllWriteSet(tx)) {
validateTasks.pushPendingSet(execTasks.getRevalidationRange(tx + 1))
}
prevWrite := lastTxIO.AllWriteSet(tx)
// Remove entries that were previously written but are no longer written
cmpMap := make(map[Key]bool)
for _, w := range res.txAllOut {
cmpMap[w.Path] = true
}
for _, v := range prevWrite {
if _, ok := cmpMap[v.Path]; !ok {
mvh.Delete(v.Path, tx)
}
}
lastTxIO.recordWrite(tx, res.txOut)
lastTxIO.recordAllWrite(tx, res.txAllOut)
}
validateTasks.pushPending(tx)
execTasks.markComplete(tx)
diagExecSuccess[tx]++
cntSuccess++
execTasks.removeDependency(tx)
} else if execErr, ok := res.err.(ErrExecAbortError); ok {
addedDependencies := false
if execErr.Dependency >= 0 {
l := len(estimateDeps[tx])
for l > 0 && estimateDeps[tx][l-1] > execErr.Dependency {
execTasks.removeDependency(estimateDeps[tx][l-1])
estimateDeps[tx] = estimateDeps[tx][:l-1]
l--
}
if txIncarnations[tx] < maxIncarnation {
addedDependencies = execTasks.addDependencies(execErr.Dependency, tx)
} else {
addedDependencies = execTasks.addDependencies(tx-1, tx)
}
} else {
estimate := 0
if len(estimateDeps[tx]) > 0 {
estimate = estimateDeps[tx][len(estimateDeps[tx])-1]
}
addedDependencies = execTasks.addDependencies(estimate, tx)
newEstimate := estimate + (estimate+tx)/2
if newEstimate >= tx {
newEstimate = tx - 1
}
estimateDeps[tx] = append(estimateDeps[tx], newEstimate)
}
execTasks.clearInProgress(tx)
if !addedDependencies {
execTasks.pushPending(tx)
}
txIncarnations[tx]++
diagExecAbort[tx]++
cntAbort++
} else {
err = res.err
break
}
// 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 validateTasks.minPending() <= maxComplete && validateTasks.minPending() >= 0 {
toValidate = append(toValidate, validateTasks.takeNextPending())
}
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)
if skipCheck[tx] || ValidateVersion(tx, lastTxIO, mvh) {
validateTasks.markComplete(tx)
} else {
log.Debug("blockstm", "* validation task FAILED", tx)
cntValidationFail++
diagExecAbort[tx]++
for _, v := range lastTxIO.AllWriteSet(tx) {
@ -222,38 +402,138 @@ func ExecuteParallel(tasks []ExecTask) (lastTxIO *TxnInputOutput, err error) {
// '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]++
addedDependencies := false
if txIncarnations[tx] >= maxIncarnation {
addedDependencies = execTasks.addDependencies(tx-1, tx)
}
execTasks.clearComplete(tx)
if !addedDependencies {
execTasks.pushPending(tx)
}
preValidated[tx] = false
txIncarnations[tx]++
}
}
// if we didn't queue work previously, do check again so we keep making progress ...
if nextTx == -1 {
nextTx = execTasks.takeNextPending()
preValidateCount := 0
invalidated := []int{}
i := sort.SearchInts(validateTasks.pending, maxComplete+1)
for i < len(validateTasks.pending) && preValidateCount < preValidateLimit {
tx := validateTasks.pending[i]
if !preValidated[tx] {
cntTotalValidations++
if !ValidateVersion(tx, lastTxIO, mvh) {
cntValidationFail++
diagExecAbort[tx]++
invalidated = append(invalidated, tx)
if execTasks.checkComplete(tx) {
execTasks.clearComplete(tx)
}
if !execTasks.checkInProgress(tx) {
for _, v := range lastTxIO.AllWriteSet(tx) {
mvh.MarkEstimate(v.Path, tx)
}
validateTasks.pushPendingSet(execTasks.getRevalidationRange(tx + 1))
addedDependencies := false
if txIncarnations[tx] >= maxIncarnation {
addedDependencies = execTasks.addDependencies(tx-1, tx)
}
if !addedDependencies {
execTasks.pushPending(tx)
}
}
txIncarnations[tx]++
preValidated[tx] = false
} else {
preValidated[tx] = true
}
preValidateCount++
}
i++
}
for _, tx := range invalidated {
validateTasks.clearPending(tx)
}
// Settle transactions that have been validated to be correct and that won't be re-executed again
maxValidated := validateTasks.maxAllComplete()
for lastSettled < maxValidated {
lastSettled++
if execTasks.checkInProgress(lastSettled) || execTasks.checkPending(lastSettled) || execTasks.blockCount[lastSettled] >= 0 {
lastSettled--
break
}
chSettle <- lastSettled
}
if validateTasks.countComplete() == len(tasks) && execTasks.countComplete() == len(tasks) {
log.Debug("blockstm exec summary", "execs", cntExec, "success", cntSuccess, "aborts", cntAbort, "validations", cntTotalValidations, "failures", cntValidationFail, "#tasks/#execs", fmt.Sprintf("%.2f%%", float64(len(tasks))/float64(cntExec)*100))
break
}
// Send the next immediate pending transaction to be executed
if execTasks.minPending() != -1 && execTasks.minPending() == maxValidated+1 {
nextTx := execTasks.takeNextPending()
if nextTx != -1 {
cntExec++
log.Debug("blockstm", "# tx queued up", nextTx)
skipCheck[nextTx] = true
chTasks <- ExecVersionView{ver: Version{nextTx, txIncarnations[nextTx]}, et: tasks[nextTx], mvh: mvh, sender: tasks[nextTx].Sender()}
}
}
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
// Send speculative tasks
for execTasks.peekPendingGE(maxValidated+3) != -1 || len(execTasks.inProgress) == 0 {
// We skip the next transaction to avoid the case where they all have conflicts and could not be
// scheduled for re-execution immediately even when it's their time to run, because they are already in
// speculative queue.
nextTx := execTasks.takePendingGE(maxValidated + 3)
if nextTx == -1 {
nextTx = execTasks.takeNextPending()
}
if nextTx != -1 {
cntExec++
task := ExecVersionView{ver: Version{nextTx, txIncarnations[nextTx]}, et: tasks[nextTx], mvh: mvh, sender: tasks[nextTx].Sender()}
specTaskQueue.Push(nextTx, task)
chSpeculativeTasks <- struct{}{}
}
}
}
for i := 0; i < numGoProcs; i++ {
chDone <- true
}
close(chTasks)
close(chSpeculativeTasks)
workerWg.Wait()
close(chResults)
settleWg.Wait()
close(chSettle)
return
var dag DAG
if profile {
dag = BuildDAG(*lastTxIO)
}
return ParallelExecutionResult{lastTxIO, &stats, &dag}, err
}

View file

@ -0,0 +1,470 @@
package blockstm
import (
"fmt"
"math/big"
"math/rand"
"os"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/log"
)
type OpType int
const readType = 0
const writeType = 1
const otherType = 2
type Op struct {
key Key
duration time.Duration
opType OpType
val int
}
type testExecTask struct {
txIdx int
ops []Op
readMap map[Key]ReadDescriptor
writeMap map[Key]WriteDescriptor
sender common.Address
nonce int
}
type PathGenerator func(addr common.Address, j int, total int) Key
type TaskRunner func(numTx int, numRead int, numWrite int, numNonIO int) (time.Duration, time.Duration)
type Timer func(txIdx int, opIdx int) time.Duration
type Sender func(int) common.Address
func NewTestExecTask(txIdx int, ops []Op, sender common.Address, nonce int) *testExecTask {
return &testExecTask{
txIdx: txIdx,
ops: ops,
readMap: make(map[Key]ReadDescriptor),
writeMap: make(map[Key]WriteDescriptor),
sender: sender,
nonce: nonce,
}
}
func sleep(i time.Duration) {
start := time.Now()
for time.Since(start) < i {
}
}
func (t *testExecTask) Execute(mvh *MVHashMap, incarnation int) error {
// Sleep for 50 microsecond to simulate setup time
sleep(time.Microsecond * 50)
version := Version{TxnIndex: t.txIdx, Incarnation: incarnation}
t.readMap = make(map[Key]ReadDescriptor)
t.writeMap = make(map[Key]WriteDescriptor)
deps := -1
for i, op := range t.ops {
k := op.key
switch op.opType {
case readType:
if _, ok := t.writeMap[k]; ok {
sleep(op.duration)
continue
}
result := mvh.Read(k, t.txIdx)
val := result.Value()
if i == 0 && val != nil && (val.(int) != t.nonce) {
return ErrExecAbortError{}
}
if result.Status() == MVReadResultDependency {
if result.depIdx > deps {
deps = result.depIdx
}
}
var readKind int
if result.Status() == MVReadResultDone {
readKind = ReadKindMap
} else if result.Status() == MVReadResultNone {
readKind = ReadKindStorage
}
sleep(op.duration)
t.readMap[k] = ReadDescriptor{k, readKind, Version{TxnIndex: result.depIdx, Incarnation: result.incarnation}}
case writeType:
t.writeMap[k] = WriteDescriptor{k, version, op.val}
case otherType:
sleep(op.duration)
default:
panic(fmt.Sprintf("Unknown op type: %d", op.opType))
}
}
if deps != -1 {
return ErrExecAbortError{deps}
}
return nil
}
func (t *testExecTask) MVWriteList() []WriteDescriptor {
return t.MVFullWriteList()
}
func (t *testExecTask) MVFullWriteList() []WriteDescriptor {
writes := make([]WriteDescriptor, 0, len(t.writeMap))
for _, v := range t.writeMap {
writes = append(writes, v)
}
return writes
}
func (t *testExecTask) MVReadList() []ReadDescriptor {
reads := make([]ReadDescriptor, 0, len(t.readMap))
for _, v := range t.readMap {
reads = append(reads, v)
}
return reads
}
func (t *testExecTask) Settle() {}
func (t *testExecTask) Sender() common.Address {
return t.sender
}
func randTimeGenerator(min time.Duration, max time.Duration) func(txIdx int, opIdx int) time.Duration {
return func(txIdx int, opIdx int) time.Duration {
return time.Duration(rand.Int63n(int64(max-min))) + min
}
}
func longTailTimeGenerator(min time.Duration, max time.Duration, i int, j int) func(txIdx int, opIdx int) time.Duration {
return func(txIdx int, opIdx int) time.Duration {
if txIdx%i == 0 && opIdx == j {
return max * 100
} else {
return time.Duration(rand.Int63n(int64(max-min))) + min
}
}
}
var randomPathGenerator = func(sender common.Address, j int, total int) Key {
return NewStateKey(sender, common.BigToHash((big.NewInt(int64(total)))))
}
var dexPathGenerator = func(sender common.Address, j int, total int) Key {
if j == total-1 || j == 2 {
return NewSubpathKey(common.BigToAddress(big.NewInt(int64(0))), 1)
} else {
return NewSubpathKey(common.BigToAddress(big.NewInt(int64(j))), 1)
}
}
var readTime = randTimeGenerator(4*time.Microsecond, 12*time.Microsecond)
var writeTime = randTimeGenerator(2*time.Microsecond, 6*time.Microsecond)
var nonIOTime = randTimeGenerator(1*time.Microsecond, 2*time.Microsecond)
func taskFactory(numTask int, sender Sender, readsPerT int, writesPerT int, nonIOPerT int, pathGenerator PathGenerator, readTime Timer, writeTime Timer, nonIOTime Timer) ([]ExecTask, time.Duration) {
exec := make([]ExecTask, 0, numTask)
var serialDuration time.Duration
senderNonces := make(map[common.Address]int)
for i := 0; i < numTask; i++ {
s := sender(i)
// Set first two ops to always read and write nonce
ops := make([]Op, 0, readsPerT+writesPerT+nonIOPerT)
ops = append(ops, Op{opType: readType, key: NewSubpathKey(s, 2), duration: readTime(i, 0), val: senderNonces[s]})
senderNonces[s]++
ops = append(ops, Op{opType: writeType, key: NewSubpathKey(s, 2), duration: writeTime(i, 1), val: senderNonces[s]})
for j := 0; j < readsPerT-1; j++ {
ops = append(ops, Op{opType: readType})
}
for j := 0; j < nonIOPerT; j++ {
ops = append(ops, Op{opType: otherType})
}
for j := 0; j < writesPerT-1; j++ {
ops = append(ops, Op{opType: writeType})
}
// shuffle ops except for the first three (read nonce, write nonce, another read) ops and last write op.
// This enables random path generator to generate deterministic paths for these "special" ops.
for j := 3; j < len(ops)-1; j++ {
k := rand.Intn(len(ops)-j-1) + j
ops[j], ops[k] = ops[k], ops[j]
}
// Generate time and key path for each op except first two that are always read and write nonce
for j := 2; j < len(ops); j++ {
if ops[j].opType == readType {
ops[j].key = pathGenerator(s, j, len(ops))
ops[j].duration = readTime(i, j)
} else if ops[j].opType == writeType {
ops[j].key = pathGenerator(s, j, len(ops))
ops[j].duration = writeTime(i, j)
} else {
ops[j].duration = nonIOTime(i, j)
}
serialDuration += ops[j].duration
}
if ops[len(ops)-1].opType != writeType {
panic("Last op must be a write")
}
t := NewTestExecTask(i, ops, s, senderNonces[s]-1)
exec = append(exec, t)
}
return exec, serialDuration
}
func testExecutorComb(t *testing.T, totalTxs []int, numReads []int, numWrites []int, numNonIO []int, taskRunner TaskRunner) {
t.Helper()
log.Root().SetHandler(log.LvlFilterHandler(log.LvlDebug, log.StreamHandler(os.Stderr, log.TerminalFormat(false))))
improved := 0
total := 0
totalExecDuration := time.Duration(0)
totalSerialDuration := time.Duration(0)
for _, numTx := range totalTxs {
for _, numRead := range numReads {
for _, numWrite := range numWrites {
for _, numNonIO := range numNonIO {
log.Info("Executing block", "numTx", numTx, "numRead", numRead, "numWrite", numWrite, "numNonIO", numNonIO)
execDuration, expectedSerialDuration := taskRunner(numTx, numRead, numWrite, numNonIO)
if execDuration < expectedSerialDuration {
improved++
}
total++
performance := "✅"
if execDuration >= expectedSerialDuration {
performance = "❌"
}
fmt.Printf("exec duration %v, serial duration %v, time reduced %v %.2f%%, %v \n", execDuration, expectedSerialDuration, expectedSerialDuration-execDuration, float64(expectedSerialDuration-execDuration)/float64(expectedSerialDuration)*100, performance)
totalExecDuration += execDuration
totalSerialDuration += expectedSerialDuration
}
}
}
}
fmt.Println("Improved: ", improved, "Total: ", total, "success rate: ", float64(improved)/float64(total)*100)
fmt.Printf("Total exec duration: %v, total serial duration: %v, time reduced: %v, time reduced percent: %.2f%%\n", totalExecDuration, totalSerialDuration, totalSerialDuration-totalExecDuration, float64(totalSerialDuration-totalExecDuration)/float64(totalSerialDuration)*100)
}
func runParallel(t *testing.T, tasks []ExecTask, validation func(TxnInputOutput) bool) time.Duration {
t.Helper()
start := time.Now()
results, _ := ExecuteParallel(tasks, false)
txio := results.TxIO
// Need to apply the final write set to storage
finalWriteSet := make(map[Key]time.Duration)
for _, task := range tasks {
task := task.(*testExecTask)
for _, op := range task.ops {
if op.opType == writeType {
finalWriteSet[op.key] = op.duration
}
}
}
for _, v := range finalWriteSet {
sleep(v)
}
duration := time.Since(start)
if validation != nil {
assert.True(t, validation(*txio))
}
return duration
}
func TestLessConflicts(t *testing.T) {
t.Parallel()
rand.Seed(0)
totalTxs := []int{10, 50, 100, 200, 300}
numReads := []int{20, 100, 200}
numWrites := []int{20, 100, 200}
numNonIO := []int{100, 500}
taskRunner := func(numTx int, numRead int, numWrite int, numNonIO int) (time.Duration, time.Duration) {
sender := func(i int) common.Address {
randomness := rand.Intn(10) + 10
return common.BigToAddress(big.NewInt(int64(i % randomness)))
}
tasks, serialDuration := taskFactory(numTx, sender, numRead, numWrite, numNonIO, randomPathGenerator, readTime, writeTime, nonIOTime)
return runParallel(t, tasks, nil), serialDuration
}
testExecutorComb(t, totalTxs, numReads, numWrites, numNonIO, taskRunner)
}
func TestAlternatingTx(t *testing.T) {
t.Parallel()
rand.Seed(0)
totalTxs := []int{200}
numReads := []int{20}
numWrites := []int{20}
numNonIO := []int{100}
taskRunner := func(numTx int, numRead int, numWrite int, numNonIO int) (time.Duration, time.Duration) {
sender := func(i int) common.Address { return common.BigToAddress(big.NewInt(int64(i % 2))) }
tasks, serialDuration := taskFactory(numTx, sender, numRead, numWrite, numNonIO, randomPathGenerator, readTime, writeTime, nonIOTime)
return runParallel(t, tasks, nil), serialDuration
}
testExecutorComb(t, totalTxs, numReads, numWrites, numNonIO, taskRunner)
}
func TestMoreConflicts(t *testing.T) {
t.Parallel()
rand.Seed(0)
totalTxs := []int{10, 50, 100, 200, 300}
numReads := []int{20, 100, 200}
numWrites := []int{20, 100, 200}
numNonIO := []int{100, 500}
taskRunner := func(numTx int, numRead int, numWrite int, numNonIO int) (time.Duration, time.Duration) {
sender := func(i int) common.Address {
randomness := rand.Intn(10) + 10
return common.BigToAddress(big.NewInt(int64(i / randomness)))
}
tasks, serialDuration := taskFactory(numTx, sender, numRead, numWrite, numNonIO, randomPathGenerator, readTime, writeTime, nonIOTime)
return runParallel(t, tasks, nil), serialDuration
}
testExecutorComb(t, totalTxs, numReads, numWrites, numNonIO, taskRunner)
}
func TestRandomTx(t *testing.T) {
t.Parallel()
rand.Seed(0)
totalTxs := []int{10, 50, 100, 200, 300}
numReads := []int{20, 100, 200}
numWrites := []int{20, 100, 200}
numNonIO := []int{100, 500}
taskRunner := func(numTx int, numRead int, numWrite int, numNonIO int) (time.Duration, time.Duration) {
// Randomly assign this tx to one of 10 senders
sender := func(i int) common.Address { return common.BigToAddress(big.NewInt(int64(rand.Intn(10)))) }
tasks, serialDuration := taskFactory(numTx, sender, numRead, numWrite, numNonIO, randomPathGenerator, readTime, writeTime, nonIOTime)
return runParallel(t, tasks, nil), serialDuration
}
testExecutorComb(t, totalTxs, numReads, numWrites, numNonIO, taskRunner)
}
func TestTxWithLongTailRead(t *testing.T) {
t.Parallel()
rand.Seed(0)
totalTxs := []int{10, 50, 100, 200, 300}
numReads := []int{20, 100, 200}
numWrites := []int{20, 100, 200}
numNonIO := []int{100, 500}
taskRunner := func(numTx int, numRead int, numWrite int, numNonIO int) (time.Duration, time.Duration) {
sender := func(i int) common.Address {
randomness := rand.Intn(10) + 10
return common.BigToAddress(big.NewInt(int64(i / randomness)))
}
longTailReadTimer := longTailTimeGenerator(4*time.Microsecond, 12*time.Microsecond, 7, 10)
tasks, serialDuration := taskFactory(numTx, sender, numRead, numWrite, numNonIO, randomPathGenerator, longTailReadTimer, writeTime, nonIOTime)
return runParallel(t, tasks, nil), serialDuration
}
testExecutorComb(t, totalTxs, numReads, numWrites, numNonIO, taskRunner)
}
func TestDexScenario(t *testing.T) {
t.Parallel()
rand.Seed(0)
totalTxs := []int{10, 50, 100, 200, 300}
numReads := []int{20, 100, 200}
numWrites := []int{20, 100, 200}
numNonIO := []int{100, 500}
validation := func(txio TxnInputOutput) bool {
for i, inputs := range txio.inputs {
foundDep := false
for _, input := range inputs {
if input.V.TxnIndex == i-1 {
foundDep = true
}
}
if !foundDep {
return false
}
}
return true
}
taskRunner := func(numTx int, numRead int, numWrite int, numNonIO int) (time.Duration, time.Duration) {
sender := func(i int) common.Address { return common.BigToAddress(big.NewInt(int64(i))) }
tasks, serialDuration := taskFactory(numTx, sender, numRead, numWrite, numNonIO, dexPathGenerator, readTime, writeTime, nonIOTime)
return runParallel(t, tasks, validation), serialDuration
}
testExecutorComb(t, totalTxs, numReads, numWrites, numNonIO, taskRunner)
}

View file

@ -6,22 +6,79 @@ import (
"github.com/emirpasic/gods/maps/treemap"
"github.com/ethereum/go-ethereum/log"
"github.com/ethereum/go-ethereum/common"
)
const FlagDone = 0
const FlagEstimate = 1
const addressType = 1
const stateType = 2
const subpathType = 3
const KeyLength = common.AddressLength + common.HashLength + 2
type Key [KeyLength]byte
func (k *Key) IsAddress() bool {
return k[KeyLength-1] == addressType
}
func (k *Key) IsState() bool {
return k[KeyLength-1] == stateType
}
func (k *Key) IsSubpath() bool {
return k[KeyLength-1] == subpathType
}
func (k *Key) GetAddress() common.Address {
return common.BytesToAddress(k[:common.AddressLength])
}
func (k *Key) GetStateKey() common.Hash {
return common.BytesToHash(k[common.AddressLength : KeyLength-2])
}
func (k *Key) GetSubpath() byte {
return k[KeyLength-2]
}
func newKey(addr common.Address, hash common.Hash, subpath byte, keyType byte) Key {
var k Key
copy(k[:common.AddressLength], addr.Bytes())
copy(k[common.AddressLength:KeyLength-2], hash.Bytes())
k[KeyLength-2] = subpath
k[KeyLength-1] = keyType
return k
}
func NewAddressKey(addr common.Address) Key {
return newKey(addr, common.Hash{}, 0, addressType)
}
func NewStateKey(addr common.Address, hash common.Hash) Key {
k := newKey(addr, hash, 0, stateType)
if !k.IsState() {
panic(fmt.Errorf("key is not a state key"))
}
return k
}
func NewSubpathKey(addr common.Address, subpath byte) Key {
return newKey(addr, common.Hash{}, subpath, subpathType)
}
type MVHashMap struct {
rw sync.RWMutex
m map[string]*TxnIndexCells // TODO: might want a more efficient key representation
m sync.Map
s sync.Map
}
func MakeMVHashMap() *MVHashMap {
return &MVHashMap{
rw: sync.RWMutex{},
m: make(map[string]*TxnIndexCells),
}
return &MVHashMap{}
}
type WriteCell struct {
@ -40,80 +97,86 @@ type Version struct {
Incarnation int
}
func (mv *MVHashMap) getKeyCells(k []byte, fNoKey func(kenc string) *TxnIndexCells) (cells *TxnIndexCells) {
kenc := string(k)
var ok bool
mv.rw.RLock()
cells, ok = mv.m[kenc]
mv.rw.RUnlock()
func (mv *MVHashMap) getKeyCells(k Key, fNoKey func(kenc Key) *TxnIndexCells) (cells *TxnIndexCells) {
val, ok := mv.m.Load(k)
if !ok {
cells = fNoKey(kenc)
cells = fNoKey(k)
} else {
cells = val.(*TxnIndexCells)
}
return
}
func (mv *MVHashMap) Write(k []byte, v Version, data interface{}) {
cells := mv.getKeyCells(k, func(kenc string) (cells *TxnIndexCells) {
func (mv *MVHashMap) Write(k Key, v Version, data interface{}) {
cells := mv.getKeyCells(k, func(kenc Key) (cells *TxnIndexCells) {
n := &TxnIndexCells{
rw: sync.RWMutex{},
tm: treemap.NewWithIntComparator(),
}
var ok bool
mv.rw.Lock()
if cells, ok = mv.m[kenc]; !ok {
mv.m[kenc] = n
cells = n
}
mv.rw.Unlock()
cells = n
val, _ := mv.m.LoadOrStore(kenc, n)
cells = val.(*TxnIndexCells)
return
})
// TODO: could probably have a scheme where this only generally requires a read lock since any given transaction transaction
// should only have one incarnation executing at a time...
cells.rw.Lock()
defer cells.rw.Unlock()
cells.rw.RLock()
ci, ok := cells.tm.Get(v.TxnIndex)
cells.rw.RUnlock()
if ok {
if ci.(*WriteCell).incarnation > v.Incarnation {
panic(fmt.Errorf("existing transaction value does not have lower incarnation: %v, %v",
string(k), v.TxnIndex))
} else if ci.(*WriteCell).flag == FlagEstimate {
log.Debug("mvhashmap marking previous estimate as done", "tx index", v.TxnIndex, "incarnation", v.Incarnation)
k, v.TxnIndex))
}
ci.(*WriteCell).flag = FlagDone
ci.(*WriteCell).incarnation = v.Incarnation
ci.(*WriteCell).data = data
} else {
cells.tm.Put(v.TxnIndex, &WriteCell{
flag: FlagDone,
incarnation: v.Incarnation,
data: data,
})
cells.rw.Lock()
if ci, ok = cells.tm.Get(v.TxnIndex); !ok {
cells.tm.Put(v.TxnIndex, &WriteCell{
flag: FlagDone,
incarnation: v.Incarnation,
data: data,
})
} else {
ci.(*WriteCell).flag = FlagDone
ci.(*WriteCell).incarnation = v.Incarnation
ci.(*WriteCell).data = data
}
cells.rw.Unlock()
}
}
func (mv *MVHashMap) MarkEstimate(k []byte, txIdx int) {
cells := mv.getKeyCells(k, func(_ string) *TxnIndexCells {
func (mv *MVHashMap) ReadStorage(k Key, fallBack func() any) any {
data, ok := mv.s.Load(string(k[:]))
if !ok {
data = fallBack()
data, _ = mv.s.LoadOrStore(string(k[:]), data)
}
return data
}
func (mv *MVHashMap) MarkEstimate(k Key, txIdx int) {
cells := mv.getKeyCells(k, func(_ Key) *TxnIndexCells {
panic(fmt.Errorf("path must already exist"))
})
cells.rw.RLock()
if ci, ok := cells.tm.Get(txIdx); !ok {
panic("should not happen - cell should be present for path")
panic(fmt.Sprintf("should not happen - cell should be present for path. TxIdx: %v, path, %x, cells keys: %v", txIdx, k, cells.tm.Keys()))
} else {
ci.(*WriteCell).flag = FlagEstimate
}
cells.rw.RUnlock()
}
func (mv *MVHashMap) Delete(k []byte, txIdx int) {
cells := mv.getKeyCells(k, func(_ string) *TxnIndexCells {
func (mv *MVHashMap) Delete(k Key, txIdx int) {
cells := mv.getKeyCells(k, func(_ Key) *TxnIndexCells {
panic(fmt.Errorf("path must already exist"))
})
@ -158,11 +221,11 @@ func (mvr MVReadResult) Status() int {
return MVReadResultNone
}
func (mv *MVHashMap) Read(k []byte, txIdx int) (res MVReadResult) {
func (mv *MVHashMap) Read(k Key, txIdx int) (res MVReadResult) {
res.depIdx = -1
res.incarnation = -1
cells := mv.getKeyCells(k, func(_ string) *TxnIndexCells {
cells := mv.getKeyCells(k, func(_ Key) *TxnIndexCells {
return nil
})
if cells == nil {
@ -170,9 +233,10 @@ func (mv *MVHashMap) Read(k []byte, txIdx int) (res MVReadResult) {
}
cells.rw.RLock()
defer cells.rw.RUnlock()
fk, fv := cells.tm.Floor(txIdx - 1)
cells.rw.RUnlock()
if fk, fv := cells.tm.Floor(txIdx - 1); fk != nil && fv != nil {
if fk != nil && fv != nil {
c := fv.(*WriteCell)
switch c.flag {
case FlagEstimate:

View file

@ -11,6 +11,13 @@ func makeStatusManager(numTasks int) (t taskStatusManager) {
t.pending[i] = i
}
t.dependency = make(map[int]map[int]bool, numTasks)
t.blockCount = make(map[int]int, numTasks)
for i := 0; i < numTasks; i++ {
t.blockCount[i] = -1
}
return
}
@ -18,6 +25,8 @@ type taskStatusManager struct {
pending []int
inProgress []int
complete []int
dependency map[int]map[int]bool
blockCount map[int]int
}
func insertInList(l []int, v int) []int {
@ -47,6 +56,35 @@ func (m *taskStatusManager) takeNextPending() int {
return x
}
func (m *taskStatusManager) peekPendingGE(n int) int {
x := sort.SearchInts(m.pending, n)
if x >= len(m.pending) {
return -1
}
return m.pending[x]
}
// Take a pending task whose transaction index is greater than or equal to the given tx index
func (m *taskStatusManager) takePendingGE(n int) int {
x := sort.SearchInts(m.pending, n)
if x >= len(m.pending) {
return -1
}
v := m.pending[x]
if x < len(m.pending)-1 {
m.pending = append(m.pending[:x], m.pending[x+1:]...)
} else {
m.pending = m.pending[:x]
}
m.inProgress = insertInList(m.inProgress, v)
return v
}
func hasNoGap(l []int) bool {
return l[0]+len(l) == l[len(l)-1]+1
}
@ -68,7 +106,11 @@ func (m taskStatusManager) maxAllComplete() int {
}
func (m *taskStatusManager) pushPending(tx int) {
m.pending = insertInList(m.pending, tx)
if !m.checkComplete(tx) && !m.checkInProgress(tx) {
m.pending = insertInList(m.pending, tx)
} else {
panic(fmt.Errorf("should not happen - clear complete or inProgress before pushing pending"))
}
}
func removeFromList(l []int, v int, expect bool) []int {
@ -108,19 +150,52 @@ 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) addDependencies(blocker int, dependent int) bool {
if blocker < 0 || blocker >= dependent {
return false
}
curBlocker := m.blockCount[dependent]
if curBlocker > blocker {
return true
}
if m.checkComplete(blocker) {
// Blocking blocker has already completed
m.blockCount[dependent] = -1
return false
}
if _, ok := m.dependency[blocker]; !ok {
m.dependency[blocker] = make(map[int]bool)
}
m.dependency[blocker][dependent] = true
m.blockCount[dependent] = blocker
return true
}
func (m *taskStatusManager) removeDependency(tx int) {
if deps, ok := m.dependency[tx]; ok && len(deps) > 0 {
for k := range deps {
if m.blockCount[k] == tx {
m.blockCount[k] = -1
if !m.checkComplete(k) && !m.checkPending(k) && !m.checkInProgress(k) {
m.pushPending(k)
}
}
}
delete(m.dependency, 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 {
@ -139,8 +214,18 @@ func (m *taskStatusManager) checkPending(tx int) bool {
return false
}
func (m *taskStatusManager) checkComplete(tx int) bool {
x := sort.SearchInts(m.complete, tx)
if x < len(m.complete) && m.complete[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
//
// '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++ {
@ -154,10 +239,20 @@ func (m *taskStatusManager) getRevalidationRange(txFrom int) (ret []int) {
func (m *taskStatusManager) pushPendingSet(set []int) {
for _, v := range set {
m.pushPending(v)
if m.checkComplete(v) {
m.clearComplete(v)
}
if !m.checkInProgress(v) {
m.pushPending(v)
}
}
}
func (m *taskStatusManager) clearComplete(tx int) {
m.complete = removeFromList(m.complete, tx, false)
}
func (m *taskStatusManager) clearPending(tx int) {
m.pending = removeFromList(m.pending, tx, false)
}

View file

@ -1,21 +1,18 @@
//nolint: unused
package blockstm
import "encoding/base64"
const (
ReadKindMap = 0
ReadKindStorage = 1
)
type ReadDescriptor struct {
Path []byte
Path Key
Kind int
V Version
}
type WriteDescriptor struct {
Path []byte
Path Key
V Version
Val interface{}
}
@ -31,14 +28,14 @@ func (txo TxnOutput) hasNewWrite(cmpSet []WriteDescriptor) bool {
return true
}
cmpMap := map[string]bool{base64.StdEncoding.EncodeToString(cmpSet[0].Path): true}
cmpMap := map[Key]bool{cmpSet[0].Path: true}
for i := 1; i < len(cmpSet); i++ {
cmpMap[base64.StdEncoding.EncodeToString(cmpSet[i].Path)] = true
cmpMap[cmpSet[i].Path] = true
}
for _, v := range txo {
if !cmpMap[base64.StdEncoding.EncodeToString(v.Path)] {
if !cmpMap[v.Path] {
return true
}
}

View file

@ -19,6 +19,7 @@ package core
import (
"fmt"
"math/big"
"time"
"github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/consensus"
@ -58,16 +59,21 @@ type ExecutionTask struct {
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.
finalStateDB *state.StateDB // The final statedb.
header *types.Header
blockChain *BlockChain
evmConfig vm.Config
result *ExecutionResult
shouldDelayFeeCal *bool
shouldRerunWithoutFeeDelay bool
sender common.Address
totalUsedGas *uint64
receipts *types.Receipts
allLogs *[]*types.Log
}
func (task *ExecutionTask) Execute(mvh *blockstm.MVHashMap, incarnation int) (err error) {
@ -76,7 +82,9 @@ func (task *ExecutionTask) Execute(mvh *blockstm.MVHashMap, incarnation int) (er
task.statedb.SetMVHashmap(mvh)
task.statedb.SetIncarnation(incarnation)
evm := vm.NewEVM(task.blockContext, vm.TxContext{}, task.statedb, task.config, task.evmConfig)
blockContext := NewEVMBlockContext(task.header, task.blockChain, nil)
evm := vm.NewEVM(blockContext, vm.TxContext{}, task.statedb, task.config, task.evmConfig)
// Create a new context to be used in the EVM environment.
txContext := NewEVMTxContext(task.msg)
@ -85,9 +93,9 @@ func (task *ExecutionTask) Execute(mvh *blockstm.MVHashMap, incarnation int) (er
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)
log.Debug("Recovered from EVM failure.", "Error:", r)
err = blockstm.ErrExecAbort
err = blockstm.ErrExecAbortError{Dependency: task.statedb.DepTxIndex()}
return
}
@ -97,11 +105,21 @@ func (task *ExecutionTask) Execute(mvh *blockstm.MVHashMap, incarnation int) (er
if *task.shouldDelayFeeCal {
task.result, err = ApplyMessageNoFeeBurnOrTip(evm, task.msg, new(GasPool).AddGas(task.gasLimit))
if _, ok := task.statedb.MVReadMap()[string(task.blockContext.Coinbase.Bytes())]; ok {
if task.result == nil || err != nil {
return blockstm.ErrExecAbortError{Dependency: task.statedb.DepTxIndex()}
}
reads := task.statedb.MVReadMap()
if _, ok := reads[blockstm.NewSubpathKey(blockContext.Coinbase, state.BalancePath)]; ok {
log.Info("Coinbase is in MVReadMap", "address", blockContext.Coinbase)
task.shouldRerunWithoutFeeDelay = true
}
if _, ok := task.statedb.MVReadMap()[string(task.result.BurntContractAddress.Bytes())]; ok {
if _, ok := reads[blockstm.NewSubpathKey(task.result.BurntContractAddress, state.BalancePath)]; ok {
log.Info("BurntContractAddress is in MVReadMap", "address", task.result.BurntContractAddress)
task.shouldRerunWithoutFeeDelay = true
}
} else {
@ -109,11 +127,11 @@ func (task *ExecutionTask) Execute(mvh *blockstm.MVHashMap, incarnation int) (er
}
if task.statedb.HadInvalidRead() || err != nil {
err = blockstm.ErrExecAbort
err = blockstm.ErrExecAbortError{Dependency: task.statedb.DepTxIndex()}
return
}
task.statedb.Finalise(false)
task.statedb.Finalise(task.config.IsEIP158(task.blockNumber))
return
}
@ -134,6 +152,87 @@ func (task *ExecutionTask) Sender() common.Address {
return task.sender
}
func (task *ExecutionTask) Settle() {
task.finalStateDB.Prepare(task.tx.Hash(), task.index)
coinbase, _ := task.blockChain.Engine().Author(task.header)
coinbaseBalance := task.finalStateDB.GetBalance(coinbase)
task.finalStateDB.ApplyMVWriteSet(task.statedb.MVWriteList())
for _, l := range task.statedb.GetLogs(task.tx.Hash(), task.blockHash) {
task.finalStateDB.AddLog(l)
}
if *task.shouldDelayFeeCal {
if task.config.IsLondon(task.blockNumber) {
task.finalStateDB.AddBalance(task.result.BurntContractAddress, task.result.FeeBurnt)
}
task.finalStateDB.AddBalance(coinbase, task.result.FeeTipped)
output1 := new(big.Int).SetBytes(task.result.SenderInitBalance.Bytes())
output2 := new(big.Int).SetBytes(coinbaseBalance.Bytes())
// Deprecating transfer log and will be removed in future fork. PLEASE DO NOT USE this transfer log going forward. Parameters won't get updated as expected going forward with EIP1559
// add transfer log
AddFeeTransferLog(
task.finalStateDB,
task.msg.From(),
coinbase,
task.result.FeeTipped,
task.result.SenderInitBalance,
coinbaseBalance,
output1.Sub(output1, task.result.FeeTipped),
output2.Add(output2, task.result.FeeTipped),
)
}
for k, v := range task.statedb.Preimages() {
task.finalStateDB.AddPreimage(k, v)
}
// Update the state with pending changes.
var root []byte
if task.config.IsByzantium(task.blockNumber) {
task.finalStateDB.Finalise(true)
} else {
root = task.finalStateDB.IntermediateRoot(task.config.IsEIP158(task.blockNumber)).Bytes()
}
*task.totalUsedGas += 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: *task.totalUsedGas}
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 = task.finalStateDB.GetLogs(task.tx.Hash(), task.blockHash)
receipt.Bloom = types.CreateBloom(types.Receipts{receipt})
receipt.BlockHash = task.blockHash
receipt.BlockNumber = task.blockNumber
receipt.TransactionIndex = uint(task.finalStateDB.TxIndex())
*task.receipts = append(*task.receipts, receipt)
*task.allLogs = append(*task.allLogs, receipt.Logs...)
}
// 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.
@ -141,6 +240,7 @@ func (task *ExecutionTask) Sender() common.Address {
// 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.
// nolint:gocognit
func (p *ParallelStateProcessor) Process(block *types.Block, statedb *state.StateDB, cfg vm.Config) (types.Receipts, []*types.Log, uint64, error) {
var (
receipts types.Receipts
@ -150,6 +250,7 @@ func (p *ParallelStateProcessor) Process(block *types.Block, statedb *state.Stat
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)
@ -159,6 +260,8 @@ func (p *ParallelStateProcessor) Process(block *types.Block, statedb *state.Stat
shouldDelayFeeCal := true
coinbase, _ := p.bc.Engine().Author(header)
// 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)
@ -167,11 +270,9 @@ func (p *ParallelStateProcessor) Process(block *types.Block, statedb *state.Stat
return nil, nil, 0, fmt.Errorf("could not apply tx %d [%v]: %w", i, tx.Hash().Hex(), err)
}
bc := NewEVMBlockContext(header, p.bc, nil)
cleansdb := statedb.Copy()
if msg.From() == bc.Coinbase {
if msg.From() == coinbase {
shouldDelayFeeCal = false
}
@ -184,22 +285,42 @@ func (p *ParallelStateProcessor) Process(block *types.Block, statedb *state.Stat
tx: tx,
index: i,
cleanStateDB: cleansdb,
blockContext: bc,
finalStateDB: statedb,
blockChain: p.bc,
header: header,
evmConfig: cfg,
shouldDelayFeeCal: &shouldDelayFeeCal,
sender: msg.From(),
totalUsedGas: usedGas,
receipts: &receipts,
allLogs: &allLogs,
}
tasks = append(tasks, task)
}
_, err := blockstm.ExecuteParallel(tasks)
backupStateDB := statedb.Copy()
_, err := blockstm.ExecuteParallel(tasks, false)
for _, task := range tasks {
task := task.(*ExecutionTask)
if task.shouldRerunWithoutFeeDelay {
shouldDelayFeeCal = false
_, err = blockstm.ExecuteParallel(tasks)
*statedb = *backupStateDB
allLogs = []*types.Log{}
receipts = types.Receipts{}
usedGas = new(uint64)
for _, t := range tasks {
t := t.(*ExecutionTask)
t.finalStateDB = backupStateDB
t.allLogs = &allLogs
t.receipts = &receipts
t.totalUsedGas = usedGas
}
_, err = blockstm.ExecuteParallel(tasks, false)
break
}
@ -210,90 +331,14 @@ func (p *ParallelStateProcessor) Process(block *types.Block, statedb *state.Stat
return nil, nil, 0, err
}
london := p.config.IsLondon(blockNumber)
statedb.Finalise(p.config.IsEIP158(blockNumber))
for _, task := range tasks {
task := task.(*ExecutionTask)
statedb.Prepare(task.tx.Hash(), task.index)
coinbaseBalance := statedb.GetBalance(task.blockContext.Coinbase)
statedb.ApplyMVWriteSet(task.MVWriteList())
for _, l := range task.statedb.GetLogs(task.tx.Hash(), blockHash) {
statedb.AddLog(l)
}
if shouldDelayFeeCal {
if london {
statedb.AddBalance(task.result.BurntContractAddress, task.result.FeeBurnt)
}
statedb.AddBalance(task.blockContext.Coinbase, task.result.FeeTipped)
output1 := new(big.Int).SetBytes(task.result.SenderInitBalance.Bytes())
output2 := new(big.Int).SetBytes(coinbaseBalance.Bytes())
// Deprecating transfer log and will be removed in future fork. PLEASE DO NOT USE this transfer log going forward. Parameters won't get updated as expected going forward with EIP1559
// add transfer log
AddFeeTransferLog(
statedb,
task.msg.From(),
task.blockContext.Coinbase,
task.result.FeeTipped,
task.result.SenderInitBalance,
coinbaseBalance,
output1.Sub(output1, task.result.FeeTipped),
output2.Add(output2, task.result.FeeTipped),
)
}
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...)
}
start := time.Now()
// Finalize the block, applying any consensus engine specific extras (e.g. block rewards)
p.engine.Finalize(p.bc, header, statedb, block.Transactions(), block.Uncles())
fmt.Println("Finalize time of parallel execution:", time.Since(start))
return receipts, allLogs, *usedGas, nil
}

View file

@ -20,6 +20,7 @@ import (
"math/big"
"github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/core/blockstm"
)
// journalEntry is a modification entry in the state change journal that can be
@ -143,7 +144,7 @@ type (
func (ch createObjectChange) revert(s *StateDB) {
delete(s.stateObjects, *ch.account)
delete(s.stateObjectsDirty, *ch.account)
MVWrite(s, ch.account.Bytes())
MVWrite(s, blockstm.NewAddressKey(*ch.account))
}
func (ch createObjectChange) dirtied() *common.Address {
@ -152,7 +153,7 @@ func (ch createObjectChange) dirtied() *common.Address {
func (ch resetObjectChange) revert(s *StateDB) {
s.setStateObject(ch.prev)
MVWrite(s, ch.prev.address.Bytes())
MVWrite(s, blockstm.NewAddressKey(ch.prev.address))
if !ch.prevdestruct && s.snap != nil {
delete(s.snapDestructs, ch.prev.addrHash)
}
@ -167,8 +168,8 @@ func (ch suicideChange) revert(s *StateDB) {
if obj != nil {
obj.suicided = ch.prev
obj.setBalance(ch.prevbalance)
MVWrite(s, subPath(ch.account.Bytes(), suicidePath))
MVWrite(s, subPath(ch.account.Bytes(), balancePath))
MVWrite(s, blockstm.NewSubpathKey(*ch.account, SuicidePath))
MVWrite(s, blockstm.NewSubpathKey(*ch.account, BalancePath))
}
}
@ -187,7 +188,7 @@ func (ch touchChange) dirtied() *common.Address {
func (ch balanceChange) revert(s *StateDB) {
s.getStateObject(*ch.account).setBalance(ch.prev)
MVWrite(s, subPath(ch.account.Bytes(), balancePath))
MVWrite(s, blockstm.NewSubpathKey(*ch.account, BalancePath))
}
func (ch balanceChange) dirtied() *common.Address {
@ -196,7 +197,7 @@ func (ch balanceChange) dirtied() *common.Address {
func (ch nonceChange) revert(s *StateDB) {
s.getStateObject(*ch.account).setNonce(ch.prev)
MVWrite(s, subPath(ch.account.Bytes(), noncePath))
MVWrite(s, blockstm.NewSubpathKey(*ch.account, NoncePath))
}
func (ch nonceChange) dirtied() *common.Address {
@ -205,7 +206,7 @@ func (ch nonceChange) dirtied() *common.Address {
func (ch codeChange) revert(s *StateDB) {
s.getStateObject(*ch.account).setCode(common.BytesToHash(ch.prevhash), ch.prevcode)
MVWrite(s, subPath(ch.account.Bytes(), codePath))
MVWrite(s, blockstm.NewSubpathKey(*ch.account, CodePath))
}
func (ch codeChange) dirtied() *common.Address {
@ -214,7 +215,7 @@ func (ch codeChange) dirtied() *common.Address {
func (ch storageChange) revert(s *StateDB) {
s.getStateObject(*ch.account).setState(ch.key, ch.prevalue)
MVWrite(s, append(ch.account.Bytes(), ch.key.Bytes()...))
MVWrite(s, blockstm.NewStateKey(*ch.account, ch.key))
}
func (ch storageChange) dirtied() *common.Address {

View file

@ -83,10 +83,10 @@ type StateDB struct {
// Block-stm related fields
mvHashmap *blockstm.MVHashMap
incarnation int
readMap map[string]blockstm.ReadDescriptor
writeMap map[string]blockstm.WriteDescriptor
readMap map[blockstm.Key]blockstm.ReadDescriptor
writeMap map[blockstm.Key]blockstm.WriteDescriptor
newStateObjects map[common.Address]struct{}
invalidRead bool
dep int
// DB error.
// State objects are used by the consensus core and VM which are
@ -169,21 +169,23 @@ func NewWithMVHashmap(root common.Hash, db Database, snaps *snapshot.Tree, mvhm
return nil, err
} else {
sdb.mvHashmap = mvhm
sdb.dep = -1
return sdb, nil
}
}
func (sdb *StateDB) SetMVHashmap(mvhm *blockstm.MVHashMap) {
sdb.mvHashmap = mvhm
sdb.dep = -1
}
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 {
if !v.Path.IsAddress() {
writes = append(writes, v)
} else if _, ok := s.newStateObjects[common.BytesToAddress(v.Path)]; ok {
} else if _, ok := s.newStateObjects[common.BytesToAddress(v.Path[:common.AddressLength])]; ok {
writes = append(writes, v)
}
}
@ -201,7 +203,7 @@ func (s *StateDB) MVFullWriteList() []blockstm.WriteDescriptor {
return writes
}
func (s *StateDB) MVReadMap() map[string]blockstm.ReadDescriptor {
func (s *StateDB) MVReadMap() map[blockstm.Key]blockstm.ReadDescriptor {
return s.readMap
}
@ -217,25 +219,33 @@ func (s *StateDB) MVReadList() []blockstm.ReadDescriptor {
func (s *StateDB) ensureReadMap() {
if s.readMap == nil {
s.readMap = make(map[string]blockstm.ReadDescriptor)
s.readMap = make(map[blockstm.Key]blockstm.ReadDescriptor)
}
}
func (s *StateDB) ensureWriteMap() {
if s.writeMap == nil {
s.writeMap = make(map[string]blockstm.WriteDescriptor)
s.writeMap = make(map[blockstm.Key]blockstm.WriteDescriptor)
}
}
func (s *StateDB) HadInvalidRead() bool {
return s.invalidRead
return s.dep >= 0
}
func (s *StateDB) DepTxIndex() int {
return s.dep
}
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) {
type StorageVal[T any] struct {
Value *T
}
func MVRead[T any](s *StateDB, k blockstm.Key, defaultV T, readStorage func(s *StateDB) T) (v T) {
if s.mvHashmap == nil {
return readStorage(s)
}
@ -243,7 +253,7 @@ func MVRead[T any](s *StateDB, k []byte, defaultV T, readStorage func(s *StateDB
s.ensureReadMap()
if s.writeMap != nil {
if _, ok := s.writeMap[string(k)]; ok {
if _, ok := s.writeMap[k]; ok {
return readStorage(s)
}
}
@ -267,8 +277,12 @@ func MVRead[T any](s *StateDB, k []byte, defaultV T, readStorage func(s *StateDB
}
case blockstm.MVReadResultDependency:
{
s.invalidRead = true
return defaultV
if res.DepIdx() > s.dep {
s.dep = res.DepIdx()
}
// Return immediate to executor when we found a dependency
panic("Found dependency")
}
case blockstm.MVReadResultNone:
{
@ -279,20 +293,19 @@ func MVRead[T any](s *StateDB, k []byte, defaultV T, readStorage func(s *StateDB
return defaultV
}
mk := string(k)
// TODO: I assume we don't want to overwrite an existing read because this could - for example - change a storage
// read to map if the same value is read multiple times.
if _, ok := s.readMap[mk]; !ok {
s.readMap[mk] = rd
if _, ok := s.readMap[k]; !ok {
s.readMap[k] = rd
}
return
}
func MVWrite(s *StateDB, k []byte) {
func MVWrite(s *StateDB, k blockstm.Key) {
if s.mvHashmap != nil {
s.ensureWriteMap()
s.writeMap[string(k)] = blockstm.WriteDescriptor{
s.writeMap[k] = blockstm.WriteDescriptor{
Path: k,
V: s.Version(),
Val: s,
@ -300,12 +313,12 @@ func MVWrite(s *StateDB, k []byte) {
}
}
func MVWritten(s *StateDB, k []byte) bool {
func MVWritten(s *StateDB, k blockstm.Key) bool {
if s.mvHashmap == nil || s.writeMap == nil {
return false
}
_, ok := s.writeMap[string(k)]
_, ok := s.writeMap[k]
return ok
}
@ -324,30 +337,27 @@ func (sw *StateDB) ApplyMVWriteSet(writes []blockstm.WriteDescriptor) {
path := writes[i].Path
sr := writes[i].Val.(*StateDB)
keyLength := len(path)
if keyLength == common.AddressLength {
sw.GetOrNewStateObject(common.BytesToAddress(path))
} else if keyLength == (common.AddressLength + common.HashLength) {
addr := common.BytesToAddress(path[:common.AddressLength])
subPath := common.BytesToHash(path[common.AddressLength:])
sw.SetState(addr, subPath, sr.GetState(addr, subPath))
if path.IsState() {
addr := path.GetAddress()
stateKey := path.GetStateKey()
state := sr.GetState(addr, stateKey)
sw.SetState(addr, stateKey, state)
} else {
addr := common.BytesToAddress(path[:common.AddressLength])
switch path[keyLength-1] {
case balancePath:
addr := path.GetAddress()
switch path.GetSubpath() {
case BalancePath:
sw.SetBalance(addr, sr.GetBalance(addr))
case noncePath:
case NoncePath:
sw.SetNonce(addr, sr.GetNonce(addr))
case codePath:
case CodePath:
sw.SetCode(addr, sr.GetCode(addr))
case suicidePath:
case SuicidePath:
stateObject := sr.getDeletedStateObject(addr)
if stateObject != nil && stateObject.deleted {
sw.Suicide(addr)
}
default:
panic(fmt.Errorf("unknown key type: %d", path[keyLength-1]))
panic(fmt.Errorf("unknown key type: %d", path.GetSubpath()))
}
}
}
@ -373,7 +383,7 @@ func (s *StateDB) GetReadMapDump() []DumpStruct {
TxInc: s.incarnation,
VerIdx: val.V.TxnIndex,
VerInc: val.V.Incarnation,
Path: val.Path,
Path: val.Path[:],
Op: "Read\n",
}
res = append(res, *temp)
@ -393,7 +403,7 @@ func (s *StateDB) GetWriteMapDump() []DumpStruct {
TxInc: s.incarnation,
VerIdx: val.V.TxnIndex,
VerInc: val.V.Incarnation,
Path: val.Path,
Path: val.Path[:],
Op: "Write\n",
}
res = append(res, *temp)
@ -512,17 +522,17 @@ func (s *StateDB) Empty(addr common.Address) bool {
}
// Create a unique path for special fields (e.g. balance, code) in a state object.
func subPath(prefix []byte, s uint8) []byte {
path := append(prefix, common.Hash{}.Bytes()...) // append a full empty hash to avoid collision with storage state
path = append(path, s) // append the special field identifier
// func subPath(prefix []byte, s uint8) [blockstm.KeyLength]byte {
// path := append(prefix, common.Hash{}.Bytes()...) // append a full empty hash to avoid collision with storage state
// path = append(path, s) // append the special field identifier
return path
}
// return path
// }
const balancePath = 1
const noncePath = 2
const codePath = 3
const suicidePath = 4
const BalancePath = 1
const NoncePath = 2
const CodePath = 3
const SuicidePath = 4
// GetBalance retrieves the balance from the given address or 0 if object not found
func (s *StateDB) GetBalance(addr common.Address) *big.Int {
@ -530,7 +540,7 @@ func (s *StateDB) GetBalance(addr common.Address) *big.Int {
return common.Big0
}
return MVRead(s, subPath(addr.Bytes(), balancePath), common.Big0, func(s *StateDB) *big.Int {
return MVRead(s, blockstm.NewSubpathKey(addr, BalancePath), common.Big0, func(s *StateDB) *big.Int {
stateObject := s.getStateObject(addr)
if stateObject != nil {
return stateObject.Balance()
@ -545,7 +555,7 @@ func (s *StateDB) GetNonce(addr common.Address) uint64 {
return 0
}
return MVRead(s, subPath(addr.Bytes(), noncePath), 0, func(s *StateDB) uint64 {
return MVRead(s, blockstm.NewSubpathKey(addr, NoncePath), 0, func(s *StateDB) uint64 {
stateObject := s.getStateObject(addr)
if stateObject != nil {
return stateObject.Nonce()
@ -572,7 +582,7 @@ func (s *StateDB) GetCode(addr common.Address) []byte {
return nil
}
return MVRead(s, subPath(addr.Bytes(), codePath), nil, func(s *StateDB) []byte {
return MVRead(s, blockstm.NewSubpathKey(addr, CodePath), nil, func(s *StateDB) []byte {
stateObject := s.getStateObject(addr)
if stateObject != nil {
return stateObject.Code(s.db)
@ -586,7 +596,7 @@ func (s *StateDB) GetCodeSize(addr common.Address) int {
return 0
}
return MVRead(s, subPath(addr.Bytes(), codePath), 0, func(s *StateDB) int {
return MVRead(s, blockstm.NewSubpathKey(addr, CodePath), 0, func(s *StateDB) int {
stateObject := s.getStateObject(addr)
if stateObject != nil {
return stateObject.CodeSize(s.db)
@ -600,7 +610,7 @@ func (s *StateDB) GetCodeHash(addr common.Address) common.Hash {
return common.Hash{}
}
return MVRead(s, subPath(addr.Bytes(), codePath), common.Hash{}, func(s *StateDB) common.Hash {
return MVRead(s, blockstm.NewSubpathKey(addr, CodePath), common.Hash{}, func(s *StateDB) common.Hash {
stateObject := s.getStateObject(addr)
if stateObject == nil {
return common.Hash{}
@ -615,7 +625,7 @@ func (s *StateDB) GetState(addr common.Address, hash common.Hash) common.Hash {
return common.Hash{}
}
return MVRead(s, append(addr.Bytes(), hash.Bytes()...), common.Hash{}, func(s *StateDB) common.Hash {
return MVRead(s, blockstm.NewStateKey(addr, hash), common.Hash{}, func(s *StateDB) common.Hash {
stateObject := s.getStateObject(addr)
if stateObject != nil {
return stateObject.GetState(s.db, hash)
@ -653,7 +663,7 @@ func (s *StateDB) GetCommittedState(addr common.Address, hash common.Hash) commo
return common.Hash{}
}
return MVRead(s, append(addr.Bytes(), hash.Bytes()...), common.Hash{}, func(s *StateDB) common.Hash {
return MVRead(s, blockstm.NewStateKey(addr, hash), common.Hash{}, func(s *StateDB) common.Hash {
stateObject := s.getStateObject(addr)
if stateObject != nil {
return stateObject.GetCommittedState(s.db, hash)
@ -684,7 +694,7 @@ func (s *StateDB) HasSuicided(addr common.Address) bool {
return false
}
return MVRead(s, subPath(addr.Bytes(), suicidePath), false, func(s *StateDB) bool {
return MVRead(s, blockstm.NewSubpathKey(addr, SuicidePath), false, func(s *StateDB) bool {
stateObject := s.getStateObject(addr)
if stateObject != nil {
return stateObject.suicided
@ -709,7 +719,7 @@ func (s *StateDB) AddBalance(addr common.Address, amount *big.Int) {
if stateObject != nil {
stateObject = s.mvRecordWritten(stateObject)
stateObject.AddBalance(amount)
MVWrite(s, subPath(addr.Bytes(), balancePath))
MVWrite(s, blockstm.NewSubpathKey(addr, BalancePath))
}
}
@ -725,7 +735,7 @@ func (s *StateDB) SubBalance(addr common.Address, amount *big.Int) {
if stateObject != nil {
stateObject = s.mvRecordWritten(stateObject)
stateObject.SubBalance(amount)
MVWrite(s, subPath(addr.Bytes(), balancePath))
MVWrite(s, blockstm.NewSubpathKey(addr, BalancePath))
}
}
@ -734,7 +744,7 @@ func (s *StateDB) SetBalance(addr common.Address, amount *big.Int) {
if stateObject != nil {
stateObject = s.mvRecordWritten(stateObject)
stateObject.SetBalance(amount)
MVWrite(s, subPath(addr.Bytes(), balancePath))
MVWrite(s, blockstm.NewSubpathKey(addr, BalancePath))
}
}
@ -743,7 +753,7 @@ func (s *StateDB) SetNonce(addr common.Address, nonce uint64) {
if stateObject != nil {
stateObject = s.mvRecordWritten(stateObject)
stateObject.SetNonce(nonce)
MVWrite(s, subPath(addr.Bytes(), noncePath))
MVWrite(s, blockstm.NewSubpathKey(addr, NoncePath))
}
}
@ -752,7 +762,7 @@ func (s *StateDB) SetCode(addr common.Address, code []byte) {
if stateObject != nil {
stateObject = s.mvRecordWritten(stateObject)
stateObject.SetCode(crypto.Keccak256Hash(code), code)
MVWrite(s, subPath(addr.Bytes(), codePath))
MVWrite(s, blockstm.NewSubpathKey(addr, CodePath))
}
}
@ -761,7 +771,7 @@ func (s *StateDB) SetState(addr common.Address, key, value common.Hash) {
if stateObject != nil {
stateObject = s.mvRecordWritten(stateObject)
stateObject.SetState(s.db, key, value)
MVWrite(s, append(addr.Bytes(), key.Bytes()...))
MVWrite(s, blockstm.NewStateKey(addr, key))
}
}
@ -794,8 +804,8 @@ func (s *StateDB) Suicide(addr common.Address) bool {
stateObject.markSuicided()
stateObject.data.Balance = new(big.Int)
MVWrite(s, subPath(addr.Bytes(), suicidePath))
MVWrite(s, subPath(addr.Bytes(), balancePath))
MVWrite(s, blockstm.NewSubpathKey(addr, SuicidePath))
MVWrite(s, blockstm.NewSubpathKey(addr, BalancePath))
return true
}
@ -853,7 +863,7 @@ func (s *StateDB) getStateObject(addr common.Address) *stateObject {
// flag set. This is needed by the state journal to revert to the correct s-
// destructed object instead of wiping all knowledge about the state object.
func (s *StateDB) getDeletedStateObject(addr common.Address) *stateObject {
return MVRead(s, addr.Bytes(), nil, func(s *StateDB) *stateObject {
return MVRead(s, blockstm.NewAddressKey(addr), nil, func(s *StateDB) *stateObject {
// Prefer live objects if any is available
if obj := s.stateObjects[addr]; obj != nil {
return obj
@ -932,16 +942,16 @@ func (s *StateDB) mvRecordWritten(object *stateObject) *stateObject {
return object
}
addrPath := object.Address().Bytes()
addrKey := blockstm.NewAddressKey(object.Address())
if MVWritten(s, addrPath) {
if MVWritten(s, addrKey) {
return object
}
// Deepcopy is needed to ensure that objects are not written by multiple transactions at the same time, because
// the input state object can come from a different transaction.
s.setStateObject(object.deepCopy(s))
MVWrite(s, addrPath)
MVWrite(s, addrKey)
return s.stateObjects[object.Address()]
}
@ -967,7 +977,7 @@ func (s *StateDB) createObject(addr common.Address) (newobj, prev *stateObject)
s.setStateObject(newobj)
s.newStateObjects[addr] = struct{}{}
MVWrite(s, addr.Bytes())
MVWrite(s, blockstm.NewAddressKey(addr))
if prev != nil && !prev.deleted {
return newobj, prev
}
@ -988,7 +998,7 @@ func (s *StateDB) CreateAccount(addr common.Address) {
newObj, prev := s.createObject(addr)
if prev != nil {
newObj.setBalance(prev.data.Balance)
MVWrite(s, subPath(addr.Bytes(), balancePath))
MVWrite(s, blockstm.NewSubpathKey(addr, BalancePath))
}
}

View file

@ -659,16 +659,21 @@ func TestMVHashMapMarkEstimate(t *testing.T) {
assert.Equal(t, balance, b)
// Tx1 mark estimate
for _, v := range states[1].writeMap {
for _, v := range states[1].MVWriteList() {
mvhm.MarkEstimate(v.Path, 1)
}
// Tx2 read again should get default (empty) vals because its dependency Tx1 is marked as estimate
v = states[2].GetState(addr, key)
b = states[2].GetBalance(addr)
defer func() {
if r := recover(); r == nil {
t.Errorf("The code did not panic")
} else {
t.Log("Recovered in f", r)
}
}()
assert.Equal(t, common.Hash{}, v)
assert.Equal(t, common.Big0, b)
// Tx2 read again should get default (empty) vals because its dependency Tx1 is marked as estimate
states[2].GetState(addr, key)
states[2].GetBalance(addr)
// Tx1 read again should get Tx0 vals
v = states[1].GetState(addr, key)

1
go.mod
View file

@ -37,6 +37,7 @@ require (
github.com/hashicorp/go-bexpr v0.1.10
github.com/hashicorp/golang-lru v0.5.5-0.20210104140557-80c98217689d
github.com/hashicorp/hcl/v2 v2.10.1
github.com/heimdalr/dag v1.2.1
github.com/holiman/bloomfilter/v2 v2.0.3
github.com/holiman/uint256 v1.2.0
github.com/huin/goupnp v1.0.3-0.20220313090229-ca81a64b4204

5
go.sum
View file

@ -187,8 +187,9 @@ github.com/go-sourcemap/sourcemap v2.1.3+incompatible/go.mod h1:F8jJfvm2KbVjc5Nq
github.com/go-sql-driver/mysql v1.4.1/go.mod h1:zAC/RDZ24gD3HViQzih4MyKcchzm+sOG5ZlKdlhCg5w=
github.com/go-stack/stack v1.8.0 h1:5SgMzNM5HxrEjV0ww2lTmX6E2Izsfxas4+YHWRs3Lsk=
github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/melR3HDY=
github.com/go-test/deep v1.0.3 h1:ZrJSEWsXzPOxaZnFteGEfooLba+ju3FYIbOrS+rQd68=
github.com/go-test/deep v1.0.3/go.mod h1:wGDj63lr65AM2AQyKZd/NYHGb0R+1RLqB8NKt3aSFNA=
github.com/go-test/deep v1.0.7 h1:/VSMRlnY/JSyqxQUzQLKVMAskpY/NZKFA5j2P+0pP2M=
github.com/go-test/deep v1.0.7/go.mod h1:QV8Hv/iy04NyLBxAdO9njL0iVPN1S4d/A3NVv1V36o8=
github.com/gofrs/uuid v3.3.0+incompatible/go.mod h1:b2aQJv3Z4Fp6yNu3cdSllBxTCLRxnplIgP/c0N/04lM=
github.com/gogo/protobuf v1.1.1/go.mod h1:r8qH/GZQm5c6nD/R0oafs1akxWv10x8SbQlK7atdtwQ=
github.com/gogo/protobuf v1.3.1/go.mod h1:SlYgWuQ5SjCEi6WLHjHCa1yvBfUnHcTbrrZtXPKa29o=
@ -275,6 +276,8 @@ github.com/hashicorp/golang-lru v0.5.5-0.20210104140557-80c98217689d h1:dg1dEPuW
github.com/hashicorp/golang-lru v0.5.5-0.20210104140557-80c98217689d/go.mod h1:iADmTwqILo4mZ8BN3D2Q6+9jd8WM5uGBxy+E8yxSoD4=
github.com/hashicorp/hcl/v2 v2.10.1 h1:h4Xx4fsrRE26ohAk/1iGF/JBqRQbyUqu5Lvj60U54ys=
github.com/hashicorp/hcl/v2 v2.10.1/go.mod h1:FwWsfWEjyV/CMj8s/gqAuiviY72rJ1/oayI9WftqcKg=
github.com/heimdalr/dag v1.2.1 h1:XJOMaoWqJK1UKdp+4zaO2uwav9GFbHMGCirdViKMRIQ=
github.com/heimdalr/dag v1.2.1/go.mod h1:Of/wUB7Yoj4dwiOcGOOYIq6MHlPF/8/QMBKFJpwg+yc=
github.com/holiman/bloomfilter/v2 v2.0.3 h1:73e0e/V0tCydx14a0SCYS/EWCxgwLZ18CZcZKVu0fao=
github.com/holiman/bloomfilter/v2 v2.0.3/go.mod h1:zpoh+gs7qcpqrHr3dB55AMiJwo0iURXE7ZOP9L9hSkA=
github.com/holiman/uint256 v1.2.0 h1:gpSYcPLWGv4sG43I2mVLiDZCNDh/EpGjSk8tmtxitHM=