fix: defer WebSocket logs until handshake delivery

This commit is contained in:
MiMoCode
2026-07-10 18:48:56 +08:00
parent 92daa9b566
commit 6fe0a8415d
3 changed files with 83 additions and 1 deletions
+53
View File
@@ -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()