WIP: Flatten the "EpCall" to be a single pointer

This commit is contained in:
Alex Forshtat 2024-08-11 19:10:43 +02:00
parent 1c42b62be1
commit f5aef8cceb

View file

@ -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)
}