Files
token_thief/tests/proxy/websocket_test.go
T

300 lines
8.2 KiB
Go

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)
}
}