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
20 changes: 17 additions & 3 deletions src/net/http/httputil/reverseproxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -837,23 +837,33 @@ func (p *ReverseProxy) handleUpgradeResponse(rw http.ResponseWriter, req *http.R
reqUpType := upgradeType(req.Header)
resUpType := upgradeType(res.Header)
if !ascii.IsPrint(resUpType) { // We know reqUpType is ASCII, it's checked by the caller.
if res.Body != nil {
res.Body.Close()
}
p.getErrorHandler()(rw, req, fmt.Errorf("backend tried to switch to invalid protocol %q", resUpType))
return
}
if !ascii.EqualFold(reqUpType, resUpType) {
if res.Body != nil {
res.Body.Close()
}
p.getErrorHandler()(rw, req, fmt.Errorf("backend tried to switch protocol %q when %q was requested", resUpType, reqUpType))
return
}

backConn, ok := res.Body.(io.ReadWriteCloser)
if !ok {
if res.Body != nil {
res.Body.Close()
}
p.getErrorHandler()(rw, req, fmt.Errorf("internal error: 101 switching protocols response with non-writable body"))
return
}

rc := http.NewResponseController(rw)
conn, brw, hijackErr := rc.Hijack()
if errors.Is(hijackErr, http.ErrNotSupported) {
backConn.Close()
p.getErrorHandler()(rw, req, fmt.Errorf("can't switch protocols using non-Hijacker ResponseWriter type %T", rw))
return
}
Expand All @@ -877,6 +887,9 @@ func (p *ReverseProxy) handleUpgradeResponse(rw http.ResponseWriter, req *http.R
defer conn.Close()

copyHeader(rw.Header(), res.Header)
removeHopByHopHeaders(rw.Header())
rw.Header().Set("Connection", "Upgrade")
rw.Header().Set("Upgrade", resUpType)

res.Header = rw.Header()
res.Body = nil // so res.Write only writes the headers; we have res.Body in backConn above
Expand All @@ -888,8 +901,8 @@ func (p *ReverseProxy) handleUpgradeResponse(rw http.ResponseWriter, req *http.R
p.getErrorHandler()(rw, req, fmt.Errorf("response flush: %v", err))
return
}
errc := make(chan error, 1)
spc := switchProtocolCopier{user: conn, backend: backConn}
errc := make(chan error, 2)
spc := switchProtocolCopier{user: conn, userReader: brw.Reader, backend: backConn}
go spc.copyToBackend(errc)
go spc.copyFromBackend(errc)

Expand All @@ -907,6 +920,7 @@ var errCopyDone = errors.New("hijacked connection copy complete")
// forth have nice names in stacks.
type switchProtocolCopier struct {
user, backend io.ReadWriter
userReader io.Reader
}

func (c switchProtocolCopier) copyFromBackend(errc chan<- error) {
Expand All @@ -925,7 +939,7 @@ func (c switchProtocolCopier) copyFromBackend(errc chan<- error) {
}

func (c switchProtocolCopier) copyToBackend(errc chan<- error) {
if _, err := io.Copy(c.backend, c.user); err != nil {
if _, err := io.Copy(c.backend, c.userReader); err != nil {
errc <- err
return
}
Expand Down
63 changes: 63 additions & 0 deletions src/net/http/httputil/reverseproxy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2258,3 +2258,66 @@ func (rc *testReadWriteCloser) Close() error {
}
return nil
}


func TestReverseProxyWebSocketDataLossAndHopByHop(t *testing.T) {
backendServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
c, _, err := w.(http.Hijacker).Hijack()
if err != nil {
t.Fatal(err)
}
defer c.Close()

// Send 101 with a hop-by-hop header that should be removed by proxy
io.WriteString(c, "HTTP/1.1 101 Switching Protocols\r\nConnection: upgrade\r\nUpgrade: WebSocket\r\nKeep-Alive: timeout=5\r\n\r\n")

bs := bufio.NewScanner(c)
if !bs.Scan() {
t.Errorf("backend failed to read: %v", bs.Err())
return
}
fmt.Fprintf(c, "echo: %s\n", bs.Text())
}))
defer backendServer.Close()

backURL, _ := url.Parse(backendServer.URL)
rproxy := NewSingleHostReverseProxy(backURL)

frontendProxy := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, req *http.Request) {
rproxy.ServeHTTP(rw, req)
}))
defer frontendProxy.Close()

u, _ := url.Parse(frontendProxy.URL)
conn, err := net.Dial("tcp", u.Host)
if err != nil {
t.Fatal(err)
}
defer conn.Close()

// Write upgrade request AND immediate WebSocket payload in one batch to test bufio buffering
rawReq := fmt.Sprintf("GET / HTTP/1.1\r\nHost: %s\r\nConnection: Upgrade\r\nUpgrade: websocket\r\n\r\nearly client payload\n", u.Host)
if _, err := conn.Write([]byte(rawReq)); err != nil {
t.Fatal(err)
}

br := bufio.NewReader(conn)
resp, err := http.ReadResponse(br, &http.Request{Method: "GET"})
if err != nil {
t.Fatal(err)
}
if resp.StatusCode != 101 {
t.Fatalf("expected status 101, got %d", resp.StatusCode)
}
if resp.Header.Get("Keep-Alive") != "" {
t.Errorf("expected Keep-Alive to be stripped, got %q", resp.Header.Get("Keep-Alive"))
}

line, err := br.ReadString('\n')
if err != nil {
t.Fatalf("failed reading echo: %v", err)
}
if want := "echo: early client payload\n"; line != want {
t.Errorf("got %q, want %q", line, want)
}
}