Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 12 additions & 4 deletions cmd/rdpgw/protocol/gateway_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -73,19 +73,22 @@ 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" +
"\r\n",
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" +
Expand All @@ -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" +
Expand All @@ -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" +
Expand All @@ -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)
Expand Down
Loading