From 8069a32ac80bb91409270c5b510d3756d189f14a Mon Sep 17 00:00:00 2001 From: Feng Ruohang Date: Fri, 31 Jul 2026 19:21:44 +0800 Subject: [PATCH] fix: track implicit HTTP response commits Mark trackingResponseWriter committed when Write or an effective Flush implicitly sends a 200 response. This keeps duplicate-response detection aligned with the actual writer chain while preserving no-op Flush behavior when unsupported. Add direct and gzip-streaming regression coverage for implicit headers, Flush delegation, and suppression of a second response. Co-authored-by: ChatGPT Co-authored-by: Claude --- cmd/api-response.go | 24 +++++-- cmd/api-response_test.go | 143 +++++++++++++++++++++++++++++++++++---- 2 files changed, 147 insertions(+), 20 deletions(-) diff --git a/cmd/api-response.go b/cmd/api-response.go index 2caf71354..cf25fd980 100644 --- a/cmd/api-response.go +++ b/cmd/api-response.go @@ -1026,7 +1026,7 @@ type unwrapper interface { Unwrap() http.ResponseWriter } -// headersAlreadyWritten returns true if the headers have already been written +// headersAlreadyWritten returns true if an HTTP status has already been written // to this response writer. It will unwrap the ResponseWriter if possible to try // and find a trackingResponseWriter. func headersAlreadyWritten(w http.ResponseWriter) bool { @@ -1041,14 +1041,18 @@ func headersAlreadyWritten(w http.ResponseWriter) bool { } } -// trackingResponseWriter wraps a ResponseWriter and notes when WriterHeader has -// been called. This allows high level request handlers to check if something -// has already sent the header. +// trackingResponseWriter wraps a ResponseWriter and records when an HTTP status +// has been written, explicitly or implicitly by Write or an effective Flush. +// +// Informational responses are treated as final. internal/http.ResponseRecorder +// has the same limitation, so 1xx support must be fixed in both layers. type trackingResponseWriter struct { http.ResponseWriter headerWritten bool } +var _ http.Flusher = (*trackingResponseWriter)(nil) + func (w *trackingResponseWriter) WriteHeader(statusCode int) { if !w.headerWritten { w.headerWritten = true @@ -1057,13 +1061,21 @@ func (w *trackingResponseWriter) WriteHeader(statusCode int) { } func (w *trackingResponseWriter) Write(b []byte) (int, error) { + if !w.headerWritten { + w.WriteHeader(http.StatusOK) + } return w.ResponseWriter.Write(b) } func (w *trackingResponseWriter) Flush() { - if f, ok := w.ResponseWriter.(http.Flusher); ok { - f.Flush() + f, ok := w.ResponseWriter.(http.Flusher) + if !ok { + return } + if !w.headerWritten { + w.WriteHeader(http.StatusOK) + } + f.Flush() } func (w *trackingResponseWriter) Unwrap() http.ResponseWriter { diff --git a/cmd/api-response_test.go b/cmd/api-response_test.go index c2d91e8a5..acfecd4d8 100644 --- a/cmd/api-response_test.go +++ b/cmd/api-response_test.go @@ -18,6 +18,7 @@ package cmd import ( + "compress/gzip" "io" "net/http" "net/http/httptest" @@ -128,6 +129,22 @@ func TestGetURLScheme(t *testing.T) { } } +type writeHeaderSpy struct { + http.ResponseWriter + codes []int +} + +func (r *writeHeaderSpy) WriteHeader(code int) { + r.codes = append(r.codes, code) + r.ResponseWriter.WriteHeader(code) +} + +func (r *writeHeaderSpy) Flush() { + if f, ok := r.ResponseWriter.(http.Flusher); ok { + f.Flush() + } +} + func TestTrackingResponseWriter(t *testing.T) { rw := httptest.NewRecorder() trw := &trackingResponseWriter{ResponseWriter: rw} @@ -140,8 +157,9 @@ func TestTrackingResponseWriter(t *testing.T) { if err != nil { t.Fatalf("Write unexpectedly failed: %v", err) } + xhttp.Flush(trw) - // Check that WriteHeader and Write were called on the underlying response writer + // Check that WriteHeader, Write, and Flush were called on the underlying response writer. resp := rw.Result() if resp.StatusCode != 299 { t.Fatalf("unexpected status: %v", resp.StatusCode) @@ -153,6 +171,9 @@ func TestTrackingResponseWriter(t *testing.T) { if string(body) != "hello" { t.Fatalf("response body incorrect: %v", string(body)) } + if !rw.Flushed { + t.Fatal("underlying ResponseRecorder was not flushed") + } // Check that Unwrap works if trw.Unwrap() != rw { @@ -160,26 +181,119 @@ func TestTrackingResponseWriter(t *testing.T) { } } +func TestTrackingResponseWriterWriteImplicitHeader(t *testing.T) { + testCases := []struct { + name string + body []byte + }{ + {name: "non-empty", body: []byte("hello")}, + {name: "empty", body: nil}, + } + + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + rec := httptest.NewRecorder() + rw := &writeHeaderSpy{ResponseWriter: rec} + trw := &trackingResponseWriter{ResponseWriter: rw} + + n, err := trw.Write(testCase.body) + if err != nil { + t.Fatalf("Write unexpectedly failed: %v", err) + } + if n != len(testCase.body) { + t.Fatalf("unexpected bytes written: got %d, want %d", n, len(testCase.body)) + } + if !trw.headerWritten { + t.Fatal("Write did not set headerWritten") + } + if len(rw.codes) != 1 || rw.codes[0] != http.StatusOK { + t.Fatalf("unexpected WriteHeader calls: got %v, want [%d]", rw.codes, http.StatusOK) + } + if got := rec.Body.String(); got != string(testCase.body) { + t.Fatalf("unexpected body: got %q, want %q", got, testCase.body) + } + }) + } +} + func TestTrackingResponseWriterFlush(t *testing.T) { - rw := httptest.NewRecorder() + rec := httptest.NewRecorder() + rw := &writeHeaderSpy{ResponseWriter: rec} + trw := &trackingResponseWriter{ResponseWriter: rw} + + xhttp.Flush(trw) + if !trw.headerWritten { + t.Fatal("Flush did not set headerWritten") + } + if len(rw.codes) != 1 || rw.codes[0] != http.StatusOK { + t.Fatalf("unexpected WriteHeader calls: got %v, want [%d]", rw.codes, http.StatusOK) + } + if !rec.Flushed { + t.Fatal("underlying ResponseRecorder was not flushed") + } +} + +func TestTrackingResponseWriterFlushUnsupported(t *testing.T) { + rw := struct{ http.ResponseWriter }{ResponseWriter: httptest.NewRecorder()} trw := &trackingResponseWriter{ResponseWriter: rw} trw.Flush() if trw.headerWritten { - t.Fatal("Flush() should not set headerWritten") + t.Fatal("unsupported Flush set headerWritten") } +} - // Simulate the ListenNotificationHandler flow: WriteHeader, Write, Flush - trw.WriteHeader(http.StatusOK) - _, err := trw.Write([]byte("event data")) +func TestTrackingResponseWriterGzipStreaming(t *testing.T) { + const ( + eventPayload = "event data" + sentinel = "" + ) + + rw := httptest.NewRecorder() + trw := &trackingResponseWriter{ResponseWriter: rw} + var ( + committed bool + writeErr error + ) + handler := gzipHandler(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + setEventStreamHeaders(w) + _, writeErr = w.Write([]byte(eventPayload)) + if writeErr != nil { + return + } + xhttp.Flush(w) + committed = headersAlreadyWritten(w) + writeResponse(w, http.StatusInternalServerError, []byte(sentinel), mimeXML) + })) + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.Header.Set("Accept-Encoding", "gzip") + + handler.ServeHTTP(trw, req) + + if writeErr != nil { + t.Fatalf("Write unexpectedly failed: %v", writeErr) + } + if !committed { + t.Fatal("headersAlreadyWritten returned false after Write and Flush") + } + resp := rw.Result() + if resp.StatusCode != http.StatusOK { + t.Fatalf("unexpected status: got %d, want %d", resp.StatusCode, http.StatusOK) + } + if got := resp.Header.Get("Content-Encoding"); got != "gzip" { + t.Fatalf("unexpected Content-Encoding: got %q, want %q", got, "gzip") + } + zr, err := gzip.NewReader(resp.Body) if err != nil { - t.Fatalf("Write failed: %v", err) + t.Fatalf("creating gzip reader failed: %v", err) } - - xhttp.Flush(trw) - - if !rw.Flushed { - t.Fatalf("xhttp.Flush should have flushed the underlying ResponseRecorder via trackingResponseWriter.Flush()") + defer zr.Close() + body, err := io.ReadAll(zr) + if err != nil { + t.Fatalf("reading gzip response body failed: %v", err) + } + if got := string(body); got != eventPayload { + t.Fatalf("unexpected response body: got %q, want %q (sentinel %q must be suppressed)", got, eventPayload, sentinel) } } @@ -191,7 +305,7 @@ func TestHeadersAlreadyWritten(t *testing.T) { t.Fatal("headers have not been written yet") } - trw.WriteHeader(123) + trw.WriteHeader(299) if !headersAlreadyWritten(trw) { t.Fatal("headers were written") } @@ -207,7 +321,8 @@ func TestHeadersAlreadyWrittenWrapped(t *testing.T) { t.Fatal("headers have not been written yet") } - wrap2.WriteHeader(123) + // Pin the current stack-wide 1xx limitation documented on trackingResponseWriter. + wrap2.WriteHeader(http.StatusContinue) if !headersAlreadyWritten(wrap2) { t.Fatal("headers were written") }