validateAccountAndPaymaster called by process, tracing

This commit is contained in:
Dror Tirosh 2024-08-13 11:23:33 +03:00
parent 2ce87f0003
commit 3ae9c9657f
3 changed files with 113 additions and 90 deletions

View file

@ -1,6 +1,7 @@
package core package core
import ( import (
"bytes"
"errors" "errors"
"fmt" "fmt"
"github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common"
@ -16,10 +17,19 @@ import (
var AA_ENTRY_POINT = common.HexToAddress("0x0000000000000000000000000000000000007560") var AA_ENTRY_POINT = common.HexToAddress("0x0000000000000000000000000000000000007560")
var AA_SENDER_CREATOR = common.HexToAddress("0x00000000000000000000000000000000ffff7560") var AA_SENDER_CREATOR = common.HexToAddress("0x00000000000000000000000000000000ffff7560")
type EntryPointCallEntry struct {
From common.Address
Input []byte
}
type EntryPointCall struct { type EntryPointCall struct {
OnEnterSuper tracing.EnterHook entries []EntryPointCallEntry
Input []byte err error
err error }
func NewEntryPointCall() *EntryPointCall {
return &EntryPointCall{
entries: make([]EntryPointCallEntry, 0),
}
} }
type ValidationPhaseResult struct { type ValidationPhaseResult struct {
@ -151,6 +161,49 @@ func CheckNonceRip7560(tx *types.Rip7560AccountAbstractionTx, st *state.StateDB)
return nil return nil
} }
// finalize validation return data from account and paymaster
func ValidateAccountAndPaymaster(time uint64, sender, paymaster *common.Address, epc *EntryPointCall) (*AcceptAccountData, *AcceptPaymasterData, error) {
if epc.err != nil {
return nil, nil, epc.err
}
if len(epc.entries) == 0 || epc.entries[0].From.Cmp(*sender) != 0 {
return nil, nil, errors.New("account validation did not call the EntryPoint 'acceptAccount' callback")
}
if paymaster == nil {
if len(epc.entries) > 1 {
return nil, nil, errors.New("EntryPoint callback called more than once")
}
} else {
if len(epc.entries) < 2 || epc.entries[1].From.Cmp(*paymaster) != 0 {
return nil, nil, errors.New("paymaster validation did not call the EntryPoint 'acceptPaymaster' callback")
}
if len(epc.entries) > 2 {
return nil, nil, errors.New("EntryPoint callback called too many times")
}
}
aad, err := ValidateAccountEntryPointCall(sender, epc.entries[0].Input)
if err != nil {
return nil, nil, err
}
err = validateValidityTimeRange(time, aad.ValidAfter.Uint64(), aad.ValidUntil.Uint64())
if err != nil {
return nil, nil, err
}
var apd *AcceptPaymasterData
if paymaster != nil {
apd, err = validatePaymasterEntryPointCall(paymaster, epc.entries[1].Input)
if err != nil {
return nil, nil, err
}
err = validateValidityTimeRange(time, apd.ValidAfter.Uint64(), apd.ValidUntil.Uint64())
if err != nil {
return nil, nil, err
}
}
return aad, apd, nil
}
func ApplyRip7560ValidationPhases(chainConfig *params.ChainConfig, bc ChainContext, author *common.Address, gp *GasPool, statedb *state.StateDB, header *types.Header, tx *types.Transaction, cfg vm.Config) (*ValidationPhaseResult, error) { func ApplyRip7560ValidationPhases(chainConfig *params.ChainConfig, bc ChainContext, author *common.Address, gp *GasPool, statedb *state.StateDB, header *types.Header, tx *types.Transaction, cfg vm.Config) (*ValidationPhaseResult, error) {
aatx := tx.Rip7560TransactionData() aatx := tx.Rip7560TransactionData()
err := CheckNonceRip7560(aatx, statedb) err := CheckNonceRip7560(aatx, statedb)
@ -171,21 +224,19 @@ func ApplyRip7560ValidationPhases(chainConfig *params.ChainConfig, bc ChainConte
blockContext := NewEVMBlockContext(header, bc, author) blockContext := NewEVMBlockContext(header, bc, author)
sender := tx.Rip7560TransactionData().Sender sender := tx.Rip7560TransactionData().Sender
paymaster := tx.Rip7560TransactionData().Paymaster
txContext := vm.TxContext{ txContext := vm.TxContext{
Origin: *sender, Origin: *sender,
GasPrice: gasPrice, GasPrice: gasPrice,
} }
evm := vm.NewEVM(blockContext, txContext, statedb, chainConfig, cfg) evm := vm.NewEVM(blockContext, txContext, statedb, chainConfig, cfg)
epc := &EntryPointCall{}
var epc *EntryPointCall
if evm.Config.Tracer == nil { if evm.Config.Tracer == nil {
epc = NewEntryPointCall()
evm.Config.Tracer = &tracing.Hooks{ evm.Config.Tracer = &tracing.Hooks{
OnEnter: epc.OnEnter, OnEnter: epc.OnEnter,
} }
} else {
// keep the original tracer's OnEnter hook
epc.OnEnterSuper = evm.Config.Tracer.OnEnter
evm.Config.Tracer.OnEnter = epc.OnEnter
} }
if evm.Config.Tracer.OnTxStart != nil { if evm.Config.Tracer.OnTxStart != nil {
@ -228,70 +279,54 @@ func ApplyRip7560ValidationPhases(chainConfig *params.ChainConfig, bc ChainConte
if resultAccountValidation.Err != nil { if resultAccountValidation.Err != nil {
return nil, resultAccountValidation.Err return nil, resultAccountValidation.Err
} }
aad, err := validateAccountEntryPointCall(epc)
if err != nil {
return nil, err
}
// clear the EntryPoint calls array after parsing
epc.err = nil
epc.Input = nil
err = validateValidityTimeRange(header.Time, aad.ValidAfter.Uint64(), aad.ValidUntil.Uint64())
if err != nil {
return nil, err
}
vpr := &ValidationPhaseResult{} vpr := &ValidationPhaseResult{}
paymasterContext, pmValidationUsedGas, pmValidAfter, pmValidUntil, err := applyPaymasterValidationFrame(epc, tx, chainConfig, signingHash, evm, gp, statedb, header) pmValidationUsedGas, err := applyPaymasterValidationFrame(tx, chainConfig, signingHash, evm, gp, statedb, header)
if err != nil { if err != nil {
return nil, err return nil, err
} }
aad, apd, err := ValidateAccountAndPaymaster(header.Time, sender, paymaster, epc)
if err != nil {
return nil, err
}
vpr.Tx = tx vpr.Tx = tx
vpr.TxHash = tx.Hash() vpr.TxHash = tx.Hash()
vpr.PreCharge = preCharge vpr.PreCharge = preCharge
vpr.EffectiveGasPrice = gasPriceUint256 vpr.EffectiveGasPrice = gasPriceUint256
vpr.PaymasterContext = paymasterContext vpr.PaymasterContext = apd.Context
vpr.DeploymentUsedGas = deploymentUsedGas vpr.DeploymentUsedGas = deploymentUsedGas
vpr.ValidationUsedGas = resultAccountValidation.UsedGas vpr.ValidationUsedGas = resultAccountValidation.UsedGas
vpr.PmValidationUsedGas = pmValidationUsedGas vpr.PmValidationUsedGas = pmValidationUsedGas
vpr.SenderValidAfter = aad.ValidAfter.Uint64() vpr.SenderValidAfter = aad.ValidAfter.Uint64()
vpr.SenderValidUntil = aad.ValidUntil.Uint64() vpr.SenderValidUntil = aad.ValidUntil.Uint64()
vpr.PmValidAfter = pmValidAfter vpr.PmValidAfter = apd.ValidAfter.Uint64()
vpr.PmValidUntil = pmValidUntil vpr.PmValidUntil = apd.ValidUntil.Uint64()
statedb.Finalise(true) statedb.Finalise(true)
return vpr, nil return vpr, nil
} }
func applyPaymasterValidationFrame(epc *EntryPointCall, tx *types.Transaction, chainConfig *params.ChainConfig, signingHash common.Hash, evm *vm.EVM, gp *GasPool, statedb *state.StateDB, header *types.Header) ([]byte, uint64, uint64, uint64, error) { func applyPaymasterValidationFrame(tx *types.Transaction, chainConfig *params.ChainConfig, signingHash common.Hash, evm *vm.EVM, gp *GasPool, statedb *state.StateDB, header *types.Header) (uint64, error) {
/*** Paymaster Validation Frame ***/ /*** Paymaster Validation Frame ***/
var pmValidationUsedGas uint64 var pmValidationUsedGas uint64
paymasterMsg, err := preparePaymasterValidationMessage(tx, chainConfig, signingHash) paymasterMsg, err := preparePaymasterValidationMessage(tx, chainConfig, signingHash)
if paymasterMsg == nil || err != nil { if paymasterMsg == nil || err != nil {
return nil, 0, 0, 0, err return 0, err
} }
resultPm, err := ApplyMessage(evm, paymasterMsg, gp) resultPm, err := ApplyMessage(evm, paymasterMsg, gp)
if err != nil { if err != nil {
return nil, 0, 0, 0, err return 0, err
} }
if resultPm.Failed() { if resultPm.Failed() {
return nil, 0, 0, 0, resultPm.Err return 0, resultPm.Err
} }
if resultPm.Failed() { if resultPm.Failed() {
return nil, 0, 0, 0, errors.New("paymaster validation failed - invalid transaction") return 0, errors.New("paymaster validation failed - invalid transaction")
} }
pmValidationUsedGas = resultPm.UsedGas pmValidationUsedGas = resultPm.UsedGas
apd, err := validatePaymasterEntryPointCall(epc) return pmValidationUsedGas, nil
if err != nil {
return nil, 0, 0, 0, err
}
err = validateValidityTimeRange(header.Time, apd.ValidAfter.Uint64(), apd.ValidUntil.Uint64())
if err != nil {
return nil, 0, 0, 0, err
}
return apd.Context, pmValidationUsedGas, apd.ValidAfter.Uint64(), apd.ValidUntil.Uint64(), nil
} }
func applyPaymasterPostOpFrame(vpr *ValidationPhaseResult, executionResult *ExecutionResult, evm *vm.EVM, gp *GasPool, statedb *state.StateDB, header *types.Header) (*ExecutionResult, error) { func applyPaymasterPostOpFrame(vpr *ValidationPhaseResult, executionResult *ExecutionResult, evm *vm.EVM, gp *GasPool, statedb *state.StateDB, header *types.Header) (*ExecutionResult, error) {
@ -466,31 +501,18 @@ func preparePostOpMessage(vpr *ValidationPhaseResult, chainConfig *params.ChainC
}, nil }, nil
} }
func validateAccountEntryPointCall(epc *EntryPointCall) (*AcceptAccountData, error) { func ValidateAccountEntryPointCall(sender *common.Address, input []byte) (*AcceptAccountData, error) {
if epc.err != nil { if len(input) != 68 {
return nil, epc.err
}
if epc.Input == nil {
return nil, errors.New("account validation did not call the EntryPoint 'acceptAccount' callback")
}
if len(epc.Input) != 68 {
return nil, errors.New("invalid account return data length") return nil, errors.New("invalid account return data length")
} }
return abiDecodeAcceptAccount(epc.Input) return abiDecodeAcceptAccount(input)
} }
func validatePaymasterEntryPointCall(epc *EntryPointCall) (*AcceptPaymasterData, error) { func validatePaymasterEntryPointCall(paymaster *common.Address, input []byte) (*AcceptPaymasterData, error) {
if epc.err != nil { if len(input) < 100 {
return nil, epc.err
}
if epc.Input == nil {
return nil, errors.New("paymaster validation did not call the EntryPoint 'acceptPaymaster' callback")
}
if len(epc.Input) < 100 {
return nil, errors.New("invalid paymaster callback data length") return nil, errors.New("invalid paymaster callback data length")
} }
apd, err := abiDecodeAcceptPaymaster(epc.Input) apd, err := abiDecodeAcceptPaymaster(input)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@ -514,25 +536,18 @@ func validateValidityTimeRange(time uint64, validAfter uint64, validUntil uint64
} }
func (epc *EntryPointCall) OnEnter(depth int, typ byte, from common.Address, to common.Address, input []byte, gas uint64, value *big.Int) { func (epc *EntryPointCall) OnEnter(depth int, typ byte, from common.Address, to common.Address, input []byte, gas uint64, value *big.Int) {
if epc.OnEnterSuper != nil {
epc.OnEnterSuper(depth, typ, from, to, input, gas, value)
}
isRip7560EntryPoint := to.Cmp(AA_ENTRY_POINT) == 0 isRip7560EntryPoint := to.Cmp(AA_ENTRY_POINT) == 0
if !isRip7560EntryPoint { if !isRip7560EntryPoint {
return return
} }
if depth != 1 { if depth != 1 {
println("ONENTER WITH WRONG DEPTH!") epc.err = errors.New("called EntryPoint not from top level account")
epc.err = errors.New("same") println(epc.err)
return return
} }
if epc.Input != nil { epc.entries = append(epc.entries, EntryPointCallEntry{
println("repeated call to ep callback") From: from,
epc.err = errors.New("same") Input: bytes.Clone(input),
return })
}
epc.Input = make([]byte, len(input))
copy(epc.Input, input)
} }

View file

@ -9,6 +9,7 @@ import (
"github.com/ethereum/go-ethereum/core/state" "github.com/ethereum/go-ethereum/core/state"
"github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/types"
"github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/core/vm"
"github.com/ethereum/go-ethereum/eth/tracers/native"
"github.com/ethereum/go-ethereum/internal/ethapi" "github.com/ethereum/go-ethereum/internal/ethapi"
"github.com/ethereum/go-ethereum/rpc" "github.com/ethereum/go-ethereum/rpc"
"math/big" "math/big"
@ -132,7 +133,9 @@ func (api *Rip7560API) traceTx(ctx context.Context, tx *types.Transaction, txctx
gp := new(core.GasPool).AddGas(10000000) gp := new(core.GasPool).AddGas(10000000)
// TODO: this is added to allow our bundler checking the 'TraceValidation' API is supported on Geth // TODO: this is added to allow our bundler checking the 'TraceValidation' API is supported on Geth
if tx.Rip7560TransactionData().Sender.Cmp(common.HexToAddress("0x0000000000000000000000000000000000000000")) == 0 { sender := tx.Rip7560TransactionData().Sender
paymaster := tx.Rip7560TransactionData().Paymaster
if sender.Cmp(common.HexToAddress("0x0000000000000000000000000000000000000000")) == 0 {
return tracer.GetResult() return tracer.GetResult()
} }
@ -141,5 +144,7 @@ func (api *Rip7560API) traceTx(ctx context.Context, tx *types.Transaction, txctx
if err != nil { if err != nil {
return nil, fmt.Errorf("tracing failed: %w", err) return nil, fmt.Errorf("tracing failed: %w", err)
} }
x:= tracer.(*native.Rip7560ValidationTracer).EntryPointCall.entries
core.ValidateAccountAndPaymaster(0, sender, paymaster, tracer.)
return tracer.GetResult() return tracer.GetResult()
} }

View file

@ -4,6 +4,7 @@ import (
"encoding/json" "encoding/json"
"github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/common/hexutil" "github.com/ethereum/go-ethereum/common/hexutil"
"github.com/ethereum/go-ethereum/core"
"github.com/ethereum/go-ethereum/core/tracing" "github.com/ethereum/go-ethereum/core/tracing"
"github.com/ethereum/go-ethereum/core/types" "github.com/ethereum/go-ethereum/core/types"
"github.com/ethereum/go-ethereum/core/vm" "github.com/ethereum/go-ethereum/core/vm"
@ -69,7 +70,7 @@ func newRip7560Tracer(ctx *tracers.Context, cfg json.RawMessage) (*tracers.Trace
return nil, err return nil, err
} }
// TODO FIX mock fields // TODO FIX mock fields
t := &rip7560ValidationTracer{ t := &Rip7560ValidationTracer{
TraceResults: make([]stateMap, ValidationFramesMaxCount), TraceResults: make([]stateMap, ValidationFramesMaxCount),
UsedOpcodes: make([]map[string]bool, ValidationFramesMaxCount), UsedOpcodes: make([]map[string]bool, ValidationFramesMaxCount),
Created: make([]map[common.Address]bool, ValidationFramesMaxCount), Created: make([]map[common.Address]bool, ValidationFramesMaxCount),
@ -119,13 +120,13 @@ type logsItem struct {
} }
// Array fields contain of all access details of all validation frames // Array fields contain of all access details of all validation frames
type rip7560ValidationTracer struct { type Rip7560ValidationTracer struct {
//rip7560TxData *types.Rip7560AccountAbstractionTx //rip7560TxData *types.Rip7560AccountAbstractionTx
EntryPointCall core.EntryPointCall
env *tracing.VMContext env *tracing.VMContext
TraceResults []stateMap `json:"traceResults"` TraceResults []stateMap `json:"traceResults"`
UsedOpcodes []map[string]bool `json:"usedOpcodes"` UsedOpcodes []map[string]bool `json:"usedOpcodes"`
Created []map[common.Address]bool `json:"created"` Created []map[common.Address]bool `json:"created"`
//Deleted []map[common.Address]bool `json:"deleted"` //Deleted []map[common.Address]bool `json:"deleted"`
lastThreeOpCodes []*lastThreeOpCodesItem lastThreeOpCodes []*lastThreeOpCodesItem
@ -142,7 +143,9 @@ type rip7560ValidationTracer struct {
//reason error // Textual reason for the interruption //reason error // Textual reason for the interruption
} }
func (b *rip7560ValidationTracer) OnEnter(depth int, typ byte, from common.Address, to common.Address, input []byte, gas uint64, value *big.Int) { func (b *Rip7560ValidationTracer) OnEnter(depth int, typ byte, from common.Address, to common.Address, input []byte, gas uint64, value *big.Int) {
b.entryPointCall.OnEnter(depth, typ, from, to, input, gas, value)
if depth == 0 { if depth == 0 {
b.createNewTopLevelFrame(to) b.createNewTopLevelFrame(to)
} }
@ -157,7 +160,7 @@ func (b *rip7560ValidationTracer) OnEnter(depth int, typ byte, from common.Addre
}) })
} }
func (b *rip7560ValidationTracer) OnExit(depth int, output []byte, gasUsed uint64, err error, reverted bool) { func (b *Rip7560ValidationTracer) OnExit(depth int, output []byte, gasUsed uint64, err error, reverted bool) {
typ := "RETURN" typ := "RETURN"
if err != nil { if err != nil {
typ = "REVERT" typ = "REVERT"
@ -169,12 +172,12 @@ func (b *rip7560ValidationTracer) OnExit(depth int, output []byte, gasUsed uint6
}) })
} }
func (b *rip7560ValidationTracer) OnTxStart(env *tracing.VMContext, tx *types.Transaction, from common.Address) { func (b *Rip7560ValidationTracer) OnTxStart(env *tracing.VMContext, tx *types.Transaction, from common.Address) {
b.env = env b.env = env
//b.rip7560TxData = tx.Rip7560TransactionData() //b.rip7560TxData = tx.Rip7560TransactionData()
} }
func (b *rip7560ValidationTracer) createNewTopLevelFrame(addr common.Address) { func (b *Rip7560ValidationTracer) createNewTopLevelFrame(addr common.Address) {
b.CurrentLevel = &entryPointCall{ b.CurrentLevel = &entryPointCall{
TopLevelTargetAddress: addr, TopLevelTargetAddress: addr,
Access: map[common.Address]*access{}, Access: map[common.Address]*access{},
@ -188,10 +191,10 @@ func (b *rip7560ValidationTracer) createNewTopLevelFrame(addr common.Address) {
return return
} }
func (b *rip7560ValidationTracer) OnTxEnd(receipt *types.Receipt, err error) { func (b *Rip7560ValidationTracer) OnTxEnd(receipt *types.Receipt, err error) {
} }
func (b *rip7560ValidationTracer) OnOpcode(pc uint64, op byte, gas, cost uint64, scope tracing.OpContext, rData []byte, depth int, err error) { func (b *Rip7560ValidationTracer) OnOpcode(pc uint64, op byte, gas, cost uint64, scope tracing.OpContext, rData []byte, depth int, err error) {
opcode := vm.OpCode(op).String() opcode := vm.OpCode(op).String()
stackSize := len(scope.StackData()) stackSize := len(scope.StackData())
@ -344,7 +347,7 @@ func StackBack(stackData []uint256.Int, n int) *uint256.Int {
return &stackData[len(stackData)-n-1] return &stackData[len(stackData)-n-1]
} }
func (b *rip7560ValidationTracer) isEXTorCALL(opcode string) bool { func (b *Rip7560ValidationTracer) isEXTorCALL(opcode string) bool {
return strings.HasPrefix(opcode, "EXT") || return strings.HasPrefix(opcode, "EXT") ||
opcode == "CALL" || opcode == "CALL" ||
opcode == "CALLCODE" || opcode == "CALLCODE" ||
@ -354,22 +357,22 @@ func (b *rip7560ValidationTracer) isEXTorCALL(opcode string) bool {
// not using 'isPrecompiled' to only allow the ones defined by the ERC-7562 as stateless precompiles // not using 'isPrecompiled' to only allow the ones defined by the ERC-7562 as stateless precompiles
// [OP-062] // [OP-062]
func (b *rip7560ValidationTracer) isAllowedPrecompile(addr common.Address) bool { func (b *Rip7560ValidationTracer) isAllowedPrecompile(addr common.Address) bool {
addrInt := addr.Big() addrInt := addr.Big()
return addrInt.Cmp(big.NewInt(0)) == 1 && addrInt.Cmp(big.NewInt(10)) == -1 return addrInt.Cmp(big.NewInt(0)) == 1 && addrInt.Cmp(big.NewInt(10)) == -1
} }
func (b *rip7560ValidationTracer) incrementCount(m map[string]uint64, k string) { func (b *Rip7560ValidationTracer) incrementCount(m map[string]uint64, k string) {
if _, ok := m[k]; !ok { if _, ok := m[k]; !ok {
m[k] = 0 m[k] = 0
} }
m[k]++ m[k]++
} }
func (b *rip7560ValidationTracer) GetResult() (json.RawMessage, error) { func (b *Rip7560ValidationTracer) GetResult() (json.RawMessage, error) {
jsonResult, err := json.MarshalIndent(*b, "", " ") jsonResult, err := json.MarshalIndent(*b, "", " ")
return jsonResult, err return jsonResult, err
} }
func (b *rip7560ValidationTracer) Stop(err error) { func (b *Rip7560ValidationTracer) Stop(err error) {
} }