package proxy_test import ( "bufio" "context" "encoding/json" "errors" "io" "net" "net/http" "net/http/httptest" "net/url" "strings" "sync" "testing" "time" "git.misaka.ren/M1saka/token_thief/config" "git.misaka.ren/M1saka/token_thief/logger" "git.misaka.ren/M1saka/token_thief/proxy" ) type noopSubmitter struct{} func (noopSubmitter) Submit(*logger.LogEntry) {} 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") } type closeNotifyConn struct { net.Conn closed chan struct{} once sync.Once } func (c *closeNotifyConn) Close() error { err := c.Conn.Close() c.once.Do(func() { close(c.closed) }) return err } type blockingHijackResponseWriter struct { *hijackResponseWriter hijacked chan struct{} release chan struct{} } func (w *blockingHijackResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { close(w.hijacked) <-w.release return w.hijackResponseWriter.Hijack() } // fakeUpstream 模拟一个最简 WebSocket 升级: // 收到 GET + Upgrade: websocket 后回 101,然后做字节回声直到对端关闭。 func fakeUpstream(t *testing.T) *httptest.Server { t.Helper() srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if strings.ToLower(r.Header.Get("Upgrade")) != "websocket" { http.Error(w, "expected websocket upgrade", http.StatusBadRequest) return } hj, ok := w.(http.Hijacker) if !ok { http.Error(w, "no hijack", http.StatusInternalServerError) return } conn, brw, err := hj.Hijack() if err != nil { t.Errorf("upstream hijack: %v", err) return } defer conn.Close() // 直接回 101 握手响应(简化版,不做真正 Sec-WebSocket-Accept 计算) _, _ = brw.WriteString("HTTP/1.1 101 Switching Protocols\r\n" + "Upgrade: websocket\r\n" + "Connection: Upgrade\r\n" + "\r\n") _ = brw.Flush() // echo buf := make([]byte, 1024) for { n, err := conn.Read(buf) if err != nil { return } if _, err := conn.Write(buf[:n]); err != nil { return } } })) return srv } func TestWebSocketHandshakeCapturedAndShutdownClosesConnection(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) proxySrv := httptest.NewServer(h) defer proxySrv.Close() pu, _ := url.Parse(proxySrv.URL) conn, err := net.DialTimeout("tcp", pu.Host, time.Second) if err != nil { t.Fatal(err) } defer conn.Close() _, _ = io.WriteString(conn, "GET /v1/realtime HTTP/1.1\r\nHost: "+pu.Host+"\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n") br := bufio.NewReader(conn) resp, err := http.ReadResponse(br, nil) if err != nil { t.Fatal(err) } if resp.StatusCode != http.StatusSwitchingProtocols { t.Fatalf("status=%d", resp.StatusCode) } select { case entry := <-entries: 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") } payload := []byte("frame-data-must-not-be-logged") if _, err := conn.Write(payload); err != nil { t.Fatal(err) } echo := make([]byte, len(payload)) if _, err := io.ReadFull(br, echo); err != nil { t.Fatalf("read websocket payload: %v", err) } if string(echo) != string(payload) { t.Fatalf("echo mismatch: got %q want %q", echo, payload) } select { case entry := <-entries: t.Fatalf("websocket frame produced an extra log entry: %+v", entry) case <-time.After(50 * time.Millisecond): } ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() if err := h.Shutdown(ctx); err != nil { t.Fatal(err) } _ = conn.SetReadDeadline(time.Now().Add(time.Second)) if _, err := conn.Read(make([]byte, 1)); err == nil { t.Fatal("connection remains open after Shutdown") } if err := h.Shutdown(ctx); err != nil { t.Fatalf("second Shutdown: %v", err) } } func TestShutdownClosesConnectionHijackedBeforeRegistration(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) } h := proxy.New(u, filter, noopSubmitter{}, 1024) downstream, peer := net.Pipe() defer peer.Close() closed := make(chan struct{}) conn := &closeNotifyConn{Conn: downstream, closed: closed} w := &blockingHijackResponseWriter{ hijackResponseWriter: &hijackResponseWriter{header: make(http.Header), conn: conn}, hijacked: make(chan struct{}), release: make(chan struct{}), } req := httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/realtime", nil) req.Header.Set("Upgrade", "websocket") req.Header.Set("Connection", "Upgrade") serveDone := make(chan struct{}) var releaseOnce sync.Once releaseHijack := func() { releaseOnce.Do(func() { close(w.release) }) } go func() { h.ServeHTTP(w, req) close(serveDone) }() defer func() { releaseHijack() ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() _ = h.Shutdown(ctx) <-serveDone }() select { case <-w.hijacked: case <-time.After(time.Second): t.Fatal("downstream connection was not hijacked") } ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() if err := h.Shutdown(ctx); err != nil { t.Fatal(err) } releaseHijack() select { case <-closed: case <-time.After(100 * time.Millisecond): t.Fatal("connection hijacked during Shutdown remained open") } } 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() u, _ := url.Parse(upstream.URL) filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, noopSubmitter{}, 1024) proxySrv := httptest.NewServer(h) defer proxySrv.Close() pu, _ := url.Parse(proxySrv.URL) conn, err := net.DialTimeout("tcp", pu.Host, time.Second) if err != nil { t.Fatal(err) } _, _ = io.WriteString(conn, "GET /v1/realtime HTTP/1.1\r\nHost: "+pu.Host+"\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n") resp, err := http.ReadResponse(bufio.NewReader(conn), nil) if err != nil { t.Fatal(err) } if resp.StatusCode != http.StatusSwitchingProtocols { t.Fatalf("status=%d", resp.StatusCode) } if err := conn.Close(); err != nil { t.Fatal(err) } deadline := time.Now().Add(time.Second) for { ctx, cancel := context.WithCancel(context.Background()) cancel() err := h.Shutdown(ctx) if err == nil { break } if time.Now().After(deadline) { t.Fatalf("connection was not unregistered after natural close: %v", err) } time.Sleep(10 * time.Millisecond) } } func TestNonWebSocketUpgradeIsManagedButNotLogged(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { hj := w.(http.Hijacker) conn, brw, err := hj.Hijack() if err != nil { t.Errorf("upstream hijack: %v", err) return } defer conn.Close() _, _ = brw.WriteString("HTTP/1.1 101 Switching Protocols\r\nUpgrade: test-protocol\r\nConnection: Upgrade\r\n\r\n") _ = brw.Flush() _, _ = io.Copy(io.Discard, conn) })) 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) proxySrv := httptest.NewServer(h) defer proxySrv.Close() pu, _ := url.Parse(proxySrv.URL) conn, err := net.DialTimeout("tcp", pu.Host, time.Second) if err != nil { t.Fatal(err) } defer conn.Close() _, _ = io.WriteString(conn, "GET /upgrade HTTP/1.1\r\nHost: "+pu.Host+"\r\nUpgrade: test-protocol\r\nConnection: Upgrade\r\n\r\n") resp, err := http.ReadResponse(bufio.NewReader(conn), nil) if err != nil { t.Fatal(err) } if resp.StatusCode != http.StatusSwitchingProtocols { t.Fatalf("status=%d", resp.StatusCode) } select { case entry := <-entries: t.Fatalf("non-WebSocket upgrade was logged: %+v", entry) case <-time.After(50 * time.Millisecond): } ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() if err := h.Shutdown(ctx); err != nil { t.Fatal(err) } } // TestWebSocketProxyPassthrough 验证 Upgrade 请求能正确透传, // 确认 captureWriter 的 Hijacker 实现没破坏 ReverseProxy 的 WS 行为。 func TestWebSocketProxyPassthrough(t *testing.T) { upstream := fakeUpstream(t) defer upstream.Close() u, err := url.Parse(upstream.URL) if err != nil { t.Fatal(err) } // disabled 模式下 ShouldLog 返回 true(全量记录),会进入捕获分支。 filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, noopSubmitter{}, 1024) proxySrv := httptest.NewServer(h) defer proxySrv.Close() // 建立到代理的 TCP 连接,手写 Upgrade 请求 pu, _ := url.Parse(proxySrv.URL) d := net.Dialer{Timeout: 3 * time.Second} conn, err := d.DialContext(context.Background(), "tcp", pu.Host) if err != nil { t.Fatal(err) } defer conn.Close() req := "GET /v1/realtime HTTP/1.1\r\n" + "Host: " + pu.Host + "\r\n" + "Upgrade: websocket\r\n" + "Connection: Upgrade\r\n" + "Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n" + "Sec-WebSocket-Version: 13\r\n" + "\r\n" if _, err := io.WriteString(conn, req); err != nil { t.Fatal(err) } _ = conn.SetReadDeadline(time.Now().Add(3 * time.Second)) br := bufio.NewReader(conn) resp, err := http.ReadResponse(br, nil) if err != nil { t.Fatalf("read upgrade response: %v", err) } if resp.StatusCode != http.StatusSwitchingProtocols { t.Fatalf("expected 101, got %d", resp.StatusCode) } if !strings.EqualFold(resp.Header.Get("Upgrade"), "websocket") { t.Fatalf("expected Upgrade: websocket, got %q", resp.Header.Get("Upgrade")) } // echo 测试 payload := "hello-websocket" if _, err := io.WriteString(conn, payload); err != nil { t.Fatal(err) } got := make([]byte, len(payload)) if _, err := io.ReadFull(br, got); err != nil { t.Fatalf("read echo: %v", err) } if string(got) != payload { t.Fatalf("echo mismatch: got %q want %q", got, payload) } }