diff --git a/proxy/proxy.go b/proxy/proxy.go index deea9d3..312b63d 100644 --- a/proxy/proxy.go +++ b/proxy/proxy.go @@ -209,7 +209,12 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { tracked := h.trackConn(conn) if upgrade := st.upgrade.Load(); upgrade != nil { cw.SetHijackedResponse(upgrade.status, upgrade.header) - submitOnce.Do(func() { submit(time.Now()) }) + return &headerDeliveryConn{ + Conn: tracked, + onDelivered: func() { + submitOnce.Do(func() { submit(time.Now()) }) + }, + } } return tracked }) diff --git a/proxy/writer.go b/proxy/writer.go index 4db1dd8..b7388d5 100644 --- a/proxy/writer.go +++ b/proxy/writer.go @@ -97,6 +97,7 @@ func (c *captureWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { c.hijacked = true if c.onHijack != nil { conn = c.onHijack(conn) + rw = bufio.NewReadWriter(rw.Reader, bufio.NewWriter(conn)) } return conn, rw, nil } @@ -131,3 +132,26 @@ func (c *trackedConn) Close() error { c.closeOnce.Do(c.onClose) return err } + +type headerDeliveryConn struct { + net.Conn + tail []byte + onDelivered func() + delivered bool +} + +func (c *headerDeliveryConn) Write(p []byte) (int, error) { + n, err := c.Conn.Write(p) + if c.delivered || err != nil || n == 0 { + return n, err + } + c.tail = append(c.tail, p[:n]...) + if bytes.Contains(c.tail, []byte("\r\n\r\n")) { + c.delivered = true + c.tail = nil + c.onDelivered() + } else if len(c.tail) > 3 { + c.tail = append(c.tail[:0], c.tail[len(c.tail)-3:]...) + } + return n, err +} diff --git a/tests/proxy/websocket_test.go b/tests/proxy/websocket_test.go index e780e03..caf42a1 100644 --- a/tests/proxy/websocket_test.go +++ b/tests/proxy/websocket_test.go @@ -3,6 +3,8 @@ package proxy_test import ( "bufio" "context" + "encoding/json" + "errors" "io" "net" "net/http" @@ -25,6 +27,24 @@ type chanSubmitter chan *logger.LogEntry func (c chanSubmitter) Submit(e *logger.LogEntry) { c <- e } +type hijackResponseWriter struct { + header http.Header + conn net.Conn +} + +func (w *hijackResponseWriter) Header() http.Header { return w.header } +func (*hijackResponseWriter) WriteHeader(int) {} +func (*hijackResponseWriter) Write(p []byte) (int, error) { + return len(p), nil +} +func (w *hijackResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + return w.conn, bufio.NewReadWriter(bufio.NewReader(w.conn), bufio.NewWriter(w.conn)), nil +} + +type failingWriteConn struct{ net.Conn } + +func (failingWriteConn) Write([]byte) (int, error) { return 0, errors.New("handshake write failed") } + // fakeUpstream 模拟一个最简 WebSocket 升级: // 收到 GET + Upgrade: websocket 后回 101,然后做字节回声直到对端关闭。 func fakeUpstream(t *testing.T) *httptest.Server { @@ -99,6 +119,13 @@ func TestWebSocketHandshakeCapturedAndShutdownClosesConnection(t *testing.T) { if entry.StatusCode != http.StatusSwitchingProtocols || len(entry.ResponseBody) != 0 { t.Fatalf("handshake entry status=%d body=%q", entry.StatusCode, entry.ResponseBody) } + var headers http.Header + if err := json.Unmarshal(entry.ResponseHeaders, &headers); err != nil { + t.Fatalf("decode response headers: %v", err) + } + if !strings.EqualFold(headers.Get("Upgrade"), "websocket") || headers.Get("X-Request-Id") == "" { + t.Fatalf("handshake metadata headers=%v", headers) + } case <-time.After(time.Second): t.Fatal("websocket handshake was not captured") } @@ -134,6 +161,32 @@ func TestWebSocketHandshakeCapturedAndShutdownClosesConnection(t *testing.T) { } } +func TestWebSocketHandshakeWriteFailureIsNotCaptured(t *testing.T) { + upstream := fakeUpstream(t) + defer upstream.Close() + u, _ := url.Parse(upstream.URL) + filter, err := config.NewFilter(config.FilterDisabled, nil) + if err != nil { + t.Fatal(err) + } + entries := make(chanSubmitter, 1) + h := proxy.New(u, filter, entries, 1024) + downstream, peer := net.Pipe() + defer peer.Close() + w := &hijackResponseWriter{header: make(http.Header), conn: failingWriteConn{downstream}} + req := httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/realtime", nil) + req.Header.Set("Upgrade", "websocket") + req.Header.Set("Connection", "Upgrade") + + h.ServeHTTP(w, req) + + select { + case entry := <-entries: + t.Fatalf("failed WebSocket handshake was captured: %+v", entry) + case <-time.After(50 * time.Millisecond): + } +} + func TestWebSocketNaturalCloseUnregistersConnection(t *testing.T) { upstream := fakeUpstream(t) defer upstream.Close()