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
36 changes: 36 additions & 0 deletions internal/httputils/wrap_writer.go
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,12 @@ func NewWrapResponseWriter(w http.ResponseWriter, protoMajor int) WrapResponseWr
if fl && hj && rf {
return &httpFancyWriter{bw}
}
if fl && hj {
return &flushHijackWriter{bw}
}
if hj {
return &hijackWriter{bw}
}
}
if fl {
return &flushWriter{bw}
Expand Down Expand Up @@ -143,6 +149,36 @@ func (f *flushWriter) Flush() {

var _ http.Flusher = &flushWriter{}

// flushHijackWriter is a HTTP writer that additionally satisfies http.Flusher
// and http.Hijacker, for writers that do not implement io.ReaderFrom.
type flushHijackWriter struct {
basicWriter
}

func (f *flushHijackWriter) Flush() {
f.wroteHeader = true
f.ResponseWriter.(http.Flusher).Flush()
}

func (f *flushHijackWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
return f.ResponseWriter.(http.Hijacker).Hijack()
}

var _ http.Flusher = &flushHijackWriter{}
var _ http.Hijacker = &flushHijackWriter{}

// hijackWriter is a HTTP writer that additionally satisfies http.Hijacker, for
// writers that implement neither http.Flusher nor io.ReaderFrom.
type hijackWriter struct {
basicWriter
}

func (f *hijackWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
return f.ResponseWriter.(http.Hijacker).Hijack()
}

var _ http.Hijacker = &hijackWriter{}

// httpFancyWriter is a HTTP writer that additionally satisfies
// http.Flusher, http.Hijacker, and io.ReaderFrom. It exists for the common case
// of wrapping the http.ResponseWriter that package http gives you, in order to
Expand Down
80 changes: 80 additions & 0 deletions internal/httputils/wrap_writer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,30 @@ func (c *CustomResponseWriter) ReadFrom(r io.Reader) (n int64, err error) {
return
}

// FlusherHijacker for testing a writer that supports http.Flusher and
// http.Hijacker, but not io.ReaderFrom, like most compressing middlewares.
type FlusherHijacker struct {
*httptest.ResponseRecorder
}

func (c *FlusherHijacker) Hijack() (net.Conn, *bufio.ReadWriter, error) {
return nil, nil, errors.New("hijack not supported in tests")
}

func (c *FlusherHijacker) Flush() {
c.ResponseRecorder.Flush()
}

// HijackerOnly for testing a writer that supports http.Hijacker, but neither
// http.Flusher nor io.ReaderFrom.
type HijackerOnly struct {
http.ResponseWriter
}

func (c *HijackerOnly) Hijack() (net.Conn, *bufio.ReadWriter, error) {
return nil, nil, errors.New("hijack not supported in tests")
}

func TestHttpFancyWriterRemembersWroteHeaderWhenFlushed(t *testing.T) {
f := &httpFancyWriter{basicWriter: basicWriter{ResponseWriter: httptest.NewRecorder()}}
f.Flush()
Expand Down Expand Up @@ -95,6 +119,62 @@ func TestNewWrapResponseWriter(t *testing.T) {
}
}

func TestNewWrapResponseWriterKeepsHijacker(t *testing.T) {
// Flusher, Hijacker and io.ReaderFrom
w1 := NewWrapResponseWriter(&CustomResponseWriter{httptest.NewRecorder()}, 1)
if _, ok := w1.(*httpFancyWriter); !ok {
t.Fatalf("expected httpFancyWriter, got %T", w1)
}
if _, ok := w1.(http.Hijacker); !ok {
t.Fatalf("%T does not implement http.Hijacker", w1)
}

// Flusher and Hijacker
w2 := NewWrapResponseWriter(&FlusherHijacker{httptest.NewRecorder()}, 1)
if _, ok := w2.(*flushHijackWriter); !ok {
t.Fatalf("expected flushHijackWriter, got %T", w2)
}
if _, ok := w2.(http.Hijacker); !ok {
t.Fatalf("%T does not implement http.Hijacker", w2)
}

// Hijacker only
w3 := NewWrapResponseWriter(&HijackerOnly{httptest.NewRecorder()}, 1)
if _, ok := w3.(*hijackWriter); !ok {
t.Fatalf("expected hijackWriter, got %T", w3)
}
if _, ok := w3.(http.Hijacker); !ok {
t.Fatalf("%T does not implement http.Hijacker", w3)
}
}

func TestFlushHijackWriterRemembersWroteHeaderWhenFlushed(t *testing.T) {
f := &flushHijackWriter{basicWriter{ResponseWriter: &FlusherHijacker{httptest.NewRecorder()}}}
f.Flush()

if !f.wroteHeader {
t.Fatal("want Flush to have set wroteHeader=true")
}
}

func TestFlushHijackWriterHijack(t *testing.T) {
f := &flushHijackWriter{basicWriter{ResponseWriter: &FlusherHijacker{httptest.NewRecorder()}}}

_, _, err := f.Hijack()
if err == nil {
t.Fatal("expected error, got nil")
}
}

func TestHijackWriterHijack(t *testing.T) {
f := &hijackWriter{basicWriter{ResponseWriter: &HijackerOnly{httptest.NewRecorder()}}}

_, _, err := f.Hijack()
if err == nil {
t.Fatal("expected error, got nil")
}
}

func TestBasicWriterWriteHeader(t *testing.T) {
rec := httptest.NewRecorder()
bw := &basicWriter{ResponseWriter: rec}
Expand Down