add standard error write tests

This commit is contained in:
Aneurin Barker Snook
2023-10-09 20:58:16 +01:00
parent a7a430a188
commit 7655bdf73d
2 changed files with 114 additions and 52 deletions
+8 -6
View File
@@ -99,18 +99,20 @@ func (e Error) WithValue(name string, value any) Error {
// Write writes the HTTP error to an HTTP response as plain text. // Write writes the HTTP error to an HTTP response as plain text.
// Additional data is omitted. // Additional data is omitted.
func (e Error) Write(w http.ResponseWriter) { func (e Error) Write(w http.ResponseWriter) (int, error) {
if e.StatusCode == 0 {
e.StatusCode = 200
}
w.WriteHeader(e.StatusCode) w.WriteHeader(e.StatusCode)
w.Write([]byte(e.Message)) return w.Write([]byte(e.Message))
} }
// WriteJSON writes the HTTP error to an HTTP response as JSON. // WriteJSON writes the HTTP error to an HTTP response as JSON.
func (e Error) WriteJSON(w http.ResponseWriter) error { func (e Error) WriteJSON(w http.ResponseWriter) error {
statusCode := e.StatusCode if e.StatusCode == 0 {
if statusCode == 0 { e.StatusCode = 200
statusCode = 200
} }
return WriteResponseJSON(w, statusCode, e) return WriteResponseJSON(w, e.StatusCode, e)
} }
// NewError creates a new REST API error. // NewError creates a new REST API error.
+82 -22
View File
@@ -7,49 +7,109 @@ import (
"testing" "testing"
) )
func TestErrorWriteJSON(t *testing.T) { type ErrorTestCase struct {
type TestCase struct {
Input Error Input Error
C int Code int
Output string Str string
JSON string
Err error Err error
} }
testCases := []TestCase{ var errorTestCases = []ErrorTestCase{
// Empty error // Empty error
{Input: Err, C: 200, Output: `{"message":""}`}, {
Input: Err,
Code: 200,
Str: "",
JSON: `{"message":""}`,
},
// Standard errors // Standard errors
{Input: ErrPermanentRedirect, C: 308, Output: `{"message":"Permanent Redirect"}`}, {
{Input: ErrNotFound, C: 404, Output: `{"message":"Not Found"}`}, Input: ErrPermanentRedirect,
{Input: ErrInternalServerError, C: 500, Output: `{"message":"Internal Server Error"}`}, Code: 308,
Str: "Permanent Redirect",
JSON: `{"message":"Permanent Redirect"}`,
},
{
Input: ErrNotFound,
Code: 404,
Str: "Not Found",
JSON: `{"message":"Not Found"}`,
},
{
Input: ErrInternalServerError,
Code: 500,
Str: "Internal Server Error",
JSON: `{"message":"Internal Server Error"}`,
},
// Error with changed message // Error with changed message
{Input: ErrBadRequest.WithMessage("Invalid Recipe"), C: 400, Output: `{"message":"Invalid Recipe"}`}, {
Input: ErrBadRequest.WithMessage("Invalid Recipe"),
Code: 400,
Str: "Invalid Recipe",
JSON: `{"message":"Invalid Recipe"}`,
},
// Error with data // Error with data
{ {
Input: ErrGatewayTimeout.WithData(map[string]any{"service": "RecipeDatabase"}), Input: ErrGatewayTimeout.WithData(map[string]any{"service": "RecipeDatabase"}),
C: 504, Code: 504,
Output: `{"message":"Gateway Timeout","data":{"service":"RecipeDatabase"}}`, Str: "Gateway Timeout",
JSON: `{"message":"Gateway Timeout","data":{"service":"RecipeDatabase"}}`,
}, },
// Error with value // Error with value
{ {
Input: ErrGatewayTimeout.WithValue("service", "RecipeDatabase"), Input: ErrGatewayTimeout.WithValue("service", "RecipeDatabase"),
C: 504, Code: 504,
Output: `{"message":"Gateway Timeout","data":{"service":"RecipeDatabase"}}`, Str: "Gateway Timeout",
JSON: `{"message":"Gateway Timeout","data":{"service":"RecipeDatabase"}}`,
}, },
// Error with error // Error with error
{ {
Input: ErrInternalServerError.WithError(errors.New("recipe is too delicious")), Input: ErrInternalServerError.WithError(errors.New("recipe is too delicious")),
C: 500, Code: 500,
Output: `{"message":"Internal Server Error","data":{"error":"recipe is too delicious"}}`, Str: "Internal Server Error",
JSON: `{"message":"Internal Server Error","data":{"error":"recipe is too delicious"}}`,
}, },
}
func TestErrorWrite(t *testing.T) {
for i, tc := range errorTestCases {
t.Logf("(%d) Testing %v", i, tc.Input)
rec := httptest.NewRecorder()
_, err := tc.Input.Write(rec)
if err != tc.Err {
t.Errorf("Expected error %v, got %v", tc.Err, err)
}
if err != nil {
continue
} }
for i, tc := range testCases { res := rec.Result()
if res.StatusCode != tc.Code {
t.Errorf("Expected status code %d, got %d", tc.Code, res.StatusCode)
}
body, err := io.ReadAll(res.Body)
if err != nil {
t.Errorf("Unexpected error reading response body: %v", err)
continue
}
if string(body) != tc.Str {
t.Errorf("Expected body %q, got %q", tc.Str, string(body))
}
}
}
func TestErrorWriteJSON(t *testing.T) {
for i, tc := range errorTestCases {
t.Logf("(%d) Testing %v", i, tc.Input) t.Logf("(%d) Testing %v", i, tc.Input)
rec := httptest.NewRecorder() rec := httptest.NewRecorder()
@@ -63,8 +123,8 @@ func TestErrorWriteJSON(t *testing.T) {
} }
res := rec.Result() res := rec.Result()
if res.StatusCode != tc.C { if res.StatusCode != tc.Code {
t.Errorf("Expected status code %d, got %d", tc.C, res.StatusCode) t.Errorf("Expected status code %d, got %d", tc.Code, res.StatusCode)
} }
body, err := io.ReadAll(res.Body) body, err := io.ReadAll(res.Body)
@@ -73,8 +133,8 @@ func TestErrorWriteJSON(t *testing.T) {
continue continue
} }
if string(body) != tc.Output { if string(body) != tc.JSON {
t.Errorf("Expected body %q, got %q", tc.Output, string(body)) t.Errorf("Expected body %q, got %q", tc.JSON, string(body))
} }
} }
} }