diff --git a/internal/debugmiddleware/debug_middleware.go b/internal/debugmiddleware/debug_middleware.go index 64e270f..9b2575a 100644 --- a/internal/debugmiddleware/debug_middleware.go +++ b/internal/debugmiddleware/debug_middleware.go @@ -17,8 +17,8 @@ type ( const redactedPlaceholder = "" -// Headers known to contain sensitive information like an API key. Note that this exclude `Authorization`, -// which is handled specially in `redactRequest` below. +// Headers known to contain sensitive information like an API key. Authorization +// headers are handled separately so their authentication scheme remains visible. var sensitiveHeaders = []string{ "api-key", "x-api-key", @@ -55,7 +55,9 @@ func (m *RequestLogger) Middleware() Middleware { return resp, err } - if respBytes, err := httputil.DumpResponse(resp, false); err == nil { + loggedResponse := *resp + loggedResponse.Header = m.redactHeaders(resp.Header) + if respBytes, err := httputil.DumpResponse(&loggedResponse, false); err == nil { m.logger.Printf("Response Content:\n%s\n", respBytes) } @@ -68,43 +70,39 @@ func (m *RequestLogger) Middleware() Middleware { // the original and that clone is returned. As a small optimization, the // original is request is returned unchanged if no redaction is necessary. func (m *RequestLogger) redactRequest(req *http.Request) (*http.Request, error) { - redactedHeaders := req.Header.Clone() + redactedHeaders := m.redactHeaders(req.Header) + if reflect.DeepEqual(req.Header, redactedHeaders) { + return req, nil + } - // Notably, the clauses below are written so they can redact multiple - // headers of the same name if necessary. - if values := redactedHeaders.Values("Authorization"); len(values) > 0 { - redactedHeaders.Del("Authorization") + redacted := req.Clone(req.Context()) + redacted.Header = redactedHeaders + return redacted, nil +} - for _, value := range values { - // In case we're using something like a bearer token (e.g. `Bearer - // `), keep the `Bearer` part for more debugging - // information. - if authKind, _, ok := strings.Cut(value, " "); ok { - redactedHeaders.Add("Authorization", authKind+" "+redactedPlaceholder) - } else { - redactedHeaders.Add("Authorization", redactedPlaceholder) +// redactHeaders returns an independent copy with sensitive values removed. +func (m *RequestLogger) redactHeaders(headers http.Header) http.Header { + redacted := headers.Clone() + for header, values := range redacted { + if strings.EqualFold(header, "Authorization") || strings.EqualFold(header, "Proxy-Authorization") { + for i, value := range values { + if authKind, _, ok := strings.Cut(value, " "); ok { + values[i] = authKind + " " + redactedPlaceholder + } else { + values[i] = redactedPlaceholder + } } - } - } - - for _, header := range m.sensitiveHeaders { - values := redactedHeaders.Values(header) - if len(values) == 0 { continue } - redactedHeaders.Del(header) - - for range values { - redactedHeaders.Add(header, redactedPlaceholder) + for _, sensitiveHeader := range m.sensitiveHeaders { + if strings.EqualFold(header, sensitiveHeader) { + for i := range values { + values[i] = redactedPlaceholder + } + break + } } } - - if reflect.DeepEqual(req.Header, redactedHeaders) { - return req, nil - } - - redacted := req.Clone(req.Context()) - redacted.Header = redactedHeaders - return redacted, nil + return redacted } diff --git a/internal/debugmiddleware/debug_middleware_test.go b/internal/debugmiddleware/debug_middleware_test.go index 11f8920..3cec690 100644 --- a/internal/debugmiddleware/debug_middleware_test.go +++ b/internal/debugmiddleware/debug_middleware_test.go @@ -200,6 +200,43 @@ func TestDebugMiddleware(t *testing.T) { require.NotContains(t, logBuf.String(), bodyContent) }) + t.Run("RedactsSensitiveResponseHeaders", func(t *testing.T) { + t.Parallel() + + middleware, logBuf := setup() + responseHeaders := http.Header{ + "sEt-CoOkIe": {"session=" + secretToken + "1", "csrf=" + secretToken + "2"}, + "aUtHoRiZaTiOn": {"Bearer " + secretToken + "3"}, + "pRoXy-AuThOrIzAtIoN": {"Basic " + secretToken + "4"}, + "X-Api-Key": {secretToken + "5"}, + "X-Request-Id": {"request-id"}, + } + originalHeaders := responseHeaders.Clone() + response := &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Header: responseHeaders, + Body: http.NoBody, + } + + req := httptest.NewRequest("GET", "https://example.com", nil) + resp, err := middleware.Middleware()(req, func(req *http.Request) (*http.Response, error) { + return response, nil + }) + require.NoError(t, err) + + logged := logBuf.String() + for _, suffix := range []string{"1", "2", "3", "4", "5"} { + require.NotContains(t, logged, secretToken+suffix) + } + require.Equal(t, 5, strings.Count(logged, redactedPlaceholder)) + require.Contains(t, logged, "Bearer "+redactedPlaceholder) + require.Contains(t, logged, "Basic "+redactedPlaceholder) + require.Contains(t, logged, "X-Request-Id: request-id") + require.Same(t, response, resp) + require.Equal(t, originalHeaders, resp.Header) + }) + t.Run("DoesNotLogOrConsumeResponseBody", func(t *testing.T) { t.Parallel()