fix: restore reviewable migration evidence
This commit is contained in:
@@ -0,0 +1,299 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user