From 62132b9680cbd8cbe6e50a346cc71922812dac00 Mon Sep 17 00:00:00 2001 From: Stefan Majewsky Date: Sat, 27 Jun 2026 23:09:47 +0200 Subject: add func assert.ErrsEqual --- assert/errequal.go | 77 +++++++++++++++++++++++++++++++++++++++++------------- 1 file changed, 59 insertions(+), 18 deletions(-) (limited to 'assert/errequal.go') 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 "" +} + +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()) + } +} -- cgit v1.3.1