cloudflare/cloudflared

Public

mirrored from https://github.com/cloudflare/cloudflaredAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
2024.9.1

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

carrier/websocket_test.go

123lines · modeblame

3ad99b24Igor Postelnik5 years ago1package carrier
2
3import (
4"context"
5"crypto/tls"
6"crypto/x509"
7"fmt"
8"math/rand"
9"testing"
10"time"
11
12gws "github.com/gorilla/websocket"
13"github.com/rs/zerolog"
14"github.com/stretchr/testify/assert"
15"github.com/stretchr/testify/require"
16"golang.org/x/net/websocket"
17
18"github.com/cloudflare/cloudflared/hello"
19"github.com/cloudflare/cloudflared/tlsconfig"
20cfwebsocket "github.com/cloudflare/cloudflared/websocket"
21)
22
23func websocketClientTLSConfig(t *testing.T) *tls.Config {
24certPool := x509.NewCertPool()
25helloCert, err := tlsconfig.GetHelloCertificateX509()
26assert.NoError(t, err)
27certPool.AddCert(helloCert)
28assert.NotNil(t, certPool)
29return &tls.Config{RootCAs: certPool}
30}
31
32func TestWebsocketHeaders(t *testing.T) {
33req := testRequest(t, "http://example.com", nil)
34wsHeaders := websocketHeaders(req)
35for _, header := range stripWebsocketHeaders {
36assert.Empty(t, wsHeaders[header])
37}
38assert.Equal(t, "curl/7.59.0", wsHeaders.Get("User-Agent"))
39}
40
41func TestServe(t *testing.T) {
42log := zerolog.Nop()
43shutdownC := make(chan struct{})
44errC := make(chan error)
45listener, err := hello.CreateTLSListener("localhost:1111")
46assert.NoError(t, err)
47defer listener.Close()
48
49go func() {
50errC <- hello.StartHelloWorldServer(&log, listener, shutdownC)
51}()
52
53req := testRequest(t, "https://localhost:1111/ws", nil)
54
55tlsConfig := websocketClientTLSConfig(t)
56assert.NotNil(t, tlsConfig)
57d := gws.Dialer{TLSClientConfig: tlsConfig}
58conn, resp, err := clientConnect(req, &d)
59assert.NoError(t, err)
60assert.Equal(t, "websocket", resp.Header.Get("Upgrade"))
61
62for i := 0; i < 1000; i++ {
63messageSize := rand.Int()%2048 + 1
64clientMessage := make([]byte, messageSize)
65// rand.Read always returns len(clientMessage) and a nil error
66rand.Read(clientMessage)
67err = conn.WriteMessage(websocket.BinaryFrame, clientMessage)
68assert.NoError(t, err)
69
70messageType, message, err := conn.ReadMessage()
71assert.NoError(t, err)
72assert.Equal(t, websocket.BinaryFrame, messageType)
73assert.Equal(t, clientMessage, message)
74}
75
76_ = conn.Close()
77close(shutdownC)
78<-errC
79}
80
81func TestWebsocketWrapper(t *testing.T) {
82listener, err := hello.CreateTLSListener("localhost:0")
83require.NoError(t, err)
84
85serverErrorChan := make(chan error)
86helloSvrCtx, cancelHelloSvr := context.WithCancel(context.Background())
87defer func() { <-serverErrorChan }()
88defer cancelHelloSvr()
89go func() {
90log := zerolog.Nop()
91serverErrorChan <- hello.StartHelloWorldServer(&log, listener, helloSvrCtx.Done())
92}()
93
94tlsConfig := websocketClientTLSConfig(t)
95d := gws.Dialer{TLSClientConfig: tlsConfig, HandshakeTimeout: time.Minute}
96testAddr := fmt.Sprintf("https://%s/ws", listener.Addr().String())
97req := testRequest(t, testAddr, nil)
98conn, resp, err := clientConnect(req, &d)
99require.NoError(t, err)
100assert.Equal(t, "websocket", resp.Header.Get("Upgrade"))
101
102// Websocket now connected to test server so lets check our wrapper
103wrapper := cfwebsocket.GorillaConn{Conn: conn}
104buf := make([]byte, 100)
105wrapper.Write([]byte("abc"))
106n, err := wrapper.Read(buf)
107require.NoError(t, err)
108require.Equal(t, n, 3)
109require.Equal(t, "abc", string(buf[:n]))
110
111// Test partial read, read 1 of 3 bytes in one read and the other 2 in another read
112wrapper.Write([]byte("abc"))
113buf = buf[:1]
114n, err = wrapper.Read(buf)
115require.NoError(t, err)
116require.Equal(t, n, 1)
117require.Equal(t, "a", string(buf[:n]))
118buf = buf[:cap(buf)]
119n, err = wrapper.Read(buf)
120require.NoError(t, err)
121require.Equal(t, n, 2)
122require.Equal(t, "bc", string(buf[:n]))
123}