Skip to content
File

Blob: integrations/tunnel_test.go

go497 lines
1package integrations
2 
3import (
4 "bufio"
5 "bytes"
6 "context"
7 cRand "crypto/rand"
8 "crypto/tls"
9 "fmt"
10 "io"
11 "net"
12 "net/http"
13 "net/http/httptest"
14 "net/url"
15 "os"
16 "strings"
17 "sync/atomic"
18 "testing"
19 "time"
20 
21 "go.miragespace.co/specter/cmd/client"
22 "go.miragespace.co/specter/cmd/server"
23 "go.miragespace.co/specter/util/bufconn"
24 "go.miragespace.co/specter/util/testcond"
25 
26 "github.com/go-chi/chi/v5"
27 "github.com/gorilla/websocket"
28 "github.com/quic-go/quic-go"
29 "github.com/quic-go/quic-go/http3"
30 "github.com/stretchr/testify/require"
31 "github.com/urfave/cli/v3"
32 "go.uber.org/zap"
33 "go.uber.org/zap/zaptest/observer"
34)
35 
36const (
37 testBinaryLength = 16
38 testBody = "yay"
39 testRespBody = "cool"
40 serverApex = "dev.con.nect.sh"
41 yamlTemplate = `version: 2
42apex: 127.0.0.1:%d
43tunnels:
44 - target: %s
45 insecure: true
46 - target: %s
47`
48)
49 
50var (
51 serverPorts = []int{21948, 21949, 21950, 21951, 21952}
52 serverHttpPorts = []int{21848, 21849, 21850, 21851, 21852}
53)
54 
55type TestWsMsg struct {
56 Message string
57}
58 
59func compileApp(cmd *cli.Command) (*cli.Command, *observer.ObservedLogs) {
60 observedZapCore, observedLogs := observer.New(zap.DebugLevel)
61 observedLogger := zap.New(observedZapCore)
62 cmd.HideHelp = true
63 return &cli.Command{
64 Name: "specter",
65 HideHelp: true,
66 HideVersion: true,
67 Commands: []*cli.Command{
68 cmd,
69 },
70 Before: func(ctx context.Context, cmd *cli.Command) (context.Context, error) {
71 cmd.Root().Metadata["logger"] = observedLogger
72 return ctx, nil
73 },
74 Metadata: make(map[string]any),
75 }, observedLogs
76}
77 
78func TestIntegrationTunnel(t *testing.T) {
79 if os.Getenv("GO_INTEGRATION_TUNNEL") == "" {
80 t.Skip("skipping integration tests")
81 }
82 
83 as := require.New(t)
84 
85 t.Logf("Creating HTTP server for forwarding target\n")
86 
87 var upgrader = websocket.Upgrader{
88 ReadBufferSize: 1024,
89 WriteBufferSize: 1024,
90 }
91 mux := chi.NewRouter()
92 mux.Get("/", func(w http.ResponseWriter, r *http.Request) {
93 w.Write([]byte(testBody))
94 })
95 mux.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
96 conn, err := upgrader.Upgrade(w, r, nil)
97 as.NoError(err)
98 defer conn.Close()
99 t := &TestWsMsg{}
100 as.NoError(conn.ReadJSON(t))
101 as.Equal(testBody, t.Message)
102 t.Message = testRespBody
103 as.NoError(conn.WriteJSON(t))
104 })
105 
106 ts := httptest.NewUnstartedServer(mux)
107 ts.EnableHTTP2 = true
108 ts.StartTLS()
109 defer ts.Close()
110 
111 t.Logf("Creating TCP server for forwarding target\n")
112 
113 listener, err := net.Listen("tcp", "127.0.0.1:0")
114 as.NoError(err)
115 defer listener.Close()
116 
117 var (
118 connectionAccepted atomic.Int32
119 )
120 defer func() {
121 as.Equal(int32(len(serverPorts)+len(serverHttpPorts)), connectionAccepted.Load(), "tcp target should have accepted all connections")
122 }()
123 go func() {
124 for {
125 c, err := listener.Accept()
126 if err != nil {
127 return
128 }
129 connectionAccepted.Add(1)
130 go io.Copy(c, c)
131 }
132 }()
133 
134 t.Logf("Generating client config\n")
135 
136 tcpTarget := listener.Addr().String()
137 
138 file, err := os.CreateTemp("", "client")
139 as.NoError(err)
140 defer os.Remove(file.Name())
141 
142 _, err = file.WriteString(fmt.Sprintf(yamlTemplate, serverPorts[0], ts.URL, fmt.Sprintf("tcp://%s", tcpTarget)))
143 as.NoError(err)
144 as.NoError(file.Close())
145 
146 ctx := t.Context()
147 
148 serverReturnCtx, serverStopped := context.WithCancel(ctx)
149 defer serverStopped()
150 
151 serverLogs := make([]*observer.ObservedLogs, len(serverPorts))
152 serverArgs := make([][]string, len(serverPorts))
153 
154 for i, port := range serverPorts {
155 dir, err := os.MkdirTemp("", fmt.Sprintf("integration-%d", port))
156 as.NoError(err)
157 defer os.RemoveAll(dir)
158 
159 serverArgs[i] = []string{
160 "specter",
161 "server",
162 "--cert-dir",
163 "../certs",
164 "--data-dir",
165 dir,
166 "--listen",
167 fmt.Sprintf("127.0.0.1:%d", port),
168 "--listen-http",
169 fmt.Sprintf("%d", serverHttpPorts[i]),
170 "--apex",
171 serverApex,
172 }
173 
174 if i != 0 {
175 serverArgs[i] = append(serverArgs[i], []string{
176 "--join",
177 fmt.Sprintf("127.0.0.1:%d", serverPorts[0]),
178 }...)
179 }
180 }
181 
182 defer func() {
183 for i, port := range serverPorts {
184 logs := serverLogs[i]
185 for _, entry := range logs.All() {
186 t.Logf("%d: %s\n", port, entry.Message)
187 // for _, x := range entry.Context {
188 // if x.Interface != nil {
189 // t.Logf(" %s: %v\n", x.Key, x.Interface)
190 // } else if x.String != "" {
191 // t.Logf(" %s: %v\n", x.Key, x.String)
192 // } else {
193 // t.Logf(" %s: %v\n", x.Key, x.Integer)
194 // }
195 // }
196 }
197 }
198 }()
199 
200 t.Run("starting servers", func(t *testing.T) {
201 as := require.New(t)
202 for i, args := range serverArgs {
203 sApp, sLogs := compileApp(server.Generate())
204 serverLogs[i] = sLogs
205 args := args
206 go func(app *cli.Command) {
207 if err := app.Run(ctx, args); err != nil {
208 as.NoError(err)
209 }
210 serverStopped()
211 }(sApp)
212 }
213 
214 as.NoError(testcond.WaitForCondition(func() bool {
215 select {
216 case <-serverReturnCtx.Done():
217 as.FailNow("server returned unexpectedly")
218 return false
219 default:
220 started := 0
221 for _, sLogs := range serverLogs {
222 serverLogs := sLogs.All()
223 for _, l := range serverLogs {
224 if strings.Contains(l.Message, "server started") {
225 started++
226 }
227 }
228 }
229 return started == (2 * len(serverPorts))
230 }
231 }, time.Millisecond*100, time.Second*30), "timeout expecting specter and gateway servers started")
232 })
233 
234 var hostMap map[string]string
235 t.Run("starting client", func(t *testing.T) {
236 as := require.New(t)
237 
238 clientReturn := make(chan struct{})
239 clientArgs := []string{
240 "specter",
241 "client",
242 "--insecure",
243 "tunnel",
244 "--config",
245 file.Name(),
246 }
247 
248 cApp, cLogs := compileApp(client.Generate())
249 cApp.Metadata["apexOverride"] = serverApex
250 go func() {
251 if err := cApp.Run(ctx, clientArgs); err != nil {
252 as.NoError(err)
253 }
254 close(clientReturn)
255 }()
256 
257 t.Logf("Waiting for client to publish tunnels\n")
258 
259 as.NoError(testcond.WaitForCondition(func() bool {
260 select {
261 case <-clientReturn:
262 as.FailNow("client returned unexpectedly")
263 return false
264 default:
265 hostMap = make(map[string]string)
266 clientLogs := cLogs.All()
267 for _, l := range clientLogs {
268 if !strings.Contains(l.Message, "published") {
269 continue
270 }
271 var hostname string
272 var proto string
273 for _, f := range l.Context {
274 switch f.Key {
275 case "hostname":
276 hostname = f.String
277 case "target":
278 if strings.Contains(f.String, "http") {
279 proto = "http"
280 }
281 if strings.Contains(f.String, "tcp") {
282 proto = "tcp"
283 }
284 }
285 }
286 if hostname != "" && proto != "" {
287 hostMap[proto] = hostname
288 }
289 }
290 return len(hostMap) == 2
291 }
292 }, time.Millisecond*100, time.Second*10), "timeout waiting for hostname of tunnels in client log")
293 
294 t.Logf("Found hostnames %v\n", hostMap)
295 })
296 
297 t.Logf("Start integration test\n")
298 
299 for serverIndex, serverPort := range serverPorts {
300 t.Run(fmt.Sprintf("with %d as endpoint", serverPort), func(t *testing.T) {
301 as := require.New(t)
302 
303 req, err := http.NewRequest("GET", fmt.Sprintf("https://%s:%d/", hostMap["http"], serverPort), nil)
304 as.NoError(err)
305 baseCfg := &tls.Config{
306 ServerName: hostMap["http"],
307 InsecureSkipVerify: true,
308 NextProtos: []string{"h2"},
309 }
310 httpClient := &http.Client{
311 Timeout: time.Second * 2,
312 }
313 
314 t.Run("HTTP over TCP", func(t *testing.T) {
315 as := require.New(t)
316 
317 cfg := baseCfg.Clone()
318 
319 httpClient.Transport = &http.Transport{
320 ForceAttemptHTTP2: true,
321 TLSClientConfig: cfg,
322 DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
323 dialer := &tls.Dialer{
324 Config: cfg,
325 }
326 return dialer.DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", serverPort))
327 },
328 }
329 
330 resp, err := httpClient.Do(req)
331 as.NoError(err)
332 defer resp.Body.Close()
333 
334 as.True(resp.ProtoAtLeast(2, 0))
335 
336 var buf bytes.Buffer
337 _, err = buf.ReadFrom(resp.Body)
338 as.NoError(err)
339 as.Equal(testBody, buf.String())
340 })
341 
342 t.Run("HTTP over QUIC", func(t *testing.T) {
343 as := require.New(t)
344 
345 h3Cfg := baseCfg.Clone()
346 h3Cfg.NextProtos = []string{"h3"}
347 httpClient.Transport = &http3.Transport{
348 TLSClientConfig: h3Cfg,
349 Dial: func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
350 return quic.DialAddr(ctx, fmt.Sprintf("127.0.0.1:%d", serverPort), h3Cfg, nil)
351 },
352 }
353 resp, err := httpClient.Do(req)
354 
355 as.NoError(err)
356 defer resp.Body.Close()
357 
358 as.True(resp.ProtoAtLeast(2, 0))
359 
360 var buf bytes.Buffer
361 _, err = buf.ReadFrom(resp.Body)
362 as.NoError(err)
363 as.Equal(testBody, buf.String())
364 })
365 
366 t.Run("WebSocket", func(t *testing.T) {
367 as := require.New(t)
368 
369 wsCfg := baseCfg.Clone()
370 wsCfg.NextProtos = []string{"http/1.1"}
371 wsDialer := &websocket.Dialer{
372 TLSClientConfig: wsCfg,
373 NetDialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
374 dialer := &tls.Dialer{
375 Config: wsCfg,
376 }
377 return dialer.DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", serverPort))
378 },
379 }
380 
381 wsConn, _, err := wsDialer.Dial(fmt.Sprintf("wss://%s:%d/ws", hostMap["http"], serverPort), nil)
382 as.NoError(err)
383 
384 defer wsConn.Close()
385 
386 testMsg := &TestWsMsg{
387 Message: testBody,
388 }
389 as.NoError(wsConn.WriteJSON(testMsg))
390 as.NoError(wsConn.ReadJSON(testMsg))
391 as.Equal(testRespBody, testMsg.Message)
392 })
393 
394 t.Run("TCP Tunnel using specter client", func(t *testing.T) {
395 as := require.New(t)
396 
397 connectArgs := []string{
398 "specter",
399 "client",
400 "--insecure",
401 "connect",
402 fmt.Sprintf("127.0.0.1:%d", serverPort),
403 }
404 
405 connectReturnCtx, connectReturn := context.WithCancel(ctx)
406 defer connectReturn()
407 
408 xApp, xLogs := compileApp(client.Generate())
409 xApp.Metadata["connectOverride"] = hostMap["tcp"]
410 
411 leftConn, rightConn := bufconn.BufferedPipe(8192)
412 xApp.Reader = leftConn
413 xApp.Writer = leftConn
414 
415 go func() {
416 if err := xApp.Run(ctx, connectArgs); err != nil {
417 as.NoError(err)
418 }
419 connectReturn()
420 }()
421 
422 as.NoError(testcond.WaitForCondition(func() bool {
423 select {
424 case <-connectReturnCtx.Done():
425 as.FailNow("connect returned unexpectedly")
426 return false
427 default:
428 connected := false
429 connectLogs := xLogs.All()
430 for _, l := range connectLogs {
431 if strings.Contains(l.Message, "established") {
432 connected = true
433 }
434 }
435 return connected
436 }
437 }, time.Millisecond*100, time.Second*3), "timeout waiting for tunnel to be connected")
438 
439 write := make([]byte, testBinaryLength)
440 _, err := io.ReadFull(cRand.Reader, write)
441 as.NoError(err)
442 
443 _, err = rightConn.Write(write)
444 as.NoError(err)
445 
446 as.NoError(rightConn.SetReadDeadline(time.Now().Add(time.Second)))
447 read := make([]byte, testBinaryLength)
448 n, err := rightConn.Read(read)
449 as.NoError(err)
450 as.Equal(testBinaryLength, n)
451 as.EqualValues(write, read)
452 })
453 
454 t.Run("TCP Tunnel over http connect", func(t *testing.T) {
455 as := require.New(t)
456 
457 // HTTP Connect start
458 dialer := &net.Dialer{
459 Timeout: time.Second,
460 }
461 tcpConn, err := dialer.Dial("tcp", fmt.Sprintf("127.0.0.1:%d", serverHttpPorts[serverIndex]))
462 as.NoError(err)
463 defer tcpConn.Close()
464 
465 proxyAddr := fmt.Sprintf("%s:%d", hostMap["tcp"], 1234)
466 req := &http.Request{
467 Method: http.MethodConnect,
468 URL: &url.URL{
469 Opaque: proxyAddr,
470 },
471 Host: proxyAddr,
472 Header: make(http.Header),
473 }
474 as.NoError(req.Write(tcpConn))
475 resp, err := http.ReadResponse(bufio.NewReader(tcpConn), req)
476 as.NoError(err)
477 as.Equal(http.StatusOK, resp.StatusCode)
478 // HTTP Connect end
479 
480 write := make([]byte, testBinaryLength)
481 _, err = io.ReadFull(cRand.Reader, write)
482 as.NoError(err)
483 
484 _, err = tcpConn.Write(write)
485 as.NoError(err)
486 
487 as.NoError(tcpConn.SetReadDeadline(time.Now().Add(time.Second)))
488 read := make([]byte, testBinaryLength)
489 n, err := tcpConn.Read(read)
490 as.NoError(err)
491 as.Equal(testBinaryLength, n)
492 as.EqualValues(write, read)
493 })
494 })
495 }
496}