Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion cfn/response.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
45 changes: 45 additions & 0 deletions cfn/response_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,12 @@ package cfn

import (
"bytes"
"errors"
"fmt"
"io"
"net/http"
"testing"
"testing/iotest"

"github.com/stretchr/testify/assert"
)
Expand Down Expand Up @@ -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,
Expand Down
Loading