300 lines
8.2 KiB
Go
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)
|
|
}
|
|
}
|