summaryrefslogtreecommitdiff
path: root/assert/errequal.go
diff options
context:
space:
mode:
authorStefan Majewsky <majewsky@gmx.net>2026-06-27 23:09:47 +0200
committerStefan Majewsky <majewsky@gmx.net>2026-06-27 23:09:47 +0200
commit62132b9680cbd8cbe6e50a346cc71922812dac00 (patch)
tree6e06a1beee0f088f3834a1db6651b004e00a67e6 /assert/errequal.go
parent8e9889d3e5b2874854832d1d882d124132546f49 (diff)
downloadgo-gg-62132b9680cbd8cbe6e50a346cc71922812dac00.tar.gz
add func assert.ErrsEqual
Diffstat (limited to 'assert/errequal.go')
-rw-r--r--assert/errequal.go77
1 files changed, 59 insertions, 18 deletions
diff --git a/assert/errequal.go b/assert/errequal.go
index 86cd719..d7bd8ba 100644
--- a/assert/errequal.go
+++ b/assert/errequal.go
@@ -17,6 +17,46 @@ import (
// - If expected is of type [*regexp.Regexp], the actual error must have a message matching that regexp.
// - If expected is of any other type, ErrEqual will panic.
func ErrEqual(t TestingTB, actual error, expected any) bool {
+ t.Helper()
+ err := errEqual(t, actual, expected)
+ if err == nil {
+ return true
+ } else {
+ t.Error(err)
+ return false
+ }
+}
+
+// ErrsEqual checks if a list of actual errors matches the expectation.
+// Both lists must be of equal length, and each individual entry must match according to [ErrEqual].
+func ErrsEqual[A ~[]error, E ~[]V, V any](t TestingTB, actual A, expected E) bool {
+ ok := true
+ for idx := range max(len(actual), len(expected)) {
+ var err error
+ switch {
+ case idx >= len(actual):
+ err = errEqual(t, missingError{}, expected[idx])
+ case idx >= len(expected):
+ err = errEqual(t, actual[idx], missingError{})
+ default:
+ err = errEqual(t, actual[idx], expected[idx])
+ }
+ if err != nil {
+ t.Errorf("in actual[%d]: %s", idx, err)
+ ok = false
+ }
+ }
+ return ok
+}
+
+// missingError is used by ErrsEqual() to mark a missing error on one side of the match.
+type missingError struct{}
+
+func (missingError) Error() string {
+ return "<missing>"
+}
+
+func errEqual(t TestingTB, actual error, expected any) error {
// coerce all types that implement `error` into the interface type `error`,
// and also coerce all nil values of concrete `error` types into untyped nil
if expectedErr, ok := expected.(error); ok {
@@ -33,42 +73,43 @@ func ErrEqual(t TestingTB, actual error, expected any) bool {
switch expected := expected.(type) {
case nil:
if actual == nil {
- return true
+ return nil
} else {
- t.Errorf("expected no error, but got %q", actual.Error())
- return false
+ return fmt.Errorf("expected no error, but got %s", formatErrorMessage(actual))
}
case error:
if actual == nil {
- t.Errorf("expected %q, but got no error", expected.Error())
- return false
+ return fmt.Errorf("expected %s, but got no error", formatErrorMessage(expected))
} else if errors.Is(actual, expected) {
- return true
+ return nil
} else {
- t.Errorf("expected %q, but got %q", expected.Error(), actual.Error())
- return false
+ return fmt.Errorf("expected %s, but got %q", formatErrorMessage(expected), actual.Error())
}
case string:
if actual == nil {
- t.Errorf("expected %q, but got no error", expected)
- return false
+ return fmt.Errorf("expected %q, but got no error", expected)
} else if actual.Error() == expected {
- return true
+ return nil
} else {
- t.Errorf("expected %q, but got %q", expected, actual.Error())
- return false
+ return fmt.Errorf("expected %q, but got %s", expected, formatErrorMessage(actual))
}
case *regexp.Regexp:
if actual == nil {
- t.Errorf("expected an error matching /%s/, but got no error", expected.String())
- return false
+ return fmt.Errorf("expected an error matching /%s/, but got no error", expected.String())
} else if expected.MatchString(actual.Error()) {
- return true
+ return nil
} else {
- t.Errorf("expected an error matching /%s/, but got %q", expected.String(), actual.Error())
- return false
+ return fmt.Errorf("expected an error matching /%s/, but got %s", expected.String(), formatErrorMessage(actual))
}
default:
panic(fmt.Sprintf("cannot handle `expected` of type %T", expected))
}
}
+
+func formatErrorMessage(err error) string {
+ if _, ok := err.(missingError); ok { //nolint:errorlint // this error is never wrapped, so errors.Is() is unnecessary
+ return err.Error()
+ } else {
+ return fmt.Sprintf("%q", err.Error())
+ }
+}