mirror of
https://github.com/pgsty/minio.git
synced 2026-08-08 23:33:30 +03:00
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 <noreply@openai.com> Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
+18
-6
@@ -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 {
|
||||
|
||||
+129
-14
@@ -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 = "<sentinel-error/>"
|
||||
)
|
||||
|
||||
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")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user