fix: defer WebSocket logs until handshake delivery
This commit is contained in:
@@ -209,7 +209,12 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|||||||
tracked := h.trackConn(conn)
|
tracked := h.trackConn(conn)
|
||||||
if upgrade := st.upgrade.Load(); upgrade != nil {
|
if upgrade := st.upgrade.Load(); upgrade != nil {
|
||||||
cw.SetHijackedResponse(upgrade.status, upgrade.header)
|
cw.SetHijackedResponse(upgrade.status, upgrade.header)
|
||||||
|
return &headerDeliveryConn{
|
||||||
|
Conn: tracked,
|
||||||
|
onDelivered: func() {
|
||||||
submitOnce.Do(func() { submit(time.Now()) })
|
submitOnce.Do(func() { submit(time.Now()) })
|
||||||
|
},
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return tracked
|
return tracked
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -97,6 +97,7 @@ func (c *captureWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
|||||||
c.hijacked = true
|
c.hijacked = true
|
||||||
if c.onHijack != nil {
|
if c.onHijack != nil {
|
||||||
conn = c.onHijack(conn)
|
conn = c.onHijack(conn)
|
||||||
|
rw = bufio.NewReadWriter(rw.Reader, bufio.NewWriter(conn))
|
||||||
}
|
}
|
||||||
return conn, rw, nil
|
return conn, rw, nil
|
||||||
}
|
}
|
||||||
@@ -131,3 +132,26 @@ func (c *trackedConn) Close() error {
|
|||||||
c.closeOnce.Do(c.onClose)
|
c.closeOnce.Do(c.onClose)
|
||||||
return err
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ package proxy_test
|
|||||||
import (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -25,6 +27,24 @@ type chanSubmitter chan *logger.LogEntry
|
|||||||
|
|
||||||
func (c chanSubmitter) Submit(e *logger.LogEntry) { c <- e }
|
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 升级:
|
// fakeUpstream 模拟一个最简 WebSocket 升级:
|
||||||
// 收到 GET + Upgrade: websocket 后回 101,然后做字节回声直到对端关闭。
|
// 收到 GET + Upgrade: websocket 后回 101,然后做字节回声直到对端关闭。
|
||||||
func fakeUpstream(t *testing.T) *httptest.Server {
|
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 {
|
if entry.StatusCode != http.StatusSwitchingProtocols || len(entry.ResponseBody) != 0 {
|
||||||
t.Fatalf("handshake entry status=%d body=%q", entry.StatusCode, entry.ResponseBody)
|
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):
|
case <-time.After(time.Second):
|
||||||
t.Fatal("websocket handshake was not captured")
|
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) {
|
func TestWebSocketNaturalCloseUnregistersConnection(t *testing.T) {
|
||||||
upstream := fakeUpstream(t)
|
upstream := fakeUpstream(t)
|
||||||
defer upstream.Close()
|
defer upstream.Close()
|
||||||
|
|||||||
Reference in New Issue
Block a user