core, internal: fix storage override

This commit is contained in:
Gary Rong 2024-07-18 12:11:34 +08:00
parent f59d013e40
commit 47bb7da7bf
2 changed files with 53 additions and 10 deletions

View file

@ -471,20 +471,28 @@ func (s *StateDB) SetState(addr common.Address, key, value common.Hash) {
// storage. This function should only be used for debugging and the mutations // storage. This function should only be used for debugging and the mutations
// must be discarded afterwards. // must be discarded afterwards.
func (s *StateDB) SetStorage(addr common.Address, storage map[common.Hash]common.Hash) { func (s *StateDB) SetStorage(addr common.Address, storage map[common.Hash]common.Hash) {
// SetStorage needs to wipe existing storage. We achieve this by pretending // SetStorage needs to wipe the existing storage. We achieve this by marking
// that the account self-destructed earlier in this block, by flagging // the account as self-destructed in this block. The effect is that storage
// it in stateObjectsDestruct. The effect of doing so is that storage lookups // lookups will not hit the disk, as it is assumed that the disk data belongs
// will not hit disk, since it is assumed that the disk-data is belonging
// to a previous incarnation of the object. // to a previous incarnation of the object.
// //
// TODO(rjl493456442) this function should only be supported by 'unwritable' // TODO (rjl493456442): This function should only be supported by 'unwritable'
// state and all mutations made should all be discarded afterwards. // state, and all mutations made should be discarded afterward.
obj := s.getStateObject(addr)
if obj != nil {
if _, ok := s.stateObjectsDestruct[addr]; !ok { if _, ok := s.stateObjectsDestruct[addr]; !ok {
s.stateObjectsDestruct[addr] = nil s.stateObjectsDestruct[addr] = obj
} }
stateObject := s.getOrNewStateObject(addr) }
newObj := s.createObject(addr)
for k, v := range storage { for k, v := range storage {
stateObject.SetState(k, v) newObj.SetState(k, v)
}
// Inherit the metadata of original object if it was existent
if obj != nil {
newObj.SetCode(common.BytesToHash(obj.CodeHash()), obj.code)
newObj.SetNonce(obj.Nonce())
newObj.SetBalance(obj.Balance(), tracing.BalanceChangeUnspecified)
} }
} }

View file

@ -781,15 +781,24 @@ func TestEstimateGas(t *testing.T) {
func TestCall(t *testing.T) { func TestCall(t *testing.T) {
t.Parallel() t.Parallel()
// Initialize test accounts // Initialize test accounts
var ( var (
accounts = newAccounts(3) accounts = newAccounts(3)
dad = common.HexToAddress("0x0000000000000000000000000000000000000dad")
genesis = &core.Genesis{ genesis = &core.Genesis{
Config: params.MergedTestChainConfig, Config: params.MergedTestChainConfig,
Alloc: types.GenesisAlloc{ Alloc: types.GenesisAlloc{
accounts[0].addr: {Balance: big.NewInt(params.Ether)}, accounts[0].addr: {Balance: big.NewInt(params.Ether)},
accounts[1].addr: {Balance: big.NewInt(params.Ether)}, accounts[1].addr: {Balance: big.NewInt(params.Ether)},
accounts[2].addr: {Balance: big.NewInt(params.Ether)}, accounts[2].addr: {Balance: big.NewInt(params.Ether)},
dad: {
Balance: big.NewInt(params.Ether),
Nonce: 1,
Storage: map[common.Hash]common.Hash{
common.Hash{}: common.HexToHash("0x0000000000000000000000000000000000000000000000000000000000000001"),
},
},
}, },
} }
genBlocks = 10 genBlocks = 10
@ -949,6 +958,32 @@ func TestCall(t *testing.T) {
}, },
want: "0x0122000000000000000000000000000000000000000000000000000000000000", want: "0x0122000000000000000000000000000000000000000000000000000000000000",
}, },
// Clear the entire storage set
{
blockNumber: rpc.LatestBlockNumber,
call: TransactionArgs{
From: &accounts[1].addr,
// Yul:
// object "Test" {
// code {
// let dad := 0x0000000000000000000000000000000000000dad
// if eq(balance(dad), 0) {
// revert(0, 0)
// }
// let slot := sload(0)
// mstore(0, slot)
// return(0, 32)
// }
// }
Input: hex2Bytes("610dad6000813103600f57600080fd5b6000548060005260206000f3"),
},
overrides: StateOverride{
dad: OverrideAccount{
State: &map[common.Hash]common.Hash{},
},
},
want: "0x0000000000000000000000000000000000000000000000000000000000000000",
},
} }
for i, tc := range testSuite { for i, tc := range testSuite {
result, err := api.Call(context.Background(), tc.call, &rpc.BlockNumberOrHash{BlockNumber: &tc.blockNumber}, &tc.overrides, &tc.blockOverrides) result, err := api.Call(context.Background(), tc.call, &rpc.BlockNumberOrHash{BlockNumber: &tc.blockNumber}, &tc.overrides, &tc.blockOverrides)