Skip to content
File

Blob: gateway/gateway_test.go

go770 lines
1package gateway
2 
3import (
4 "context"
5 "crypto/rand"
6 "crypto/rsa"
7 "crypto/tls"
8 "crypto/x509"
9 "encoding/pem"
10 "fmt"
11 "io"
12 "math/big"
13 "net"
14 "net/http"
15 "os"
16 "testing"
17 "time"
18 
19 "go.miragespace.co/specter/overlay"
20 "go.miragespace.co/specter/spec/chord"
21 "go.miragespace.co/specter/spec/cipher"
22 mocks "go.miragespace.co/specter/spec/mocks"
23 "go.miragespace.co/specter/spec/protocol"
24 "go.miragespace.co/specter/spec/rpc"
25 "go.miragespace.co/specter/spec/tun"
26 "go.miragespace.co/specter/util/bufconn"
27 
28 "github.com/go-chi/chi/v5"
29 "github.com/libp2p/go-yamux/v4"
30 "github.com/quic-go/quic-go"
31 "github.com/quic-go/quic-go/http3"
32 "github.com/stretchr/testify/mock"
33 "github.com/stretchr/testify/require"
34 "go.uber.org/zap"
35 "go.uber.org/zap/zaptest"
36 "golang.org/x/net/http2"
37 "golang.org/x/net/http2/h2c"
38)
39 
40const (
41 testDomain = "a.b.c.d.com"
42)
43 
44func generateTLSConfig(protos []string) *tls.Config {
45 key, err := rsa.GenerateKey(rand.Reader, 1024)
46 if err != nil {
47 panic(err)
48 }
49 template := x509.Certificate{
50 SerialNumber: big.NewInt(1),
51 }
52 certDER, err := x509.CreateCertificate(rand.Reader, &template, &template, &key.PublicKey, key)
53 if err != nil {
54 panic(err)
55 }
56 keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)})
57 certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
58 
59 tlsCert, err := tls.X509KeyPair(certPEM, keyPEM)
60 if err != nil {
61 panic(err)
62 }
63 return &tls.Config{
64 InsecureSkipVerify: true,
65 Certificates: []tls.Certificate{tlsCert},
66 NextProtos: protos,
67 }
68}
69 
70func getTCPListener(as *require.Assertions) (net.Listener, int) {
71 l, err := net.Listen("tcp", "127.0.0.1:0")
72 as.NoError(err)
73 
74 return l, l.Addr().(*net.TCPAddr).Port
75}
76 
77func getUDPListener(as *require.Assertions) (net.PacketConn, int) {
78 l, err := net.ListenPacket("udp", "127.0.0.1:0")
79 as.NoError(err)
80 
81 return l, l.LocalAddr().(*net.UDPAddr).Port
82}
83 
84func getH2Listener(as *require.Assertions) (net.Listener, int) {
85 l, err := tls.Listen("tcp", "127.0.0.1:0", generateTLSConfig([]string{
86 tun.ALPN(protocol.Link_HTTP2),
87 tun.ALPN(protocol.Link_HTTP),
88 tun.ALPN(protocol.Link_TCP),
89 tun.ALPN(protocol.Link_UNKNOWN),
90 }))
91 as.NoError(err)
92 
93 return l, l.Addr().(*net.TCPAddr).Port
94}
95 
96func getDialer(proto string, sn string) *tls.Dialer {
97 sni := testDomain
98 if sn != "" {
99 sni = sn + "." + testDomain
100 }
101 dialer := &tls.Dialer{
102 Config: &tls.Config{
103 ServerName: sni,
104 InsecureSkipVerify: true,
105 },
106 }
107 if proto != "" {
108 dialer.Config.NextProtos = []string{proto}
109 }
110 return dialer
111}
112 
113func getQuicDialer(proto string, sn string) func(context.Context, string) (*quic.Conn, error) {
114 return func(ctx context.Context, addr string) (*quic.Conn, error) {
115 return quic.DialAddr(ctx, addr, getDialer(proto, sn).Config, nil)
116 }
117}
118 
119func getH1Client(host string, port int) *http.Client {
120 return &http.Client{
121 Timeout: time.Second,
122 Transport: &http.Transport{
123 ForceAttemptHTTP2: false,
124 DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
125 return getDialer("http/1.1", host).DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", port))
126 },
127 },
128 }
129}
130 
131func getHTTPClient(port int) *http.Client {
132 dialer := &net.Dialer{}
133 return &http.Client{
134 Timeout: time.Second,
135 Transport: &http.Transport{
136 DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
137 return dialer.DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", port))
138 },
139 },
140 }
141}
142 
143func getH2Client(host string, port int) *http.Client {
144 return &http.Client{
145 Timeout: time.Second,
146 Transport: &http.Transport{
147 ForceAttemptHTTP2: true,
148 DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
149 return getDialer(tun.ALPN(protocol.Link_HTTP2), host).DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", port))
150 },
151 },
152 }
153}
154 
155func getH3Client(host string, port int) *http.Client {
156 return &http.Client{
157 Timeout: time.Second,
158 Transport: &http3.Transport{
159 TLSClientConfig: getDialer("h3", host).Config,
160 Dial: func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
161 return quic.DialAddr(ctx, fmt.Sprintf("127.0.0.1:%d", port), tlsCfg, cfg)
162 },
163 },
164 }
165}
166 
167type configOption func(*GatewayConfig)
168 
169func withExtraRootDomains(domains []string) configOption {
170 return func(gc *GatewayConfig) {
171 gc.RootDomains = append(gc.RootDomains, domains...)
172 }
173}
174 
175func setupGateway(t *testing.T, as *require.Assertions, httpListener net.Listener, options ...configOption) (udpPort int, tcpPort int, mockS *mocks.TunnelServer, done func()) {
176 logger := zaptest.NewLogger(t, zaptest.WrapOptions(zap.AddCaller()))
177 
178 var q net.PacketConn
179 var h2 net.Listener
180 
181 q, udpPort = getUDPListener(as)
182 
183 h2, tcpPort = getH2Listener(as)
184 
185 ss := generateTLSConfig([]string{})
186 alpnMux, err := overlay.NewMux(&quic.Transport{Conn: q})
187 as.NoError(err)
188 
189 h3 := alpnMux.With(cipher.GetGatewayTLSConfig(func(chi *tls.ClientHelloInfo) (*tls.Certificate, error) {
190 return &ss.Certificates[0], nil
191 }, nil), append(cipher.H3Protos, tun.ALPN(protocol.Link_TCP))...)
192 
193 mockS = new(mocks.TunnelServer)
194 
195 fakeStats := chi.NewRouter()
196 fakeStats.Get("/stats", func(w http.ResponseWriter, r *http.Request) {
197 w.WriteHeader(http.StatusOK)
198 })
199 
200 conf := GatewayConfig{
201 Logger: logger,
202 TunnelServer: mockS,
203 HTTPListener: httpListener,
204 H2Listener: h2,
205 H3Listener: h3,
206 RootDomains: []string{testDomain},
207 GatewayPort: udpPort,
208 AdminUser: os.Getenv("INTERNAL_USER"),
209 AdminPass: os.Getenv("INTERNAL_PASS"),
210 Handlers: InternalHandlers{
211 Chord: fakeStats,
212 },
213 Options: Options{
214 TransportBufferSize: 1024 * 8,
215 ProxyBufferSize: 1024 * 8,
216 },
217 }
218 
219 for _, option := range options {
220 option(&conf)
221 }
222 
223 g := New(conf)
224 
225 ctx, cancel := context.WithCancel(context.Background())
226 go alpnMux.Accept(ctx)
227 g.MustStart(ctx)
228 
229 return udpPort, tcpPort, mockS, func() {
230 cancel()
231 h2.Close()
232 h3.Close()
233 alpnMux.Close()
234 g.Close()
235 }
236}
237 
238func TestH2HTTPNotFound(t *testing.T) {
239 as := require.New(t)
240 
241 _, tcpPort, mockS, done := setupGateway(t, as, nil)
242 defer done()
243 
244 testHost := "hello"
245 
246 mockS.On("Identity").Return(&protocol.Node{
247 Id: chord.Random(),
248 Address: "127.0.0.1:1234",
249 })
250 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
251 return l.GetAlpn() == protocol.Link_HTTP && l.GetHostname() == testHost
252 })).Return(nil, tun.ErrDestinationNotFound)
253 
254 c := getH2Client(testHost, tcpPort)
255 
256 resp, err := c.Get(fmt.Sprintf("https://%s.%s", testHost, testDomain))
257 as.NoError(err)
258 defer resp.Body.Close()
259 
260 b, err := io.ReadAll(resp.Body)
261 as.NoError(err)
262 
263 as.Contains(string(b), "not found")
264 as.NotEmpty(resp.Header.Get("alt-svc"))
265 as.Equal("false", resp.Header.Get("http3"))
266 
267 mockS.AssertExpectations(t)
268}
269 
270func TestH3HTTPNotFound(t *testing.T) {
271 as := require.New(t)
272 
273 udpPort, _, mockS, done := setupGateway(t, as, nil)
274 defer done()
275 
276 testHost := "hello"
277 
278 mockS.On("Identity").Return(&protocol.Node{
279 Id: chord.Random(),
280 Address: "127.0.0.1:1234",
281 })
282 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
283 return l.GetAlpn() == protocol.Link_HTTP && l.GetHostname() == testHost
284 })).Return(nil, tun.ErrDestinationNotFound)
285 
286 c := getH3Client(testHost, udpPort)
287 
288 resp, err := c.Get(fmt.Sprintf("https://%s.%s", testHost, testDomain))
289 as.NoError(err)
290 defer resp.Body.Close()
291 
292 b, err := io.ReadAll(resp.Body)
293 as.NoError(err)
294 
295 as.Contains(string(b), "not found")
296 as.NotEmpty(resp.Header.Get("alt-svc"))
297 as.Equal("true", resp.Header.Get("http3"))
298 
299 mockS.AssertExpectations(t)
300}
301 
302func TestHTTPNotConnected(t *testing.T) {
303 as := require.New(t)
304 
305 udpPort, _, mockS, done := setupGateway(t, as, nil)
306 defer done()
307 
308 testHost := "hello"
309 
310 mockS.On("Identity").Return(&protocol.Node{
311 Id: chord.Random(),
312 Address: "127.0.0.1:1234",
313 })
314 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
315 return l.GetAlpn() == protocol.Link_HTTP && l.GetHostname() == testHost
316 })).Return(nil, tun.ErrTunnelClientNotConnected)
317 
318 c := getH3Client(testHost, udpPort)
319 
320 resp, err := c.Get(fmt.Sprintf("https://%s.%s", testHost, testDomain))
321 as.NoError(err)
322 defer resp.Body.Close()
323 
324 b, err := io.ReadAll(resp.Body)
325 as.NoError(err)
326 
327 as.Contains(string(b), "not connected")
328 as.NotEmpty(resp.Header.Get("alt-svc"))
329 as.Equal("true", resp.Header.Get("http3"))
330 
331 mockS.AssertExpectations(t)
332}
333 
334type miniClient struct {
335 c chan net.Conn
336 net.Listener
337}
338 
339func (b *miniClient) Accept() (net.Conn, error) {
340 c := <-b.c
341 if c == nil {
342 return nil, net.ErrClosed
343 }
344 return c, nil
345}
346 
347func (b *miniClient) Close() error {
348 return nil
349}
350 
351func serveMiniClient(as *require.Assertions, ch chan net.Conn, resp string) {
352 h := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
353 as.Equal("https", r.Header.Get("X-Forwarded-Proto"))
354 as.NotEmpty(r.Header.Get("X-Forwarded-Host"))
355 w.Write([]byte(resp))
356 })
357 h2s := &http2.Server{}
358 h1 := &http.Server{
359 Handler: h2c.NewHandler(h, h2s),
360 }
361 h1.Serve(&miniClient{c: ch})
362}
363 
364func TestRejectInvalidHostnames(t *testing.T) {
365 as := require.New(t)
366 logger := zaptest.NewLogger(t, zaptest.WrapOptions(zap.AddCaller()))
367 
368 conf := GatewayConfig{
369 Logger: logger,
370 RootDomains: []string{testDomain},
371 AdminUser: os.Getenv("INTERNAL_USER"),
372 AdminPass: os.Getenv("INTERNAL_PASS"),
373 }
374 g := New(conf)
375 
376 hostnames := []string{
377 "bleh",
378 "bleh.com",
379 "192.168.1.1",
380 }
381 
382 for _, hostname := range hostnames {
383 _, _, err := g.parseAddr(hostname)
384 as.Error(err)
385 }
386}
387 
388func TestH1HTTPFound(t *testing.T) {
389 as := require.New(t)
390 
391 _, tcpPort, mockS, done := setupGateway(t, as, nil)
392 defer done()
393 
394 testHost := "hello"
395 testResponse := "this is fine"
396 
397 c1, c2 := bufconn.BufferedPipe(8192)
398 ch := make(chan net.Conn, 1)
399 go serveMiniClient(as, ch, testResponse)
400 defer close(ch)
401 ch <- c2
402 
403 mockS.On("Identity").Return(&protocol.Node{
404 Id: chord.Random(),
405 Address: "127.0.0.1:1234",
406 })
407 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
408 return l.GetAlpn() == protocol.Link_HTTP && l.GetHostname() == testHost
409 })).Return(c1, nil)
410 
411 c := getH1Client(testHost, tcpPort)
412 
413 req, err := http.NewRequest("GET", fmt.Sprintf("https://%s.%s", testHost, testDomain), nil)
414 as.NoError(err)
415 
416 req.Host = "fail.com"
417 
418 resp, err := c.Do(req)
419 as.NoError(err)
420 defer resp.Body.Close()
421 
422 b, err := io.ReadAll(resp.Body)
423 as.NoError(err)
424 
425 as.Contains(string(b), testResponse)
426 as.NotEmpty(resp.Header.Get("alt-svc"))
427 as.Equal("false", resp.Header.Get("http3"))
428 
429 mockS.AssertExpectations(t)
430}
431 
432func TestH2HTTPFound(t *testing.T) {
433 as := require.New(t)
434 
435 _, tcpPort, mockS, done := setupGateway(t, as, nil)
436 defer done()
437 
438 testHost := "hello"
439 testResponse := "this is fine"
440 
441 c1, c2 := bufconn.BufferedPipe(8192)
442 ch := make(chan net.Conn, 1)
443 go serveMiniClient(as, ch, testResponse)
444 defer close(ch)
445 ch <- c2
446 
447 mockS.On("Identity").Return(&protocol.Node{
448 Id: chord.Random(),
449 Address: "127.0.0.1:1234",
450 })
451 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
452 return l.GetAlpn() == protocol.Link_HTTP && l.GetHostname() == testHost
453 })).Return(c1, nil)
454 
455 c := getH2Client(testHost, tcpPort)
456 
457 req, err := http.NewRequest("GET", fmt.Sprintf("https://%s.%s", testHost, testDomain), nil)
458 as.NoError(err)
459 
460 resp, err := c.Do(req)
461 as.NoError(err)
462 defer resp.Body.Close()
463 
464 b, err := io.ReadAll(resp.Body)
465 as.NoError(err)
466 
467 as.Contains(string(b), testResponse)
468 as.NotEmpty(resp.Header.Get("alt-svc"))
469 as.Equal("false", resp.Header.Get("http3"))
470 
471 mockS.AssertExpectations(t)
472}
473 
474func TestH3HTTPFound(t *testing.T) {
475 as := require.New(t)
476 
477 udpPort, _, mockS, done := setupGateway(t, as, nil)
478 defer done()
479 
480 testHost := "hello"
481 testResponse := "this is fine from h3"
482 
483 c1, c2 := bufconn.BufferedPipe(8192)
484 ch := make(chan net.Conn, 1)
485 go serveMiniClient(as, ch, testResponse)
486 defer close(ch)
487 ch <- c2
488 
489 mockS.On("Identity").Return(&protocol.Node{
490 Id: chord.Random(),
491 Address: "127.0.0.1:1234",
492 })
493 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
494 return l.GetAlpn() == protocol.Link_HTTP && l.GetHostname() == testHost
495 })).Return(c1, nil)
496 
497 c := getH3Client(testHost, udpPort)
498 
499 req, err := http.NewRequest("GET", fmt.Sprintf("https://%s.%s", testHost, testDomain), nil)
500 as.NoError(err)
501 
502 resp, err := c.Do(req)
503 as.NoError(err)
504 defer resp.Body.Close()
505 
506 b, err := io.ReadAll(resp.Body)
507 as.NoError(err)
508 
509 as.Contains(string(b), testResponse)
510 as.NotEmpty(resp.Header.Get("alt-svc"))
511 as.Equal("true", resp.Header.Get("http3"))
512 
513 mockS.AssertExpectations(t)
514}
515 
516func TestH2TCPNotFound(t *testing.T) {
517 as := require.New(t)
518 
519 _, tcpPort, mockS, done := setupGateway(t, as, nil)
520 defer done()
521 
522 testHost := "hello"
523 
524 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
525 return l.GetAlpn() == protocol.Link_TCP && l.GetHostname() == testHost
526 })).Return(nil, tun.ErrDestinationNotFound)
527 
528 dialer := getDialer(tun.ALPN(protocol.Link_TCP), testHost)
529 conn, err := dialer.DialContext(context.Background(), "tcp", fmt.Sprintf("127.0.0.1:%d", tcpPort))
530 as.NoError(err)
531 
532 cfg := yamux.DefaultConfig()
533 cfg.LogOutput = io.Discard
534 session, err := yamux.Client(conn, cfg, nil)
535 as.NoError(err)
536 
537 ctx := t.Context()
538 stream, err := session.OpenStream(ctx)
539 as.NoError(err)
540 
541 status := &protocol.TunnelStatus{}
542 err = rpc.Send(stream, status)
543 as.NoError(err)
544 err = rpc.BoundedReceive(stream, status, 1024)
545 as.NoError(err)
546 as.NotEqual(protocol.TunnelStatusCode_STATUS_OK, status.GetStatus())
547 
548 <-time.After(time.Millisecond * 100)
549 
550 mockS.AssertExpectations(t)
551}
552 
553func TestH3TCPNotFound(t *testing.T) {
554 as := require.New(t)
555 
556 udpPort, _, mockS, done := setupGateway(t, as, nil)
557 defer done()
558 
559 testHost := "hello"
560 
561 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
562 return l.GetAlpn() == protocol.Link_TCP && l.GetHostname() == testHost
563 })).Return(nil, tun.ErrDestinationNotFound)
564 
565 dial := getQuicDialer(tun.ALPN(protocol.Link_TCP), testHost)
566 conn, err := dial(context.Background(), fmt.Sprintf("127.0.0.1:%d", udpPort))
567 as.NoError(err)
568 
569 ctx, cancel := context.WithTimeout(context.Background(), time.Second)
570 defer cancel()
571 b, err := conn.OpenStreamSync(ctx)
572 as.NoError(err)
573 
574 status := &protocol.TunnelStatus{}
575 err = rpc.Send(b, status)
576 as.NoError(err)
577 err = rpc.BoundedReceive(b, status, 1024)
578 as.NoError(err)
579 as.NotEqual(protocol.TunnelStatusCode_STATUS_OK, status.GetStatus())
580 
581 <-time.After(time.Millisecond * 100)
582 
583 mockS.AssertExpectations(t)
584}
585 
586func TestTCPNotConnected(t *testing.T) {
587 as := require.New(t)
588 
589 udpPort, _, mockS, done := setupGateway(t, as, nil)
590 defer done()
591 
592 testHost := "hello"
593 
594 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
595 return l.GetAlpn() == protocol.Link_TCP && l.GetHostname() == testHost
596 })).Return(nil, tun.ErrTunnelClientNotConnected)
597 
598 dial := getQuicDialer(tun.ALPN(protocol.Link_TCP), testHost)
599 conn, err := dial(context.Background(), fmt.Sprintf("127.0.0.1:%d", udpPort))
600 as.NoError(err)
601 
602 ctx, cancel := context.WithTimeout(context.Background(), time.Second)
603 defer cancel()
604 b, err := conn.OpenStreamSync(ctx)
605 as.NoError(err)
606 
607 status := &protocol.TunnelStatus{}
608 err = rpc.Send(b, status)
609 as.NoError(err)
610 err = rpc.BoundedReceive(b, status, 1024)
611 as.NoError(err)
612 as.Equal(protocol.TunnelStatusCode_NO_DIRECT, status.GetStatus())
613 
614 <-time.After(time.Millisecond * 100)
615 
616 mockS.AssertExpectations(t)
617}
618 
619func TestH2TCPFound(t *testing.T) {
620 as := require.New(t)
621 
622 _, tcpPort, mockS, done := setupGateway(t, as, nil)
623 defer done()
624 
625 testHost := "hello"
626 bufLength := 20
627 
628 c1, c2 := bufconn.BufferedPipe(8192)
629 
630 go func() {
631 tun.SendStatusProto(c2, nil)
632 
633 buf := make([]byte, bufLength)
634 n, err := io.ReadFull(c2, buf)
635 as.NoError(err)
636 as.Equal(bufLength, n)
637 
638 rand.Read(buf)
639 n, err = c2.Write(buf)
640 as.NoError(err)
641 as.Equal(bufLength, n)
642 }()
643 
644 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
645 return l.GetAlpn() == protocol.Link_TCP && l.GetHostname() == testHost
646 })).Return(c1, nil)
647 
648 dialer := getDialer(tun.ALPN(protocol.Link_TCP), testHost)
649 conn, err := dialer.DialContext(context.Background(), "tcp", fmt.Sprintf("127.0.0.1:%d", tcpPort))
650 as.NoError(err)
651 
652 cfg := yamux.DefaultConfig()
653 cfg.LogOutput = io.Discard
654 session, err := yamux.Client(conn, cfg, nil)
655 as.NoError(err)
656 
657 ctx := t.Context()
658 stream, err := session.OpenStream(ctx)
659 as.NoError(err)
660 
661 status := &protocol.TunnelStatus{}
662 err = rpc.Send(stream, status)
663 as.NoError(err)
664 err = rpc.BoundedReceive(stream, status, 1024)
665 as.NoError(err)
666 as.Equal(protocol.TunnelStatusCode_STATUS_OK, status.GetStatus())
667 
668 buf := make([]byte, bufLength)
669 
670 rand.Read(buf)
671 n, err := stream.Write(buf)
672 as.NoError(err)
673 as.Equal(bufLength, n)
674 
675 n, err = io.ReadFull(stream, buf)
676 as.NoError(err)
677 as.Equal(bufLength, n)
678 
679 mockS.AssertExpectations(t)
680}
681 
682func TestH3TCPFound(t *testing.T) {
683 as := require.New(t)
684 
685 udpPort, _, mockS, done := setupGateway(t, as, nil)
686 defer done()
687 
688 testHost := "hello"
689 bufLength := 20
690 
691 c1, c2 := bufconn.BufferedPipe(8192)
692 
693 go func() {
694 tun.SendStatusProto(c2, nil)
695 
696 buf := make([]byte, bufLength)
697 n, err := io.ReadFull(c2, buf)
698 as.NoError(err)
699 as.Equal(bufLength, n)
700 
701 rand.Read(buf)
702 n, err = c2.Write(buf)
703 as.NoError(err)
704 as.Equal(bufLength, n)
705 }()
706 
707 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
708 return l.GetAlpn() == protocol.Link_TCP && l.GetHostname() == testHost
709 })).Return(c1, nil)
710 
711 dial := getQuicDialer(tun.ALPN(protocol.Link_TCP), testHost)
712 conn, err := dial(context.Background(), fmt.Sprintf("127.0.0.1:%d", udpPort))
713 as.NoError(err)
714 
715 ctx, cancel := context.WithTimeout(context.Background(), time.Second)
716 defer cancel()
717 stream, err := conn.OpenStreamSync(ctx)
718 as.NoError(err)
719 
720 status := &protocol.TunnelStatus{}
721 err = rpc.Send(stream, status)
722 as.NoError(err)
723 err = rpc.BoundedReceive(stream, status, 1024)
724 as.NoError(err)
725 as.Equal(protocol.TunnelStatusCode_STATUS_OK, status.GetStatus())
726 
727 buf := make([]byte, bufLength)
728 
729 rand.Read(buf)
730 n, err := stream.Write(buf)
731 as.NoError(err)
732 as.Equal(bufLength, n)
733 
734 n, err = io.ReadFull(stream, buf)
735 as.NoError(err)
736 as.Equal(bufLength, n)
737 
738 mockS.AssertExpectations(t)
739}
740 
741func TestH2RejectALPN(t *testing.T) {
742 as := require.New(t)
743 
744 _, tcpPort, mockS, done := setupGateway(t, as, nil)
745 defer done()
746 
747 testHost := "hello"
748 
749 dialer := getDialer("h3", testHost)
750 _, err := dialer.DialContext(context.Background(), "tcp", fmt.Sprintf("127.0.0.1:%d", tcpPort))
751 as.Error(err)
752 
753 mockS.AssertExpectations(t)
754}
755 
756func TestH3RejectALPN(t *testing.T) {
757 as := require.New(t)
758 
759 udpPort, _, mockS, done := setupGateway(t, as, nil)
760 defer done()
761 
762 testHost := "hello"
763 
764 dial := getQuicDialer(tun.ALPN(protocol.Link_HTTP2), testHost)
765 _, err := dial(context.Background(), fmt.Sprintf("127.0.0.1:%d", udpPort))
766 as.Error(err)
767 
768 mockS.AssertExpectations(t)
769}