package session import ( "bufio" "context" "errors" "io" "net" "net/http" "net/http/httptest" "strings" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) // TestForwarder_OpenStream_NoProxyHost 驗證 baseURL 為空時直接拒絕。 func TestForwarder_OpenStream_NoProxyHost(t *testing.T) { f := NewForwarder("", nil) _, err := f.OpenStream(context.Background(), "vAc_x") assert.Error(t, err) } // TestForwarder_OpenStream_EmptyToken 驗證空 token 拒絕。 func TestForwarder_OpenStream_EmptyToken(t *testing.T) { f := NewForwarder("http://localhost:9999", nil) _, err := f.OpenStream(context.Background(), "") assert.Error(t, err) } // TestForwarder_ForwardWebSocket_NilReq 驗證 nil req 直接拒絕。 func TestForwarder_ForwardWebSocket_NilReq(t *testing.T) { f := NewForwarder("http://localhost:9999", nil) _, err := f.ForwardWebSocket(context.Background(), "vAc_x", nil) assert.Error(t, err) } // TestForwarder_ForwardWebSocket_NoProxyHost 驗證 baseURL 為空時 OpenStream 就失敗。 func TestForwarder_ForwardWebSocket_NoProxyHost(t *testing.T) { f := NewForwarder("", nil) req, _ := http.NewRequest(http.MethodGet, "/ws/devices/x/inference", nil) _, err := f.ForwardWebSocket(context.Background(), "vAc_x", req) assert.Error(t, err) } // fakeRawProxy 起一個假的 remote-proxy /internal/forward/raw: // hijack 後回 "200 Connected",接著 handoff 給 onStream 模擬 local agent 端行為。 // // 回傳 baseURL 供 NewForwarder。 func fakeRawProxy(t *testing.T, onStream func(conn net.Conn)) string { t.Helper() ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { hj, ok := w.(http.Hijacker) require.True(t, ok, "test server must support hijack") conn, _, err := hj.Hijack() require.NoError(t, err) // 回 raw forward 握手 if _, err := conn.Write([]byte("HTTP/1.1 200 Connected\r\n\r\n")); err != nil { _ = conn.Close() return } onStream(conn) // 模擬 tunnel 另一端(local agent) })) t.Cleanup(ts.Close) return ts.URL } // TestForwarder_ForwardWebSocket_101_Success 驗證 happy path: // forwarder 寫 upgrade request → 假 local agent 回 101 → forwarder 回 WebSocketConn。 func TestForwarder_ForwardWebSocket_101_Success(t *testing.T) { const framePayload = "hello-ws-frame" baseURL := fakeRawProxy(t, func(conn net.Conn) { defer conn.Close() // 讀 upgrade request req, err := http.ReadRequest(bufio.NewReader(conn)) if err != nil { return } // 驗證 upgrade header 有被帶過來 assert.True(t, strings.EqualFold(req.Header.Get("Upgrade"), "websocket"), "local agent 應收到 Upgrade: websocket") // 回 101 + 一段 frame bytes(模擬升級後 local agent 主動推的資料) _, _ = conn.Write([]byte( "HTTP/1.1 101 Switching Protocols\r\n" + "Upgrade: websocket\r\n" + "Connection: Upgrade\r\n" + "Sec-WebSocket-Accept: fake-accept\r\n\r\n" + framePayload, )) }) f := NewForwarder(baseURL, nil) req, _ := http.NewRequest(http.MethodGet, "/ws/devices/dev1/inference", nil) req.Header.Set("Upgrade", "websocket") req.Header.Set("Connection", "Upgrade") req.Header.Set("Sec-WebSocket-Key", "dGhlIHNhbXBsZSBub25jZQ==") req.Header.Set("Sec-WebSocket-Version", "13") wsConn, err := f.ForwardWebSocket(context.Background(), "vAc_x", req) require.NoError(t, err) require.NotNil(t, wsConn) defer wsConn.Conn.Close() assert.Equal(t, http.StatusSwitchingProtocols, wsConn.Resp.StatusCode) assert.Equal(t, "fake-accept", wsConn.Resp.Header.Get("Sec-WebSocket-Accept")) // 驗證 101 後緊跟的 frame bytes 沒有被 bufio 預讀吞掉(prefixConn 接回)。 buf := make([]byte, len(framePayload)) n, err := io.ReadFull(wsConn.Conn, buf) require.NoError(t, err) assert.Equal(t, framePayload, string(buf[:n]), "101 後預讀的 frame bytes 應完整保留") } // TestForwarder_ForwardWebSocket_Non101_Rejected 驗證 local agent 回非 101 時, // 回傳 wsUpgradeError 且能透過 AsWSUpgradeError 取出原 response。 func TestForwarder_ForwardWebSocket_Non101_Rejected(t *testing.T) { baseURL := fakeRawProxy(t, func(conn net.Conn) { defer conn.Close() _, _ = http.ReadRequest(bufio.NewReader(conn)) // local agent 拒絕 upgrade:回 403 _, _ = conn.Write([]byte( "HTTP/1.1 403 Forbidden\r\n" + "Content-Type: text/plain\r\n" + "Content-Length: 7\r\n\r\n" + "denied!", )) }) f := NewForwarder(baseURL, nil) req, _ := http.NewRequest(http.MethodGet, "/ws/devices/dev1/inference", nil) req.Header.Set("Upgrade", "websocket") _, err := f.ForwardWebSocket(context.Background(), "vAc_x", req) require.Error(t, err) resp, ok := AsWSUpgradeError(err) require.True(t, ok, "應為 wsUpgradeError") assert.Equal(t, http.StatusForbidden, resp.StatusCode) body, _ := io.ReadAll(resp.Body) _ = resp.Body.Close() assert.Equal(t, "denied!", string(body)) } // TestForwarder_ForwardWebSocket_NoSession 驗證無 session(remote-proxy 回 502)→ // 映射成 ErrSessionNotFound(沿用 OpenStream 的錯誤映射)。 func TestForwarder_ForwardWebSocket_NoSession(t *testing.T) { ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusBadGateway) _, _ = w.Write([]byte(`{"error":{"code":"TUNNEL_DISCONNECTED"}}`)) })) defer ts.Close() f := NewForwarder(ts.URL, nil) req, _ := http.NewRequest(http.MethodGet, "/ws/devices/dev1/inference", nil) req.Header.Set("Upgrade", "websocket") _, err := f.ForwardWebSocket(context.Background(), "vAc_dead", req) assert.ErrorIs(t, err, ErrSessionNotFound) } // TestForwarder_OpenStream_502_TreatedAsNotFound 驗證當 remote-proxy 回 502 // (session 不存在時的雛形行為)→ 包裝成 ErrSessionNotFound。 // // 用 httptest 起一個假的 internal endpoint,回 502 JSON。 func TestForwarder_OpenStream_502_TreatedAsNotFound(t *testing.T) { ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusBadGateway) _, _ = w.Write([]byte(`{"error":{"code":"TUNNEL_DISCONNECTED","message":"session not connected"}}`)) })) defer ts.Close() f := NewForwarder(ts.URL, nil) _, err := f.OpenStream(context.Background(), "vAc_dead") if !errors.Is(err, ErrSessionNotFound) { t.Fatalf("expected ErrSessionNotFound, got %v", err) } } // TestForwarder_OpenStream_HandshakeRead 驗證能正確讀「HTTP/1.1 200 Connected\r\n\r\n」 // 握手;用一個假 server 回正確握手後立刻 close — 期望我們的 OpenStream 成功, // 後續 Read 拿 EOF(這對 forwarder 而言是合法情境,由 caller 處理)。 // // 此 case 直接驗證 happy-path 握手解析;真正的端對端轉發由 integration test 涵蓋。 func TestForwarder_OpenStream_HandshakeRead(t *testing.T) { // 為了保證 server 端在 200 Connected 後不再寫 body(讓 forwarder 結束 header 讀 // 不被預讀干擾),用一個 raw TCP listener 而非 httptest.NewServer。 // 但 raw listener 會增加測試複雜度;在 unit test 用 httptest 已足以驗證 // 「能 parse 200 Connected + 兩個 \r\n」的路徑——讀 body 結束會回 EOF, // 後續 caller 用該 conn 才會發現問題,這裡僅驗證 OpenStream 不 error。 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // 不能直接寫 raw "HTTP/1.1 200 Connected\r\n\r\n" — httptest 會額外加 // content-length 等 header。改用 hijack 模擬真實 raw forward 行為。 hj, ok := w.(http.Hijacker) if !ok { http.Error(w, "no hijacker", 500) return } conn, _, err := hj.Hijack() if err != nil { return } defer conn.Close() _, _ = conn.Write([]byte("HTTP/1.1 200 Connected\r\n\r\n")) // 不再寫;讓 forwarder 拿到 conn 後若 read 會 EOF })) defer ts.Close() f := NewForwarder(ts.URL, nil) conn, err := f.OpenStream(context.Background(), "vAc_x") if err != nil { t.Fatalf("OpenStream should succeed: %v", err) } defer conn.Close() // 不再做 read 驗證(行為由 integration test 涵蓋) }