diff --git a/cmd/rdpgw/protocol/gateway_test.go b/cmd/rdpgw/protocol/gateway_test.go index 779966f..df270fe 100644 --- a/cmd/rdpgw/protocol/gateway_test.go +++ b/cmd/rdpgw/protocol/gateway_test.go @@ -73,11 +73,13 @@ func TestHandleGatewayProtocolRouting(t *testing.T) { cases := []struct { name string + connectionID string request string wantStatusLine string }{ { - name: "RDG_OUT_DATA without upgrade headers routes to legacy", + name: "RDG_OUT_DATA without upgrade headers routes to legacy", + connectionID: "test-legacy", request: "RDG_OUT_DATA /remoteDesktopGateway/ HTTP/1.1\r\n" + "Host: " + addr + "\r\n" + "Rdg-Connection-Id: test-legacy\r\n" + @@ -85,7 +87,8 @@ func TestHandleGatewayProtocolRouting(t *testing.T) { wantStatusLine: "HTTP/1.1 200 OK", }, { - name: "RDG_OUT_DATA with upgrade headers routes to websocket", + name: "RDG_OUT_DATA with upgrade headers routes to websocket", + connectionID: "test-ws", request: "RDG_OUT_DATA /remoteDesktopGateway/ HTTP/1.1\r\n" + "Host: " + addr + "\r\n" + "Rdg-Connection-Id: test-ws\r\n" + @@ -97,7 +100,8 @@ func TestHandleGatewayProtocolRouting(t *testing.T) { wantStatusLine: "HTTP/1.1 101 Switching Protocols", }, { - name: "RDG_OUT_DATA with Connection token list still routes to websocket", + name: "RDG_OUT_DATA with Connection token list still routes to websocket", + connectionID: "test-ws-list", request: "RDG_OUT_DATA /remoteDesktopGateway/ HTTP/1.1\r\n" + "Host: " + addr + "\r\n" + "Rdg-Connection-Id: test-ws-list\r\n" + @@ -109,7 +113,8 @@ func TestHandleGatewayProtocolRouting(t *testing.T) { wantStatusLine: "HTTP/1.1 101 Switching Protocols", }, { - name: "RDG_OUT_DATA with partially matching headers routes to legacy", + name: "RDG_OUT_DATA with partially matching headers routes to legacy", + connectionID: "test-partial", request: "RDG_OUT_DATA /remoteDesktopGateway/ HTTP/1.1\r\n" + "Host: " + addr + "\r\n" + "Rdg-Connection-Id: test-partial\r\n" + @@ -121,6 +126,9 @@ func TestHandleGatewayProtocolRouting(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { + c.Delete(tc.connectionID) + t.Cleanup(func() { c.Delete(tc.connectionID) }) + conn, err := net.DialTimeout("tcp", addr, 2*time.Second) if err != nil { t.Fatalf("dial: %v", err)