Skip to content
File

Blob: gateway/extra_test.go

go147 lines
1package gateway
2 
3import (
4 "context"
5 "fmt"
6 "io"
7 "net"
8 "net/http"
9 "testing"
10 
11 "go.miragespace.co/specter/spec/chord"
12 "go.miragespace.co/specter/spec/protocol"
13 "go.miragespace.co/specter/spec/tun"
14 
15 "github.com/quic-go/quic-go/http3"
16 "github.com/stretchr/testify/mock"
17 "github.com/stretchr/testify/require"
18)
19 
20const (
21 extraRootDomain = "x.y.z.net"
22)
23 
24func TestH2ExtraApexIndex(t *testing.T) {
25 as := require.New(t)
26 
27 _, tcpPort, mockS, done := setupGateway(t, as, nil, withExtraRootDomains([]string{extraRootDomain}))
28 defer done()
29 
30 c := getH2Client("", tcpPort)
31 dialer := getDialer(tun.ALPN(protocol.Link_HTTP2), "")
32 dialer.Config.ServerName = extraRootDomain
33 c.Transport.(*http.Transport).DialTLSContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
34 return dialer.DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", tcpPort))
35 }
36 
37 resp, err := c.Get(fmt.Sprintf("https://%s/", extraRootDomain))
38 as.NoError(err)
39 defer resp.Body.Close()
40 
41 b, err := io.ReadAll(resp.Body)
42 as.NoError(err)
43 
44 as.Contains(string(b), extraRootDomain)
45 as.NotEmpty(resp.Header.Get("alt-svc"))
46 as.Equal("false", resp.Header.Get("http3"))
47 
48 mockS.AssertExpectations(t)
49}
50 
51func TestH3ExtraApexIndex(t *testing.T) {
52 as := require.New(t)
53 
54 udpPort, _, mockS, done := setupGateway(t, as, nil, withExtraRootDomains([]string{extraRootDomain}))
55 defer done()
56 
57 c := getH3Client("", udpPort)
58 dialer := getDialer("h3", "")
59 dialer.Config.ServerName = extraRootDomain
60 c.Transport.(*http3.Transport).TLSClientConfig = dialer.Config
61 
62 resp, err := c.Get(fmt.Sprintf("https://%s/", extraRootDomain))
63 as.NoError(err)
64 defer resp.Body.Close()
65 
66 b, err := io.ReadAll(resp.Body)
67 as.NoError(err)
68 
69 as.Contains(string(b), extraRootDomain)
70 as.NotEmpty(resp.Header.Get("alt-svc"))
71 as.Equal("true", resp.Header.Get("http3"))
72 
73 mockS.AssertExpectations(t)
74}
75 
76func TestH2HTTPNotFoundExtra(t *testing.T) {
77 as := require.New(t)
78 
79 _, tcpPort, mockS, done := setupGateway(t, as, nil, withExtraRootDomains([]string{extraRootDomain}))
80 defer done()
81 
82 testHost := "hello"
83 
84 mockS.On("Identity").Return(&protocol.Node{
85 Id: chord.Random(),
86 Address: "127.0.0.1:1234",
87 })
88 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
89 return l.GetAlpn() == protocol.Link_HTTP && l.GetHostname() == testHost
90 })).Return(nil, tun.ErrDestinationNotFound)
91 
92 c := getH2Client(testHost, tcpPort)
93 dialer := getDialer(tun.ALPN(protocol.Link_HTTP2), "")
94 dialer.Config.ServerName = testHost + "." + extraRootDomain
95 c.Transport.(*http.Transport).DialTLSContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
96 return dialer.DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", tcpPort))
97 }
98 
99 resp, err := c.Get(fmt.Sprintf("https://%s.%s", testHost, extraRootDomain))
100 as.NoError(err)
101 defer resp.Body.Close()
102 
103 b, err := io.ReadAll(resp.Body)
104 as.NoError(err)
105 
106 as.Contains(string(b), "not found")
107 as.NotEmpty(resp.Header.Get("alt-svc"))
108 as.Equal("false", resp.Header.Get("http3"))
109 
110 mockS.AssertExpectations(t)
111}
112 
113func TestH3HTTPNotFoundExtra(t *testing.T) {
114 as := require.New(t)
115 
116 udpPort, _, mockS, done := setupGateway(t, as, nil, withExtraRootDomains([]string{extraRootDomain}))
117 defer done()
118 
119 testHost := "hello"
120 
121 mockS.On("Identity").Return(&protocol.Node{
122 Id: chord.Random(),
123 Address: "127.0.0.1:1234",
124 })
125 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
126 return l.GetAlpn() == protocol.Link_HTTP && l.GetHostname() == testHost
127 })).Return(nil, tun.ErrDestinationNotFound)
128 
129 c := getH3Client(testHost, udpPort)
130 dialer := getDialer("h3", "")
131 dialer.Config.ServerName = testHost + "." + extraRootDomain
132 c.Transport.(*http3.Transport).TLSClientConfig = dialer.Config
133 
134 resp, err := c.Get(fmt.Sprintf("https://%s.%s", testHost, extraRootDomain))
135 as.NoError(err)
136 defer resp.Body.Close()
137 
138 b, err := io.ReadAll(resp.Body)
139 as.NoError(err)
140 
141 as.Contains(string(b), "not found")
142 as.NotEmpty(resp.Header.Get("alt-svc"))
143 as.Equal("true", resp.Header.Get("http3"))
144 
145 mockS.AssertExpectations(t)
146}