diff --git a/docs/pages/release_notes.rst b/docs/pages/release_notes.rst index 49e224974..0cb178ef8 100644 --- a/docs/pages/release_notes.rst +++ b/docs/pages/release_notes.rst @@ -2,6 +2,14 @@ Release notes ############# +**************** +Unreleased +**************** + +## Minor fixes/changes + +- #4620: HTTP client: the response body is now closed when a response exceeds the maximum size, so an oversized response no longer holds a connection until the client timeout. By @stevenvegt in https://github.com/nuts-foundation/nuts-node/pull/4621 + **************** Peanut (v6.2.14) **************** diff --git a/http/client/client.go b/http/client/client.go index a6e6e0c7c..96428f853 100644 --- a/http/client/client.go +++ b/http/client/client.go @@ -165,7 +165,9 @@ func (s *StrictHTTPClient) Do(req *http.Request) (*http.Response, error) { return nil, err } if result.Body != nil { - body, err := limitedReadAll(result.Body) + originalBody := result.Body + defer originalBody.Close() + body, err := limitedReadAll(originalBody) if err != nil { return nil, err } diff --git a/http/client/client_test.go b/http/client/client_test.go index 3ab853913..a4c5d1b36 100644 --- a/http/client/client_test.go +++ b/http/client/client_test.go @@ -21,8 +21,10 @@ package client import ( "crypto/tls" "fmt" + "io" "net/http" "net/http/httptest" + "strconv" "strings" "sync" "sync/atomic" @@ -177,6 +179,65 @@ func TestLimitedReadAll(t *testing.T) { }) } +func TestStrictHTTPClient_ClosesResponseBody(t *testing.T) { + oldStrictMode := StrictMode + StrictMode = false + t.Cleanup(func() { StrictMode = oldStrictMode }) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + size, _ := strconv.Atoi(request.URL.Query().Get("size")) + _, _ = writer.Write([]byte(strings.Repeat("a", size))) + })) + t.Cleanup(server.Close) + + t.Run("response exceeds limit", func(t *testing.T) { + transport := &closeRecordingTransport{base: SafeHttpTransport} + client := &StrictHTTPClient{client: &http.Client{Transport: transport}} + request, _ := http.NewRequest(http.MethodGet, server.URL+"?size="+strconv.Itoa(DefaultMaxHttpResponseSize+1), nil) + + _, err := client.Do(request) + + assert.EqualError(t, err, "data to read exceeds max. safety limit of 1048576 bytes") + assert.Equal(t, int32(1), transport.closed.Load()) + }) + t.Run("response within limit", func(t *testing.T) { + transport := &closeRecordingTransport{base: SafeHttpTransport} + client := &StrictHTTPClient{client: &http.Client{Transport: transport}} + request, _ := http.NewRequest(http.MethodGet, server.URL+"?size=10", nil) + + response, err := client.Do(request) + + require.NoError(t, err) + data, _ := io.ReadAll(response.Body) + assert.Len(t, data, 10) + assert.Equal(t, int32(1), transport.closed.Load()) + }) +} + +// closeRecordingTransport wraps response bodies to count how often they are closed. +type closeRecordingTransport struct { + base http.RoundTripper + closed atomic.Int32 +} + +func (c *closeRecordingTransport) RoundTrip(request *http.Request) (*http.Response, error) { + response, err := c.base.RoundTrip(request) + if err != nil { + return nil, err + } + response.Body = &closeRecordingBody{ReadCloser: response.Body, closed: &c.closed} + return response, nil +} + +type closeRecordingBody struct { + io.ReadCloser + closed *atomic.Int32 +} + +func (c *closeRecordingBody) Close() error { + c.closed.Add(1) + return c.ReadCloser.Close() +} + func TestMaxConns(t *testing.T) { oldStrictMode := StrictMode StrictMode = false