diff --git a/core/state_processor_rip7560.go b/core/state_processor_rip7560.go index 2fa12e0feb..6a45aab67b 100644 --- a/core/state_processor_rip7560.go +++ b/core/state_processor_rip7560.go @@ -25,8 +25,9 @@ func PackValidationData(authorizerMagic uint64, validUntil, validAfter uint64) [ } type EntryPointCall struct { - caller common.Address - input []byte + OnEnterSuper tracing.EnterHook + Input []byte + err error } type ValidationPhaseResult struct { @@ -43,10 +44,6 @@ type ValidationPhaseResult struct { SenderValidUntil uint64 PmValidAfter uint64 PmValidUntil uint64 - // tracking the calls to the EntryPoint precompile - PmUsed bool - EpCalls []*EntryPointCall - OnEnterSuper tracing.EnterHook } // HandleRip7560Transactions apply state changes of all sequential RIP-7560 transactions and return @@ -187,19 +184,16 @@ func ApplyRip7560ValidationPhases(chainConfig *params.ChainConfig, bc ChainConte GasPrice: gasPrice, } evm := vm.NewEVM(blockContext, txContext, statedb, chainConfig, cfg) - vpr := &ValidationPhaseResult{ - PmUsed: false, - EpCalls: make([]*EntryPointCall, 0), - } + epc := &EntryPointCall{} if evm.Config.Tracer == nil { evm.Config.Tracer = &tracing.Hooks{ - OnEnter: vpr.OnEnter, + OnEnter: epc.OnEnter, } } else { // keep the original tracer's OnEnter hook - vpr.OnEnterSuper = evm.Config.Tracer.OnEnter - evm.Config.Tracer.OnEnter = vpr.OnEnter + epc.OnEnterSuper = evm.Config.Tracer.OnEnter + evm.Config.Tracer.OnEnter = epc.OnEnter } if evm.Config.Tracer.OnTxStart != nil { @@ -242,9 +236,11 @@ func ApplyRip7560ValidationPhases(chainConfig *params.ChainConfig, bc ChainConte if resultAccountValidation.Err != nil { return nil, resultAccountValidation.Err } - aad, err := vpr.validateAccountEntryPointCall() + aad, err := validateAccountEntryPointCall(epc) // clear the EntryPoint calls array after parsing - vpr.EpCalls = make([]*EntryPointCall, 0) + epc.err = nil + epc.Input = nil + if err != nil { return nil, err } @@ -253,7 +249,8 @@ func ApplyRip7560ValidationPhases(chainConfig *params.ChainConfig, bc ChainConte return nil, err } - paymasterContext, pmValidationUsedGas, pmValidAfter, pmValidUntil, err := applyPaymasterValidationFrame(vpr, tx, chainConfig, signingHash, evm, gp, statedb, header) + vpr := &ValidationPhaseResult{} + paymasterContext, pmValidationUsedGas, pmValidAfter, pmValidUntil, err := applyPaymasterValidationFrame(epc, tx, chainConfig, signingHash, evm, gp, statedb, header) if err != nil { return nil, err } @@ -275,14 +272,13 @@ func ApplyRip7560ValidationPhases(chainConfig *params.ChainConfig, bc ChainConte return vpr, nil } -func applyPaymasterValidationFrame(vpr *ValidationPhaseResult, 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(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) { /*** Paymaster Validation Frame ***/ var pmValidationUsedGas uint64 paymasterMsg, err := preparePaymasterValidationMessage(tx, chainConfig, signingHash) if paymasterMsg == nil || err != nil { return nil, 0, 0, 0, err } - vpr.PmUsed = true resultPm, err := ApplyMessage(evm, paymasterMsg, gp) if err != nil { return nil, 0, 0, 0, err @@ -294,7 +290,7 @@ func applyPaymasterValidationFrame(vpr *ValidationPhaseResult, tx *types.Transac return nil, 0, 0, 0, errors.New("paymaster validation failed - invalid transaction") } pmValidationUsedGas = resultPm.UsedGas - apd, err := vpr.validatePaymasterEntryPointCall() + apd, err := validatePaymasterEntryPointCall(epc) if err != nil { return nil, 0, 0, 0, err } @@ -477,34 +473,31 @@ func preparePostOpMessage(vpr *ValidationPhaseResult, chainConfig *params.ChainC }, nil } -func (vpr *ValidationPhaseResult) validateAccountEntryPointCall() (*AcceptAccountData, error) { - if len(vpr.EpCalls) == 0 { +func validateAccountEntryPointCall(epc *EntryPointCall) (*AcceptAccountData, error) { + if epc.err != nil { + return nil, epc.err + } + if epc.Input == nil { return nil, errors.New("account validation did not call the EntryPoint callback") } - if len(vpr.EpCalls) > 1 { - return nil, errors.New("account validation illegally called the EntryPoint callback multiple times") - } - epCall := vpr.EpCalls[0] - - if len(epCall.input) != 68 { + if len(epc.Input) != 68 { return nil, errors.New("invalid account return data length") } - return abiDecodeAcceptAccount(epCall.input) + return abiDecodeAcceptAccount(epc.Input) } -func (vpr *ValidationPhaseResult) validatePaymasterEntryPointCall() (*AcceptPaymasterData, error) { - if len(vpr.EpCalls) == 0 { +func validatePaymasterEntryPointCall(epc *EntryPointCall) (*AcceptPaymasterData, error) { + if epc.err != nil { + return nil, epc.err + } + if epc.Input == nil { return nil, errors.New("paymaster validation did not call the EntryPoint callback") } - if vpr.PmUsed && len(vpr.EpCalls) > 1 { - return nil, errors.New("paymaster validation illegally called the EntryPoint callback multiple times") - } - epCall := vpr.EpCalls[0] - if len(epCall.input) < 100 { + if len(epc.Input) < 100 { return nil, errors.New("invalid paymaster callback data length") } - apd, err := abiDecodeAcceptPaymaster(epCall.input) + apd, err := abiDecodeAcceptPaymaster(epc.Input) if err != nil { return nil, err } @@ -527,17 +520,26 @@ func validateValidityTimeRange(time uint64, validAfter uint64, validUntil uint64 return nil } -func (vpr *ValidationPhaseResult) OnEnter(depth int, typ byte, from common.Address, to common.Address, input []byte, gas uint64, value *big.Int) { - if vpr.OnEnterSuper != nil { - vpr.OnEnterSuper(depth, typ, from, to, input, gas, value) +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 - if isRip7560EntryPoint { - inputBytes := make([]byte, len(input)) - copy(inputBytes, input) - vpr.EpCalls = append(vpr.EpCalls, &EntryPointCall{ - caller: from, - input: inputBytes, - }) + if !isRip7560EntryPoint { + return } + + if depth != 1 { + println("ONENTER WITH WRONG DEPTH!") + epc.err = errors.New("same") + return + } + if epc.Input != nil { + println("repeated call to ep callback") + epc.err = errors.New("same") + return + } + + epc.Input = make([]byte, len(input)) + copy(epc.Input, input) }