diff --git a/cfn/response.go b/cfn/response.go index 5b786f5f..6d7cd768 100644 --- a/cfn/response.go +++ b/cfn/response.go @@ -66,12 +66,12 @@ func (r *Response) sendWith(client httpClient) error { if err != nil { return err } + defer res.Body.Close() body, err = io.ReadAll(res.Body) if err != nil { return err } - res.Body.Close() if res.StatusCode != 200 { log.Printf("StatusCode: %d\nBody: %v\n", res.StatusCode, string(body)) diff --git a/cfn/response_test.go b/cfn/response_test.go index c833c2b8..b3fa7668 100644 --- a/cfn/response_test.go +++ b/cfn/response_test.go @@ -4,10 +4,12 @@ package cfn import ( "bytes" + "errors" "fmt" "io" "net/http" "testing" + "testing/iotest" "github.com/stretchr/testify/assert" ) @@ -44,6 +46,49 @@ type nopCloser struct { func (nopCloser) Close() error { return nil } +type trackingReadCloser struct { + io.Reader + closeCount int + closeErr error +} + +func (r *trackingReadCloser) Close() error { + r.closeCount++ + return r.closeErr +} + +func TestResponseBodyClosed(t *testing.T) { + readErr := errors.New("response body read failed") + closeErr := errors.New("response body close failed") + for _, test := range []struct { + name string + statusCode int + reader io.Reader + closeErr error + wantErr error + }{ + {"success", http.StatusOK, bytes.NewBufferString(""), nil, nil}, + {"HTTP error", http.StatusForbidden, bytes.NewBufferString("forbidden"), nil, fmt.Errorf("invalid status code. got: %d", http.StatusForbidden)}, + {"read error", http.StatusOK, iotest.ErrReader(readErr), nil, readErr}, + {"partial read error", http.StatusOK, io.MultiReader(bytes.NewBufferString("partial body"), iotest.ErrReader(readErr)), nil, readErr}, + {"close error", http.StatusOK, bytes.NewBufferString(""), closeErr, nil}, + {"read and close errors", http.StatusOK, iotest.ErrReader(readErr), closeErr, readErr}, + } { + t.Run(test.name, func(t *testing.T) { + body := &trackingReadCloser{Reader: test.reader, closeErr: test.closeErr} + client := &mockClient{ + DoFunc: func(req *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: test.statusCode, Body: body}, nil + }, + } + r := &Response{Status: StatusSuccess, url: "http://pre-signed-S3-url-for-response"} + + assert.Equal(t, test.wantErr, r.sendWith(client)) + assert.Equal(t, 1, body.closeCount) + }) + } +} + func TestRequestSentCorrectly(t *testing.T) { r := &Response{ Status: StatusSuccess,