chore(all): enable errorlint and fix all errors

- Ignore fmt.Errorf not using wrapping with `%w`
- Use `errors.Is` instead of `==`
- Use `errors.As` instead of direct type assertions
This commit is contained in:
Quentin Mc Gaw 2025-01-19 22:48:49 +01:00
parent 04a336aee8
commit d245d7e868
No known key found for this signature in database
GPG key ID: 6B26BAFFE648CAFB
118 changed files with 428 additions and 316 deletions

View file

@ -7,6 +7,7 @@ run:
linters: linters:
disable-all: true disable-all: true
enable: enable:
- errorlint
- goimports - goimports
- gosimple - gosimple
- govet - govet
@ -32,12 +33,13 @@ linters:
# - errcheck #lot of false positives # - errcheck #lot of false positives
# - contextcheck # - contextcheck
# - errchkjson # lots of false positives # - errchkjson # lots of false positives
# - errorlint # this check crashes
# - exhaustive # silly check # - exhaustive # silly check
# - makezero # false positives # - makezero # false positives
# - nilerr # several intentional # - nilerr # several intentional
linters-settings: linters-settings:
errorlint:
errorf: false
gofmt: gofmt:
simplify: true simplify: true
revive: revive:

View file

@ -89,7 +89,7 @@ func TestWaitDeployed(t *testing.T) {
select { select {
case <-mined: case <-mined:
if err != test.wantErr { if !errors.Is(err, test.wantErr) {
t.Errorf("test %q: error mismatch: want %q, got %q", name, test.wantErr, err) t.Errorf("test %q: error mismatch: want %q, got %q", name, test.wantErr, err)
} }
if address != test.wantAddress { if address != test.wantAddress {

View file

@ -19,6 +19,7 @@ package abi
import ( import (
"bytes" "bytes"
"encoding/hex" "encoding/hex"
"errors"
"fmt" "fmt"
"math" "math"
"math/big" "math/big"
@ -1105,7 +1106,7 @@ func TestPackAndUnpackIncompatibleNumber(t *testing.T) {
{Type: ty}, {Type: ty},
} }
decoded, err := decodeABI.Unpack(packed) decoded, err := decodeABI.Unpack(packed)
if err != testCase.err { if !errors.Is(err, testCase.err) {
t.Fatalf("Expected error %v, actual error %v. case %d", testCase.err, err, i) t.Fatalf("Expected error %v, actual error %v. case %d", testCase.err, err, i)
} }
if err != nil { if err != nil {

View file

@ -307,7 +307,7 @@ func TestCacheFind(t *testing.T) {
} }
for i, test := range tests { for i, test := range tests {
a, err := cache.find(test.Query) a, err := cache.find(test.Query)
if !reflect.DeepEqual(err, test.WantError) { if !errors.Is(err, test.WantError) {
t.Errorf("test %d: error mismatch for query %v\ngot %q\nwant %q", i, test.Query, err, test.WantError) t.Errorf("test %d: error mismatch for query %v\ngot %q\nwant %q", i, test.Query, err, test.WantError)
continue continue
} }

View file

@ -17,6 +17,7 @@
package keystore package keystore
import ( import (
"errors"
"math/rand" "math/rand"
"os" "os"
"runtime" "runtime"
@ -127,7 +128,7 @@ func TestTimedUnlock(t *testing.T) {
// Signing without passphrase fails because account is locked // Signing without passphrase fails because account is locked
_, err = ks.SignHash(accounts.Account{Address: a1.Address}, testSigData) _, 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) 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 // Signing fails again after automatic locking
time.Sleep(250 * time.Millisecond) time.Sleep(250 * time.Millisecond)
_, err = ks.SignHash(accounts.Account{Address: a1.Address}, testSigData) _, 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) 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 // Signing fails again after automatic locking
time.Sleep(250 * time.Millisecond) time.Sleep(250 * time.Millisecond)
_, err = ks.SignHash(accounts.Account{Address: a1.Address}, testSigData) _, 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) 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) end := time.Now().Add(500 * time.Millisecond)
for time.Now().Before(end) { 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 return
} else if err != nil { } else if err != nil {
t.Errorf("Sign error: %v", err) t.Errorf("Sign error: %v", err)

View file

@ -19,6 +19,7 @@ package keystore
import ( import (
"crypto/rand" "crypto/rand"
"encoding/hex" "encoding/hex"
"errors"
"fmt" "fmt"
"path/filepath" "path/filepath"
"reflect" "reflect"
@ -90,7 +91,7 @@ func TestKeyStorePassphraseDecryptionFail(t *testing.T) {
if err != nil { if err != nil {
t.Fatal(err) 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) t.Fatalf("wrong error for invalid password\ngot %q\nwant %q", err, ErrDecrypt)
} }
} }

View file

@ -119,7 +119,7 @@ func (w *ledgerDriver) Open(device io.ReadWriter, passphrase string) error {
_, err := w.ledgerDerive(accounts.DefaultBaseDerivationPath) _, err := w.ledgerDerive(accounts.DefaultBaseDerivationPath)
if err != nil { if err != nil {
// Ethereum app is not running or in browser mode, nothing more to do, return // 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 w.browser = true
} }
return nil return nil
@ -141,7 +141,7 @@ func (w *ledgerDriver) Close() error {
// Heartbeat implements usbwallet.driver, performing a sanity check against the // Heartbeat implements usbwallet.driver, performing a sanity check against the
// Ledger to see if it's still online. // Ledger to see if it's still online.
func (w *ledgerDriver) Heartbeat() error { 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 w.failure = err
return err return err
} }

View file

@ -18,6 +18,7 @@ package light
import ( import (
"crypto/rand" "crypto/rand"
"errors"
"testing" "testing"
"time" "time"
@ -259,13 +260,13 @@ func (c *committeeChainTest) setClockPeriod(period float64) {
} }
func (c *committeeChainTest) addFixedCommitteeRoot(tc *testCommitteeChain, period uint64, expErr error) { 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) 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) { 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) 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 { if addCommittee {
committee = tc.periods[period+1].committee 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) c.t.Errorf("Incorrect error output from InsertUpdate at period %d (expected %v, got %v)", period, expErr, err)
} }
} }

View file

@ -17,6 +17,7 @@
package sync package sync
import ( import (
"errors"
"sort" "sort"
"github.com/ethereum/go-ethereum/beacon/light" "github.com/ethereum/go-ethereum/beacon/light"
@ -380,12 +381,12 @@ func (s *ForwardUpdateSync) Process(requester request.Requester, events []reques
func (s *ForwardUpdateSync) processResponse(requester request.Requester, u updateResponse) (success bool) { func (s *ForwardUpdateSync) processResponse(requester request.Requester, u updateResponse) (success bool) {
for i, update := range u.response.Updates { for i, update := range u.response.Updates {
if err := s.chain.InsertUpdate(update, u.response.Committees[i]); err != nil { if err := s.chain.InsertUpdate(update, u.response.Committees[i]); err != nil {
if err == light.ErrInvalidPeriod { if errors.Is(err, light.ErrInvalidPeriod) {
// there is a gap in the update periods; stop processing without // there is a gap in the update periods; stop processing without
// failing and try again next time // failing and try again next time
return return
} }
if err == light.ErrInvalidUpdate || err == light.ErrWrongCommitteeRoot || err == light.ErrCannotReorg { if errors.Is(err, light.ErrInvalidUpdate) || errors.Is(err, light.ErrWrongCommitteeRoot) || errors.Is(err, light.ErrCannotReorg) {
requester.Fail(u.sid.Server, "invalid update received") requester.Fail(u.sid.Server, "invalid update received")
} else { } else {
log.Error("Unexpected InsertUpdate error", "error", err) log.Error("Unexpected InsertUpdate error", "error", err)

View file

@ -923,13 +923,13 @@ func testExternalUI(api *core.SignerAPI) {
} }
} }
expectApprove := func(testcase string, err error) { expectApprove := func(testcase string, err error) {
if err == nil || err == accounts.ErrUnknownAccount { if err == nil || errors.Is(err, accounts.ErrUnknownAccount) {
return return
} }
addErr(fmt.Sprintf("%v: expected no error, got %v", testcase, err.Error())) addErr(fmt.Sprintf("%v: expected no error, got %v", testcase, err.Error()))
} }
expectDeny := func(testcase string, err error) { expectDeny := func(testcase string, err error) {
if err == nil || err != core.ErrRequestDenied { if err == nil || !errors.Is(err, core.ErrRequestDenied) {
addErr(fmt.Sprintf("%v: expected ErrRequestDenied, got %v", testcase, err)) addErr(fmt.Sprintf("%v: expected ErrRequestDenied, got %v", testcase, err))
} }
} }

View file

@ -275,7 +275,7 @@ func blocksFromFile(chainfile string, gblock *types.Block) ([]*types.Block, erro
blocks[0] = gblock blocks[0] = gblock
for i := 0; ; i++ { for i := 0; ; i++ {
var b types.Block var b types.Block
if err := stream.Decode(&b); err == io.EOF { if err := stream.Decode(&b); errors.Is(err, io.EOF) {
break break
} else if err != nil { } else if err != nil {
return nil, fmt.Errorf("at block index %d: %v", i, err) return nil, fmt.Errorf("at block index %d: %v", i, err)

View file

@ -18,6 +18,7 @@ package v5test
import ( import (
"bytes" "bytes"
"errors"
"net" "net"
"slices" "slices"
"sync" "sync"
@ -96,7 +97,7 @@ func (s *Suite) TestPingLargeRequestID(t *utesting.T) {
case *v5wire.Pong: case *v5wire.Pong:
t.Errorf("PONG response with unknown request ID %x", resp.ReqID) t.Errorf("PONG response with unknown request ID %x", resp.ReqID)
case *readError: case *readError:
if resp.err == v5wire.ErrInvalidReqID { if errors.Is(resp.err, v5wire.ErrInvalidReqID) {
t.Error("response with oversized request ID") t.Error("response with oversized request ID")
} else if !netutil.IsTimeout(resp.err) { } else if !netutil.IsTimeout(resp.err) {
t.Error(resp) t.Error(resp)

View file

@ -20,6 +20,7 @@ import (
"bytes" "bytes"
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"math/big" "math/big"
@ -166,7 +167,7 @@ func timedExec(bench bool, execFunc func() ([]byte, uint64, error)) ([]byte, exe
if haveGasUsed != gasUsed { if haveGasUsed != gasUsed {
panic(fmt.Sprintf("gas differs, have %v want %v", haveGasUsed, gasUsed)) panic(fmt.Sprintf("gas differs, have %v want %v", haveGasUsed, gasUsed))
} }
if haveErr != err { if !errors.Is(haveErr, err) {
panic(fmt.Sprintf("err differs, have %v want %v", haveErr, err)) panic(fmt.Sprintf("err differs, have %v want %v", haveErr, err))
} }
} }

View file

@ -120,7 +120,8 @@ func loadConfig(file string, cfg *gethConfig) error {
err = tomlSettings.NewDecoder(bufio.NewReader(f)).Decode(cfg) err = tomlSettings.NewDecoder(bufio.NewReader(f)).Decode(cfg)
// Add file name to errors that have a line number. // Add file name to errors that have a line number.
if _, ok := err.(*toml.LineError); ok { lineErr := new(toml.LineError)
if ok := errors.As(err, &lineErr); ok {
err = errors.New(file + ", " + err.Error()) err = errors.New(file + ", " + err.Error())
} }
return err return err

View file

@ -22,6 +22,7 @@ import (
"bytes" "bytes"
"container/list" "container/list"
"encoding/hex" "encoding/hex"
"errors"
"flag" "flag"
"fmt" "fmt"
"io" "io"
@ -106,7 +107,7 @@ func rlpToText(in *inStream, out io.Writer) error {
stream := rlp.NewStream(in, 0) stream := rlp.NewStream(in, 0)
for { for {
if err := dump(in, stream, 0, out); err != nil { if err := dump(in, stream, 0, out); err != nil {
if err != io.EOF { if !errors.Is(err, io.EOF) {
return err return err
} }
break break
@ -149,7 +150,7 @@ func dump(in *inStream, s *rlp.Stream, depth int, out io.Writer) error {
if i > 0 { if i > 0 {
fmt.Fprint(out, ",\n") fmt.Fprint(out, ",\n")
} }
if err := dump(in, s, depth+1, out); err == rlp.EOL { if err := dump(in, s, depth+1, out); errors.Is(err, rlp.EOL) {
break break
} else if err != nil { } else if err != nil {
return err return err

View file

@ -195,7 +195,7 @@ func ImportChain(chain *core.BlockChain, fn string) error {
i := 0 i := 0
for ; i < importBatchSize; i++ { for ; i < importBatchSize; i++ {
var b types.Block var b types.Block
if err := stream.Decode(&b); err == io.EOF { if err := stream.Decode(&b); errors.Is(err, io.EOF) {
break break
} else if err != nil { } else if err != nil {
return fmt.Errorf("at block %d: %v", n, err) return fmt.Errorf("at block %d: %v", n, err)
@ -515,7 +515,7 @@ func ImportPreimages(db ethdb.Database, fn string) error {
var blob []byte var blob []byte
if err := stream.Decode(&blob); err != nil { if err := stream.Decode(&blob); err != nil {
if err == io.EOF { if errors.Is(err, io.EOF) {
break break
} }
return err return err
@ -725,7 +725,7 @@ func ImportLDBData(db ethdb.Database, f string, startIndex int64, interrupt chan
key, val []byte key, val []byte
) )
if err := stream.Decode(&op); err != nil { if err := stream.Decode(&op); err != nil {
if err == io.EOF { if errors.Is(err, io.EOF) {
break break
} }
return err return err

View file

@ -18,6 +18,7 @@ package bitutil
import ( import (
"bytes" "bytes"
"errors"
"fmt" "fmt"
"math/rand" "math/rand"
"testing" "testing"
@ -107,7 +108,7 @@ func TestDecodingCycle(t *testing.T) {
data := hexutil.MustDecode(tt.input) data := hexutil.MustDecode(tt.input)
orig, err := bitsetDecodeBytes(data, tt.size) orig, err := bitsetDecodeBytes(data, tt.size)
if err != tt.fail { if !errors.Is(err, tt.fail) {
t.Errorf("test %d: failure mismatch: have %v, want %v", i, err, tt.fail) t.Errorf("test %d: failure mismatch: have %v, want %v", i, err, tt.fail)
} }
if err != nil { if err != nil {
@ -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) 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 // 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) t.Errorf("decoding error mismatch for long data: have %v, want %v", err, errExceededTarget)
} }
} }

View file

@ -32,6 +32,7 @@ package hexutil
import ( import (
"encoding/hex" "encoding/hex"
"errors"
"fmt" "fmt"
"math/big" "math/big"
"strconv" "strconv"
@ -223,18 +224,21 @@ func decodeNibble(in byte) uint64 {
} }
func mapError(err error) error { func mapError(err error) error {
if err, ok := err.(*strconv.NumError); ok { numErr := new(strconv.NumError)
switch err.Err { if ok := errors.As(err, &numErr); ok {
case strconv.ErrRange: switch {
case errors.Is(numErr.Err, strconv.ErrRange):
return ErrUint64Range return ErrUint64Range
case strconv.ErrSyntax: case errors.Is(numErr.Err, strconv.ErrSyntax):
return ErrSyntax return ErrSyntax
} }
} }
if _, ok := err.(hex.InvalidByteError); ok {
var invalidByteErr hex.InvalidByteError
if ok := errors.As(err, &invalidByteErr); ok {
return ErrSyntax return ErrSyntax
} }
if err == hex.ErrLength { if errors.Is(err, hex.ErrLength) {
return ErrOddLength return ErrOddLength
} }
return err return err

View file

@ -19,6 +19,7 @@ package hexutil
import ( import (
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"math/big" "math/big"
"reflect" "reflect"
@ -355,7 +356,7 @@ func (b *Uint) UnmarshalJSON(input []byte) error {
func (b *Uint) UnmarshalText(input []byte) error { func (b *Uint) UnmarshalText(input []byte) error {
var u64 Uint64 var u64 Uint64
err := u64.UnmarshalText(input) err := u64.UnmarshalText(input)
if u64 > Uint64(^uint(0)) || err == ErrUint64Range { if u64 > Uint64(^uint(0)) || errors.Is(err, ErrUint64Range) {
return ErrUintRange return ErrUintRange
} else if err != nil { } else if err != nil {
return err return err
@ -410,7 +411,8 @@ func checkNumberText(input []byte) (raw []byte, err error) {
} }
func wrapTypeError(err error, typ reflect.Type) error { func wrapTypeError(err error, typ reflect.Type) error {
if _, ok := err.(*decError); ok { decErr := new(decError)
if ok := errors.As(err, &decErr); ok {
return &json.UnmarshalTypeError{Value: err.Error(), Type: typ} return &json.UnmarshalTypeError{Value: err.Error(), Type: typ}
} }
return err return err

View file

@ -18,6 +18,7 @@ package common
import ( import (
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"os" "os"
) )
@ -29,7 +30,8 @@ func LoadJSON(file string, val interface{}) error {
return err return err
} }
if err := json.Unmarshal(content, val); err != nil { if err := json.Unmarshal(content, val); err != nil {
if syntaxerr, ok := err.(*json.SyntaxError); ok { syntaxerr := new(json.SyntaxError)
if ok := errors.As(err, &syntaxerr); ok {
line := findLine(content, syntaxerr.Offset) line := findLine(content, syntaxerr.Offset)
return fmt.Errorf("JSON syntax error at %v:%v: %v", file, line, err) return fmt.Errorf("JSON syntax error at %v:%v: %v", file, line, err)
} }

View file

@ -19,6 +19,7 @@ package clique
import ( import (
"bytes" "bytes"
"crypto/ecdsa" "crypto/ecdsa"
"errors"
"fmt" "fmt"
"math/big" "math/big"
"slices" "slices"
@ -469,7 +470,7 @@ func (tt *cliqueTest) run(t *testing.T) {
t.Fatalf("failed to import batch %d, block %d: %v", j, k, err) t.Fatalf("failed to import batch %d, block %d: %v", j, k, err)
} }
} }
if _, err = chain.InsertChain(batches[len(batches)-1]); err != tt.failure { if _, err = chain.InsertChain(batches[len(batches)-1]); !errors.Is(err, tt.failure) {
t.Errorf("failure mismatch: have %v, want %v", err, tt.failure) t.Errorf("failure mismatch: have %v, want %v", err, tt.failure)
} }
if tt.failure != nil { if tt.failure != nil {

View file

@ -169,11 +169,13 @@ func (b *bridge) Send(call jsre.Call) (goja.Value, error) {
} else { } else {
code := -32603 code := -32603
var data interface{} var data interface{}
if err, ok := err.(rpc.Error); ok { var rcpErr rpc.Error
code = err.ErrorCode() if ok := errors.As(err, &rcpErr); ok {
code = rcpErr.ErrorCode()
} }
if err, ok := err.(rpc.DataError); ok { var rcpDataErr rpc.DataError
data = err.ErrorData() if ok := errors.As(err, &rcpDataErr); ok {
data = rcpDataErr.ErrorData()
} }
setError(resp, code, err.Error(), data) setError(resp, code, err.Error(), data)
} }

View file

@ -148,7 +148,8 @@ func (c *Console) init(preload []string) error {
for _, path := range preload { for _, path := range preload {
if err := c.jsre.Exec(path); err != nil { if err := c.jsre.Exec(path); err != nil {
failure := err.Error() failure := err.Error()
if gojaErr, ok := err.(*goja.Exception); ok { gojaErr := new(goja.Exception)
if ok := errors.As(err, &gojaErr); ok {
failure = gojaErr.String() failure = gojaErr.String()
} }
return fmt.Errorf("%s: %v", path, failure) return fmt.Errorf("%s: %v", path, failure)
@ -205,7 +206,8 @@ func (c *Console) initExtensions() error {
const methodNotFound = -32601 const methodNotFound = -32601
apis, err := c.client.SupportedModules() apis, err := c.client.SupportedModules()
if err != nil { if err != nil {
if rpcErr, ok := err.(rpc.Error); ok && rpcErr.ErrorCode() == methodNotFound { var rpcErr rpc.Error
if ok := errors.As(err, &rpcErr); ok && rpcErr.ErrorCode() == methodNotFound {
log.Warn("Server does not support method rpc_modules, using default API list.") log.Warn("Server does not support method rpc_modules, using default API list.")
apis = defaultAPIs apis = defaultAPIs
} else { } else {
@ -423,7 +425,7 @@ func (c *Console) Interactive() {
return return
case err := <-inputErr: case err := <-inputErr:
if err == liner.ErrPromptAborted { if errors.Is(err, liner.ErrPromptAborted) {
// When prompting for multi-line input, the first Ctrl-C resets // When prompting for multi-line input, the first Ctrl-C resets
// the multi-line state. // the multi-line state.
prompt, indents, input = c.prompt, 0, "" prompt, indents, input = c.prompt, 0, ""

View file

@ -1388,7 +1388,7 @@ func (bc *BlockChain) InsertReceiptChain(blockChain types.Blocks, receiptChain [
// Write downloaded chain data and corresponding receipt chain data // Write downloaded chain data and corresponding receipt chain data
if len(ancientBlocks) > 0 { if len(ancientBlocks) > 0 {
if n, err := writeAncient(ancientBlocks, ancientReceipts); err != nil { if n, err := writeAncient(ancientBlocks, ancientReceipts); err != nil {
if err == errInsertionInterrupted { if errors.Is(err, errInsertionInterrupted) {
return 0, nil return 0, nil
} }
return n, err return n, err
@ -1396,7 +1396,7 @@ func (bc *BlockChain) InsertReceiptChain(blockChain types.Blocks, receiptChain [
} }
if len(liveBlocks) > 0 { if len(liveBlocks) > 0 {
if n, err := writeLive(liveBlocks, liveReceipts); err != nil { if n, err := writeLive(liveBlocks, liveReceipts); err != nil {
if err == errInsertionInterrupted { if errors.Is(err, errInsertionInterrupted) {
return 0, nil return 0, nil
} }
return n, err return n, err

View file

@ -158,7 +158,7 @@ func testBlockChainImport(chain types.Blocks, blockchain *BlockChain) error {
err = blockchain.validator.ValidateBody(block) err = blockchain.validator.ValidateBody(block)
} }
if err != nil { if err != nil {
if err == ErrKnownBlock { if errors.Is(err, ErrKnownBlock) {
continue continue
} }
return err return err

View file

@ -18,6 +18,7 @@ package forkid
import ( import (
"bytes" "bytes"
"errors"
"hash/crc32" "hash/crc32"
"math" "math"
"math/big" "math/big"
@ -340,7 +341,7 @@ func TestValidation(t *testing.T) {
genesis := core.DefaultGenesisBlock().ToBlock() genesis := core.DefaultGenesisBlock().ToBlock()
for i, tt := range tests { for i, tt := range tests {
filter := newFilter(tt.config, genesis, func() (uint64, uint64) { return tt.head, tt.time }) filter := newFilter(tt.config, genesis, func() (uint64, uint64) { return tt.head, tt.time })
if err := filter(tt.id); err != tt.err { if err := filter(tt.id); !errors.Is(err, tt.err) {
t.Errorf("test %d: validation error mismatch: have %v, want %v", i, err, tt.err) t.Errorf("test %d: validation error mismatch: have %v, want %v", i, err, tt.err)
} }
} }

View file

@ -19,6 +19,7 @@ package rawdb
import ( import (
"bytes" "bytes"
"encoding/binary" "encoding/binary"
"errors"
"fmt" "fmt"
"math/rand" "math/rand"
"os" "os"
@ -68,7 +69,7 @@ func TestFreezerBasics(t *testing.T) {
} }
// Check that we cannot read too far // Check that we cannot read too far
_, err = f.Retrieve(uint64(255)) _, err = f.Retrieve(uint64(255))
if err != errOutOfBounds { if !errors.Is(err, errOutOfBounds) {
t.Fatal(err) t.Fatal(err)
} }
} }
@ -878,7 +879,7 @@ func checkRetrieveError(t *testing.T, f *freezerTable, items map[uint64]error) {
if err == nil { if err == nil {
t.Fatalf("unexpected value %x for item %d, want error %v", item, value, wantError) t.Fatalf("unexpected value %x for item %d, want error %v", item, value, wantError)
} }
if err != wantError { if !errors.Is(err, wantError) {
t.Fatalf("wrong error for item %d: %v", item, err) t.Fatalf("wrong error for item %d: %v", item, err)
} }
} }
@ -1361,7 +1362,8 @@ func runRandTest(rt randTest) bool {
func TestRandom(t *testing.T) { func TestRandom(t *testing.T) {
if err := quick.Check(runRandTest, nil); err != nil { if err := quick.Check(runRandTest, nil); err != nil {
if cerr, ok := err.(*quick.CheckError); ok { cerr := new(quick.CheckError)
if ok := errors.As(err, &cerr); ok {
t.Fatalf("random test iteration %d failed: %s", cerr.Count, spew.Sdump(cerr.In)) t.Fatalf("random test iteration %d failed: %s", cerr.Count, spew.Sdump(cerr.In))
} }
t.Fatal(err) t.Fatal(err)

View file

@ -104,7 +104,7 @@ func TestFreezerModifyRollback(t *testing.T) {
require.NoError(t, op.AppendRaw("test", 2, make([]byte, 2048))) require.NoError(t, op.AppendRaw("test", 2, make([]byte, 2048)))
return theError return theError
}) })
if err != theError { if !errors.Is(err, theError) {
t.Errorf("ModifyAncients returned wrong error %q", err) t.Errorf("ModifyAncients returned wrong error %q", err)
} }
checkAncientCount(t, f, "test", 0) checkAncientCount(t, f, "test", 0)
@ -372,7 +372,7 @@ func checkAncientCount(t *testing.T, f *Freezer, kind string, n uint64) {
} }
if _, err := f.Ancient(kind, index); err == nil { if _, err := f.Ancient(kind, index); err == nil {
t.Errorf("Ancient(%q, %d) didn't return expected error", kind, index) 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) t.Errorf("Ancient(%q, %d) returned unexpected error %q", kind, index, err)
} }
} }

View file

@ -18,6 +18,7 @@ package snapshot
import ( import (
"bytes" "bytes"
"errors"
"testing" "testing"
"github.com/VictoriaMetrics/fastcache" "github.com/VictoriaMetrics/fastcache"
@ -313,7 +314,7 @@ func TestDiskPartialMerge(t *testing.T) {
assertAccount := func(account common.Hash, data []byte) { assertAccount := func(account common.Hash, data []byte) {
t.Helper() t.Helper()
blob, err := base.AccountRLP(account) 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) 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) { if bytes.Compare(account[:], genMarker) <= 0 && !bytes.Equal(blob, data) {
@ -329,7 +330,7 @@ func TestDiskPartialMerge(t *testing.T) {
assertStorage := func(account common.Hash, slot common.Hash, data []byte) { assertStorage := func(account common.Hash, slot common.Hash, data []byte) {
t.Helper() t.Helper()
blob, err := base.Storage(account, slot) 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) 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) { if bytes.Compare(append(account[:], slot[:]...), genMarker) <= 0 && !bytes.Equal(blob, data) {

View file

@ -670,7 +670,8 @@ func (dl *diskLayer) generate(stats *generatorStats) {
if err := generateAccounts(ctx, dl, accMarker); err != nil { if err := generateAccounts(ctx, dl, accMarker); err != nil {
// Extract the received interruption signal if exists // Extract the received interruption signal if exists
if aerr, ok := err.(*abortErr); ok { aerr := new(abortErr)
if ok := errors.As(err, &aerr); ok {
abort = aerr.abort abort = aerr.abort
} }
// Aborted by internal error, wait the signal // Aborted by internal error, wait the signal

View file

@ -19,6 +19,7 @@ package snapshot
import ( import (
crand "crypto/rand" crand "crypto/rand"
"encoding/binary" "encoding/binary"
"errors"
"fmt" "fmt"
"math/rand" "math/rand"
"testing" "testing"
@ -118,10 +119,10 @@ func TestDiskLayerExternalInvalidationFullFlatten(t *testing.T) {
t.Fatalf("failed to merge diff layer onto disk: %v", err) 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 // 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) 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) t.Errorf("stale reference returned storage slot: %#x (err: %v)", slot, err)
} }
if n := len(snaps.layers); n != 1 { 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) 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 // 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) 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) t.Errorf("stale reference returned storage slot: %#x (err: %v)", slot, err)
} }
if n := len(snaps.layers); n != 2 { 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) 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 // 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) 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) t.Errorf("stale reference returned storage slot: %#x (err: %v)", slot, err)
} }
if n := len(snaps.layers); n != 3 { if n := len(snaps.layers); n != 3 {

View file

@ -437,7 +437,8 @@ func (test *stateTest) verify(root common.Hash, next common.Hash, db *triedb.Dat
func TestStateChanges(t *testing.T) { func TestStateChanges(t *testing.T) {
config := &quick.Config{MaxCount: 1000} config := &quick.Config{MaxCount: 1000}
err := quick.Check((*stateTest).run, config) err := quick.Check((*stateTest).run, config)
if cerr, ok := err.(*quick.CheckError); ok { cerr := new(quick.CheckError)
if ok := errors.As(err, &cerr); ok {
test := cerr.In[0].(*stateTest) test := cerr.In[0].(*stateTest)
t.Errorf("%v:\n%s", test.err, test) t.Errorf("%v:\n%s", test.err, test)
} else if err != nil { } else if err != nil {

View file

@ -19,6 +19,7 @@ package state
import ( import (
"bytes" "bytes"
"encoding/binary" "encoding/binary"
"errors"
"fmt" "fmt"
"maps" "maps"
"math" "math"
@ -304,7 +305,8 @@ func TestCopyObjectState(t *testing.T) {
func TestSnapshotRandom(t *testing.T) { func TestSnapshotRandom(t *testing.T) {
config := &quick.Config{MaxCount: 1000} config := &quick.Config{MaxCount: 1000}
err := quick.Check((*snapshotTest).run, config) err := quick.Check((*snapshotTest).run, config)
if cerr, ok := err.(*quick.CheckError); ok { cerr := new(quick.CheckError)
if ok := errors.As(err, &cerr); ok {
test := cerr.In[0].(*snapshotTest) test := cerr.In[0].(*snapshotTest)
t.Errorf("%v:\n%s", test.err, test) t.Errorf("%v:\n%s", test.err, test)
} else if err != nil { } else if err != nil {

View file

@ -17,6 +17,7 @@
package core package core
import ( import (
"errors"
"fmt" "fmt"
"math" "math"
"math/big" "math/big"
@ -60,7 +61,7 @@ func (result *ExecutionResult) Return() []byte {
// Revert returns the concrete revert reason if the execution is aborted by `REVERT` // Revert returns the concrete revert reason if the execution is aborted by `REVERT`
// opcode. Note the reason can be nil if no data supplied with revert opcode. // opcode. Note the reason can be nil if no data supplied with revert opcode.
func (result *ExecutionResult) Revert() []byte { func (result *ExecutionResult) Revert() []byte {
if result.Err != vm.ErrExecutionReverted { if !errors.Is(result.Err, vm.ErrExecutionReverted) {
return nil return nil
} }
return common.CopyBytes(result.ReturnData) return common.CopyBytes(result.ReturnData)

View file

@ -96,7 +96,7 @@ func (journal *journal) load(add func([]*types.Transaction) []error) error {
// Parse the next transaction and terminate on error // Parse the next transaction and terminate on error
tx := new(types.Transaction) tx := new(types.Transaction)
if err = stream.Decode(tx); err != nil { if err = stream.Decode(tx); err != nil {
if err != io.EOF { if !errors.Is(err, io.EOF) {
failure = err failure = err
} }
if batch.Len() > 0 { if batch.Len() > 0 {

View file

@ -421,7 +421,7 @@ func TestNegativeValue(t *testing.T) {
tx, _ := types.SignTx(types.NewTransaction(0, common.Address{}, big.NewInt(-1), 100, big.NewInt(1), nil), types.HomesteadSigner{}, key) tx, _ := types.SignTx(types.NewTransaction(0, common.Address{}, big.NewInt(-1), 100, big.NewInt(1), nil), types.HomesteadSigner{}, key)
from, _ := deriveSender(tx) from, _ := deriveSender(tx)
testAddBalance(pool, from, big.NewInt(1)) testAddBalance(pool, from, big.NewInt(1))
if err := pool.addRemote(tx); err != txpool.ErrNegativeValue { if err := pool.addRemote(tx); !errors.Is(err, txpool.ErrNegativeValue) {
t.Error("expected", txpool.ErrNegativeValue, "got", err) t.Error("expected", txpool.ErrNegativeValue, "got", err)
} }
} }
@ -434,7 +434,7 @@ func TestTipAboveFeeCap(t *testing.T) {
tx := dynamicFeeTx(0, 100, big.NewInt(1), big.NewInt(2), key) tx := dynamicFeeTx(0, 100, big.NewInt(1), big.NewInt(2), key)
if err := pool.addRemote(tx); err != core.ErrTipAboveFeeCap { if err := pool.addRemote(tx); !errors.Is(err, core.ErrTipAboveFeeCap) {
t.Error("expected", core.ErrTipAboveFeeCap, "got", err) t.Error("expected", core.ErrTipAboveFeeCap, "got", err)
} }
} }
@ -449,12 +449,12 @@ func TestVeryHighValues(t *testing.T) {
veryBigNumber.Lsh(veryBigNumber, 300) veryBigNumber.Lsh(veryBigNumber, 300)
tx := dynamicFeeTx(0, 100, big.NewInt(1), veryBigNumber, key) tx := dynamicFeeTx(0, 100, big.NewInt(1), veryBigNumber, key)
if err := pool.addRemote(tx); err != core.ErrTipVeryHigh { if err := pool.addRemote(tx); !errors.Is(err, core.ErrTipVeryHigh) {
t.Error("expected", core.ErrTipVeryHigh, "got", err) t.Error("expected", core.ErrTipVeryHigh, "got", err)
} }
tx2 := dynamicFeeTx(0, 100, veryBigNumber, big.NewInt(1), key) tx2 := dynamicFeeTx(0, 100, veryBigNumber, big.NewInt(1), key)
if err := pool.addRemote(tx2); err != core.ErrFeeCapVeryHigh { if err := pool.addRemote(tx2); !errors.Is(err, core.ErrFeeCapVeryHigh) {
t.Error("expected", core.ErrFeeCapVeryHigh, "got", err) t.Error("expected", core.ErrFeeCapVeryHigh, "got", err)
} }
} }
@ -1804,7 +1804,7 @@ func TestUnderpricing(t *testing.T) {
t.Fatalf("failed to add well priced transaction: %v", err) t.Fatalf("failed to add well priced transaction: %v", err)
} }
// Ensure that replacing a pending transaction with a future transaction fails // Ensure that replacing a pending transaction with a future transaction fails
if err := pool.addRemote(pricedTransaction(5, 100000, big.NewInt(6), keys[1])); err != txpool.ErrFutureReplacePending { if err := pool.addRemote(pricedTransaction(5, 100000, big.NewInt(6), keys[1])); !errors.Is(err, txpool.ErrFutureReplacePending) {
t.Fatalf("adding future replace transaction error mismatch: have %v, want %v", err, txpool.ErrFutureReplacePending) t.Fatalf("adding future replace transaction error mismatch: have %v, want %v", err, txpool.ErrFutureReplacePending)
} }
pending, queued = pool.Stats() pending, queued = pool.Stats()
@ -2174,7 +2174,7 @@ func TestReplacement(t *testing.T) {
if err := pool.addRemoteSync(pricedTransaction(0, 100000, big.NewInt(1), key)); err != nil { if err := pool.addRemoteSync(pricedTransaction(0, 100000, big.NewInt(1), key)); err != nil {
t.Fatalf("failed to add original cheap pending transaction: %v", err) t.Fatalf("failed to add original cheap pending transaction: %v", err)
} }
if err := pool.addRemote(pricedTransaction(0, 100001, big.NewInt(1), key)); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(pricedTransaction(0, 100001, big.NewInt(1), key)); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original cheap pending transaction replacement error mismatch: have %v, want %v", err, txpool.ErrReplaceUnderpriced) t.Fatalf("original cheap pending transaction replacement error mismatch: have %v, want %v", err, txpool.ErrReplaceUnderpriced)
} }
if err := pool.addRemote(pricedTransaction(0, 100000, big.NewInt(2), key)); err != nil { if err := pool.addRemote(pricedTransaction(0, 100000, big.NewInt(2), key)); err != nil {
@ -2187,7 +2187,7 @@ func TestReplacement(t *testing.T) {
if err := pool.addRemoteSync(pricedTransaction(0, 100000, big.NewInt(price), key)); err != nil { if err := pool.addRemoteSync(pricedTransaction(0, 100000, big.NewInt(price), key)); err != nil {
t.Fatalf("failed to add original proper pending transaction: %v", err) t.Fatalf("failed to add original proper pending transaction: %v", err)
} }
if err := pool.addRemote(pricedTransaction(0, 100001, big.NewInt(threshold-1), key)); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(pricedTransaction(0, 100001, big.NewInt(threshold-1), key)); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original proper pending transaction replacement error mismatch: have %v, want %v", err, txpool.ErrReplaceUnderpriced) t.Fatalf("original proper pending transaction replacement error mismatch: have %v, want %v", err, txpool.ErrReplaceUnderpriced)
} }
if err := pool.addRemote(pricedTransaction(0, 100000, big.NewInt(threshold), key)); err != nil { if err := pool.addRemote(pricedTransaction(0, 100000, big.NewInt(threshold), key)); err != nil {
@ -2201,7 +2201,7 @@ func TestReplacement(t *testing.T) {
if err := pool.addRemote(pricedTransaction(2, 100000, big.NewInt(1), key)); err != nil { if err := pool.addRemote(pricedTransaction(2, 100000, big.NewInt(1), key)); err != nil {
t.Fatalf("failed to add original cheap queued transaction: %v", err) t.Fatalf("failed to add original cheap queued transaction: %v", err)
} }
if err := pool.addRemote(pricedTransaction(2, 100001, big.NewInt(1), key)); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(pricedTransaction(2, 100001, big.NewInt(1), key)); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original cheap queued transaction replacement error mismatch: have %v, want %v", err, txpool.ErrReplaceUnderpriced) t.Fatalf("original cheap queued transaction replacement error mismatch: have %v, want %v", err, txpool.ErrReplaceUnderpriced)
} }
if err := pool.addRemote(pricedTransaction(2, 100000, big.NewInt(2), key)); err != nil { if err := pool.addRemote(pricedTransaction(2, 100000, big.NewInt(2), key)); err != nil {
@ -2211,7 +2211,7 @@ func TestReplacement(t *testing.T) {
if err := pool.addRemote(pricedTransaction(2, 100000, big.NewInt(price), key)); err != nil { if err := pool.addRemote(pricedTransaction(2, 100000, big.NewInt(price), key)); err != nil {
t.Fatalf("failed to add original proper queued transaction: %v", err) t.Fatalf("failed to add original proper queued transaction: %v", err)
} }
if err := pool.addRemote(pricedTransaction(2, 100001, big.NewInt(threshold-1), key)); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(pricedTransaction(2, 100001, big.NewInt(threshold-1), key)); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original proper queued transaction replacement error mismatch: have %v, want %v", err, txpool.ErrReplaceUnderpriced) t.Fatalf("original proper queued transaction replacement error mismatch: have %v, want %v", err, txpool.ErrReplaceUnderpriced)
} }
if err := pool.addRemote(pricedTransaction(2, 100000, big.NewInt(threshold), key)); err != nil { if err := pool.addRemote(pricedTransaction(2, 100000, big.NewInt(threshold), key)); err != nil {
@ -2275,7 +2275,7 @@ func TestReplacementDynamicFee(t *testing.T) {
} }
// 2. Don't bump tip or feecap => discard // 2. Don't bump tip or feecap => discard
tx = dynamicFeeTx(nonce, 100001, big.NewInt(2), big.NewInt(1), key) tx = dynamicFeeTx(nonce, 100001, big.NewInt(2), big.NewInt(1), key)
if err := pool.addRemote(tx); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(tx); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original cheap %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced) t.Fatalf("original cheap %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced)
} }
// 3. Bump both more than min => accept // 3. Bump both more than min => accept
@ -2298,22 +2298,22 @@ func TestReplacementDynamicFee(t *testing.T) {
} }
// 6. Bump tip max allowed so it's still underpriced => discard // 6. Bump tip max allowed so it's still underpriced => discard
tx = dynamicFeeTx(nonce, 100000, big.NewInt(gasFeeCap), big.NewInt(tipThreshold-1), key) tx = dynamicFeeTx(nonce, 100000, big.NewInt(gasFeeCap), big.NewInt(tipThreshold-1), key)
if err := pool.addRemote(tx); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(tx); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original proper %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced) t.Fatalf("original proper %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced)
} }
// 7. Bump fee cap max allowed so it's still underpriced => discard // 7. Bump fee cap max allowed so it's still underpriced => discard
tx = dynamicFeeTx(nonce, 100000, big.NewInt(feeCapThreshold-1), big.NewInt(gasTipCap), key) tx = dynamicFeeTx(nonce, 100000, big.NewInt(feeCapThreshold-1), big.NewInt(gasTipCap), key)
if err := pool.addRemote(tx); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(tx); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original proper %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced) t.Fatalf("original proper %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced)
} }
// 8. Bump tip min for acceptance => accept // 8. Bump tip min for acceptance => accept
tx = dynamicFeeTx(nonce, 100000, big.NewInt(gasFeeCap), big.NewInt(tipThreshold), key) tx = dynamicFeeTx(nonce, 100000, big.NewInt(gasFeeCap), big.NewInt(tipThreshold), key)
if err := pool.addRemote(tx); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(tx); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original proper %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced) t.Fatalf("original proper %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced)
} }
// 9. Bump fee cap min for acceptance => accept // 9. Bump fee cap min for acceptance => accept
tx = dynamicFeeTx(nonce, 100000, big.NewInt(feeCapThreshold), big.NewInt(gasTipCap), key) tx = dynamicFeeTx(nonce, 100000, big.NewInt(feeCapThreshold), big.NewInt(gasTipCap), key)
if err := pool.addRemote(tx); err != txpool.ErrReplaceUnderpriced { if err := pool.addRemote(tx); !errors.Is(err, txpool.ErrReplaceUnderpriced) {
t.Fatalf("original proper %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced) t.Fatalf("original proper %s transaction replacement error mismatch: have %v, want %v", stage, err, txpool.ErrReplaceUnderpriced)
} }
// 10. Check events match expected (3 new executable txs during pending, 0 during queue) // 10. Check events match expected (3 new executable txs during pending, 0 during queue)

View file

@ -19,6 +19,7 @@ package types
import ( import (
"bytes" "bytes"
"encoding/json" "encoding/json"
"errors"
"math" "math"
"math/big" "math/big"
"reflect" "reflect"
@ -300,7 +301,7 @@ func TestDecodeEmptyTypedReceipt(t *testing.T) {
input := []byte{0x80} input := []byte{0x80}
var r Receipt var r Receipt
err := rlp.DecodeBytes(input, &r) err := rlp.DecodeBytes(input, &r)
if err != errShortTypedReceipt { if !errors.Is(err, errShortTypedReceipt) {
t.Fatal("wrong error:", err) t.Fatal("wrong error:", err)
} }
} }

View file

@ -76,7 +76,7 @@ func TestDecodeEmptyTypedTx(t *testing.T) {
input := []byte{0x80} input := []byte{0x80}
var tx Transaction var tx Transaction
err := rlp.DecodeBytes(input, &tx) err := rlp.DecodeBytes(input, &tx)
if err != errShortTypedTx { if !errors.Is(err, errShortTypedTx) {
t.Fatal("wrong error:", err) t.Fatal("wrong error:", err)
} }
} }
@ -569,7 +569,7 @@ func TestYParityJSONUnmarshalling(t *testing.T) {
// Unmarshal the tx // Unmarshal the tx
var tx Transaction var tx Transaction
err = tx.UnmarshalJSON(jsonBytes) err = tx.UnmarshalJSON(jsonBytes)
if err != test.wantErr { if !errors.Is(err, test.wantErr) {
t.Fatalf("wrong error: got %v, want %v", err, test.wantErr) t.Fatalf("wrong error: got %v, want %v", err, test.wantErr)
} }
}) })

View file

@ -18,6 +18,7 @@ package vm
import ( import (
"encoding/hex" "encoding/hex"
"errors"
"reflect" "reflect"
"testing" "testing"
@ -66,7 +67,7 @@ func TestEOFMarshaling(t *testing.T) {
got Container got Container
) )
t.Logf("b: %#x", b) t.Logf("b: %#x", b)
if err := got.UnmarshalBinary(b, true); err != nil && err != test.err { if err := got.UnmarshalBinary(b, true); err != nil && !errors.Is(err, test.err) {
t.Fatalf("test %d: got error \"%v\", want \"%v\"", i, err, test.err) t.Fatalf("test %d: got error \"%v\", want \"%v\"", i, err, test.err)
} }
if !reflect.DeepEqual(got, test.want) { if !reflect.DeepEqual(got, test.want) {

View file

@ -235,7 +235,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. // when we're in homestead this also counts for code storage gas errors.
if err != nil { if err != nil {
evm.StateDB.RevertToSnapshot(snapshot) evm.StateDB.RevertToSnapshot(snapshot)
if err != ErrExecutionReverted { if !errors.Is(err, ErrExecutionReverted) {
if evm.Config.Tracer != nil && evm.Config.Tracer.OnGasChange != nil { if evm.Config.Tracer != nil && evm.Config.Tracer.OnGasChange != nil {
evm.Config.Tracer.OnGasChange(gas, 0, tracing.GasChangeCallFailedExecution) evm.Config.Tracer.OnGasChange(gas, 0, tracing.GasChangeCallFailedExecution)
} }
@ -291,7 +291,7 @@ func (evm *EVM) CallCode(caller ContractRef, addr common.Address, input []byte,
} }
if err != nil { if err != nil {
evm.StateDB.RevertToSnapshot(snapshot) evm.StateDB.RevertToSnapshot(snapshot)
if err != ErrExecutionReverted { if !errors.Is(err, ErrExecutionReverted) {
if evm.Config.Tracer != nil && evm.Config.Tracer.OnGasChange != nil { if evm.Config.Tracer != nil && evm.Config.Tracer.OnGasChange != nil {
evm.Config.Tracer.OnGasChange(gas, 0, tracing.GasChangeCallFailedExecution) evm.Config.Tracer.OnGasChange(gas, 0, tracing.GasChangeCallFailedExecution)
} }
@ -338,7 +338,7 @@ func (evm *EVM) DelegateCall(caller ContractRef, addr common.Address, input []by
} }
if err != nil { if err != nil {
evm.StateDB.RevertToSnapshot(snapshot) evm.StateDB.RevertToSnapshot(snapshot)
if err != ErrExecutionReverted { if !errors.Is(err, ErrExecutionReverted) {
if evm.Config.Tracer != nil && evm.Config.Tracer.OnGasChange != nil { if evm.Config.Tracer != nil && evm.Config.Tracer.OnGasChange != nil {
evm.Config.Tracer.OnGasChange(gas, 0, tracing.GasChangeCallFailedExecution) evm.Config.Tracer.OnGasChange(gas, 0, tracing.GasChangeCallFailedExecution)
} }
@ -396,7 +396,7 @@ func (evm *EVM) StaticCall(caller ContractRef, addr common.Address, input []byte
} }
if err != nil { if err != nil {
evm.StateDB.RevertToSnapshot(snapshot) evm.StateDB.RevertToSnapshot(snapshot)
if err != ErrExecutionReverted { if !errors.Is(err, ErrExecutionReverted) {
if evm.Config.Tracer != nil && evm.Config.Tracer.OnGasChange != nil { if evm.Config.Tracer != nil && evm.Config.Tracer.OnGasChange != nil {
evm.Config.Tracer.OnGasChange(gas, 0, tracing.GasChangeCallFailedExecution) evm.Config.Tracer.OnGasChange(gas, 0, tracing.GasChangeCallFailedExecution)
} }
@ -509,9 +509,9 @@ func (evm *EVM) create(caller ContractRef, codeAndHash *codeAndHash, gas uint64,
contract.IsDeployment = true contract.IsDeployment = true
ret, err = evm.initNewContract(contract, address, value) ret, err = evm.initNewContract(contract, address, value)
if err != nil && (evm.chainRules.IsHomestead || err != ErrCodeStoreOutOfGas) { if err != nil && (evm.chainRules.IsHomestead || !errors.Is(err, ErrCodeStoreOutOfGas)) {
evm.StateDB.RevertToSnapshot(snapshot) evm.StateDB.RevertToSnapshot(snapshot)
if err != ErrExecutionReverted { if !errors.Is(err, ErrExecutionReverted) {
contract.UseGas(contract.Gas, evm.Config.Tracer, tracing.GasChangeCallFailedExecution) contract.UseGas(contract.Gas, evm.Config.Tracer, tracing.GasChangeCallFailedExecution)
} }
} }

View file

@ -43,8 +43,8 @@ func TestMemoryGasCost(t *testing.T) {
} }
for i, tt := range tests { for i, tt := range tests {
v, err := memoryGasCost(&Memory{}, tt.size) v, err := memoryGasCost(&Memory{}, tt.size)
if (err == ErrGasUintOverflow) != tt.overflow { if errors.Is(err, ErrGasUintOverflow) != tt.overflow {
t.Errorf("test %d: overflow mismatch: have %v, want %v", i, 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 { if v != tt.cost {
t.Errorf("test %d: gas cost mismatch: have %v, want %v", i, v, tt.cost) t.Errorf("test %d: gas cost mismatch: have %v, want %v", i, v, tt.cost)

View file

@ -17,6 +17,7 @@
package vm package vm
import ( import (
"errors"
"math" "math"
"github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common"
@ -679,9 +680,9 @@ func opCreate(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([]b
// homestead we must check for CodeStoreOutOfGasError (homestead only // homestead we must check for CodeStoreOutOfGasError (homestead only
// rule) and treat as an error, if the ruleset is frontier we must // rule) and treat as an error, if the ruleset is frontier we must
// ignore this error and pretend the operation was successful. // 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() stackvalue.Clear()
} else if suberr != nil && suberr != ErrCodeStoreOutOfGas { } else if suberr != nil && !errors.Is(suberr, ErrCodeStoreOutOfGas) {
stackvalue.Clear() stackvalue.Clear()
} else { } else {
stackvalue.SetBytes(addr.Bytes()) stackvalue.SetBytes(addr.Bytes())
@ -690,7 +691,7 @@ func opCreate(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([]b
scope.Contract.RefundGas(returnGas, interpreter.evm.Config.Tracer, tracing.GasChangeCallLeftOverRefunded) scope.Contract.RefundGas(returnGas, interpreter.evm.Config.Tracer, tracing.GasChangeCallLeftOverRefunded)
if suberr == ErrExecutionReverted { if errors.Is(suberr, ErrExecutionReverted) {
interpreter.returnData = res // set REVERT data to return data buffer interpreter.returnData = res // set REVERT data to return data buffer
return res, nil return res, nil
} }
@ -726,7 +727,7 @@ func opCreate2(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([]
scope.Stack.push(&stackvalue) scope.Stack.push(&stackvalue)
scope.Contract.RefundGas(returnGas, interpreter.evm.Config.Tracer, tracing.GasChangeCallLeftOverRefunded) scope.Contract.RefundGas(returnGas, interpreter.evm.Config.Tracer, tracing.GasChangeCallLeftOverRefunded)
if suberr == ErrExecutionReverted { if errors.Is(suberr, ErrExecutionReverted) {
interpreter.returnData = res // set REVERT data to return data buffer interpreter.returnData = res // set REVERT data to return data buffer
return res, nil return res, nil
} }
@ -760,7 +761,7 @@ func opCall(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([]byt
temp.SetOne() temp.SetOne()
} }
stack.push(&temp) stack.push(&temp)
if err == nil || err == ErrExecutionReverted { if err == nil || errors.Is(err, ErrExecutionReverted) {
scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret) scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret)
} }
@ -793,7 +794,7 @@ func opCallCode(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext) ([
temp.SetOne() temp.SetOne()
} }
stack.push(&temp) stack.push(&temp)
if err == nil || err == ErrExecutionReverted { if err == nil || errors.Is(err, ErrExecutionReverted) {
scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret) scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret)
} }
@ -822,7 +823,7 @@ func opDelegateCall(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext
temp.SetOne() temp.SetOne()
} }
stack.push(&temp) stack.push(&temp)
if err == nil || err == ErrExecutionReverted { if err == nil || errors.Is(err, ErrExecutionReverted) {
scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret) scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret)
} }
@ -851,7 +852,7 @@ func opStaticCall(pc *uint64, interpreter *EVMInterpreter, scope *ScopeContext)
temp.SetOne() temp.SetOne()
} }
stack.push(&temp) stack.push(&temp)
if err == nil || err == ErrExecutionReverted { if err == nil || errors.Is(err, ErrExecutionReverted) {
scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret) scope.Memory.Set(retOffset.Uint64(), retSize.Uint64(), ret)
} }

View file

@ -17,6 +17,7 @@
package vm package vm
import ( import (
"errors"
"fmt" "fmt"
"github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common"
@ -322,7 +323,7 @@ func (in *EVMInterpreter) Run(contract *Contract, input []byte, readOnly bool) (
pc++ pc++
} }
if err == errStopToken { if errors.Is(err, errStopToken) {
err = nil // clear stop token error err = nil // clear stop token error
} }

View file

@ -8,6 +8,7 @@ import (
"bytes" "bytes"
"encoding" "encoding"
"encoding/hex" "encoding/hex"
"errors"
"fmt" "fmt"
"hash" "hash"
"io" "io"
@ -166,7 +167,7 @@ func testHashes2X(t *testing.T) {
if _, err := h.Read(sum); err != nil { if _, err := h.Read(sum); err != nil {
t.Fatalf("#%d (single write): error from Read: %v", i, err) t.Fatalf("#%d (single write): error from Read: %v", i, err)
} }
if n, err := h.Read(sum); n != 0 || err != io.EOF { if n, err := h.Read(sum); n != 0 || !errors.Is(err, io.EOF) {
t.Fatalf("#%d (single write): Read did not return (0, io.EOF) after exhaustion, got (%v, %v)", i, n, err) t.Fatalf("#%d (single write): Read did not return (0, io.EOF) after exhaustion, got (%v, %v)", i, n, err)
} }
if gotHex := fmt.Sprintf("%x", sum); gotHex != expectedHex { if gotHex := fmt.Sprintf("%x", sum); gotHex != expectedHex {

View file

@ -191,7 +191,8 @@ func FromECDSAPub(pub *ecdsa.PublicKey) []byte {
// HexToECDSA parses a secp256k1 private key. // HexToECDSA parses a secp256k1 private key.
func HexToECDSA(hexkey string) (*ecdsa.PrivateKey, error) { func HexToECDSA(hexkey string) (*ecdsa.PrivateKey, error) {
b, err := hex.DecodeString(hexkey) b, err := hex.DecodeString(hexkey)
if byteErr, ok := err.(hex.InvalidByteError); ok { var byteErr hex.InvalidByteError
if ok := errors.As(err, &byteErr); ok {
return nil, fmt.Errorf("invalid hex character %q in private key", byte(byteErr)) return nil, fmt.Errorf("invalid hex character %q in private key", byte(byteErr))
} else if err != nil { } else if err != nil {
return nil, errors.New("invalid hex data for private key") return nil, errors.New("invalid hex data for private key")
@ -228,7 +229,7 @@ func readASCII(buf []byte, r *bufio.Reader) (n int, err error) {
for ; n < len(buf); n++ { for ; n < len(buf); n++ {
buf[n], err = r.ReadByte() buf[n], err = r.ReadByte()
switch { switch {
case err == io.EOF || buf[n] < '!': case errors.Is(err, io.EOF) || buf[n] < '!':
return n, nil return n, nil
case err != nil: case err != nil:
return n, err return n, err
@ -242,7 +243,7 @@ func checkKeyFileEnd(r *bufio.Reader) error {
for i := 0; ; i++ { for i := 0; ; i++ {
b, err := r.ReadByte() b, err := r.ReadByte()
switch { switch {
case err == io.EOF: case errors.Is(err, io.EOF):
return nil return nil
case err != nil: case err != nil:
return err return err

View file

@ -20,6 +20,7 @@ import (
"bytes" "bytes"
"crypto/ecdsa" "crypto/ecdsa"
"encoding/hex" "encoding/hex"
"errors"
"math/big" "math/big"
"os" "os"
"reflect" "reflect"
@ -66,11 +67,11 @@ func BenchmarkSha3(b *testing.B) {
func TestUnmarshalPubkey(t *testing.T) { func TestUnmarshalPubkey(t *testing.T) {
key, err := UnmarshalPubkey(nil) 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) t.Fatalf("expected error, got %v, %v", err, key)
} }
key, err = UnmarshalPubkey([]byte{1, 2, 3}) 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) t.Fatalf("expected error, got %v, %v", err, key)
} }

View file

@ -152,12 +152,12 @@ func TestTooBigSharedKey(t *testing.T) {
} }
_, err = prv1.GenerateShared(&prv2.PublicKey, 32, 32) _, 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") t.Fatal("ecdh: shared key should be too large for curve")
} }
_, err = prv2.GenerateShared(&prv1.PublicKey, 32, 32) _, 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") t.Fatal("ecdh: shared key should be too large for curve")
} }
} }
@ -355,7 +355,7 @@ func TestBasicKeyValidation(t *testing.T) {
for _, b := range badBytes { for _, b := range badBytes {
ct[0] = b ct[0] = b
_, err := prv.Decrypt(ct, nil, nil) _, err := prv.Decrypt(ct, nil, nil)
if err != ErrInvalidPublicKey { if !errors.Is(err, ErrInvalidPublicKey) {
t.Fatal("ecies: validated an invalid key") t.Fatal("ecies: validated an invalid key")
} }
} }

View file

@ -12,6 +12,7 @@ import (
"crypto/ecdsa" "crypto/ecdsa"
"crypto/rand" "crypto/rand"
"encoding/hex" "encoding/hex"
"errors"
"io" "io"
"testing" "testing"
) )
@ -91,7 +92,7 @@ func TestInvalidRecoveryID(t *testing.T) {
sig, _ := Sign(msg, seckey) sig, _ := Sign(msg, seckey)
sig[64] = 99 sig[64] = 99
_, err := RecoverPubkey(msg, sig) _, err := RecoverPubkey(msg, sig)
if err != ErrInvalidRecoveryID { if !errors.Is(err, ErrInvalidRecoveryID) {
t.Fatalf("got %q, want %q", err, ErrInvalidRecoveryID) t.Fatalf("got %q, want %q", err, ErrInvalidRecoveryID)
} }
} }

View file

@ -113,7 +113,7 @@ func (api *AdminAPI) ImportChain(file string) (bool, error) {
// Load a batch of blocks from the input file // Load a batch of blocks from the input file
for len(blocks) < cap(blocks) { for len(blocks) < cap(blocks) {
block := new(types.Block) block := new(types.Block)
if err := stream.Decode(block); err == io.EOF { if err := stream.Decode(block); errors.Is(err, io.EOF) {
break break
} else if err != nil { } else if err != nil {
return false, fmt.Errorf("block %d: failed to parse: %v", index, err) return false, fmt.Errorf("block %d: failed to parse: %v", index, err)

View file

@ -1295,7 +1295,9 @@ func TestNilWithdrawals(t *testing.T) {
status, err = api.NewPayloadV2(*execData.ExecutionPayload) status, err = api.NewPayloadV2(*execData.ExecutionPayload)
} }
if err != nil { if err != nil {
t.Fatalf("error validating payload: %v", err.(*engine.EngineAPIError).ErrorData()) engineAPIErr := new(engine.EngineAPIError)
_ = errors.As(err, &engineAPIErr)
t.Fatalf("error validating payload: %v", engineAPIErr.ErrorData())
} else if status.Status != engine.VALID { } else if status.Status != engine.VALID {
t.Fatalf("invalid payload") t.Fatalf("invalid payload")
} }
@ -1644,7 +1646,9 @@ func TestParentBeaconBlockRoot(t *testing.T) {
} }
resp, err := api.ForkchoiceUpdatedV3(fcState, &blockParams) resp, err := api.ForkchoiceUpdatedV3(fcState, &blockParams)
if err != nil { if err != nil {
t.Fatalf("error preparing payload, err=%v", err.(*engine.EngineAPIError).ErrorData()) engineAPIErr := new(engine.EngineAPIError)
_ = errors.As(err, &engineAPIErr)
t.Fatalf("error preparing payload, err=%v", engineAPIErr.ErrorData())
} }
if resp.PayloadStatus.Status != engine.VALID { if resp.PayloadStatus.Status != engine.VALID {
t.Fatalf("unexpected status (got: %s, want: %s)", resp.PayloadStatus.Status, engine.VALID) t.Fatalf("unexpected status (got: %s, want: %s)", resp.PayloadStatus.Status, engine.VALID)
@ -1675,7 +1679,9 @@ func TestParentBeaconBlockRoot(t *testing.T) {
fcState.HeadBlockHash = execData.ExecutionPayload.BlockHash fcState.HeadBlockHash = execData.ExecutionPayload.BlockHash
resp, err = api.ForkchoiceUpdatedV3(fcState, nil) resp, err = api.ForkchoiceUpdatedV3(fcState, nil)
if err != nil { if err != nil {
t.Fatalf("error preparing payload, err=%v", err.(*engine.EngineAPIError).ErrorData()) engineAPIErr := new(engine.EngineAPIError)
_ = errors.As(err, &engineAPIErr)
t.Fatalf("error preparing payload, err=%v", engineAPIErr.ErrorData())
} }
if resp.PayloadStatus.Status != engine.VALID { if resp.PayloadStatus.Status != engine.VALID {
t.Fatalf("unexpected status (got: %s, want: %s)", resp.PayloadStatus.Status, engine.VALID) t.Fatalf("unexpected status (got: %s, want: %s)", resp.PayloadStatus.Status, engine.VALID)

View file

@ -566,7 +566,7 @@ func (d *Downloader) spawnSync(fetchers []func() error) error {
} }
if got := <-errc; got != nil { if got := <-errc; got != nil {
err = got err = got
if got != errCanceled { if !errors.Is(got, errCanceled) {
break // receive a meaningful error, bubble it up break // receive a meaningful error, bubble it up
} }
} }
@ -817,7 +817,7 @@ func (d *Downloader) processSnapSyncContent() error {
}() }()
closeOnErr := func(s *stateSync) { 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 d.queue.Close() // wake up Results
} }
} }

View file

@ -278,25 +278,25 @@ func (s *skeleton) startup() {
// signalling as the sync loop should never terminate (TM). // signalling as the sync loop should never terminate (TM).
newhead, err := s.sync(head) newhead, err := s.sync(head)
switch { switch {
case err == errSyncLinked: case errors.Is(err, errSyncLinked):
// Sync cycle linked up to the genesis block, or the existent chain // Sync cycle linked up to the genesis block, or the existent chain
// segment. Tear down the loop and restart it so, it can properly // segment. Tear down the loop and restart it so, it can properly
// notify the backfiller. Don't account a new head. // notify the backfiller. Don't account a new head.
head = nil head = nil
case err == errSyncMerged: case errors.Is(err, errSyncMerged):
// Subchains were merged, we just need to reinit the internal // Subchains were merged, we just need to reinit the internal
// start to continue on the tail of the merged chain. Don't // start to continue on the tail of the merged chain. Don't
// announce a new head, // announce a new head,
head = nil head = nil
case err == errSyncReorged: case errors.Is(err, errSyncReorged):
// The subchain being synced got modified at the head in a // The subchain being synced got modified at the head in a
// way that requires resyncing it. Restart sync with the new // way that requires resyncing it. Restart sync with the new
// head to force a cleanup. // head to force a cleanup.
head = newhead head = newhead
case err == errTerminated: case errors.Is(err, errTerminated):
// Sync was requested to be terminated from within, stop and // Sync was requested to be terminated from within, stop and
// return (no need to pass a message, was already done internally) // return (no need to pass a message, was already done internally)
return return

View file

@ -459,7 +459,7 @@ func TestInvalidGetRangeLogsRequest(t *testing.T) {
api = NewFilterAPI(sys) api = NewFilterAPI(sys)
) )
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) t.Errorf("Expected Logs for invalid range return error, but got: %v", err)
} }
} }

View file

@ -19,6 +19,7 @@ package filters
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"math/big" "math/big"
"strings" "strings"
"testing" "testing"
@ -382,7 +383,7 @@ func TestFilters(t *testing.T) {
if err == nil { if err == nil {
t.Fatal("expected error") t.Fatal("expected error")
} }
if err != context.DeadlineExceeded { if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("expected context.DeadlineExceeded, got %v", err) t.Fatalf("expected context.DeadlineExceeded, got %v", err)
} }
}) })

View file

@ -90,7 +90,7 @@ func TestFeeHistory(t *testing.T) {
if len(blobBaseFee) != len(baseFee) { if len(blobBaseFee) != len(baseFee) {
t.Fatalf("Test case %d: blobBaseFee array length mismatch, want %d, got %d", i, len(baseFee), len(blobBaseFee)) t.Fatalf("Test case %d: blobBaseFee array length mismatch, want %d, got %d", i, len(baseFee), len(blobBaseFee))
} }
if err != c.expErr && !errors.Is(err, c.expErr) { if !errors.Is(err, c.expErr) {
t.Fatalf("Test case %d: error mismatch, want %v, got %v", i, c.expErr, err) t.Fatalf("Test case %d: error mismatch, want %v, got %v", i, c.expErr, err)
} }
} }

View file

@ -116,16 +116,16 @@ func markError(p *Peer, err error) {
return return
} }
m := meters.get(p.Inbound()) m := meters.get(p.Inbound())
switch errors.Unwrap(err) { switch {
case errNetworkIDMismatch: case errors.Is(err, errNetworkIDMismatch):
m.networkIDMismatch.Mark(1) m.networkIDMismatch.Mark(1)
case errProtocolVersionMismatch: case errors.Is(err, errProtocolVersionMismatch):
m.protocolVersionMismatch.Mark(1) m.protocolVersionMismatch.Mark(1)
case errGenesisMismatch: case errors.Is(err, errGenesisMismatch):
m.genesisMismatch.Mark(1) m.genesisMismatch.Mark(1)
case errForkIDRejected: case errors.Is(err, errForkIDRejected):
m.forkidRejected.Mark(1) m.forkidRejected.Mark(1)
case p2p.DiscReadTimeout: case errors.Is(err, p2p.DiscReadTimeout):
m.timeoutError.Mark(1) m.timeoutError.Mark(1)
default: default:
m.peerError.Mark(1) m.peerError.Mark(1)

View file

@ -2290,11 +2290,11 @@ func (s *Syncer) processTrienodeHealResponse(res *trienodeHealResponse) {
s.trienodeHealBytes += common.StorageSize(len(node)) s.trienodeHealBytes += common.StorageSize(len(node))
err := s.healer.scheduler.ProcessNode(trie.NodeSyncResult{Path: res.paths[i], Data: node}) err := s.healer.scheduler.ProcessNode(trie.NodeSyncResult{Path: res.paths[i], Data: node})
switch err { switch {
case nil: case err == nil:
case trie.ErrAlreadyProcessed: case errors.Is(err, trie.ErrAlreadyProcessed):
s.trienodeHealDups++ s.trienodeHealDups++
case trie.ErrNotRequested: case errors.Is(err, trie.ErrNotRequested):
s.trienodeHealNops++ s.trienodeHealNops++
default: default:
log.Error("Invalid trienode processed", "hash", hash, "err", err) log.Error("Invalid trienode processed", "hash", hash, "err", err)
@ -2377,11 +2377,11 @@ func (s *Syncer) processBytecodeHealResponse(res *bytecodeHealResponse) {
s.bytecodeHealBytes += common.StorageSize(len(node)) s.bytecodeHealBytes += common.StorageSize(len(node))
err := s.healer.scheduler.ProcessCode(trie.CodeSyncResult{Hash: hash, Data: node}) err := s.healer.scheduler.ProcessCode(trie.CodeSyncResult{Hash: hash, Data: node})
switch err { switch {
case nil: case err == nil:
case trie.ErrAlreadyProcessed: case errors.Is(err, trie.ErrAlreadyProcessed):
s.bytecodeHealDups++ s.bytecodeHealDups++
case trie.ErrNotRequested: case errors.Is(err, trie.ErrNotRequested):
s.bytecodeHealNops++ s.bytecodeHealNops++
default: default:
log.Error("Invalid bytecode processed", "hash", hash, "err", err) log.Error("Invalid bytecode processed", "hash", hash, "err", err)

View file

@ -118,12 +118,12 @@ func (eth *Ethereum) hashState(ctx context.Context, block *types.Block, reexec u
} }
} }
if err != nil { if err != nil {
switch err.(type) { missingNodeErr := new(*trie.MissingNodeError)
case *trie.MissingNodeError: ok := errors.As(err, missingNodeErr)
if ok {
return nil, nil, fmt.Errorf("required historical state unavailable (reexec=%d)", reexec) return nil, nil, fmt.Errorf("required historical state unavailable (reexec=%d)", reexec)
default:
return nil, nil, err
} }
return nil, nil, err
} }
} }
// State is available at historical point, re-execute the blocks on top for // State is available at historical point, re-execute the blocks on top for

View file

@ -286,7 +286,7 @@ func testTransactionInBlock(t *testing.T, client *rpc.Client) {
} }
// Test tx in block not found. // 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") t.Fatal("error should be ethereum.NotFound")
} }

View file

@ -22,6 +22,7 @@ package leveldb
import ( import (
"bytes" "bytes"
stderrors "errors"
"fmt" "fmt"
"sync" "sync"
"time" "time"
@ -120,7 +121,8 @@ func NewCustom(file string, namespace string, customize func(options *opt.Option
// Open the db and recover any potential corruptions // Open the db and recover any potential corruptions
db, err := leveldb.OpenFile(file, options) db, err := leveldb.OpenFile(file, options)
if _, corrupted := err.(*errors.ErrCorrupted); corrupted { corruptedErr := new(errors.ErrCorrupted)
if corrupted := stderrors.As(err, &corruptedErr); corrupted {
db, err = leveldb.RecoverFile(file, nil) db, err = leveldb.RecoverFile(file, nil)
} }
if err != nil { if err != nil {

View file

@ -19,6 +19,7 @@ package pebble
import ( import (
"bytes" "bytes"
"errors"
"fmt" "fmt"
"runtime" "runtime"
"sync" "sync"
@ -288,7 +289,7 @@ func (d *Database) Has(key []byte) (bool, error) {
return false, pebble.ErrClosed return false, pebble.ErrClosed
} }
_, closer, err := d.db.Get(key) _, closer, err := d.db.Get(key)
if err == pebble.ErrNotFound { if errors.Is(err, pebble.ErrNotFound) {
return false, nil return false, nil
} else if err != nil { } else if err != nil {
return false, err return false, err

View file

@ -17,6 +17,7 @@
package event package event
import ( import (
"errors"
"math/rand" "math/rand"
"sync" "sync"
"testing" "testing"
@ -59,7 +60,7 @@ func TestMuxErrorAfterStop(t *testing.T) {
if _, isopen := <-sub.Chan(); isopen { if _, isopen := <-sub.Chan(); isopen {
t.Errorf("subscription channel was not closed") 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) t.Errorf("Post error mismatch, got: %s, expected: %s", err, ErrMuxClosed)
} }
} }

View file

@ -56,7 +56,7 @@ loop:
t.Fatalf("wrong int %d, want %d", got, want) t.Fatalf("wrong int %d, want %d", got, want)
} }
case err := <-sub.Err(): case err := <-sub.Err():
if err != errInts { if !errors.Is(err, errInts) {
t.Fatalf("wrong error: got %q, want %q", err, errInts) t.Fatalf("wrong error: got %q, want %q", err, errInts)
} }
if want != 2 { if want != 2 {

View file

@ -218,7 +218,7 @@ func extractTarball(ar io.Reader, dest string) error {
// Move to the next file header. // Move to the next file header.
header, err := tr.Next() header, err := tr.Next()
if err != nil { if err != nil {
if err == io.EOF { if errors.Is(err, io.EOF) {
return nil return nil
} }
return err return err

View file

@ -19,6 +19,7 @@ package build
import ( import (
"bufio" "bufio"
"bytes" "bytes"
"errors"
"flag" "flag"
"fmt" "fmt"
"go/parser" "go/parser"
@ -97,7 +98,8 @@ func RunGit(args ...string) string {
var stdout, stderr bytes.Buffer var stdout, stderr bytes.Buffer
cmd.Stdout, cmd.Stderr = &stdout, &stderr cmd.Stdout, cmd.Stderr = &stdout, &stderr
if err := cmd.Run(); err != nil { if err := cmd.Run(); err != nil {
if e, ok := err.(*exec.Error); ok && e.Err == exec.ErrNotFound { e := new(exec.Error)
if ok := errors.As(err, &e); ok && errors.Is(e.Err, exec.ErrNotFound) {
if !warnedAboutGit { if !warnedAboutGit {
log.Println("Warning: can't find 'git' in PATH") log.Println("Warning: can't find 'git' in PATH")
warnedAboutGit = true warnedAboutGit = true

View file

@ -19,6 +19,7 @@ package cmdtest
import ( import (
"bufio" "bufio"
"bytes" "bytes"
"errors"
"fmt" "fmt"
"io" "io"
"os" "os"
@ -206,13 +207,12 @@ func (tt *TestCmd) Interrupt() {
// It will only return a valid value after the process has finished. // It will only return a valid value after the process has finished.
func (tt *TestCmd) ExitStatus() int { func (tt *TestCmd) ExitStatus() int {
if tt.Err != nil { if tt.Err != nil {
exitErr := tt.Err.(*exec.ExitError) exitErr := new(exec.ExitError)
if exitErr != nil { _ = errors.As(tt.Err, &exitErr)
if status, ok := exitErr.Sys().(syscall.WaitStatus); ok { if status, ok := exitErr.Sys().(syscall.WaitStatus); ok {
return status.ExitStatus() return status.ExitStatus()
} }
} }
}
return 0 return 0
} }

View file

@ -110,7 +110,7 @@ func (r *Reader) ReadAt(entry *Entry, off int64) (int, error) {
n += headerSize n += headerSize
// An entry with a non-zero length should not return EOF when // An entry with a non-zero length should not return EOF when
// reading the value. // reading the value.
if err == io.EOF { if errors.Is(err, io.EOF) {
return n, io.ErrUnexpectedEOF return n, io.ErrUnexpectedEOF
} }
return n, err return n, err
@ -151,7 +151,7 @@ func (r *Reader) LengthAt(off int64) (int64, error) {
func (r *Reader) ReadMetadataAt(off int64) (typ uint16, length uint32, err error) { func (r *Reader) ReadMetadataAt(off int64) (typ uint16, length uint32, err error) {
b := make([]byte, headerSize) b := make([]byte, headerSize)
if n, err := r.r.ReadAt(b, off); err != nil { if n, err := r.r.ReadAt(b, off); err != nil {
if err == io.EOF && n > 0 { if errors.Is(err, io.EOF) && n > 0 {
return 0, 0, io.ErrUnexpectedEOF return 0, 0, io.ErrUnexpectedEOF
} }
return 0, 0, err return 0, 0, err
@ -177,7 +177,7 @@ func (r *Reader) Find(want uint16) (*Entry, error) {
) )
for { for {
typ, length, err = r.ReadMetadataAt(off) typ, length, err = r.ReadMetadataAt(off)
if err == io.EOF { if errors.Is(err, io.EOF) {
return nil, io.EOF return nil, io.EOF
} else if err != nil { } else if err != nil {
return nil, err return nil, err
@ -204,7 +204,7 @@ func (r *Reader) FindAll(want uint16) ([]*Entry, error) {
) )
for { for {
typ, length, err = r.ReadMetadataAt(off) typ, length, err = r.ReadMetadataAt(off)
if err == io.EOF { if errors.Is(err, io.EOF) {
return entries, nil return entries, nil
} else if err != nil { } else if err != nil {
return entries, err return entries, err

View file

@ -182,7 +182,7 @@ func (it *RawIterator) Number() uint64 {
// Error returns the error status of the iterator. It should be called before // Error returns the error status of the iterator. It should be called before
// reading from any of the iterator's values. // reading from any of the iterator's values.
func (it *RawIterator) Error() error { func (it *RawIterator) Error() error {
if it.err == io.EOF { if errors.Is(it.err, io.EOF) {
return nil return nil
} }
return it.err return it.err

View file

@ -17,6 +17,7 @@
package jsre package jsre
import ( import (
"errors"
"fmt" "fmt"
"io" "io"
"reflect" "reflect"
@ -60,7 +61,8 @@ func prettyPrint(vm *goja.Runtime, value goja.Value, w io.Writer) {
// prettyError writes err to standard output. // prettyError writes err to standard output.
func prettyError(vm *goja.Runtime, err error, w io.Writer) { func prettyError(vm *goja.Runtime, err error, w io.Writer) {
failure := err.Error() failure := err.Error()
if gojaErr, ok := err.(*goja.Exception); ok { gojaErr := new(goja.Exception)
if ok := errors.As(err, &gojaErr); ok {
failure = gojaErr.String() failure = gojaErr.String()
} }
fmt.Fprint(w, ErrorColor("%s", failure)) fmt.Fprint(w, ErrorColor("%s", failure))

View file

@ -20,6 +20,7 @@ package metrics
import ( import (
"bufio" "bufio"
"errors"
"fmt" "fmt"
"io" "io"
"os" "os"
@ -42,7 +43,7 @@ func ReadDiskStats(stats *DiskStats) error {
// Read the next line and split to key and value // Read the next line and split to key and value
line, err := in.ReadString('\n') line, err := in.ReadString('\n')
if err != nil { if err != nil {
if err == io.EOF { if errors.Is(err, io.EOF) {
return nil return nil
} }
return err return err

View file

@ -33,7 +33,8 @@ var (
) )
func convertFileLockError(err error) error { func convertFileLockError(err error) error {
if errno, ok := err.(syscall.Errno); ok && datadirInUseErrnos[uint(errno)] { var errno syscall.Errno
if ok := errors.As(err, &errno); ok && datadirInUseErrnos[uint(errno)] {
return ErrDatadirUsed return ErrDatadirUsed
} }
return err return err

View file

@ -56,7 +56,7 @@ func TestNodeCloseMultipleTimes(t *testing.T) {
// Ensure that a stopped node can be stopped again // Ensure that a stopped node can be stopped again
for i := 0; i < 3; i++ { 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) t.Fatalf("iter %d: stop failure mismatch: have %v, want %v", i, err, ErrNodeStopped)
} }
} }
@ -72,14 +72,14 @@ func TestNodeStartMultipleTimes(t *testing.T) {
if err := stack.Start(); err != nil { if err := stack.Start(); err != nil {
t.Fatalf("failed to start node: %v", err) 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) t.Fatalf("start failure mismatch: have %v, want %v ", err, ErrNodeRunning)
} }
// Ensure that a node can be stopped, but only once // Ensure that a node can be stopped, but only once
if err := stack.Close(); err != nil { if err := stack.Close(); err != nil {
t.Fatalf("failed to stop node: %v", err) 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) t.Fatalf("stop failure mismatch: have %v, want %v ", err, ErrNodeStopped)
} }
} }
@ -101,7 +101,7 @@ func TestNodeUsedDataDir(t *testing.T) {
// Create a second node based on the same data directory and ensure failure // Create a second node based on the same data directory and ensure failure
_, err = New(&Config{DataDir: dir}) _, 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) t.Fatalf("duplicate datadir failure mismatch: have %v, want %v", err, ErrDatadirUsed)
} }
} }
@ -297,7 +297,7 @@ func TestLifecycleStartupError(t *testing.T) {
stack.RegisterLifecycle(failer) stack.RegisterLifecycle(failer)
// Start the protocol stack and ensure all started services stop // Start the protocol stack and ensure all started services stop
if err := stack.Start(); err != failure { if err := stack.Start(); !errors.Is(err, failure) {
t.Fatalf("stack startup failure mismatch: have %v, want %v", err, failure) t.Fatalf("stack startup failure mismatch: have %v, want %v", err, failure)
} }
for id := range lifecycles { for id := range lifecycles {
@ -361,15 +361,17 @@ func TestLifecycleTerminationGuarantee(t *testing.T) {
} }
// Stop the stack, verify failure and check all terminations // Stop the stack, verify failure and check all terminations
err = stack.Close() err = stack.Close()
if err, ok := err.(*StopError); !ok { stopErr := new(StopError)
ok := errors.As(err, &stopErr)
if !ok {
t.Fatalf("termination failure mismatch: have %v, want StopError", err) t.Fatalf("termination failure mismatch: have %v, want StopError", err)
} else { } else {
failer := reflect.TypeOf(&InstrumentedService{}) failer := reflect.TypeOf(&InstrumentedService{})
if err.Services[failer] != failure { if !errors.Is(stopErr.Services[failer], failure) {
t.Fatalf("failer termination failure mismatch: have %v, want %v", err.Services[failer], failure) t.Fatalf("failer termination failure mismatch: have %v, want %v", stopErr.Services[failer], failure)
} }
if len(err.Services) != 1 { if len(stopErr.Services) != 1 {
t.Fatalf("failure count mismatch: have %d, want %d", len(err.Services), 1) t.Fatalf("failure count mismatch: have %d, want %d", len(stopErr.Services), 1)
} }
} }
for id := range lifecycles { for id := range lifecycles {

View file

@ -281,7 +281,7 @@ func (h *httpServer) doStop() {
ctx, cancel := context.WithTimeout(context.Background(), shutdownTimeout) ctx, cancel := context.WithTimeout(context.Background(), shutdownTimeout)
defer cancel() defer cancel()
err := h.server.Shutdown(ctx) err := h.server.Shutdown(ctx)
if err != nil && err == ctx.Err() { if err != nil && errors.Is(err, ctx.Err()) {
h.log.Warn("HTTP server graceful shutdown timed out") h.log.Warn("HTTP server graceful shutdown timed out")
h.server.Close() h.server.Close()
} }

View file

@ -626,7 +626,8 @@ func (t *dialTask) String() string {
} }
func cleanupDialErr(err error) error { func cleanupDialErr(err error) error {
if netErr, ok := err.(*net.OpError); ok && netErr.Op == "dial" { netErr := new(net.OpError)
if ok := errors.As(err, &netErr); ok && netErr.Op == "dial" {
return netErr.Err return netErr.Err
} }
return err return err

View file

@ -101,7 +101,7 @@ func (test *udpTest) packetInFrom(wantError error, key *ecdsa.PrivateKey, addr n
test.t.Errorf("%s encode error: %v", data.Name(), err) test.t.Errorf("%s encode error: %v", data.Name(), err)
} }
test.sent = append(test.sent, enc) test.sent = append(test.sent, enc)
if err = test.udp.handlePacket(addr, enc); err != wantError { if err = test.udp.handlePacket(addr, enc); !errors.Is(err, wantError) {
test.t.Errorf("error mismatch: got %q, want %q", err, wantError) test.t.Errorf("error mismatch: got %q, want %q", err, wantError)
} }
} }
@ -112,7 +112,7 @@ func (test *udpTest) waitPacketOut(validate interface{}) (closed bool) {
test.t.Helper() test.t.Helper()
dgram, err := test.pipe.receive() dgram, err := test.pipe.receive()
if err == errClosed { if errors.Is(err, errClosed) {
return true return true
} else if err != nil { } else if err != nil {
test.t.Error("packet receive error:", err) test.t.Error("packet receive error:", err)
@ -151,7 +151,7 @@ func TestUDPv4_pingTimeout(t *testing.T) {
key := newkey() key := newkey()
toaddr := &net.UDPAddr{IP: net.ParseIP("1.2.3.4"), Port: 2222} toaddr := &net.UDPAddr{IP: net.ParseIP("1.2.3.4"), Port: 2222}
node := enode.NewV4(&key.PublicKey, toaddr.IP, 0, toaddr.Port) 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) t.Error("expected timeout error, got", err)
} }
} }
@ -211,7 +211,7 @@ func TestUDPv4_responseTimeouts(t *testing.T) {
for i := 0; i < nReqs; i++ { for i := 0; i < nReqs; i++ {
select { select {
case err := <-timeoutErr: case err := <-timeoutErr:
if err != errTimeout { if !errors.Is(err, errTimeout) {
t.Fatalf("got non-timeout error on timeoutErr %d: %v", i, err) t.Fatalf("got non-timeout error on timeoutErr %d: %v", i, err)
} }
nTimeoutsRecv++ nTimeoutsRecv++
@ -241,7 +241,7 @@ func TestUDPv4_findnodeTimeout(t *testing.T) {
toid := enode.ID{1, 2, 3, 4} toid := enode.ID{1, 2, 3, 4}
target := v4wire.Pubkey{4, 5, 6, 7} target := v4wire.Pubkey{4, 5, 6, 7}
result, err := test.udp.findnode(toid, toaddr, target) result, err := test.udp.findnode(toid, toaddr, target)
if err != errTimeout { if !errors.Is(err, errTimeout) {
t.Error("expected timeout error, got", err) t.Error("expected timeout error, got", err)
} }
if len(result) > 0 { if len(result) > 0 {

View file

@ -20,6 +20,7 @@ import (
"bytes" "bytes"
"crypto/ecdsa" "crypto/ecdsa"
"encoding/binary" "encoding/binary"
"errors"
"fmt" "fmt"
"math/rand" "math/rand"
"net" "net"
@ -241,7 +242,7 @@ func TestUDPv5_pingCall(t *testing.T) {
done <- err done <- err
}() }()
test.waitPacketOut(func(p *v5wire.Ping, addr netip.AddrPort, _ v5wire.Nonce) {}) test.waitPacketOut(func(p *v5wire.Ping, addr netip.AddrPort, _ v5wire.Nonce) {})
if err := <-done; err != errTimeout { if err := <-done; !errors.Is(err, errTimeout) {
t.Fatalf("want errTimeout, got %q", err) t.Fatalf("want errTimeout, got %q", err)
} }
@ -266,7 +267,7 @@ func TestUDPv5_pingCall(t *testing.T) {
wrongAddr := netip.MustParseAddrPort("33.44.55.22:10101") wrongAddr := netip.MustParseAddrPort("33.44.55.22:10101")
test.packetInFrom(test.remotekey, wrongAddr, &v5wire.Pong{ReqID: p.ReqID}) 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) t.Fatalf("want errTimeout for reply from wrong IP, got %q", err)
} }
} }
@ -379,7 +380,7 @@ func TestUDPv5_multipleHandshakeRounds(t *testing.T) {
test.waitPacketOut(func(p *v5wire.Ping, addr netip.AddrPort, nonce v5wire.Nonce) { test.waitPacketOut(func(p *v5wire.Ping, addr netip.AddrPort, nonce v5wire.Nonce) {
test.packetIn(&v5wire.Whoareyou{Nonce: 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) t.Fatalf("unexpected ping error: %q", err)
} }
} }
@ -488,7 +489,7 @@ func TestUDPv5_talkRequest(t *testing.T) {
done <- err done <- err
}() }()
test.waitPacketOut(func(p *v5wire.TalkRequest, addr netip.AddrPort, _ v5wire.Nonce) {}) test.waitPacketOut(func(p *v5wire.TalkRequest, addr netip.AddrPort, _ v5wire.Nonce) {})
if err := <-done; err != errTimeout { if err := <-done; !errors.Is(err, errTimeout) {
t.Fatalf("want errTimeout, got %q", err) t.Fatalf("want errTimeout, got %q", err)
} }
@ -817,10 +818,10 @@ func (test *udpV5Test) waitPacketOut(validate interface{}) (closed bool) {
exptype := fn.Type().In(0) exptype := fn.Type().In(0)
dgram, err := test.pipe.receive() dgram, err := test.pipe.receive()
if err == errClosed { if errors.Is(err, errClosed) {
return true return true
} }
if err == errTimeout { if errors.Is(err, errTimeout) {
test.t.Fatalf("timed out waiting for %v", exptype) test.t.Fatalf("timed out waiting for %v", exptype)
return false return false
} }

View file

@ -128,7 +128,7 @@ var (
// returns false, it is pretty certain that the packet causing the error does not belong // returns false, it is pretty certain that the packet causing the error does not belong
// to discv5. // to discv5.
func IsInvalidHeader(err error) bool { 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. // Packet sizes.

View file

@ -95,7 +95,7 @@ func TestClientSyncTreeBadNode(t *testing.T) {
c := NewClient(Config{Resolver: r, Logger: testlog.Logger(t, log.LvlTrace)}) c := NewClient(Config{Resolver: r, Logger: testlog.Logger(t, log.LvlTrace)})
_, err := c.SyncTree("enrtree://AKPYQIUQIL7PSIACI32J7FGZW56E5FKHEFCCOFHILBIMW3M6LWXS2@n") _, err := c.SyncTree("enrtree://AKPYQIUQIL7PSIACI32J7FGZW56E5FKHEFCCOFHILBIMW3M6LWXS2@n")
wantErr := nameError{name: "INDMVBZEEQ4ESVYAKGIYU74EAA.n", err: entryError{typ: "enr", err: errInvalidENR}} wantErr := nameError{name: "INDMVBZEEQ4ESVYAKGIYU74EAA.n", err: entryError{typ: "enr", err: errInvalidENR}}
if err != wantErr { if !errors.Is(err, wantErr) {
t.Fatalf("expected sync error %q, got %q", wantErr, err) t.Fatalf("expected sync error %q, got %q", wantErr, err)
} }
} }

View file

@ -47,7 +47,8 @@ type nameError struct {
} }
func (err nameError) Error() string { func (err nameError) Error() string {
if ee, ok := err.err.(entryError); ok { ee := new(entryError)
if ok := errors.As(err.err, &ee); ok {
return fmt.Sprintf("invalid %s entry at %s: %v", ee.typ, err.name, ee.err) return fmt.Sprintf("invalid %s entry at %s: %v", ee.typ, err.name, ee.err)
} }
return err.name + ": " + err.err.Error() return err.name + ": " + err.err.Error()

View file

@ -17,6 +17,7 @@
package dnsdisc package dnsdisc
import ( import (
"errors"
"reflect" "reflect"
"testing" "testing"
@ -54,7 +55,7 @@ func TestParseRoot(t *testing.T) {
if !reflect.DeepEqual(e, test.e) { if !reflect.DeepEqual(e, test.e) {
t.Errorf("test %d: wrong entry %s, want %s", i, spew.Sdump(e), spew.Sdump(test.e)) t.Errorf("test %d: wrong entry %s, want %s", i, spew.Sdump(e), spew.Sdump(test.e))
} }
if err != test.err { if !errors.Is(err, test.err) {
t.Errorf("test %d: wrong error %q, want %q", i, err, test.err) t.Errorf("test %d: wrong error %q, want %q", i, err, test.err)
} }
} }
@ -131,7 +132,7 @@ func TestParseEntry(t *testing.T) {
if !reflect.DeepEqual(e, test.e) { if !reflect.DeepEqual(e, test.e) {
t.Errorf("test %d: wrong entry %s, want %s", i, spew.Sdump(e), spew.Sdump(test.e)) t.Errorf("test %d: wrong entry %s, want %s", i, spew.Sdump(e), spew.Sdump(test.e))
} }
if err != test.err { if !errors.Is(err, test.err) {
t.Errorf("test %d: wrong error %q, want %q", i, err, test.err) t.Errorf("test %d: wrong error %q, want %q", i, err, test.err)
} }
} }

View file

@ -20,6 +20,7 @@ import (
"bytes" "bytes"
"crypto/rand" "crypto/rand"
"encoding/binary" "encoding/binary"
stderrors "errors"
"fmt" "fmt"
"net/netip" "net/netip"
"os" "os"
@ -99,7 +100,8 @@ func newMemoryDB() (*DB, error) {
func newPersistentDB(path string) (*DB, error) { func newPersistentDB(path string) (*DB, error) {
opts := &opt.Options{OpenFilesCacheCapacity: 5} opts := &opt.Options{OpenFilesCacheCapacity: 5}
db, err := leveldb.OpenFile(path, opts) db, err := leveldb.OpenFile(path, opts)
if _, iscorrupted := err.(*errors.ErrCorrupted); iscorrupted { errCorrupted := new(*errors.ErrCorrupted)
if iscorrupted := stderrors.As(err, errCorrupted); iscorrupted {
db, err = leveldb.RecoverFile(path, nil) db, err = leveldb.RecoverFile(path, nil)
} }
if err != nil { if err != nil {
@ -111,15 +113,15 @@ func newPersistentDB(path string) (*DB, error) {
currentVer = currentVer[:binary.PutVarint(currentVer, int64(dbVersion))] currentVer = currentVer[:binary.PutVarint(currentVer, int64(dbVersion))]
blob, err := db.Get([]byte(dbVersionKey), nil) blob, err := db.Get([]byte(dbVersionKey), nil)
switch err { switch {
case leveldb.ErrNotFound: case stderrors.Is(err, leveldb.ErrNotFound):
// Version not found (i.e. empty cache), insert it // Version not found (i.e. empty cache), insert it
if err := db.Put([]byte(dbVersionKey), currentVer, nil); err != nil { if err := db.Put([]byte(dbVersionKey), currentVer, nil); err != nil {
db.Close() db.Close()
return nil, err return nil, err
} }
case nil: case err == nil:
// Version present, flush if different // Version present, flush if different
if !bytes.Equal(blob, currentVer) { if !bytes.Equal(blob, currentVer) {
db.Close() db.Close()

View file

@ -228,13 +228,13 @@ func decodeRecord(s *rlp.Stream) (dec Record, raw []byte, err error) {
return dec, raw, err return dec, raw, err
} }
if err = s.Decode(&dec.signature); err != nil { if err = s.Decode(&dec.signature); err != nil {
if err == rlp.EOL { if errors.Is(err, rlp.EOL) {
err = errIncompleteList err = errIncompleteList
} }
return dec, raw, err return dec, raw, err
} }
if err = s.Decode(&dec.seq); err != nil { if err = s.Decode(&dec.seq); err != nil {
if err == rlp.EOL { if errors.Is(err, rlp.EOL) {
err = errIncompleteList err = errIncompleteList
} }
return dec, raw, err return dec, raw, err
@ -244,13 +244,13 @@ func decodeRecord(s *rlp.Stream) (dec Record, raw []byte, err error) {
for i := 0; ; i++ { for i := 0; ; i++ {
var kv pair var kv pair
if err := s.Decode(&kv.k); err != nil { if err := s.Decode(&kv.k); err != nil {
if err == rlp.EOL { if errors.Is(err, rlp.EOL) {
break break
} }
return dec, raw, err return dec, raw, err
} }
if err := s.Decode(&kv.v); err != nil { if err := s.Decode(&kv.v); err != nil {
if err == rlp.EOL { if errors.Is(err, rlp.EOL) {
return dec, raw, errIncompletePair return dec, raw, errIncompletePair
} }
return dec, raw, err return dec, raw, err

View file

@ -19,6 +19,7 @@ package enr
import ( import (
"bytes" "bytes"
"encoding/binary" "encoding/binary"
"errors"
"fmt" "fmt"
"math/rand" "math/rand"
"testing" "testing"
@ -97,7 +98,8 @@ func TestLoadErrors(t *testing.T) {
// Check error for invalid keys. // Check error for invalid keys.
var list []uint var list []uint
err = r.Load(WithEntry(ip4.ENRKey(), &list)) err = r.Load(WithEntry(ip4.ENRKey(), &list))
kerr, ok := err.(*KeyError) kerr := new(KeyError)
ok := errors.As(err, &kerr)
if !ok { if !ok {
t.Fatalf("expected KeyError, got %T", err) t.Fatalf("expected KeyError, got %T", err)
} }
@ -149,7 +151,7 @@ func TestSortedGetAndSet(t *testing.T) {
func TestDirty(t *testing.T) { func TestDirty(t *testing.T) {
var r Record 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) t.Errorf("expected errEncodeUnsigned, got %#v", err)
} }
@ -164,7 +166,7 @@ func TestDirty(t *testing.T) {
if len(r.signature) != 0 { if len(r.signature) != 0 {
t.Error("signature still set after modification") 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) t.Errorf("expected errEncodeUnsigned, got %#v", err)
} }
} }
@ -248,7 +250,7 @@ func TestRecordTooBig(t *testing.T) {
// set a big value for random key, expect error // set a big value for random key, expect error
r.Set(WithEntry(key, randomString(SizeLimit))) 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) t.Fatalf("expected to get errTooBig, got %#v", err)
} }
@ -274,7 +276,7 @@ func TestDecodeIncomplete(t *testing.T) {
for _, test := range tests { for _, test := range tests {
var r Record var r Record
err := rlp.DecodeBytes(test.input, &r) err := rlp.DecodeBytes(test.input, &r)
if err != test.err { if !errors.Is(err, test.err) {
t.Errorf("wrong error for %X: %v", test.input, err) t.Errorf("wrong error for %X: %v", test.input, err)
} }
} }

View file

@ -240,7 +240,7 @@ type KeyError struct {
// Error implements error. // Error implements error.
func (err *KeyError) Error() string { 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("missing ENR key %q", err.Key)
} }
return fmt.Sprintf("ENR key %q: %v", err.Key, err.Err) return fmt.Sprintf("ENR key %q: %v", err.Key, err.Err)
@ -255,7 +255,7 @@ func (err *KeyError) Unwrap() error {
func IsNotFound(err error) bool { func IsNotFound(err error) bool {
var ke *KeyError var ke *KeyError
if errors.As(err, &ke) { if errors.As(err, &ke) {
return ke.Err == errNotFound return errors.Is(ke.Err, errNotFound)
} }
return false return false
} }

View file

@ -18,6 +18,7 @@ package p2p
import ( import (
"bytes" "bytes"
"errors"
"fmt" "fmt"
"io" "io"
"runtime" "runtime"
@ -55,7 +56,7 @@ loop:
go func() { go func() {
if err := SendItems(rw1, 1); err == nil { if err := SendItems(rw1, 1); err == nil {
t.Error("EncodeMsg returned nil error") 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) t.Errorf("EncodeMsg returned wrong error: got %v, want %v", err, ErrPipeClosed)
} }
close(done) close(done)
@ -91,7 +92,7 @@ func TestEOFSignal(t *testing.T) {
// empty reader // empty reader
eof := make(chan struct{}, 1) eof := make(chan struct{}, 1)
sig := &eofSignal{new(bytes.Buffer), 0, eof} sig := &eofSignal{new(bytes.Buffer), 0, eof}
if n, err := sig.Read(rb); n != 0 || err != io.EOF { if n, err := sig.Read(rb); n != 0 || !errors.Is(err, io.EOF) {
t.Errorf("Read returned unexpected values: (%v, %v)", n, err) t.Errorf("Read returned unexpected values: (%v, %v)", n, err)
} }
select { select {
@ -118,7 +119,7 @@ func TestEOFSignal(t *testing.T) {
if n, err := sig.Read(rb); n != 4 || err != nil { if n, err := sig.Read(rb); n != 4 || err != nil {
t.Errorf("Read returned unexpected values: (%v, %v)", n, err) t.Errorf("Read returned unexpected values: (%v, %v)", n, err)
} }
if n, err := sig.Read(rb); n != 0 || err != io.EOF { if n, err := sig.Read(rb); n != 0 || !errors.Is(err, io.EOF) {
t.Errorf("Read returned unexpected values: (%v, %v)", n, err) t.Errorf("Read returned unexpected values: (%v, %v)", n, err)
} }
select { select {

View file

@ -70,20 +70,20 @@ func markDialError(err error) {
if err2 := errors.Unwrap(err); err2 != nil { if err2 := errors.Unwrap(err); err2 != nil {
err = err2 err = err2
} }
switch err { switch {
case DiscTooManyPeers: case errors.Is(err, DiscTooManyPeers):
dialTooManyPeers.Mark(1) dialTooManyPeers.Mark(1)
case DiscAlreadyConnected: case errors.Is(err, DiscAlreadyConnected):
dialAlreadyConnected.Mark(1) dialAlreadyConnected.Mark(1)
case DiscSelf: case errors.Is(err, DiscSelf):
dialSelf.Mark(1) dialSelf.Mark(1)
case DiscUselessPeer: case errors.Is(err, DiscUselessPeer):
dialUselessPeer.Mark(1) dialUselessPeer.Mark(1)
case DiscUnexpectedIdentity: case errors.Is(err, DiscUnexpectedIdentity):
dialUnexpectedIdentity.Mark(1) dialUnexpectedIdentity.Mark(1)
case errEncHandshakeError: case errors.Is(err, errEncHandshakeError):
dialEncHandshakeError.Mark(1) dialEncHandshakeError.Mark(1)
case errProtoHandshakeError: case errors.Is(err, errProtoHandshakeError):
dialProtoHandshakeError.Mark(1) dialProtoHandshakeError.Mark(1)
} }
} }

View file

@ -16,18 +16,22 @@
package netutil package netutil
import "errors"
// IsTemporaryError checks whether the given error should be considered temporary. // IsTemporaryError checks whether the given error should be considered temporary.
func IsTemporaryError(err error) bool { func IsTemporaryError(err error) bool {
tempErr, ok := err.(interface { var tempErr interface {
Temporary() bool Temporary() bool
}) }
ok := errors.As(err, &tempErr)
return ok && tempErr.Temporary() || isPacketTooBig(err) return ok && tempErr.Temporary() || isPacketTooBig(err)
} }
// IsTimeout checks whether the given error is a timeout. // IsTimeout checks whether the given error is a timeout.
func IsTimeout(err error) bool { func IsTimeout(err error) bool {
timeoutErr, ok := err.(interface { var timeoutErr interface {
Timeout() bool Timeout() bool
}) }
ok := errors.As(err, &timeoutErr)
return ok && timeoutErr.Timeout() return ok && timeoutErr.Timeout()
} }

View file

@ -17,6 +17,7 @@
package netutil package netutil
import ( import (
"errors"
"net" "net"
"testing" "testing"
"time" "time"
@ -52,7 +53,8 @@ func TestIsPacketTooBig(t *testing.T) {
listener.SetDeadline(time.Now().Add(1 * time.Second)) listener.SetDeadline(time.Now().Add(1 * time.Second))
n, _, err := listener.ReadFrom(buf) n, _, err := listener.ReadFrom(buf)
if err != nil { if err != nil {
if nerr, ok := err.(net.Error); ok && nerr.Timeout() { var nerr net.Error
if ok := errors.As(err, &nerr); ok && nerr.Timeout() {
continue continue
} }
if !isPacketTooBig(err) { if !isPacketTooBig(err) {

View file

@ -17,6 +17,7 @@
package netutil package netutil
import ( import (
"errors"
"fmt" "fmt"
"math/rand" "math/rand"
"net" "net"
@ -178,7 +179,7 @@ func TestCheckRelayIP(t *testing.T) {
for _, test := range tests { for _, test := range tests {
err := CheckRelayIP(parseIP(test.sender), parseIP(test.addr)) err := CheckRelayIP(parseIP(test.sender), parseIP(test.addr))
if err != test.want { if !errors.Is(err, test.want) {
t.Errorf("%s from %s: got %q, want %q", test.addr, test.sender, err, test.want) t.Errorf("%s from %s: got %q, want %q", test.addr, test.sender, err, test.want)
} }
} }

View file

@ -20,6 +20,7 @@
package netutil package netutil
import ( import (
"errors"
"net" "net"
"os" "os"
"syscall" "syscall"
@ -31,8 +32,10 @@ const _WSAEMSGSIZE = syscall.Errno(10040)
// fit the receive buffer. On Windows, WSARecvFrom returns // fit the receive buffer. On Windows, WSARecvFrom returns
// code WSAEMSGSIZE and no data if this happens. // code WSAEMSGSIZE and no data if this happens.
func isPacketTooBig(err error) bool { func isPacketTooBig(err error) bool {
if opErr, ok := err.(*net.OpError); ok { opErr := new(net.OpError)
if scErr, ok := opErr.Err.(*os.SyscallError); ok { if ok := errors.As(err, &opErr); ok {
scErr := new(os.SyscallError)
if ok := errors.As(opErr.Err, &scErr); ok {
return scErr.Err == _WSAEMSGSIZE return scErr.Err == _WSAEMSGSIZE
} }
return opErr.Err == _WSAEMSGSIZE return opErr.Err == _WSAEMSGSIZE

View file

@ -272,7 +272,8 @@ loop:
} }
writeStart <- struct{}{} writeStart <- struct{}{}
case err = <-readErr: case err = <-readErr:
if r, ok := err.(DiscReason); ok { var r DiscReason
if ok := errors.As(err, &r); ok {
remoteRequested = true remoteRequested = true
reason = r reason = r
} else { } else {

View file

@ -103,13 +103,15 @@ func (d DiscReason) Error() string {
} }
func discReasonForError(err error) DiscReason { func discReasonForError(err error) DiscReason {
if reason, ok := err.(DiscReason); ok { var reason DiscReason
if ok := errors.As(err, &reason); ok {
return reason return reason
} }
if errors.Is(err, errProtocolReturned) { if errors.Is(err, errProtocolReturned) {
return DiscQuitting return DiscQuitting
} }
peerError, ok := err.(*peerError) peerError := new(peerError)
ok := errors.As(err, &peerError)
if ok { if ok {
switch peerError.code { switch peerError.code {
case errInvalidMsgCode, errInvalidMsg: case errInvalidMsgCode, errInvalidMsg:

View file

@ -138,7 +138,7 @@ func TestPeerProtoReadMsg(t *testing.T) {
select { select {
case err := <-errc: case err := <-errc:
if err != errProtocolReturned { if !errors.Is(err, errProtocolReturned) {
t.Errorf("peer returned error: %v", err) t.Errorf("peer returned error: %v", err)
} }
case <-time.After(2 * time.Second): case <-time.After(2 * time.Second):
@ -190,7 +190,7 @@ func TestPeerDisconnect(t *testing.T) {
} }
select { select {
case reason := <-disc: case reason := <-disc:
if reason != DiscQuitting { if !errors.Is(reason, DiscQuitting) {
t.Errorf("run returned wrong reason: got %v, want %v", reason, DiscQuitting) t.Errorf("run returned wrong reason: got %v, want %v", reason, DiscQuitting)
} }
case <-time.After(500 * time.Millisecond): case <-time.After(500 * time.Millisecond):

View file

@ -280,7 +280,7 @@ func TestServerAtCap(t *testing.T) {
// Try inserting a non-trusted connection. // Try inserting a non-trusted connection.
anotherID := randomID() anotherID := randomID()
c := newconn(anotherID) c := newconn(anotherID)
if err := srv.checkpoint(c, srv.checkpointPostHandshake); err != DiscTooManyPeers { if err := srv.checkpoint(c, srv.checkpointPostHandshake); !errors.Is(err, DiscTooManyPeers) {
t.Error("wrong error for insert:", err) t.Error("wrong error for insert:", err)
} }
// Try inserting a trusted connection. // Try inserting a trusted connection.
@ -295,7 +295,7 @@ func TestServerAtCap(t *testing.T) {
// Remove from trusted set and try again // Remove from trusted set and try again
srv.RemoveTrustedPeer(newNode(trustedID, "")) srv.RemoveTrustedPeer(newNode(trustedID, ""))
c = newconn(trustedID) c = newconn(trustedID)
if err := srv.checkpoint(c, srv.checkpointPostHandshake); err != DiscTooManyPeers { if err := srv.checkpoint(c, srv.checkpointPostHandshake); !errors.Is(err, DiscTooManyPeers) {
t.Error("wrong error for insert:", err) t.Error("wrong error for insert:", err)
} }
@ -345,7 +345,7 @@ func TestServerPeerLimits(t *testing.T) {
dialDest := clientnode dialDest := clientnode
conn, _ := net.Pipe() conn, _ := net.Pipe()
srv.SetupConn(conn, flags, dialDest) srv.SetupConn(conn, flags, dialDest)
if tp.closeErr != DiscTooManyPeers { if !errors.Is(tp.closeErr, DiscTooManyPeers) {
t.Errorf("unexpected close error: %q", tp.closeErr) t.Errorf("unexpected close error: %q", tp.closeErr)
} }
conn.Close() conn.Close()
@ -355,11 +355,11 @@ func TestServerPeerLimits(t *testing.T) {
// Check that server allows a trusted peer despite being full. // Check that server allows a trusted peer despite being full.
conn, _ = net.Pipe() conn, _ = net.Pipe()
srv.SetupConn(conn, flags, dialDest) srv.SetupConn(conn, flags, dialDest)
if tp.closeErr == DiscTooManyPeers { if errors.Is(tp.closeErr, DiscTooManyPeers) {
t.Errorf("failed to bypass MaxPeers with trusted node: %q", tp.closeErr) t.Errorf("failed to bypass MaxPeers with trusted node: %q", tp.closeErr)
} }
if tp.closeErr != DiscUselessPeer { if !errors.Is(tp.closeErr, DiscUselessPeer) {
t.Errorf("unexpected close error: %q", tp.closeErr) t.Errorf("unexpected close error: %q", tp.closeErr)
} }
conn.Close() conn.Close()
@ -369,7 +369,7 @@ func TestServerPeerLimits(t *testing.T) {
// Check that server is full again. // Check that server is full again.
conn, _ = net.Pipe() conn, _ = net.Pipe()
srv.SetupConn(conn, flags, dialDest) srv.SetupConn(conn, flags, dialDest)
if tp.closeErr != DiscTooManyPeers { if !errors.Is(tp.closeErr, DiscTooManyPeers) {
t.Errorf("unexpected close error: %q", tp.closeErr) t.Errorf("unexpected close error: %q", tp.closeErr)
} }
conn.Close() conn.Close()
@ -564,7 +564,7 @@ func TestServerInboundThrottle(t *testing.T) {
go func() { go func() {
conn.SetDeadline(time.Now().Add(timeout)) conn.SetDeadline(time.Now().Add(timeout))
buf := make([]byte, 10) buf := make([]byte, 10)
if n, err := conn.Read(buf); err != io.EOF || n != 0 { if n, err := conn.Read(buf); !errors.Is(err, io.EOF) || n != 0 {
t.Errorf("expected io.EOF and n == 0, got error %q and n == %d", err, n) t.Errorf("expected io.EOF and n == 0, got error %q and n == %d", err, n)
} }
connClosed <- struct{}{} connClosed <- struct{}{}

View file

@ -113,7 +113,8 @@ func (t *rlpxTransport) close(err error) {
// Tell the remote end why we're disconnecting if possible. // Tell the remote end why we're disconnecting if possible.
// We only bother doing this if the underlying connection supports // We only bother doing this if the underlying connection supports
// setting a timeout tough. // setting a timeout tough.
if reason, ok := err.(DiscReason); ok && reason != DiscNetworkError { var reason DiscReason
if ok := errors.As(err, &reason); ok && reason != DiscNetworkError {
// We do not use the WriteMsg func since we want a custom deadline // We do not use the WriteMsg func since we want a custom deadline
deadline := time.Now().Add(discWriteTimeout) deadline := time.Now().Add(discWriteTimeout)
if err := t.conn.SetWriteDeadline(deadline); err == nil { if err := t.conn.SetWriteDeadline(deadline); err == nil {

View file

@ -123,25 +123,26 @@ func (err *decodeError) Error() string {
} }
func wrapStreamError(err error, typ reflect.Type) error { func wrapStreamError(err error, typ reflect.Type) error {
switch err { switch {
case ErrCanonInt: case errors.Is(err, ErrCanonInt):
return &decodeError{msg: "non-canonical integer (leading zero bytes)", typ: typ} return &decodeError{msg: "non-canonical integer (leading zero bytes)", typ: typ}
case ErrCanonSize: case errors.Is(err, ErrCanonSize):
return &decodeError{msg: "non-canonical size information", typ: typ} return &decodeError{msg: "non-canonical size information", typ: typ}
case ErrExpectedList: case errors.Is(err, ErrExpectedList):
return &decodeError{msg: "expected input list", typ: typ} return &decodeError{msg: "expected input list", typ: typ}
case ErrExpectedString: case errors.Is(err, ErrExpectedString):
return &decodeError{msg: "expected input string or byte", typ: typ} return &decodeError{msg: "expected input string or byte", typ: typ}
case errUintOverflow: case errors.Is(err, errUintOverflow):
return &decodeError{msg: "input string too long", typ: typ} return &decodeError{msg: "input string too long", typ: typ}
case errNotAtEOL: case errors.Is(err, errNotAtEOL):
return &decodeError{msg: "input list has too many elements", typ: typ} return &decodeError{msg: "input list has too many elements", typ: typ}
} }
return err return err
} }
func addErrorContext(err error, ctx string) error { func addErrorContext(err error, ctx string) error {
if decErr, ok := err.(*decodeError); ok { decErr := new(decodeError)
if ok := errors.As(err, &decErr); ok {
decErr.ctx = append(decErr.ctx, ctx) decErr.ctx = append(decErr.ctx, ctx)
} }
return err return err
@ -326,7 +327,7 @@ func decodeSliceElems(s *Stream, val reflect.Value, elemdec decoder) error {
val.SetLen(i + 1) val.SetLen(i + 1)
} }
// decode into element // decode into element
if err := elemdec(s, val.Index(i)); err == EOL { if err := elemdec(s, val.Index(i)); errors.Is(err, EOL) {
break break
} else if err != nil { } else if err != nil {
return addErrorContext(err, fmt.Sprint("[", i, "]")) return addErrorContext(err, fmt.Sprint("[", i, "]"))
@ -345,7 +346,7 @@ func decodeListArray(s *Stream, val reflect.Value, elemdec decoder) error {
vlen := val.Len() vlen := val.Len()
i := 0 i := 0
for ; i < vlen; i++ { 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 break
} else if err != nil { } else if err != nil {
return addErrorContext(err, fmt.Sprint("[", i, "]")) return addErrorContext(err, fmt.Sprint("[", i, "]"))
@ -417,7 +418,7 @@ func makeStructDecoder(typ reflect.Type) (decoder, error) {
} }
for i, f := range fields { for i, f := range fields {
err := f.info.decoder(s, val.Field(f.index)) err := f.info.decoder(s, val.Field(f.index))
if err == EOL { if errors.Is(err, EOL) {
if f.optional { if f.optional {
// The field is optional, so reaching the end of the list before // The field is optional, so reaching the end of the list before
// reaching the last field is acceptable. All remaining undecoded // 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)) v, err := s.readUint(byte(size))
switch { switch {
case err == ErrCanonSize: case errors.Is(err, ErrCanonSize):
// Adjust error because we're not reading a size right now. // Adjust error because we're not reading a size right now.
return 0, ErrCanonInt return 0, ErrCanonInt
case err != nil: case err != nil:
@ -948,7 +949,8 @@ func (s *Stream) Decode(val interface{}) error {
} }
err = decoder(s, rval.Elem()) err = decoder(s, rval.Elem())
if decErr, ok := err.(*decodeError); ok && len(decErr.ctx) > 0 { decErr := new(decodeError)
if ok := errors.As(err, &decErr); ok && len(decErr.ctx) > 0 {
// Add decode target type to error so context has more meaning. // Add decode target type to error so context has more meaning.
decErr.ctx = append(decErr.ctx, fmt.Sprint("(", rtyp.Elem(), ")")) decErr.ctx = append(decErr.ctx, fmt.Sprint("(", rtyp.Elem(), ")"))
} }
@ -1041,10 +1043,10 @@ func (s *Stream) readKind() (kind Kind, size uint64, err error) {
if len(s.stack) == 0 { if len(s.stack) == 0 {
// At toplevel, Adjust the error to actual EOF. io.EOF is // At toplevel, Adjust the error to actual EOF. io.EOF is
// used by callers to determine when to stop decoding. // used by callers to determine when to stop decoding.
switch err { switch {
case io.ErrUnexpectedEOF: case errors.Is(err, io.ErrUnexpectedEOF):
err = io.EOF err = io.EOF
case ErrValueTooLarge: case errors.Is(err, ErrValueTooLarge):
err = io.EOF err = io.EOF
} }
} }
@ -1129,7 +1131,7 @@ func (s *Stream) readFull(buf []byte) (err error) {
nn, err = s.r.Read(buf[n:]) nn, err = s.r.Read(buf[n:])
n += nn n += nn
} }
if err == io.EOF { if errors.Is(err, io.EOF) {
if n < len(buf) { if n < len(buf) {
err = io.ErrUnexpectedEOF err = io.ErrUnexpectedEOF
} else { } else {
@ -1147,7 +1149,7 @@ func (s *Stream) readByte() (byte, error) {
return 0, err return 0, err
} }
b, err := s.r.ReadByte() b, err := s.r.ReadByte()
if err == io.EOF { if errors.Is(err, io.EOF) {
err = io.ErrUnexpectedEOF err = io.ErrUnexpectedEOF
} }
return b, err return b, err

View file

@ -251,7 +251,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) t.Errorf("Uint error mismatch, got %v, want %v", err, EOL)
} }
if err = s.ListEnd(); err != nil { if err = s.ListEnd(); err != nil {
@ -331,16 +331,16 @@ func TestStreamReadBytes(t *testing.T) {
func TestDecodeErrors(t *testing.T) { func TestDecodeErrors(t *testing.T) {
r := bytes.NewReader(nil) 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) t.Errorf("Decode(r, nil) error mismatch, got %q, want %q", err, errDecodeIntoNil)
} }
var nilptr *struct{} 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) 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) t.Errorf("Decode(r, struct{}{}) error mismatch, got %q, want %q", err, errNoPointer)
} }
@ -349,7 +349,7 @@ func TestDecodeErrors(t *testing.T) {
t.Errorf("Decode(r, new(chan bool)) error mismatch, got %q, want %q", err, expectErr) t.Errorf("Decode(r, new(chan bool)) error mismatch, got %q, want %q", err, expectErr)
} }
if err := Decode(r, new(uint)); err != io.EOF { if err := Decode(r, new(uint)); !errors.Is(err, io.EOF) {
t.Errorf("Decode(r, new(int)) error mismatch, got %q, want %q", err, io.EOF) t.Errorf("Decode(r, new(int)) error mismatch, got %q, want %q", err, io.EOF)
} }
} }

View file

@ -476,7 +476,7 @@ func TestEncodeToReaderPiecewise(t *testing.T) {
} }
n, err := r.Read(output[start:end]) n, err := r.Read(output[start:end])
end = start + n end = start + n
if err == io.EOF { if errors.Is(err, io.EOF) {
break break
} else if err != nil { } else if err != nil {
return nil, err return nil, err

View file

@ -129,7 +129,7 @@ func TestSplitUint64(t *testing.T) {
if !bytes.Equal(rest, unhex(test.rest)) { if !bytes.Equal(rest, unhex(test.rest)) {
t.Errorf("test %d: rest mismatch: got %x, want %s (input %q)", i, rest, test.rest, test.input) t.Errorf("test %d: rest mismatch: got %x, want %s (input %q)", i, rest, test.rest, test.input)
} }
if err != test.err { if !errors.Is(err, test.err) {
t.Errorf("test %d: error mismatch: got %q, want %q", i, err, test.err) t.Errorf("test %d: error mismatch: got %q, want %q", i, err, test.err)
} }
} }
@ -213,7 +213,7 @@ func TestSplit(t *testing.T) {
if !bytes.Equal(rest, unhex(test.rest)) { if !bytes.Equal(rest, unhex(test.rest)) {
t.Errorf("test %d: rest mismatch: got %x, want %s", i, rest, test.rest) t.Errorf("test %d: rest mismatch: got %x, want %s", i, rest, test.rest)
} }
if err != test.err { if !errors.Is(err, test.err) {
t.Errorf("test %d: error mismatch: got %q, want %q", i, err, test.err) t.Errorf("test %d: error mismatch: got %q, want %q", i, err, test.err)
} }
} }
@ -251,7 +251,7 @@ func TestReadSize(t *testing.T) {
for _, test := range tests { for _, test := range tests {
size, err := readSize(unhex(test.input), test.slen) size, err := readSize(unhex(test.input), test.slen)
if err != test.err { if !errors.Is(err, test.err) {
t.Errorf("readSize(%s, %d): error mismatch: got %q, want %q", test.input, test.slen, err, test.err) t.Errorf("readSize(%s, %d): error mismatch: got %q, want %q", test.input, test.slen, err, test.err)
continue continue
} }

View file

@ -17,6 +17,7 @@
package rlp package rlp
import ( import (
"errors"
"fmt" "fmt"
"maps" "maps"
"reflect" "reflect"
@ -140,7 +141,8 @@ func structFields(typ reflect.Type) (fields []field, err error) {
// Filter/validate fields. // Filter/validate fields.
structFields, structTags, err := rlpstruct.ProcessFields(allStructFields) structFields, structTags, err := rlpstruct.ProcessFields(allStructFields)
if err != nil { if err != nil {
if tagErr, ok := err.(rlpstruct.TagError); ok { tagErr := new(rlpstruct.TagError)
if ok := errors.As(err, &tagErr); ok {
tagErr.StructType = typ.String() tagErr.StructType = typ.String()
return nil, tagErr return nil, tagErr
} }

Some files were not shown because too many files have changed in this diff Show more