cloudflare/cloudflared

Public

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

CodeCommitsIssuesPull requestsActionsInsightsSecurity
e2262085e57de4f5dbd5f3858b951b3e9f9596a8

Branches

Tags

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

Clone

HTTPS

Download ZIP

connection/http2_test.go

370lines · modecode

1package connection
2
3import (
4 "context"
5 "fmt"
6 "io"
7 "io/ioutil"
8 "net"
9 "net/http"
10 "net/http/httptest"
11 "sync"
12 "testing"
13 "time"
14
15 "github.com/stretchr/testify/assert"
16
17 "github.com/cloudflare/cloudflared/tunnelrpc/pogs"
18 tunnelpogs "github.com/cloudflare/cloudflared/tunnelrpc/pogs"
19
20 "github.com/gobwas/ws/wsutil"
21 "github.com/rs/zerolog"
22 "github.com/stretchr/testify/require"
23 "golang.org/x/net/http2"
24)
25
26var (
27 testTransport = http2.Transport{}
28)
29
30func newTestHTTP2Connection() (*http2Connection, net.Conn) {
31 edgeConn, originConn := net.Pipe()
32 var connIndex = uint8(0)
33 return NewHTTP2Connection(
34 originConn,
35 testConfig,
36 &NamedTunnelConfig{},
37 &pogs.ConnectionOptions{},
38 NewObserver(&log, &log, false),
39 connIndex,
40 mockConnectedFuse{},
41 nil,
42 ), edgeConn
43}
44
45func TestServeHTTP(t *testing.T) {
46 tests := []testRequest{
47 {
48 name: "ok",
49 endpoint: "ok",
50 expectedStatus: http.StatusOK,
51 expectedBody: []byte(http.StatusText(http.StatusOK)),
52 },
53 {
54 name: "large_file",
55 endpoint: "large_file",
56 expectedStatus: http.StatusOK,
57 expectedBody: testLargeResp,
58 },
59 {
60 name: "Bad request",
61 endpoint: "400",
62 expectedStatus: http.StatusBadRequest,
63 expectedBody: []byte(http.StatusText(http.StatusBadRequest)),
64 },
65 {
66 name: "Internal server error",
67 endpoint: "500",
68 expectedStatus: http.StatusInternalServerError,
69 expectedBody: []byte(http.StatusText(http.StatusInternalServerError)),
70 },
71 {
72 name: "Proxy error",
73 endpoint: "error",
74 expectedStatus: http.StatusBadGateway,
75 expectedBody: nil,
76 isProxyError: true,
77 },
78 }
79
80 http2Conn, edgeConn := newTestHTTP2Connection()
81
82 ctx, cancel := context.WithCancel(context.Background())
83 var wg sync.WaitGroup
84 wg.Add(1)
85 go func() {
86 defer wg.Done()
87 http2Conn.Serve(ctx)
88 }()
89
90 edgeHTTP2Conn, err := testTransport.NewClientConn(edgeConn)
91 require.NoError(t, err)
92
93 for _, test := range tests {
94 endpoint := fmt.Sprintf("http://localhost:8080/%s", test.endpoint)
95 req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
96 require.NoError(t, err)
97
98 resp, err := edgeHTTP2Conn.RoundTrip(req)
99 require.NoError(t, err)
100 require.Equal(t, test.expectedStatus, resp.StatusCode)
101 if test.expectedBody != nil {
102 respBody, err := ioutil.ReadAll(resp.Body)
103 require.NoError(t, err)
104 require.Equal(t, test.expectedBody, respBody)
105 }
106 if test.isProxyError {
107 require.Equal(t, responseMetaHeaderCfd, resp.Header.Get(ResponseMetaHeaderField))
108 } else {
109 require.Equal(t, responseMetaHeaderOrigin, resp.Header.Get(ResponseMetaHeaderField))
110 }
111 }
112 cancel()
113 wg.Wait()
114}
115
116type mockNamedTunnelRPCClient struct {
117 registered chan struct{}
118 unregistered chan struct{}
119}
120
121func (mc mockNamedTunnelRPCClient) RegisterConnection(
122 c context.Context,
123 config *NamedTunnelConfig,
124 options *tunnelpogs.ConnectionOptions,
125 connIndex uint8,
126 observer *Observer,
127) error {
128 close(mc.registered)
129 return nil
130}
131
132func (mc mockNamedTunnelRPCClient) GracefulShutdown(ctx context.Context, gracePeriod time.Duration) {
133 close(mc.unregistered)
134}
135
136func (mockNamedTunnelRPCClient) Close() {}
137
138type mockRPCClientFactory struct {
139 registered chan struct{}
140 unregistered chan struct{}
141}
142
143func (mf *mockRPCClientFactory) newMockRPCClient(context.Context, io.ReadWriteCloser, *zerolog.Logger) NamedTunnelRPCClient {
144 return mockNamedTunnelRPCClient{
145 registered: mf.registered,
146 unregistered: mf.unregistered,
147 }
148}
149
150type wsRespWriter struct {
151 *httptest.ResponseRecorder
152 readPipe *io.PipeReader
153 writePipe *io.PipeWriter
154}
155
156func newWSRespWriter() *wsRespWriter {
157 readPipe, writePipe := io.Pipe()
158 return &wsRespWriter{
159 httptest.NewRecorder(),
160 readPipe,
161 writePipe,
162 }
163}
164
165func (w *wsRespWriter) RespBody() io.ReadWriter {
166 return nowriter{w.readPipe}
167}
168
169func (w *wsRespWriter) Write(data []byte) (n int, err error) {
170 return w.writePipe.Write(data)
171}
172
173func TestServeWS(t *testing.T) {
174 http2Conn, _ := newTestHTTP2Connection()
175
176 ctx, cancel := context.WithCancel(context.Background())
177 var wg sync.WaitGroup
178 wg.Add(1)
179 go func() {
180 defer wg.Done()
181 http2Conn.Serve(ctx)
182 }()
183
184 respWriter := newWSRespWriter()
185 readPipe, writePipe := io.Pipe()
186
187 req, err := http.NewRequestWithContext(ctx, http.MethodGet, "http://localhost:8080/ws", readPipe)
188 require.NoError(t, err)
189 req.Header.Set(internalUpgradeHeader, websocketUpgrade)
190
191 wg.Add(1)
192 go func() {
193 defer wg.Done()
194 http2Conn.ServeHTTP(respWriter, req)
195 }()
196
197 data := []byte("test websocket")
198 err = wsutil.WriteClientText(writePipe, data)
199 require.NoError(t, err)
200
201 respBody, err := wsutil.ReadServerText(respWriter.RespBody())
202 require.NoError(t, err)
203 require.Equal(t, data, respBody, fmt.Sprintf("Expect %s, got %s", string(data), string(respBody)))
204
205 cancel()
206 resp := respWriter.Result()
207 // http2RespWriter should rewrite status 101 to 200
208 require.Equal(t, http.StatusOK, resp.StatusCode)
209 require.Equal(t, responseMetaHeaderOrigin, resp.Header.Get(ResponseMetaHeaderField))
210
211 wg.Wait()
212}
213
214func TestServeControlStream(t *testing.T) {
215 http2Conn, edgeConn := newTestHTTP2Connection()
216
217 rpcClientFactory := mockRPCClientFactory{
218 registered: make(chan struct{}),
219 unregistered: make(chan struct{}),
220 }
221 http2Conn.newRPCClientFunc = rpcClientFactory.newMockRPCClient
222
223 ctx, cancel := context.WithCancel(context.Background())
224 var wg sync.WaitGroup
225 wg.Add(1)
226 go func() {
227 defer wg.Done()
228 http2Conn.Serve(ctx)
229 }()
230
231 req, err := http.NewRequestWithContext(ctx, http.MethodGet, "http://localhost:8080/", nil)
232 require.NoError(t, err)
233 req.Header.Set(internalUpgradeHeader, controlStreamUpgrade)
234
235 edgeHTTP2Conn, err := testTransport.NewClientConn(edgeConn)
236 require.NoError(t, err)
237
238 wg.Add(1)
239 go func() {
240 defer wg.Done()
241 edgeHTTP2Conn.RoundTrip(req)
242 }()
243
244 <-rpcClientFactory.registered
245 cancel()
246 <-rpcClientFactory.unregistered
247 assert.False(t, http2Conn.stoppedGracefully)
248
249 wg.Wait()
250}
251
252func TestGracefulShutdownHTTP2(t *testing.T) {
253 http2Conn, edgeConn := newTestHTTP2Connection()
254
255 rpcClientFactory := mockRPCClientFactory{
256 registered: make(chan struct{}),
257 unregistered: make(chan struct{}),
258 }
259 events := &eventCollectorSink{}
260 http2Conn.newRPCClientFunc = rpcClientFactory.newMockRPCClient
261 http2Conn.observer.RegisterSink(events)
262 shutdownC := make(chan struct{})
263 http2Conn.gracefulShutdownC = shutdownC
264
265 ctx, cancel := context.WithCancel(context.Background())
266 var wg sync.WaitGroup
267 wg.Add(1)
268 go func() {
269 defer wg.Done()
270 http2Conn.Serve(ctx)
271 }()
272
273 req, err := http.NewRequestWithContext(ctx, http.MethodGet, "http://localhost:8080/", nil)
274 require.NoError(t, err)
275 req.Header.Set(internalUpgradeHeader, controlStreamUpgrade)
276
277 edgeHTTP2Conn, err := testTransport.NewClientConn(edgeConn)
278 require.NoError(t, err)
279
280 wg.Add(1)
281 go func() {
282 defer wg.Done()
283 _, _ = edgeHTTP2Conn.RoundTrip(req)
284 }()
285
286 select {
287 case <-rpcClientFactory.registered:
288 break //ok
289 case <-time.Tick(time.Second):
290 t.Fatal("timeout out waiting for registration")
291 }
292
293 // signal graceful shutdown
294 close(shutdownC)
295
296 select {
297 case <-rpcClientFactory.unregistered:
298 break //ok
299 case <-time.Tick(time.Second):
300 t.Fatal("timeout out waiting for unregistered signal")
301 }
302 assert.True(t, http2Conn.stoppedGracefully)
303
304 cancel()
305 wg.Wait()
306
307 events.assertSawEvent(t, Event{
308 Index: http2Conn.connIndex,
309 EventType: Unregistering,
310 })
311}
312
313func benchmarkServeHTTP(b *testing.B, test testRequest) {
314 http2Conn, edgeConn := newTestHTTP2Connection()
315
316 ctx, cancel := context.WithCancel(context.Background())
317 var wg sync.WaitGroup
318 wg.Add(1)
319 go func() {
320 defer wg.Done()
321 http2Conn.Serve(ctx)
322 }()
323
324 endpoint := fmt.Sprintf("http://localhost:8080/%s", test.endpoint)
325 req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
326 require.NoError(b, err)
327
328 edgeHTTP2Conn, err := testTransport.NewClientConn(edgeConn)
329 require.NoError(b, err)
330
331 b.ResetTimer()
332 for i := 0; i < b.N; i++ {
333 b.StartTimer()
334 resp, err := edgeHTTP2Conn.RoundTrip(req)
335 b.StopTimer()
336 require.NoError(b, err)
337 require.Equal(b, test.expectedStatus, resp.StatusCode)
338 if test.expectedBody != nil {
339 respBody, err := ioutil.ReadAll(resp.Body)
340 require.NoError(b, err)
341 require.Equal(b, test.expectedBody, respBody)
342 }
343 resp.Body.Close()
344 }
345
346 cancel()
347 wg.Wait()
348}
349
350func BenchmarkServeHTTPSimple(b *testing.B) {
351 test := testRequest{
352 name: "ok",
353 endpoint: "ok",
354 expectedStatus: http.StatusOK,
355 expectedBody: []byte(http.StatusText(http.StatusOK)),
356 }
357
358 benchmarkServeHTTP(b, test)
359}
360
361func BenchmarkServeHTTPLargeFile(b *testing.B) {
362 test := testRequest{
363 name: "large_file",
364 endpoint: "large_file",
365 expectedStatus: http.StatusOK,
366 expectedBody: testLargeResp,
367 }
368
369 benchmarkServeHTTP(b, test)
370}
371