From d1b9c1cf1b2a54281240c9ef0d826406dc9d59ac Mon Sep 17 00:00:00 2001 From: bufrr Date: Tue, 20 Feb 2024 17:03:02 +0800 Subject: [PATCH] all: update error type checking --- accounts/keystore/keystore_test.go | 9 +++++---- accounts/keystore/plain_test.go | 3 ++- accounts/usbwallet/ledger.go | 4 ++-- beacon/light/committee_chain_test.go | 7 ++++--- cmd/evm/main.go | 4 +++- cmd/geth/accountcmd.go | 6 ++++-- cmd/geth/config.go | 3 ++- common/bitutil/compress_test.go | 3 ++- common/hexutil/hexutil.go | 15 +++++++++------ common/hexutil/json.go | 6 ++++-- common/test_utils.go | 4 +++- console/bridge.go | 10 ++++++---- console/console.go | 6 ++++-- core/blockchain.go | 10 ++++++---- core/blockchain_test.go | 2 +- core/rawdb/freezer_table_test.go | 6 ++++-- core/rawdb/freezer_test.go | 2 +- core/state/snapshot/disklayer_test.go | 5 +++-- core/state/snapshot/generate.go | 3 ++- core/state/snapshot/snapshot_test.go | 13 +++++++------ core/state/statedb_fuzz_test.go | 5 ++--- core/state/statedb_test.go | 3 ++- core/types/receipt_test.go | 3 ++- core/types/transaction_test.go | 2 +- core/vm/evm.go | 13 +++++++------ core/vm/gas_table_test.go | 5 +++-- core/vm/instructions.go | 17 +++++++++-------- core/vm/interpreter.go | 3 ++- crypto/crypto.go | 3 ++- crypto/crypto_test.go | 5 +++-- crypto/ecies/ecies_test.go | 6 +++--- crypto/secp256k1/secp256_test.go | 3 ++- eth/downloader/downloader.go | 2 +- eth/downloader/downloader_test.go | 7 ++++--- eth/downloader/skeleton.go | 8 ++++---- eth/filters/filter_system_test.go | 2 +- eth/state_accessor.go | 5 +++-- ethclient/ethclient_test.go | 2 +- ethdb/leveldb/leveldb.go | 6 ++++-- event/event_test.go | 3 ++- event/subscription_test.go | 2 +- internal/build/util.go | 4 +++- internal/cmdtest/test_cmd.go | 4 +++- internal/jsre/pretty.go | 4 +++- node/errors.go | 3 ++- node/node_test.go | 8 ++++---- p2p/dial.go | 6 ++++-- p2p/discover/v4_udp_test.go | 8 ++++---- p2p/discover/v5_udp_test.go | 13 +++++++------ p2p/discover/v5wire/encoding.go | 2 +- p2p/dnsdisc/error.go | 3 ++- p2p/enode/nodedb.go | 13 +++++++------ p2p/enr/enr_test.go | 11 ++++++----- p2p/enr/entries.go | 4 ++-- p2p/message_test.go | 3 ++- p2p/netutil/error_test.go | 4 +++- p2p/peer.go | 3 ++- p2p/peer_error.go | 6 ++++-- p2p/peer_test.go | 2 +- p2p/transport.go | 4 +++- rlp/decode.go | 14 ++++++++------ rlp/decode_test.go | 8 ++++---- rlp/typecache.go | 4 +++- rpc/client.go | 3 ++- rpc/client_test.go | 12 +++++++----- rpc/http_test.go | 8 +++++--- rpc/json.go | 6 ++++-- rpc/subscription.go | 4 ++-- tests/init_test.go | 3 ++- trie/iterator.go | 12 +++++++----- trie/node.go | 4 +++- trie/node_test.go | 7 +++++-- trie/trie_test.go | 15 +++++++++------ 73 files changed, 257 insertions(+), 174 deletions(-) diff --git a/accounts/keystore/keystore_test.go b/accounts/keystore/keystore_test.go index c9a23eddd6..c6defd945c 100644 --- a/accounts/keystore/keystore_test.go +++ b/accounts/keystore/keystore_test.go @@ -17,6 +17,7 @@ package keystore import ( + "errors" "math/rand" "os" "runtime" @@ -127,7 +128,7 @@ func TestTimedUnlock(t *testing.T) { // Signing without passphrase fails because account is locked _, err = ks.SignHash(accounts.Account{Address: a1.Address}, testSigData) - if err != ErrLocked { + if !errors.Is(err, ErrLocked) { t.Fatal("Signing should've failed with ErrLocked before unlocking, got ", err) } @@ -145,7 +146,7 @@ func TestTimedUnlock(t *testing.T) { // Signing fails again after automatic locking time.Sleep(250 * time.Millisecond) _, err = ks.SignHash(accounts.Account{Address: a1.Address}, testSigData) - if err != ErrLocked { + if !errors.Is(err, ErrLocked) { t.Fatal("Signing should've failed with ErrLocked timeout expired, got ", err) } } @@ -185,7 +186,7 @@ func TestOverrideUnlock(t *testing.T) { // Signing fails again after automatic locking time.Sleep(250 * time.Millisecond) _, err = ks.SignHash(accounts.Account{Address: a1.Address}, testSigData) - if err != ErrLocked { + if !errors.Is(err, ErrLocked) { t.Fatal("Signing should've failed with ErrLocked timeout expired, got ", err) } } @@ -206,7 +207,7 @@ func TestSignRace(t *testing.T) { } end := time.Now().Add(500 * time.Millisecond) for time.Now().Before(end) { - if _, err := ks.SignHash(accounts.Account{Address: a1.Address}, testSigData); err == ErrLocked { + if _, err := ks.SignHash(accounts.Account{Address: a1.Address}, testSigData); errors.Is(err, ErrLocked) { return } else if err != nil { t.Errorf("Sign error: %v", err) diff --git a/accounts/keystore/plain_test.go b/accounts/keystore/plain_test.go index 737eb7fd61..c90982876e 100644 --- a/accounts/keystore/plain_test.go +++ b/accounts/keystore/plain_test.go @@ -19,6 +19,7 @@ package keystore import ( "crypto/rand" "encoding/hex" + "errors" "fmt" "path/filepath" "reflect" @@ -90,7 +91,7 @@ func TestKeyStorePassphraseDecryptionFail(t *testing.T) { if err != nil { t.Fatal(err) } - if _, err = ks.GetKey(k1.Address, account.URL.Path, "bar"); err != ErrDecrypt { + if _, err = ks.GetKey(k1.Address, account.URL.Path, "bar"); !errors.Is(err, ErrDecrypt) { t.Fatalf("wrong error for invalid password\ngot %q\nwant %q", err, ErrDecrypt) } } diff --git a/accounts/usbwallet/ledger.go b/accounts/usbwallet/ledger.go index d0cb93e74e..15c8a04d9a 100644 --- a/accounts/usbwallet/ledger.go +++ b/accounts/usbwallet/ledger.go @@ -119,7 +119,7 @@ func (w *ledgerDriver) Open(device io.ReadWriter, passphrase string) error { _, err := w.ledgerDerive(accounts.DefaultBaseDerivationPath) if err != nil { // Ethereum app is not running or in browser mode, nothing more to do, return - if err == errLedgerReplyInvalidHeader { + if errors.Is(err, errLedgerReplyInvalidHeader) { w.browser = true } return nil @@ -141,7 +141,7 @@ func (w *ledgerDriver) Close() error { // Heartbeat implements usbwallet.driver, performing a sanity check against the // Ledger to see if it's still online. func (w *ledgerDriver) Heartbeat() error { - if _, err := w.ledgerVersion(); err != nil && err != errLedgerInvalidVersionReply { + if _, err := w.ledgerVersion(); err != nil && !errors.Is(err, errLedgerInvalidVersionReply) { w.failure = err return err } diff --git a/beacon/light/committee_chain_test.go b/beacon/light/committee_chain_test.go index 60ea2a0efd..fc5e5ec8a8 100644 --- a/beacon/light/committee_chain_test.go +++ b/beacon/light/committee_chain_test.go @@ -18,6 +18,7 @@ package light import ( "crypto/rand" + "errors" "testing" "time" @@ -259,13 +260,13 @@ func (c *committeeChainTest) setClockPeriod(period float64) { } func (c *committeeChainTest) addFixedCommitteeRoot(tc *testCommitteeChain, period uint64, expErr error) { - if err := c.chain.addFixedCommitteeRoot(period, tc.periods[period].committee.Root()); err != expErr { + if err := c.chain.addFixedCommitteeRoot(period, tc.periods[period].committee.Root()); !errors.Is(err, expErr) { c.t.Errorf("Incorrect error output from addFixedCommitteeRoot at period %d (expected %v, got %v)", period, expErr, err) } } func (c *committeeChainTest) addCommittee(tc *testCommitteeChain, period uint64, expErr error) { - if err := c.chain.addCommittee(period, tc.periods[period].committee); err != expErr { + if err := c.chain.addCommittee(period, tc.periods[period].committee); !errors.Is(err, expErr) { c.t.Errorf("Incorrect error output from addCommittee at period %d (expected %v, got %v)", period, expErr, err) } } @@ -275,7 +276,7 @@ func (c *committeeChainTest) insertUpdate(tc *testCommitteeChain, period uint64, if addCommittee { committee = tc.periods[period+1].committee } - if err := c.chain.InsertUpdate(tc.periods[period].update, committee); err != expErr { + if err := c.chain.InsertUpdate(tc.periods[period].update, committee); !errors.Is(err, expErr) { c.t.Errorf("Incorrect error output from InsertUpdate at period %d (expected %v, got %v)", period, expErr, err) } } diff --git a/cmd/evm/main.go b/cmd/evm/main.go index c3e6a4af91..eed94ab766 100644 --- a/cmd/evm/main.go +++ b/cmd/evm/main.go @@ -18,6 +18,7 @@ package main import ( + "errors" "fmt" "math/big" "os" @@ -248,7 +249,8 @@ func init() { func main() { if err := app.Run(os.Args); err != nil { code := 1 - if ec, ok := err.(*t8ntool.NumberedError); ok { + var ec *t8ntool.NumberedError + if errors.As(err, &ec) { code = ec.ExitCode() } fmt.Fprintln(os.Stderr, err) diff --git a/cmd/geth/accountcmd.go b/cmd/geth/accountcmd.go index cc22684e0b..7170bab3a6 100644 --- a/cmd/geth/accountcmd.go +++ b/cmd/geth/accountcmd.go @@ -17,6 +17,7 @@ package main import ( + "errors" "fmt" "os" @@ -233,11 +234,12 @@ func unlockAccount(ks *keystore.KeyStore, address string, i int, passwords []str log.Info("Unlocked account", "address", account.Address.Hex()) return account, password } - if err, ok := err.(*keystore.AmbiguousAddrError); ok { + var err *keystore.AmbiguousAddrError + if errors.As(err, &err) { log.Info("Unlocked account", "address", account.Address.Hex()) return ambiguousAddrRecovery(ks, err, password), password } - if err != keystore.ErrDecrypt { + if !errors.Is(err, keystore.ErrDecrypt) { // No need to prompt again if the error is not decryption-related. break } diff --git a/cmd/geth/config.go b/cmd/geth/config.go index 5f52f1df54..1548fb8d4a 100644 --- a/cmd/geth/config.go +++ b/cmd/geth/config.go @@ -106,7 +106,8 @@ func loadConfig(file string, cfg *gethConfig) error { err = tomlSettings.NewDecoder(bufio.NewReader(f)).Decode(cfg) // Add file name to errors that have a line number. - if _, ok := err.(*toml.LineError); ok { + var lineError *toml.LineError + if errors.As(err, &lineError) { err = errors.New(file + ", " + err.Error()) } return err diff --git a/common/bitutil/compress_test.go b/common/bitutil/compress_test.go index c6f6fe8bcf..b58c9d11df 100644 --- a/common/bitutil/compress_test.go +++ b/common/bitutil/compress_test.go @@ -18,6 +18,7 @@ package bitutil import ( "bytes" + "errors" "fmt" "math/rand" "testing" @@ -143,7 +144,7 @@ func TestCompression(t *testing.T) { t.Errorf("decoding mismatch for dense data: have %x, want %x, error %v", data, in, err) } // Check that decompressing a longer input than the target fails - if _, err := DecompressBytes([]byte{0xc0, 0x01, 0x01}, 2); err != errExceededTarget { + if _, err := DecompressBytes([]byte{0xc0, 0x01, 0x01}, 2); !errors.Is(err, errExceededTarget) { t.Errorf("decoding error mismatch for long data: have %v, want %v", err, errExceededTarget) } } diff --git a/common/hexutil/hexutil.go b/common/hexutil/hexutil.go index d3201850a8..07b3932c15 100644 --- a/common/hexutil/hexutil.go +++ b/common/hexutil/hexutil.go @@ -32,6 +32,7 @@ package hexutil import ( "encoding/hex" + "errors" "fmt" "math/big" "strconv" @@ -223,18 +224,20 @@ func decodeNibble(in byte) uint64 { } func mapError(err error) error { - if err, ok := err.(*strconv.NumError); ok { - switch err.Err { - case strconv.ErrRange: + var numErr *strconv.NumError + if errors.As(err, &numErr) { + switch { + case errors.Is(numErr.Err, strconv.ErrRange): return ErrUint64Range - case strconv.ErrSyntax: + case errors.Is(numErr.Err, strconv.ErrSyntax): return ErrSyntax } } - if _, ok := err.(hex.InvalidByteError); ok { + var invalidByteError hex.InvalidByteError + if errors.As(err, &invalidByteError) { return ErrSyntax } - if err == hex.ErrLength { + if errors.Is(err, hex.ErrLength) { return ErrOddLength } return err diff --git a/common/hexutil/json.go b/common/hexutil/json.go index e0ac98f52d..795af4b593 100644 --- a/common/hexutil/json.go +++ b/common/hexutil/json.go @@ -19,6 +19,7 @@ package hexutil import ( "encoding/hex" "encoding/json" + "errors" "fmt" "math/big" "reflect" @@ -355,7 +356,7 @@ func (b *Uint) UnmarshalJSON(input []byte) error { func (b *Uint) UnmarshalText(input []byte) error { var u64 Uint64 err := u64.UnmarshalText(input) - if u64 > Uint64(^uint(0)) || err == ErrUint64Range { + if u64 > Uint64(^uint(0)) || errors.Is(err, ErrUint64Range) { return ErrUintRange } else if err != nil { return err @@ -410,7 +411,8 @@ func checkNumberText(input []byte) (raw []byte, err error) { } func wrapTypeError(err error, typ reflect.Type) error { - if _, ok := err.(*decError); ok { + var de *decError + if errors.As(err, &de) { return &json.UnmarshalTypeError{Value: err.Error(), Type: typ} } return err diff --git a/common/test_utils.go b/common/test_utils.go index 7a175412f4..145000e23d 100644 --- a/common/test_utils.go +++ b/common/test_utils.go @@ -18,6 +18,7 @@ package common import ( "encoding/json" + "errors" "fmt" "os" ) @@ -29,7 +30,8 @@ func LoadJSON(file string, val interface{}) error { return err } if err := json.Unmarshal(content, val); err != nil { - if syntaxerr, ok := err.(*json.SyntaxError); ok { + var syntaxerr *json.SyntaxError + if errors.As(err, &syntaxerr) { line := findLine(content, syntaxerr.Offset) return fmt.Errorf("JSON syntax error at %v:%v: %v", file, line, err) } diff --git a/console/bridge.go b/console/bridge.go index 37578041ca..0fbc5c96e9 100644 --- a/console/bridge.go +++ b/console/bridge.go @@ -434,11 +434,13 @@ func (b *bridge) Send(call jsre.Call) (goja.Value, error) { } else { code := -32603 var data interface{} - if err, ok := err.(rpc.Error); ok { - code = err.ErrorCode() + var rpcErr rpc.Error + if errors.As(err, &rpcErr) { + code = rpcErr.ErrorCode() } - if err, ok := err.(rpc.DataError); ok { - data = err.ErrorData() + var dataErr rpc.DataError + if errors.As(err, &dataErr) { + data = dataErr.ErrorData() } setError(resp, code, err.Error(), data) } diff --git a/console/console.go b/console/console.go index cdee53684e..4aec967c3b 100644 --- a/console/console.go +++ b/console/console.go @@ -149,7 +149,8 @@ func (c *Console) init(preload []string) error { for _, path := range preload { if err := c.jsre.Exec(path); err != nil { failure := err.Error() - if gojaErr, ok := err.(*goja.Exception); ok { + var gojaErr *goja.Exception + if errors.As(err, &gojaErr) { failure = gojaErr.String() } return fmt.Errorf("%s: %v", path, failure) @@ -206,7 +207,8 @@ func (c *Console) initExtensions() error { const methodNotFound = -32601 apis, err := c.client.SupportedModules() if err != nil { - if rpcErr, ok := err.(rpc.Error); ok && rpcErr.ErrorCode() == methodNotFound { + var rpcErr rpc.Error + if errors.As(err, &rpcErr) && rpcErr.ErrorCode() == methodNotFound { log.Warn("Server does not support method rpc_modules, using default API list.") apis = defaultAPIs } else { diff --git a/core/blockchain.go b/core/blockchain.go index b1bbc3d598..b292adfed9 100644 --- a/core/blockchain.go +++ b/core/blockchain.go @@ -275,7 +275,8 @@ func NewBlockChain(db ethdb.Database, cacheConfig *CacheConfig, genesis *Genesis // to database if the genesis block is not present yet, or load the // stored one from database. chainConfig, genesisHash, genesisErr := SetupGenesisBlockWithOverride(db, triedb, genesis, overrides) - if _, ok := genesisErr.(*params.ConfigCompatError); genesisErr != nil && !ok { + var configCompatError *params.ConfigCompatError + if genesisErr != nil && !errors.As(genesisErr, &configCompatError) { return nil, genesisErr } log.Info("") @@ -455,7 +456,8 @@ func NewBlockChain(db ethdb.Database, cacheConfig *CacheConfig, genesis *Genesis go bc.updateFutureBlocks() // Rewind the chain in case of an incompatible config upgrade. - if compat, ok := genesisErr.(*params.ConfigCompatError); ok { + var compat *params.ConfigCompatError + if errors.As(genesisErr, &compat) { log.Warn("Rewinding chain to upgrade configuration", "err", compat) if compat.RewindToTime > 0 { bc.SetHeadWithTimestamp(compat.RewindToTime) @@ -1282,7 +1284,7 @@ func (bc *BlockChain) InsertReceiptChain(blockChain types.Blocks, receiptChain [ // Write downloaded chain data and corresponding receipt chain data if len(ancientBlocks) > 0 { if n, err := writeAncient(ancientBlocks, ancientReceipts); err != nil { - if err == errInsertionInterrupted { + if errors.Is(err, errInsertionInterrupted) { return 0, nil } return n, err @@ -1290,7 +1292,7 @@ func (bc *BlockChain) InsertReceiptChain(blockChain types.Blocks, receiptChain [ } if len(liveBlocks) > 0 { if n, err := writeLive(liveBlocks, liveReceipts); err != nil { - if err == errInsertionInterrupted { + if errors.Is(err, errInsertionInterrupted) { return 0, nil } return n, err diff --git a/core/blockchain_test.go b/core/blockchain_test.go index 876d662f74..e4054b654f 100644 --- a/core/blockchain_test.go +++ b/core/blockchain_test.go @@ -154,7 +154,7 @@ func testBlockChainImport(chain types.Blocks, blockchain *BlockChain) error { err = blockchain.validator.ValidateBody(block) } if err != nil { - if err == ErrKnownBlock { + if errors.Is(err, ErrKnownBlock) { continue } return err diff --git a/core/rawdb/freezer_table_test.go b/core/rawdb/freezer_table_test.go index 91b4943e59..cff986d81c 100644 --- a/core/rawdb/freezer_table_test.go +++ b/core/rawdb/freezer_table_test.go @@ -19,6 +19,7 @@ package rawdb import ( "bytes" "encoding/binary" + "errors" "fmt" "math/rand" "os" @@ -68,7 +69,7 @@ func TestFreezerBasics(t *testing.T) { } // Check that we cannot read too far _, err = f.Retrieve(uint64(255)) - if err != errOutOfBounds { + if !errors.Is(err, errOutOfBounds) { t.Fatal(err) } } @@ -1361,7 +1362,8 @@ func runRandTest(rt randTest) bool { func TestRandom(t *testing.T) { if err := quick.Check(runRandTest, nil); err != nil { - if cerr, ok := err.(*quick.CheckError); ok { + var cerr *quick.CheckError + if errors.As(err, &cerr) { t.Fatalf("random test iteration %d failed: %s", cerr.Count, spew.Sdump(cerr.In)) } t.Fatal(err) diff --git a/core/rawdb/freezer_test.go b/core/rawdb/freezer_test.go index b4bd6a382a..89b1251fce 100644 --- a/core/rawdb/freezer_test.go +++ b/core/rawdb/freezer_test.go @@ -373,7 +373,7 @@ func checkAncientCount(t *testing.T, f *Freezer, kind string, n uint64) { } if _, err := f.Ancient(kind, index); err == nil { t.Errorf("Ancient(%q, %d) didn't return expected error", kind, index) - } else if err != errOutOfBounds { + } else if !errors.Is(err, errOutOfBounds) { t.Errorf("Ancient(%q, %d) returned unexpected error %q", kind, index, err) } } diff --git a/core/state/snapshot/disklayer_test.go b/core/state/snapshot/disklayer_test.go index 168458c405..87b65b393f 100644 --- a/core/state/snapshot/disklayer_test.go +++ b/core/state/snapshot/disklayer_test.go @@ -18,6 +18,7 @@ package snapshot import ( "bytes" + "errors" "testing" "github.com/VictoriaMetrics/fastcache" @@ -311,7 +312,7 @@ func TestDiskPartialMerge(t *testing.T) { assertAccount := func(account common.Hash, data []byte) { t.Helper() blob, err := base.AccountRLP(account) - if bytes.Compare(account[:], genMarker) > 0 && err != ErrNotCoveredYet { + if bytes.Compare(account[:], genMarker) > 0 && !errors.Is(err, ErrNotCoveredYet) { t.Fatalf("test %d: post-marker (%x) account access (%x) succeeded: %x", i, genMarker, account, blob) } if bytes.Compare(account[:], genMarker) <= 0 && !bytes.Equal(blob, data) { @@ -327,7 +328,7 @@ func TestDiskPartialMerge(t *testing.T) { assertStorage := func(account common.Hash, slot common.Hash, data []byte) { t.Helper() blob, err := base.Storage(account, slot) - if bytes.Compare(append(account[:], slot[:]...), genMarker) > 0 && err != ErrNotCoveredYet { + if bytes.Compare(append(account[:], slot[:]...), genMarker) > 0 && !errors.Is(err, ErrNotCoveredYet) { t.Fatalf("test %d: post-marker (%x) storage access (%x:%x) succeeded: %x", i, genMarker, account, slot, blob) } if bytes.Compare(append(account[:], slot[:]...), genMarker) <= 0 && !bytes.Equal(blob, data) { diff --git a/core/state/snapshot/generate.go b/core/state/snapshot/generate.go index 8de4b134d3..c075dfe713 100644 --- a/core/state/snapshot/generate.go +++ b/core/state/snapshot/generate.go @@ -687,7 +687,8 @@ func (dl *diskLayer) generate(stats *generatorStats) { if err := generateAccounts(ctx, dl, accMarker); err != nil { // Extract the received interruption signal if exists - if aerr, ok := err.(*abortErr); ok { + var aerr *abortErr + if errors.As(err, &aerr) { abort = aerr.abort } // Aborted by internal error, wait the signal diff --git a/core/state/snapshot/snapshot_test.go b/core/state/snapshot/snapshot_test.go index a9ab3eaea3..2d549d654b 100644 --- a/core/state/snapshot/snapshot_test.go +++ b/core/state/snapshot/snapshot_test.go @@ -19,6 +19,7 @@ package snapshot import ( crand "crypto/rand" "encoding/binary" + "errors" "fmt" "math/rand" "testing" @@ -118,10 +119,10 @@ func TestDiskLayerExternalInvalidationFullFlatten(t *testing.T) { t.Fatalf("failed to merge diff layer onto disk: %v", err) } // Since the base layer was modified, ensure that data retrievals on the external reference fail - if acc, err := ref.Account(common.HexToHash("0x01")); err != ErrSnapshotStale { + if acc, err := ref.Account(common.HexToHash("0x01")); !errors.Is(err, ErrSnapshotStale) { t.Errorf("stale reference returned account: %#x (err: %v)", acc, err) } - if slot, err := ref.Storage(common.HexToHash("0xa1"), common.HexToHash("0xb1")); err != ErrSnapshotStale { + if slot, err := ref.Storage(common.HexToHash("0xa1"), common.HexToHash("0xb1")); !errors.Is(err, ErrSnapshotStale) { t.Errorf("stale reference returned storage slot: %#x (err: %v)", slot, err) } if n := len(snaps.layers); n != 1 { @@ -168,10 +169,10 @@ func TestDiskLayerExternalInvalidationPartialFlatten(t *testing.T) { t.Fatalf("failed to merge accumulator onto disk: %v", err) } // Since the base layer was modified, ensure that data retrievals on the external reference fail - if acc, err := ref.Account(common.HexToHash("0x01")); err != ErrSnapshotStale { + if acc, err := ref.Account(common.HexToHash("0x01")); !errors.Is(err, ErrSnapshotStale) { t.Errorf("stale reference returned account: %#x (err: %v)", acc, err) } - if slot, err := ref.Storage(common.HexToHash("0xa1"), common.HexToHash("0xb1")); err != ErrSnapshotStale { + if slot, err := ref.Storage(common.HexToHash("0xa1"), common.HexToHash("0xb1")); !errors.Is(err, ErrSnapshotStale) { t.Errorf("stale reference returned storage slot: %#x (err: %v)", slot, err) } if n := len(snaps.layers); n != 2 { @@ -230,10 +231,10 @@ func TestDiffLayerExternalInvalidationPartialFlatten(t *testing.T) { t.Fatalf("failed to flatten diff layer into accumulator: %v", err) } // Since the accumulator diff layer was modified, ensure that data retrievals on the external reference fail - if acc, err := ref.Account(common.HexToHash("0x01")); err != ErrSnapshotStale { + if acc, err := ref.Account(common.HexToHash("0x01")); !errors.Is(err, ErrSnapshotStale) { t.Errorf("stale reference returned account: %#x (err: %v)", acc, err) } - if slot, err := ref.Storage(common.HexToHash("0xa1"), common.HexToHash("0xb1")); err != ErrSnapshotStale { + if slot, err := ref.Storage(common.HexToHash("0xa1"), common.HexToHash("0xb1")); !errors.Is(err, ErrSnapshotStale) { t.Errorf("stale reference returned storage slot: %#x (err: %v)", slot, err) } if n := len(snaps.layers); n != 3 { diff --git a/core/state/statedb_fuzz_test.go b/core/state/statedb_fuzz_test.go index b416bcf1f3..506b436e72 100644 --- a/core/state/statedb_fuzz_test.go +++ b/core/state/statedb_fuzz_test.go @@ -384,10 +384,9 @@ func (test *stateTest) verify(root common.Hash, next common.Hash, db *triedb.Dat func TestStateChanges(t *testing.T) { config := &quick.Config{MaxCount: 1000} err := quick.Check((*stateTest).run, config) - if cerr, ok := err.(*quick.CheckError); ok { + var cerr *quick.CheckError + if errors.As(err, &cerr) { test := cerr.In[0].(*stateTest) t.Errorf("%v:\n%s", test.err, test) - } else if err != nil { - t.Error(err) } } diff --git a/core/state/statedb_test.go b/core/state/statedb_test.go index cd86a7f4b6..ddfd105a47 100644 --- a/core/state/statedb_test.go +++ b/core/state/statedb_test.go @@ -227,7 +227,8 @@ func TestCopy(t *testing.T) { func TestSnapshotRandom(t *testing.T) { config := &quick.Config{MaxCount: 1000} err := quick.Check((*snapshotTest).run, config) - if cerr, ok := err.(*quick.CheckError); ok { + var cerr *quick.CheckError + if errors.As(err, &cerr) { test := cerr.In[0].(*snapshotTest) t.Errorf("%v:\n%s", test.err, test) } else if err != nil { diff --git a/core/types/receipt_test.go b/core/types/receipt_test.go index a7b2644471..1224d1275c 100644 --- a/core/types/receipt_test.go +++ b/core/types/receipt_test.go @@ -19,6 +19,7 @@ package types import ( "bytes" "encoding/json" + "errors" "math" "math/big" "reflect" @@ -300,7 +301,7 @@ func TestDecodeEmptyTypedReceipt(t *testing.T) { input := []byte{0x80} var r Receipt err := rlp.DecodeBytes(input, &r) - if err != errShortTypedReceipt { + if !errors.Is(err, errShortTypedReceipt) { t.Fatal("wrong error:", err) } } diff --git a/core/types/transaction_test.go b/core/types/transaction_test.go index 76a010d2e5..9001e88be4 100644 --- a/core/types/transaction_test.go +++ b/core/types/transaction_test.go @@ -75,7 +75,7 @@ func TestDecodeEmptyTypedTx(t *testing.T) { input := []byte{0x80} var tx Transaction err := rlp.DecodeBytes(input, &tx) - if err != errShortTypedTx { + if !errors.Is(err, errShortTypedTx) { t.Fatal("wrong error:", err) } } diff --git a/core/vm/evm.go b/core/vm/evm.go index 16cc854908..f7934fbe1b 100644 --- a/core/vm/evm.go +++ b/core/vm/evm.go @@ -17,6 +17,7 @@ package vm import ( + "errors" "math/big" "sync/atomic" @@ -246,7 +247,7 @@ func (evm *EVM) Call(caller ContractRef, addr common.Address, input []byte, gas // when we're in homestead this also counts for code storage gas errors. if err != nil { evm.StateDB.RevertToSnapshot(snapshot) - if err != ErrExecutionReverted { + if !errors.Is(err, ErrExecutionReverted) { gas = 0 } // TODO: consider clearing up unused snapshots: @@ -299,7 +300,7 @@ func (evm *EVM) CallCode(caller ContractRef, addr common.Address, input []byte, } if err != nil { evm.StateDB.RevertToSnapshot(snapshot) - if err != ErrExecutionReverted { + if !errors.Is(err, ErrExecutionReverted) { gas = 0 } } @@ -343,7 +344,7 @@ func (evm *EVM) DelegateCall(caller ContractRef, addr common.Address, input []by } if err != nil { evm.StateDB.RevertToSnapshot(snapshot) - if err != ErrExecutionReverted { + if !errors.Is(err, ErrExecutionReverted) { gas = 0 } } @@ -399,7 +400,7 @@ func (evm *EVM) StaticCall(caller ContractRef, addr common.Address, input []byte } if err != nil { evm.StateDB.RevertToSnapshot(snapshot) - if err != ErrExecutionReverted { + if !errors.Is(err, ErrExecutionReverted) { gas = 0 } } @@ -492,9 +493,9 @@ func (evm *EVM) create(caller ContractRef, codeAndHash *codeAndHash, gas uint64, // When an error was returned by the EVM or when setting the creation code // above we revert to the snapshot and consume any gas remaining. Additionally // when we're in homestead this also counts for code storage gas errors. - if err != nil && (evm.chainRules.IsHomestead || err != ErrCodeStoreOutOfGas) { + if err != nil && (evm.chainRules.IsHomestead || !errors.Is(err, ErrCodeStoreOutOfGas)) { evm.StateDB.RevertToSnapshot(snapshot) - if err != ErrExecutionReverted { + if !errors.Is(err, ErrExecutionReverted) { contract.UseGas(contract.Gas) } } diff --git a/core/vm/gas_table_test.go b/core/vm/gas_table_test.go index 4a2545b6ed..159a2bea5e 100644 --- a/core/vm/gas_table_test.go +++ b/core/vm/gas_table_test.go @@ -18,6 +18,7 @@ package vm import ( "bytes" + "errors" "math" "math/big" "sort" @@ -43,8 +44,8 @@ func TestMemoryGasCost(t *testing.T) { } for i, tt := range tests { v, err := memoryGasCost(&Memory{}, tt.size) - if (err == ErrGasUintOverflow) != tt.overflow { - t.Errorf("test %d: overflow mismatch: have %v, want %v", i, err == ErrGasUintOverflow, tt.overflow) + if (errors.Is(err, ErrGasUintOverflow)) != tt.overflow { + t.Errorf("test %d: overflow mismatch: have %v, want %v", i, errors.Is(err, ErrGasUintOverflow), tt.overflow) } if v != tt.cost { t.Errorf("test %d: gas cost mismatch: have %v, want %v", i, v, tt.cost) diff --git a/core/vm/instructions.go b/core/vm/instructions.go index b8055de6bc..23b6fe951f 100644 --- a/core/vm/instructions.go +++ b/core/vm/instructions.go @@ -17,6 +17,7 @@ package vm import ( + "errors" "math" "github.com/ethereum/go-ethereum/common" @@ -597,9 +598,9 @@ func opCreate(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([]b // homestead we must check for CodeStoreOutOfGasError (homestead only // rule) and treat as an error, if the ruleset is frontier we must // ignore this error and pretend the operation was successful. - if interpreter.evm.chainRules.IsHomestead && suberr == ErrCodeStoreOutOfGas { + if interpreter.evm.chainRules.IsHomestead && errors.Is(suberr, ErrCodeStoreOutOfGas) { stackvalue.Clear() - } else if suberr != nil && suberr != ErrCodeStoreOutOfGas { + } else if suberr != nil && !errors.Is(suberr, ErrCodeStoreOutOfGas) { stackvalue.Clear() } else { stackvalue.SetBytes(addr.Bytes()) @@ -607,7 +608,7 @@ func opCreate(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([]b scope.Stack.push(&stackvalue) scope.Contract.Gas += returnGas - if suberr == ErrExecutionReverted { + if errors.Is(suberr, ErrExecutionReverted) { interpreter.returnData = res // set REVERT data to return data buffer return res, nil } @@ -642,7 +643,7 @@ func opCreate2(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([] scope.Stack.push(&stackvalue) scope.Contract.Gas += returnGas - if suberr == ErrExecutionReverted { + if errors.Is(suberr, ErrExecutionReverted) { interpreter.returnData = res // set REVERT data to return data buffer return res, nil } @@ -676,7 +677,7 @@ func opCall(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([]byt temp.SetOne() } stack.push(&temp) - if err == nil || err == ErrExecutionReverted { + if err == nil || errors.Is(err, ErrExecutionReverted) { scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret) } scope.Contract.Gas += returnGas @@ -708,7 +709,7 @@ func opCallCode(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([ temp.SetOne() } stack.push(&temp) - if err == nil || err == ErrExecutionReverted { + if err == nil || errors.Is(err, ErrExecutionReverted) { scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret) } scope.Contract.Gas += returnGas @@ -736,7 +737,7 @@ func opDelegateCall(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext temp.SetOne() } stack.push(&temp) - if err == nil || err == ErrExecutionReverted { + if err == nil || errors.Is(err, ErrExecutionReverted) { scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret) } scope.Contract.Gas += returnGas @@ -764,7 +765,7 @@ func opStaticCall(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) temp.SetOne() } stack.push(&temp) - if err == nil || err == ErrExecutionReverted { + if err == nil || errors.Is(err, ErrExecutionReverted) { scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret) } scope.Contract.Gas += returnGas diff --git a/core/vm/interpreter.go b/core/vm/interpreter.go index 1968289f4e..a2e9014f88 100644 --- a/core/vm/interpreter.go +++ b/core/vm/interpreter.go @@ -17,6 +17,7 @@ package vm import ( + "errors" "github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common/math" "github.com/ethereum/go-ethereum/crypto" @@ -234,7 +235,7 @@ func (in *EVMInterpreter) Run(contract *Contract, input []byte, readOnly bool) ( pc++ } - if err == errStopToken { + if errors.Is(err, errStopToken) { err = nil // clear stop token error } diff --git a/crypto/crypto.go b/crypto/crypto.go index 2492165d38..ff1bff8e2e 100644 --- a/crypto/crypto.go +++ b/crypto/crypto.go @@ -182,7 +182,8 @@ func FromECDSAPub(pub *ecdsa.PublicKey) []byte { // HexToECDSA parses a secp256k1 private key. func HexToECDSA(hexkey string) (*ecdsa.PrivateKey, error) { b, err := hex.DecodeString(hexkey) - if byteErr, ok := err.(hex.InvalidByteError); ok { + var byteErr hex.InvalidByteError + if errors.As(err, &byteErr) { return nil, fmt.Errorf("invalid hex character %q in private key", byte(byteErr)) } else if err != nil { return nil, errors.New("invalid hex data for private key") diff --git a/crypto/crypto_test.go b/crypto/crypto_test.go index da123cf980..7d348c137f 100644 --- a/crypto/crypto_test.go +++ b/crypto/crypto_test.go @@ -20,6 +20,7 @@ import ( "bytes" "crypto/ecdsa" "encoding/hex" + "errors" "math/big" "os" "reflect" @@ -66,11 +67,11 @@ func BenchmarkSha3(b *testing.B) { func TestUnmarshalPubkey(t *testing.T) { key, err := UnmarshalPubkey(nil) - if err != errInvalidPubkey || key != nil { + if !errors.Is(err, errInvalidPubkey) || key != nil { t.Fatalf("expected error, got %v, %v", err, key) } key, err = UnmarshalPubkey([]byte{1, 2, 3}) - if err != errInvalidPubkey || key != nil { + if !errors.Is(err, errInvalidPubkey) || key != nil { t.Fatalf("expected error, got %v, %v", err, key) } diff --git a/crypto/ecies/ecies_test.go b/crypto/ecies/ecies_test.go index e3da71010e..0b3a2e070c 100644 --- a/crypto/ecies/ecies_test.go +++ b/crypto/ecies/ecies_test.go @@ -152,12 +152,12 @@ func TestTooBigSharedKey(t *testing.T) { } _, err = prv1.GenerateShared(&prv2.PublicKey, 32, 32) - if err != ErrSharedKeyTooBig { + if !errors.Is(err, ErrSharedKeyTooBig) { t.Fatal("ecdh: shared key should be too large for curve") } _, err = prv2.GenerateShared(&prv1.PublicKey, 32, 32) - if err != ErrSharedKeyTooBig { + if !errors.Is(err, ErrSharedKeyTooBig) { t.Fatal("ecdh: shared key should be too large for curve") } } @@ -355,7 +355,7 @@ func TestBasicKeyValidation(t *testing.T) { for _, b := range badBytes { ct[0] = b _, err := prv.Decrypt(ct, nil, nil) - if err != ErrInvalidPublicKey { + if !errors.Is(err, ErrInvalidPublicKey) { t.Fatal("ecies: validated an invalid key") } } diff --git a/crypto/secp256k1/secp256_test.go b/crypto/secp256k1/secp256_test.go index 74408d06d2..8407cac3f2 100644 --- a/crypto/secp256k1/secp256_test.go +++ b/crypto/secp256k1/secp256_test.go @@ -13,6 +13,7 @@ import ( "crypto/elliptic" "crypto/rand" "encoding/hex" + "errors" "io" "testing" ) @@ -92,7 +93,7 @@ func TestInvalidRecoveryID(t *testing.T) { sig, _ := Sign(msg, seckey) sig[64] = 99 _, err := RecoverPubkey(msg, sig) - if err != ErrInvalidRecoveryID { + if !errors.Is(err, ErrInvalidRecoveryID) { t.Fatalf("got %q, want %q", err, ErrInvalidRecoveryID) } } diff --git a/eth/downloader/downloader.go b/eth/downloader/downloader.go index 6e7c5dcf02..ced74eb37f 100644 --- a/eth/downloader/downloader.go +++ b/eth/downloader/downloader.go @@ -1547,7 +1547,7 @@ func (d *Downloader) processSnapSyncContent() error { }() closeOnErr := func(s *stateSync) { - if err := s.Wait(); err != nil && err != errCancelStateFetch && err != errCanceled && err != snap.ErrCancelled { + if err := s.Wait(); err != nil && !errors.Is(err, errCancelStateFetch) && !errors.Is(err, errCanceled) && !errors.Is(err, snap.ErrCancelled) { d.queue.Close() // wake up Results } } diff --git a/eth/downloader/downloader_test.go b/eth/downloader/downloader_test.go index 2468e1a980..f9009f718d 100644 --- a/eth/downloader/downloader_test.go +++ b/eth/downloader/downloader_test.go @@ -17,6 +17,7 @@ package downloader import ( + "errors" "fmt" "math/big" "os" @@ -614,7 +615,7 @@ func testBoundedForkedSync(t *testing.T, protocol uint, mode SyncMode) { assertOwnChain(t, tester, len(chainA.blocks)) // Synchronise with the second peer and ensure that the fork is rejected to being too old - if err := tester.sync("rewriter", nil, mode); err != errInvalidAncestor { + if err := tester.sync("rewriter", nil, mode); !errors.Is(err, errInvalidAncestor) { t.Fatalf("sync failure mismatch: have %v, want %v", err, errInvalidAncestor) } } @@ -649,7 +650,7 @@ func testBoundedHeavyForkedSync(t *testing.T, protocol uint, mode SyncMode) { tester.newPeer("heavy-rewriter", protocol, chainB.blocks[1:]) // Synchronise with the second peer and ensure that the fork is rejected to being too old - if err := tester.sync("heavy-rewriter", nil, mode); err != errInvalidAncestor { + if err := tester.sync("heavy-rewriter", nil, mode); !errors.Is(err, errInvalidAncestor) { t.Fatalf("sync failure mismatch: have %v, want %v", err, errInvalidAncestor) } } @@ -854,7 +855,7 @@ func testHighTDStarvationAttack(t *testing.T, protocol uint, mode SyncMode) { chain := testChainBase.shorten(1) tester.newPeer("attack", protocol, chain.blocks[1:]) - if err := tester.sync("attack", big.NewInt(1000000), mode); err != errStallingPeer { + if err := tester.sync("attack", big.NewInt(1000000), mode); !errors.Is(err, errStallingPeer) { t.Fatalf("synchronisation error mismatch: have %v, want %v", err, errStallingPeer) } } diff --git a/eth/downloader/skeleton.go b/eth/downloader/skeleton.go index 873ee950b6..e4f18f974a 100644 --- a/eth/downloader/skeleton.go +++ b/eth/downloader/skeleton.go @@ -278,25 +278,25 @@ func (s *skeleton) startup() { // signalling as the sync loop should never terminate (TM). newhead, err := s.sync(head) switch { - case err == errSyncLinked: + case errors.Is(err, errSyncLinked): // Sync cycle linked up to the genesis block, or the existent chain // segment. Tear down the loop and restart it so, it can properly // notify the backfiller. Don't account a new head. head = nil - case err == errSyncMerged: + case errors.Is(err, errSyncMerged): // Subchains were merged, we just need to reinit the internal // start to continue on the tail of the merged chain. Don't // announce a new head, head = nil - case err == errSyncReorged: + case errors.Is(err, errSyncReorged): // The subchain being synced got modified at the head in a // way that requires resyncing it. Restart sync with the new // head to force a cleanup. head = newhead - case err == errTerminated: + case errors.Is(err, errTerminated): // Sync was requested to be terminated from within, stop and // return (no need to pass a message, was already done internally) return diff --git a/eth/filters/filter_system_test.go b/eth/filters/filter_system_test.go index 99c012cc84..1c590d2c29 100644 --- a/eth/filters/filter_system_test.go +++ b/eth/filters/filter_system_test.go @@ -468,7 +468,7 @@ func TestInvalidGetRangeLogsRequest(t *testing.T) { api = NewFilterAPI(sys, false) ) - if _, err := api.GetLogs(context.Background(), FilterCriteria{FromBlock: big.NewInt(2), ToBlock: big.NewInt(1)}); err != errInvalidBlockRange { + if _, err := api.GetLogs(context.Background(), FilterCriteria{FromBlock: big.NewInt(2), ToBlock: big.NewInt(1)}); !errors.Is(err, errInvalidBlockRange) { t.Errorf("Expected Logs for invalid range return error, but got: %v", err) } } diff --git a/eth/state_accessor.go b/eth/state_accessor.go index 526361a2b8..0a354bbd9d 100644 --- a/eth/state_accessor.go +++ b/eth/state_accessor.go @@ -117,8 +117,9 @@ func (eth *Ethereum) hashState(ctx context.Context, block *types.Block, reexec u } } if err != nil { - switch err.(type) { - case *trie.MissingNodeError: + var missingNodeError *trie.MissingNodeError + switch { + case errors.As(err, &missingNodeError): return nil, nil, fmt.Errorf("required historical state unavailable (reexec=%d)", reexec) default: return nil, nil, err diff --git a/ethclient/ethclient_test.go b/ethclient/ethclient_test.go index 0d2675f8d1..70fee5209b 100644 --- a/ethclient/ethclient_test.go +++ b/ethclient/ethclient_test.go @@ -398,7 +398,7 @@ func testTransactionInBlock(t *testing.T, client *rpc.Client) { } // Test tx in block not found. - if _, err := ec.TransactionInBlock(context.Background(), block.Hash(), 20); err != ethereum.NotFound { + if _, err := ec.TransactionInBlock(context.Background(), block.Hash(), 20); !errors.Is(err, ethereum.NotFound) { t.Fatal("error should be ethereum.NotFound") } diff --git a/ethdb/leveldb/leveldb.go b/ethdb/leveldb/leveldb.go index e58efbddbe..28be5207dc 100644 --- a/ethdb/leveldb/leveldb.go +++ b/ethdb/leveldb/leveldb.go @@ -21,7 +21,9 @@ package leveldb import ( + "errors" "fmt" + dberrors "github.com/syndtr/goleveldb/leveldb/errors" "strings" "sync" "time" @@ -31,7 +33,6 @@ import ( "github.com/ethereum/go-ethereum/log" "github.com/ethereum/go-ethereum/metrics" "github.com/syndtr/goleveldb/leveldb" - "github.com/syndtr/goleveldb/leveldb/errors" "github.com/syndtr/goleveldb/leveldb/filter" "github.com/syndtr/goleveldb/leveldb/opt" "github.com/syndtr/goleveldb/leveldb/util" @@ -120,7 +121,8 @@ func NewCustom(file string, namespace string, customize func(options *opt.Option // Open the db and recover any potential corruptions db, err := leveldb.OpenFile(file, options) - if _, corrupted := err.(*errors.ErrCorrupted); corrupted { + var errCorrupted *dberrors.ErrCorrupted + if errors.As(err, &errCorrupted) { db, err = leveldb.RecoverFile(file, nil) } if err != nil { diff --git a/event/event_test.go b/event/event_test.go index 84b37eca3b..66c824f047 100644 --- a/event/event_test.go +++ b/event/event_test.go @@ -17,6 +17,7 @@ package event import ( + "errors" "math/rand" "sync" "testing" @@ -59,7 +60,7 @@ func TestMuxErrorAfterStop(t *testing.T) { if _, isopen := <-sub.Chan(); isopen { t.Errorf("subscription channel was not closed") } - if err := mux.Post(testEvent(0)); err != ErrMuxClosed { + if err := mux.Post(testEvent(0)); !errors.Is(err, ErrMuxClosed) { t.Errorf("Post error mismatch, got: %s, expected: %s", err, ErrMuxClosed) } } diff --git a/event/subscription_test.go b/event/subscription_test.go index 743d0bf67d..2cd5d4e7c9 100644 --- a/event/subscription_test.go +++ b/event/subscription_test.go @@ -56,7 +56,7 @@ loop: t.Fatalf("wrong int %d, want %d", got, want) } case err := <-sub.Err(): - if err != errInts { + if !errors.Is(err, errInts) { t.Fatalf("wrong error: got %q, want %q", err, errInts) } if want != 2 { diff --git a/internal/build/util.go b/internal/build/util.go index b41014a16f..2b47783763 100644 --- a/internal/build/util.go +++ b/internal/build/util.go @@ -19,6 +19,7 @@ package build import ( "bufio" "bytes" + "errors" "flag" "fmt" "go/parser" @@ -98,7 +99,8 @@ func RunGit(args ...string) string { var stdout, stderr bytes.Buffer cmd.Stdout, cmd.Stderr = &stdout, &stderr if err := cmd.Run(); err != nil { - if e, ok := err.(*exec.Error); ok && e.Err == exec.ErrNotFound { + var e *exec.Error + if errors.As(err, &e) && errors.Is(e.Err, exec.ErrNotFound) { if !warnedAboutGit { log.Println("Warning: can't find 'git' in PATH") warnedAboutGit = true diff --git a/internal/cmdtest/test_cmd.go b/internal/cmdtest/test_cmd.go index 4890d0b7c6..6291c4cbf7 100644 --- a/internal/cmdtest/test_cmd.go +++ b/internal/cmdtest/test_cmd.go @@ -19,6 +19,7 @@ package cmdtest import ( "bufio" "bytes" + "errors" "fmt" "io" "os" @@ -206,7 +207,8 @@ func (tt *TestCmd) Interrupt() { // It will only return a valid value after the process has finished. func (tt *TestCmd) ExitStatus() int { if tt.Err != nil { - exitErr := tt.Err.(*exec.ExitError) + var exitErr *exec.ExitError + errors.As(tt.Err, &exitErr) if exitErr != nil { if status, ok := exitErr.Sys().(syscall.WaitStatus); ok { return status.ExitStatus() diff --git a/internal/jsre/pretty.go b/internal/jsre/pretty.go index bd772b4927..3efff62ee2 100644 --- a/internal/jsre/pretty.go +++ b/internal/jsre/pretty.go @@ -17,6 +17,7 @@ package jsre import ( + "errors" "fmt" "io" "reflect" @@ -60,7 +61,8 @@ func prettyPrint(vm *goja.Runtime, value goja.Value, w io.Writer) { // prettyError writes err to standard output. func prettyError(vm *goja.Runtime, err error, w io.Writer) { failure := err.Error() - if gojaErr, ok := err.(*goja.Exception); ok { + var gojaErr *goja.Exception + if errors.As(err, &gojaErr) { failure = gojaErr.String() } fmt.Fprint(w, ErrorColor("%s", failure)) diff --git a/node/errors.go b/node/errors.go index 67547bf691..f486380bcd 100644 --- a/node/errors.go +++ b/node/errors.go @@ -33,7 +33,8 @@ var ( ) func convertFileLockError(err error) error { - if errno, ok := err.(syscall.Errno); ok && datadirInUseErrnos[uint(errno)] { + var errno syscall.Errno + if errors.As(err, &errno) && datadirInUseErrnos[uint(errno)] { return ErrDatadirUsed } return err diff --git a/node/node_test.go b/node/node_test.go index 04810a815b..9e4af57aca 100644 --- a/node/node_test.go +++ b/node/node_test.go @@ -55,7 +55,7 @@ func TestNodeCloseMultipleTimes(t *testing.T) { // Ensure that a stopped node can be stopped again for i := 0; i < 3; i++ { - if err := stack.Close(); err != ErrNodeStopped { + if err := stack.Close(); !errors.Is(err, ErrNodeStopped) { t.Fatalf("iter %d: stop failure mismatch: have %v, want %v", i, err, ErrNodeStopped) } } @@ -71,14 +71,14 @@ func TestNodeStartMultipleTimes(t *testing.T) { if err := stack.Start(); err != nil { t.Fatalf("failed to start node: %v", err) } - if err := stack.Start(); err != ErrNodeRunning { + if err := stack.Start(); !errors.Is(err, ErrNodeRunning) { t.Fatalf("start failure mismatch: have %v, want %v ", err, ErrNodeRunning) } // Ensure that a node can be stopped, but only once if err := stack.Close(); err != nil { t.Fatalf("failed to stop node: %v", err) } - if err := stack.Close(); err != ErrNodeStopped { + if err := stack.Close(); !errors.Is(err, ErrNodeStopped) { t.Fatalf("stop failure mismatch: have %v, want %v ", err, ErrNodeStopped) } } @@ -100,7 +100,7 @@ func TestNodeUsedDataDir(t *testing.T) { // Create a second node based on the same data directory and ensure failure _, err = New(&Config{DataDir: dir}) - if err != ErrDatadirUsed { + if !errors.Is(err, ErrDatadirUsed) { t.Fatalf("duplicate datadir failure mismatch: have %v, want %v", err, ErrDatadirUsed) } } diff --git a/p2p/dial.go b/p2p/dial.go index 5e4ab1d50d..086e8e8bed 100644 --- a/p2p/dial.go +++ b/p2p/dial.go @@ -474,7 +474,8 @@ func (t *dialTask) run(d *dialScheduler) { err := t.dial(d, t.dest) if err != nil { // For static nodes, resolve one more time if dialing fails. - if _, ok := err.(*dialError); ok && t.flags&staticDialedConn != 0 { + var dialError *dialError + if errors.As(err, &dialError) && t.flags&staticDialedConn != 0 { if t.resolve(d) { t.dial(d, t.dest) } @@ -537,7 +538,8 @@ func (t *dialTask) String() string { } func cleanupDialErr(err error) error { - if netErr, ok := err.(*net.OpError); ok && netErr.Op == "dial" { + var netErr *net.OpError + if errors.As(err, &netErr) && netErr.Op == "dial" { return netErr.Err } return err diff --git a/p2p/discover/v4_udp_test.go b/p2p/discover/v4_udp_test.go index 361e379626..d0838f4d48 100644 --- a/p2p/discover/v4_udp_test.go +++ b/p2p/discover/v4_udp_test.go @@ -111,7 +111,7 @@ func (test *udpTest) waitPacketOut(validate interface{}) (closed bool) { test.t.Helper() dgram, err := test.pipe.receive() - if err == errClosed { + if errors.Is(err, errClosed) { return true } else if err != nil { test.t.Error("packet receive error:", err) @@ -150,7 +150,7 @@ func TestUDPv4_pingTimeout(t *testing.T) { key := newkey() toaddr := &net.UDPAddr{IP: net.ParseIP("1.2.3.4"), Port: 2222} node := enode.NewV4(&key.PublicKey, toaddr.IP, 0, toaddr.Port) - if _, err := test.udp.ping(node); err != errTimeout { + if _, err := test.udp.ping(node); !errors.Is(err, errTimeout) { t.Error("expected timeout error, got", err) } } @@ -210,7 +210,7 @@ func TestUDPv4_responseTimeouts(t *testing.T) { for i := 0; i < nReqs; i++ { select { case err := <-timeoutErr: - if err != errTimeout { + if !errors.Is(err, errTimeout) { t.Fatalf("got non-timeout error on timeoutErr %d: %v", i, err) } nTimeoutsRecv++ @@ -240,7 +240,7 @@ func TestUDPv4_findnodeTimeout(t *testing.T) { toid := enode.ID{1, 2, 3, 4} target := v4wire.Pubkey{4, 5, 6, 7} result, err := test.udp.findnode(toid, toaddr, target) - if err != errTimeout { + if !errors.Is(err, errTimeout) { t.Error("expected timeout error, got", err) } if len(result) > 0 { diff --git a/p2p/discover/v5_udp_test.go b/p2p/discover/v5_udp_test.go index eaa969ea8b..508525057a 100644 --- a/p2p/discover/v5_udp_test.go +++ b/p2p/discover/v5_udp_test.go @@ -20,6 +20,7 @@ import ( "bytes" "crypto/ecdsa" "encoding/binary" + "errors" "fmt" "math/rand" "net" @@ -239,7 +240,7 @@ func TestUDPv5_pingCall(t *testing.T) { done <- err }() test.waitPacketOut(func(p *v5wire.Ping, addr *net.UDPAddr, _ v5wire.Nonce) {}) - if err := <-done; err != errTimeout { + if err := <-done; !errors.Is(err, errTimeout) { t.Fatalf("want errTimeout, got %q", err) } @@ -264,7 +265,7 @@ func TestUDPv5_pingCall(t *testing.T) { wrongAddr := &net.UDPAddr{IP: net.IP{33, 44, 55, 22}, Port: 10101} test.packetInFrom(test.remotekey, wrongAddr, &v5wire.Pong{ReqID: p.ReqID}) }) - if err := <-done; err != errTimeout { + if err := <-done; !errors.Is(err, errTimeout) { t.Fatalf("want errTimeout for reply from wrong IP, got %q", err) } } @@ -377,7 +378,7 @@ func TestUDPv5_multipleHandshakeRounds(t *testing.T) { test.waitPacketOut(func(p *v5wire.Ping, addr *net.UDPAddr, nonce v5wire.Nonce) { test.packetIn(&v5wire.Whoareyou{Nonce: nonce}) }) - if err := <-done; err != errTimeout { + if err := <-done; !errors.Is(err, errTimeout) { t.Fatalf("unexpected ping error: %q", err) } } @@ -486,7 +487,7 @@ func TestUDPv5_talkRequest(t *testing.T) { done <- err }() test.waitPacketOut(func(p *v5wire.TalkRequest, addr *net.UDPAddr, _ v5wire.Nonce) {}) - if err := <-done; err != errTimeout { + if err := <-done; !errors.Is(err, errTimeout) { t.Fatalf("want errTimeout, got %q", err) } @@ -817,10 +818,10 @@ func (test *udpV5Test) waitPacketOut(validate interface{}) (closed bool) { exptype := fn.Type().In(0) dgram, err := test.pipe.receive() - if err == errClosed { + if errors.Is(err, errClosed) { return true } - if err == errTimeout { + if errors.Is(err, errTimeout) { test.t.Fatalf("timed out waiting for %v", exptype) return false } diff --git a/p2p/discover/v5wire/encoding.go b/p2p/discover/v5wire/encoding.go index 5108910620..3ad38dba40 100644 --- a/p2p/discover/v5wire/encoding.go +++ b/p2p/discover/v5wire/encoding.go @@ -128,7 +128,7 @@ var ( // returns false, it is pretty certain that the packet causing the error does not belong // to discv5. func IsInvalidHeader(err error) bool { - return err == errTooShort || err == errInvalidHeader || err == errMsgTooShort + return errors.Is(err, errTooShort) || errors.Is(err, errInvalidHeader) || errors.Is(err, errMsgTooShort) } // Packet sizes. diff --git a/p2p/dnsdisc/error.go b/p2p/dnsdisc/error.go index 39955cabff..4c129b28f7 100644 --- a/p2p/dnsdisc/error.go +++ b/p2p/dnsdisc/error.go @@ -47,7 +47,8 @@ type nameError struct { } func (err nameError) Error() string { - if ee, ok := err.err.(entryError); ok { + var ee entryError + if errors.As(err.err, &ee) { return fmt.Sprintf("invalid %s entry at %s: %v", ee.typ, err.name, ee.err) } return err.name + ": " + err.err.Error() diff --git a/p2p/enode/nodedb.go b/p2p/enode/nodedb.go index 7e7fb69b29..f482345a53 100644 --- a/p2p/enode/nodedb.go +++ b/p2p/enode/nodedb.go @@ -20,6 +20,7 @@ import ( "bytes" "crypto/rand" "encoding/binary" + "errors" "fmt" "net" "os" @@ -28,7 +29,7 @@ import ( "github.com/ethereum/go-ethereum/rlp" "github.com/syndtr/goleveldb/leveldb" - "github.com/syndtr/goleveldb/leveldb/errors" + dberrors "github.com/syndtr/goleveldb/leveldb/errors" "github.com/syndtr/goleveldb/leveldb/iterator" "github.com/syndtr/goleveldb/leveldb/opt" "github.com/syndtr/goleveldb/leveldb/storage" @@ -98,7 +99,8 @@ func newMemoryDB() (*DB, error) { func newPersistentDB(path string) (*DB, error) { opts := &opt.Options{OpenFilesCacheCapacity: 5} db, err := leveldb.OpenFile(path, opts) - if _, iscorrupted := err.(*errors.ErrCorrupted); iscorrupted { + var errCorrupted *dberrors.ErrCorrupted + if errors.As(err, &errCorrupted) { db, err = leveldb.RecoverFile(path, nil) } if err != nil { @@ -110,15 +112,14 @@ func newPersistentDB(path string) (*DB, error) { currentVer = currentVer[:binary.PutVarint(currentVer, int64(dbVersion))] blob, err := db.Get([]byte(dbVersionKey), nil) - switch err { - case leveldb.ErrNotFound: + switch { + case errors.Is(err, leveldb.ErrNotFound): // Version not found (i.e. empty cache), insert it if err := db.Put([]byte(dbVersionKey), currentVer, nil); err != nil { db.Close() return nil, err } - - case nil: + case err == nil: // Version present, flush if different if !bytes.Equal(blob, currentVer) { db.Close() diff --git a/p2p/enr/enr_test.go b/p2p/enr/enr_test.go index b85ee209d5..366a7e763f 100644 --- a/p2p/enr/enr_test.go +++ b/p2p/enr/enr_test.go @@ -19,6 +19,7 @@ package enr import ( "bytes" "encoding/binary" + "errors" "fmt" "math/rand" "testing" @@ -97,8 +98,8 @@ func TestLoadErrors(t *testing.T) { // Check error for invalid keys. var list []uint err = r.Load(WithEntry(ip4.ENRKey(), &list)) - kerr, ok := err.(*KeyError) - if !ok { + var kerr *KeyError + if !errors.As(err, &kerr) { t.Fatalf("expected KeyError, got %T", err) } assert.Equal(t, kerr.Key, ip4.ENRKey()) @@ -149,7 +150,7 @@ func TestSortedGetAndSet(t *testing.T) { func TestDirty(t *testing.T) { var r Record - if _, err := rlp.EncodeToBytes(r); err != errEncodeUnsigned { + if _, err := rlp.EncodeToBytes(r); !errors.Is(err, errEncodeUnsigned) { t.Errorf("expected errEncodeUnsigned, got %#v", err) } @@ -164,7 +165,7 @@ func TestDirty(t *testing.T) { if len(r.signature) != 0 { t.Error("signature still set after modification") } - if _, err := rlp.EncodeToBytes(r); err != errEncodeUnsigned { + if _, err := rlp.EncodeToBytes(r); !errors.Is(err, errEncodeUnsigned) { t.Errorf("expected errEncodeUnsigned, got %#v", err) } } @@ -248,7 +249,7 @@ func TestRecordTooBig(t *testing.T) { // set a big value for random key, expect error r.Set(WithEntry(key, randomString(SizeLimit))) - if err := signTest([]byte{5}, &r); err != errTooBig { + if err := signTest([]byte{5}, &r); !errors.Is(err, errTooBig) { t.Fatalf("expected to get errTooBig, got %#v", err) } diff --git a/p2p/enr/entries.go b/p2p/enr/entries.go index 9945a436c9..c3704d0a49 100644 --- a/p2p/enr/entries.go +++ b/p2p/enr/entries.go @@ -175,7 +175,7 @@ type KeyError struct { // Error implements error. func (err *KeyError) Error() string { - if err.Err == errNotFound { + if errors.Is(err.Err, errNotFound) { return fmt.Sprintf("missing ENR key %q", err.Key) } return fmt.Sprintf("ENR key %q: %v", err.Key, err.Err) @@ -190,7 +190,7 @@ func (err *KeyError) Unwrap() error { func IsNotFound(err error) bool { var ke *KeyError if errors.As(err, &ke) { - return ke.Err == errNotFound + return errors.Is(ke.Err, errNotFound) } return false } diff --git a/p2p/message_test.go b/p2p/message_test.go index e575c5d96e..37fee20312 100644 --- a/p2p/message_test.go +++ b/p2p/message_test.go @@ -18,6 +18,7 @@ package p2p import ( "bytes" + "errors" "fmt" "io" "runtime" @@ -55,7 +56,7 @@ loop: go func() { if err := SendItems(rw1, 1); err == nil { t.Error("EncodeMsg returned nil error") - } else if err != ErrPipeClosed { + } else if !errors.Is(err, ErrPipeClosed) { t.Errorf("EncodeMsg returned wrong error: got %v, want %v", err, ErrPipeClosed) } close(done) diff --git a/p2p/netutil/error_test.go b/p2p/netutil/error_test.go index 84d5c2c206..745fdafffa 100644 --- a/p2p/netutil/error_test.go +++ b/p2p/netutil/error_test.go @@ -17,6 +17,7 @@ package netutil import ( + "errors" "net" "testing" "time" @@ -52,7 +53,8 @@ func TestIsPacketTooBig(t *testing.T) { listener.SetDeadline(time.Now().Add(1 * time.Second)) n, _, err := listener.ReadFrom(buf) if err != nil { - if nerr, ok := err.(net.Error); ok && nerr.Timeout() { + var nerr net.Error + if errors.As(err, &nerr) && nerr.Timeout() { continue } if !isPacketTooBig(err) { diff --git a/p2p/peer.go b/p2p/peer.go index 65a7903f58..cec5ecc77b 100644 --- a/p2p/peer.go +++ b/p2p/peer.go @@ -272,7 +272,8 @@ loop: } writeStart <- struct{}{} case err = <-readErr: - if r, ok := err.(DiscReason); ok { + var r DiscReason + if errors.As(err, &r) { remoteRequested = true reason = r } else { diff --git a/p2p/peer_error.go b/p2p/peer_error.go index ebc59de251..2909b40c7d 100644 --- a/p2p/peer_error.go +++ b/p2p/peer_error.go @@ -100,13 +100,15 @@ func (d DiscReason) Error() string { } func discReasonForError(err error) DiscReason { - if reason, ok := err.(DiscReason); ok { + var reason DiscReason + if errors.As(err, &reason) { return reason } if errors.Is(err, errProtocolReturned) { return DiscQuitting } - peerError, ok := err.(*peerError) + var peerError *peerError + ok := errors.As(err, &peerError) if ok { switch peerError.code { case errInvalidMsgCode, errInvalidMsg: diff --git a/p2p/peer_test.go b/p2p/peer_test.go index 4308bbd2eb..91862f14a1 100644 --- a/p2p/peer_test.go +++ b/p2p/peer_test.go @@ -138,7 +138,7 @@ func TestPeerProtoReadMsg(t *testing.T) { select { case err := <-errc: - if err != errProtocolReturned { + if !errors.Is(err, errProtocolReturned) { t.Errorf("peer returned error: %v", err) } case <-time.After(2 * time.Second): diff --git a/p2p/transport.go b/p2p/transport.go index 4f6bb569bf..73c1e170d8 100644 --- a/p2p/transport.go +++ b/p2p/transport.go @@ -19,6 +19,7 @@ package p2p import ( "bytes" "crypto/ecdsa" + "errors" "fmt" "io" "net" @@ -113,7 +114,8 @@ func (t *rlpxTransport) close(err error) { // We only bother doing this if the underlying connection supports // setting a timeout tough. if t.conn != nil { - if r, ok := err.(DiscReason); ok && r != DiscNetworkError { + var r DiscReason + if errors.As(err, &r) && !errors.Is(r, DiscNetworkError) { deadline := time.Now().Add(discWriteTimeout) if err := t.conn.SetWriteDeadline(deadline); err == nil { // Connection supports write deadline. diff --git a/rlp/decode.go b/rlp/decode.go index 9b17d2d810..f500b09c99 100644 --- a/rlp/decode.go +++ b/rlp/decode.go @@ -141,7 +141,8 @@ func wrapStreamError(err error, typ reflect.Type) error { } func addErrorContext(err error, ctx string) error { - if decErr, ok := err.(*decodeError); ok { + var decErr *decodeError + if errors.As(err, &decErr) { decErr.ctx = append(decErr.ctx, ctx) } return err @@ -326,7 +327,7 @@ func decodeSliceElems(s *Stream, val reflect.Value, elemdec decoder) error { val.SetLen(i + 1) } // decode into element - if err := elemdec(s, val.Index(i)); err == EOL { + if err := elemdec(s, val.Index(i)); errors.Is(err, EOL) { break } else if err != nil { return addErrorContext(err, fmt.Sprint("[", i, "]")) @@ -345,7 +346,7 @@ func decodeListArray(s *Stream, val reflect.Value, elemdec decoder) error { vlen := val.Len() i := 0 for ; i < vlen; i++ { - if err := elemdec(s, val.Index(i)); err == EOL { + if err := elemdec(s, val.Index(i)); errors.Is(err, EOL) { break } else if err != nil { return addErrorContext(err, fmt.Sprint("[", i, "]")) @@ -417,7 +418,7 @@ func makeStructDecoder(typ reflect.Type) (decoder, error) { } for i, f := range fields { err := f.info.decoder(s, val.Field(f.index)) - if err == EOL { + if errors.Is(err, EOL) { if f.optional { // The field is optional, so reaching the end of the list before // reaching the last field is acceptable. All remaining undecoded @@ -757,7 +758,7 @@ func (s *Stream) uint(maxbits int) (uint64, error) { } v, err := s.readUint(byte(size)) switch { - case err == ErrCanonSize: + case errors.Is(err, ErrCanonSize): // Adjust error because we're not reading a size right now. return 0, ErrCanonInt case err != nil: @@ -948,7 +949,8 @@ func (s *Stream) Decode(val interface{}) error { } err = decoder(s, rval.Elem()) - if decErr, ok := err.(*decodeError); ok && len(decErr.ctx) > 0 { + var decErr *decodeError + if errors.As(err, &decErr) && len(decErr.ctx) > 0 { // Add decode target type to error so context has more meaning. decErr.ctx = append(decErr.ctx, fmt.Sprint("(", rtyp.Elem(), ")")) } diff --git a/rlp/decode_test.go b/rlp/decode_test.go index 07d9c579a6..67bc50dff1 100644 --- a/rlp/decode_test.go +++ b/rlp/decode_test.go @@ -250,7 +250,7 @@ func TestStreamList(t *testing.T) { } } - if _, err := s.Uint(); err != EOL { + if _, err := s.Uint(); !errors.Is(err, EOL) { t.Errorf("Uint error mismatch, got %v, want %v", err, EOL) } if err = s.ListEnd(); err != nil { @@ -331,16 +331,16 @@ func TestStreamReadBytes(t *testing.T) { func TestDecodeErrors(t *testing.T) { r := bytes.NewReader(nil) - if err := Decode(r, nil); err != errDecodeIntoNil { + if err := Decode(r, nil); !errors.Is(err, errDecodeIntoNil) { t.Errorf("Decode(r, nil) error mismatch, got %q, want %q", err, errDecodeIntoNil) } var nilptr *struct{} - if err := Decode(r, nilptr); err != errDecodeIntoNil { + if err := Decode(r, nilptr); !errors.Is(err, errDecodeIntoNil) { t.Errorf("Decode(r, nilptr) error mismatch, got %q, want %q", err, errDecodeIntoNil) } - if err := Decode(r, struct{}{}); err != errNoPointer { + if err := Decode(r, struct{}{}); !errors.Is(err, errNoPointer) { t.Errorf("Decode(r, struct{}{}) error mismatch, got %q, want %q", err, errNoPointer) } diff --git a/rlp/typecache.go b/rlp/typecache.go index 3e37c9d2fc..5a112fa53f 100644 --- a/rlp/typecache.go +++ b/rlp/typecache.go @@ -17,6 +17,7 @@ package rlp import ( + "errors" "fmt" "reflect" "sync" @@ -142,7 +143,8 @@ func structFields(typ reflect.Type) (fields []field, err error) { // Filter/validate fields. structFields, structTags, err := rlpstruct.ProcessFields(allStructFields) if err != nil { - if tagErr, ok := err.(rlpstruct.TagError); ok { + var tagErr rlpstruct.TagError + if errors.As(err, &tagErr) { tagErr.StructType = typ.String() return nil, tagErr } diff --git a/rpc/client.go b/rpc/client.go index 2b0016db8f..80f062a089 100644 --- a/rpc/client.go +++ b/rpc/client.go @@ -712,7 +712,8 @@ func (c *Client) drainRead() { func (c *Client) read(codec ServerCodec) { for { msgs, batch, err := codec.readBatch() - if _, ok := err.(*json.SyntaxError); ok { + var syntaxError *json.SyntaxError + if errors.As(err, &syntaxError) { msg := errorMessage(&parseError{err.Error()}) codec.writeJSON(context.Background(), msg, true) } diff --git a/rpc/client_test.go b/rpc/client_test.go index ac02ad33cf..224473f409 100644 --- a/rpc/client_test.go +++ b/rpc/client_test.go @@ -106,17 +106,19 @@ func TestClientErrorData(t *testing.T) { // The method handler returns an error value which implements the rpc.Error // interface, i.e. it has a custom error code. The server returns this error code. expectedCode := testError{}.ErrorCode() - if e, ok := err.(Error); !ok { + var e Error + if !errors.As(err, &e) { t.Fatalf("client did not return rpc.Error, got %#v", e) } else if e.ErrorCode() != expectedCode { t.Fatalf("wrong error code %d, want %d", e.ErrorCode(), expectedCode) } // Check data. - if e, ok := err.(DataError); !ok { + var de DataError + if !errors.As(err, &de) { t.Fatalf("client did not return rpc.DataError, got %#v", e) - } else if e.ErrorData() != (testError{}.ErrorData()) { - t.Fatalf("wrong error data %#v, want %#v", e.ErrorData(), testError{}.ErrorData()) + } else if de.ErrorData() != (testError{}.ErrorData()) { + t.Fatalf("wrong error data %#v, want %#v", de.ErrorData(), testError{}.ErrorData()) } } @@ -633,7 +635,7 @@ func TestClientNotificationStorm(t *testing.T) { t.Fatalf("(%d/%d) unexpected value %d", i, count, val) } case err := <-sub.Err(): - if wantError && err != ErrSubscriptionQueueOverflow { + if wantError && !errors.Is(err, ErrSubscriptionQueueOverflow) { t.Fatalf("(%d/%d) got error %q, want %q", i, count, err, ErrSubscriptionQueueOverflow) } else if !wantError { t.Fatalf("(%d/%d) got unexpected error %q", i, count, err) diff --git a/rpc/http_test.go b/rpc/http_test.go index ad86ca15ae..09b231c5ac 100644 --- a/rpc/http_test.go +++ b/rpc/http_test.go @@ -18,6 +18,7 @@ package rpc import ( "context" + "errors" "fmt" "net/http" "net/http/httptest" @@ -148,7 +149,8 @@ func TestHTTPErrorResponse(t *testing.T) { t.Fatal("error was expected") } - httpErr, ok := err.(HTTPError) + var httpErr HTTPError + ok := errors.As(err, &httpErr) if !ok { t.Fatalf("unexpected error type %T", err) } @@ -234,12 +236,12 @@ func TestNewContextWithHeaders(t *testing.T) { ctx3 := NewContextWithHeaders(ctx2, newHdr("key-2", "val-2")) expectedHeaders = 3 - if err := client.CallContext(ctx3, nil, "test"); err != ErrNoResult { + if err := client.CallContext(ctx3, nil, "test"); !errors.Is(err, ErrNoResult) { t.Error("call failed", err) } expectedHeaders = 2 - if err := client.CallContext(ctx2, nil, "test"); err != ErrNoResult { + if err := client.CallContext(ctx2, nil, "test"); !errors.Is(err, ErrNoResult) { t.Error("call failed:", err) } } diff --git a/rpc/json.go b/rpc/json.go index 5557a80760..b2576a6115 100644 --- a/rpc/json.go +++ b/rpc/json.go @@ -125,11 +125,13 @@ func errorMessage(err error) *jsonrpcMessage { Code: errcodeDefault, Message: err.Error(), }} - ec, ok := err.(Error) + var ec Error + ok := errors.As(err, &ec) if ok { msg.Error.Code = ec.ErrorCode() } - de, ok := err.(DataError) + var de DataError + ok = errors.As(err, &de) if ok { msg.Error.Data = de.ErrorData() } diff --git a/rpc/subscription.go b/rpc/subscription.go index 9cb0727547..e712b92947 100644 --- a/rpc/subscription.go +++ b/rpc/subscription.go @@ -310,7 +310,7 @@ func (sub *ClientSubscription) run() { // Send the error. if err != nil { - if err == ErrClientQuit { + if errors.Is(err, ErrClientQuit) { // ErrClientQuit gets here when Client.Close is called. This is reported as a // nil error because it's not an error, but we can't close sub.err here. err = nil @@ -346,7 +346,7 @@ func (sub *ClientSubscription) forward() (unsubscribeServer bool, err error) { if !recv.IsNil() { err = recv.Interface().(error) } - if err == errUnsubscribed { + if errors.Is(err, errUnsubscribed) { // Exiting because Unsubscribe was called, unsubscribe on server. return true, nil } diff --git a/tests/init_test.go b/tests/init_test.go index e9bb99dc7d..9c0eff0088 100644 --- a/tests/init_test.go +++ b/tests/init_test.go @@ -52,7 +52,8 @@ func readJSON(reader io.Reader, value interface{}) error { return fmt.Errorf("error reading JSON file: %v", err) } if err = json.Unmarshal(data, &value); err != nil { - if syntaxerr, ok := err.(*json.SyntaxError); ok { + var syntaxerr *json.SyntaxError + if errors.As(err, &syntaxerr) { line := findLine(data, syntaxerr.Offset) return fmt.Errorf("JSON syntax error at line %v: %v", line, err) } diff --git a/trie/iterator.go b/trie/iterator.go index 83ccc0740f..21a7e709c4 100644 --- a/trie/iterator.go +++ b/trie/iterator.go @@ -268,10 +268,11 @@ func (it *nodeIterator) NodeBlob() []byte { } func (it *nodeIterator) Error() error { - if it.err == errIteratorEnd { + if errors.Is(it.err, errIteratorEnd) { return nil } - if seek, ok := it.err.(seekError); ok { + var seek seekError + if errors.As(it.err, &seek) { return seek.err } return it.err @@ -282,10 +283,11 @@ func (it *nodeIterator) Error() error { // sets the Error field to the encountered failure. If `descend` is false, // skips iterating over any subnodes of the current node. func (it *nodeIterator) Next(descend bool) bool { - if it.err == errIteratorEnd { + if errors.Is(it.err, errIteratorEnd) { return false } - if seek, ok := it.err.(seekError); ok { + var seek seekError + if errors.As(it.err, &seek) { if it.err = it.seek(seek.key); it.err != nil { return false } @@ -307,7 +309,7 @@ func (it *nodeIterator) seek(prefix []byte) error { // Move forward until we're just before the closest match to key. for { state, parentIndex, path, err := it.peekSeek(key) - if err == errIteratorEnd { + if errors.Is(err, errIteratorEnd) { return errIteratorEnd } else if err != nil { return seekError{prefix, err} diff --git a/trie/node.go b/trie/node.go index 15bbf62f1c..91bd84fb1a 100644 --- a/trie/node.go +++ b/trie/node.go @@ -17,6 +17,7 @@ package trie import ( + "errors" "fmt" "io" "strings" @@ -242,7 +243,8 @@ func wrapError(err error, ctx string) error { if err == nil { return nil } - if decErr, ok := err.(*decodeError); ok { + var decErr *decodeError + if errors.As(err, &decErr) { decErr.stack = append(decErr.stack, ctx) return decErr } diff --git a/trie/node_test.go b/trie/node_test.go index 9b8b33748f..e14aa5aa29 100644 --- a/trie/node_test.go +++ b/trie/node_test.go @@ -18,6 +18,7 @@ package trie import ( "bytes" + "errors" "testing" "github.com/ethereum/go-ethereum/crypto" @@ -59,7 +60,8 @@ func TestDecodeFullNodeWrongSizeChild(t *testing.T) { rlp.Encode(buf, fullNodeData) _, err := decodeNode([]byte("testdecode"), buf.Bytes()) - if _, ok := err.(*decodeError); !ok { + var decodeError *decodeError + if !errors.As(err, &decodeError) { t.Fatalf("decodeNode returned wrong err: %v", err) } } @@ -78,7 +80,8 @@ func TestDecodeFullNodeWrongNestedFullNode(t *testing.T) { rlp.Encode(buf, fullNodeData) _, err := decodeNode([]byte("testdecode"), buf.Bytes()) - if _, ok := err.(*decodeError); !ok { + var decodeError *decodeError + if !errors.As(err, &decodeError) { t.Fatalf("decodeNode returned wrong err: %v", err) } } diff --git a/trie/trie_test.go b/trie/trie_test.go index 379a866f7e..dbe66e7536 100644 --- a/trie/trie_test.go +++ b/trie/trie_test.go @@ -76,7 +76,8 @@ func testMissingRoot(t *testing.T, scheme string) { if trie != nil { t.Error("New returned non-nil trie for invalid root") } - if _, ok := err.(*MissingNodeError); !ok { + var missingNodeError *MissingNodeError + if !errors.As(err, &missingNodeError) { t.Errorf("New returned wrong error: %v", err) } } @@ -146,11 +147,12 @@ func testMissingNode(t *testing.T, memonly bool, scheme string) { } _, err = trie.Get([]byte("120000")) - if _, ok := err.(*MissingNodeError); !ok { + var missingNodeError *MissingNodeError + if !errors.As(err, &missingNodeError) { t.Errorf("Wrong error: %v", err) } _, err = trie.Get([]byte("120099")) - if _, ok := err.(*MissingNodeError); !ok { + if !errors.As(err, &missingNodeError) { t.Errorf("Wrong error: %v", err) } _, err = trie.Get([]byte("123456")) @@ -158,11 +160,11 @@ func testMissingNode(t *testing.T, memonly bool, scheme string) { t.Errorf("Unexpected error: %v", err) } err = trie.Update([]byte("120099"), []byte("zxcv")) - if _, ok := err.(*MissingNodeError); !ok { + if !errors.As(err, &missingNodeError) { t.Errorf("Wrong error: %v", err) } err = trie.Delete([]byte("123456")) - if _, ok := err.(*MissingNodeError); !ok { + if !errors.As(err, &missingNodeError) { t.Errorf("Wrong error: %v", err) } } @@ -620,7 +622,8 @@ func runRandTest(rt randTest) error { func TestRandom(t *testing.T) { if err := quick.Check(runRandTestBool, nil); err != nil { - if cerr, ok := err.(*quick.CheckError); ok { + var cerr *quick.CheckError + if errors.As(err, &cerr) { t.Fatalf("random test iteration %d failed: %s", cerr.Count, spew.Sdump(cerr.In)) } t.Fatal(err)