package proxy_test import ( "bufio" "context" "io" "net" "net/http" "net/http/httptest" "net/url" "strings" "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 } // 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) } 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 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) } }