Skip to content
File

Blob: tun/client/tunnel_test.go

go283 lines
1package client
2 
3import (
4 "fmt"
5 "io"
6 "net"
7 "net/http"
8 "net/http/httptest"
9 "os"
10 "testing"
11 
12 "go.miragespace.co/specter/spec/chord"
13 "go.miragespace.co/specter/spec/mocks"
14 "go.miragespace.co/specter/spec/pki"
15 "go.miragespace.co/specter/spec/protocol"
16 "go.miragespace.co/specter/spec/rpc"
17 
18 "github.com/stretchr/testify/mock"
19 "github.com/stretchr/testify/require"
20 "github.com/zhangyunhao116/skipmap"
21 "go.uber.org/zap/zaptest"
22)
23 
24func TestPublishHostnameReuse(t *testing.T) {
25 as := require.New(t)
26 logger := zaptest.NewLogger(t)
27 
28 file, err := os.CreateTemp("", "client")
29 as.NoError(err)
30 defer os.Remove(file.Name())
31 
32 ctx := t.Context()
33 
34 token := &protocol.ClientToken{
35 Token: []byte("test"),
36 }
37 cl := &protocol.Node{
38 Id: chord.Random(),
39 }
40 
41 der, cert, key := makeCertificate(as, logger, cl, token, nil)
42 cfg := &Config{
43 path: file.Name(),
44 router: skipmap.NewString[route](),
45 Apex: testApex,
46 Certificate: cert,
47 PrivKey: key,
48 Tunnels: []Tunnel{
49 {
50 Target: "tcp://127.0.0.1:5432",
51 },
52 {
53 Target: "tcp://127.0.0.1:2345",
54 },
55 {
56 Target: "https://example.com",
57 },
58 },
59 }
60 as.NoError(cfg.validate())
61 
62 m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) {
63 resp := &protocol.RegisteredHostnamesResponse{
64 Hostnames: []string{testHostname, "bastion.example.com"},
65 }
66 s.On("RegisteredHostnames", mock.Anything, mock.Anything).Return(resp, nil)
67 transportHelper(t1, der)
68 }
69 
70 client, _, assertion := setupClient(t, as, ctx, logger, nil, cfg, nil, m, true, len(cfg.Tunnels))
71 defer assertion()
72 defer client.Close()
73 
74 // assert that we are not reusing the custom domain
75 for _, tunnel := range client.Configuration.Tunnels {
76 as.Equal(testHostname, tunnel.Hostname)
77 }
78}
79 
80func TestRegisterAndProxy(t *testing.T) {
81 as := require.New(t)
82 logger := zaptest.NewLogger(t)
83 
84 file, err := os.CreateTemp("", "client")
85 as.NoError(err)
86 defer os.Remove(file.Name())
87 
88 ctx := t.Context()
89 
90 tcpListener, err := net.Listen("tcp", "127.0.0.1:0")
91 as.NoError(err)
92 defer tcpListener.Close()
93 
94 go func() {
95 conn, err := tcpListener.Accept()
96 as.NoError(err)
97 conn.Write([]byte("hi"))
98 }()
99 
100 token := &protocol.ClientToken{
101 Token: []byte("test"),
102 }
103 cl := &protocol.Node{
104 Id: chord.Random(),
105 }
106 
107 // no token
108 cfg := &Config{
109 path: file.Name(),
110 router: skipmap.NewString[route](),
111 Apex: testApex,
112 Tunnels: []Tunnel{
113 {
114 Target: fmt.Sprintf("tcp://%s", tcpListener.Addr()),
115 },
116 },
117 }
118 as.NoError(cfg.validate())
119 
120 key, err := pki.UnmarshalPrivateKey([]byte(cfg.PrivKey))
121 as.NoError(err)
122 der, cert, _ := makeCertificate(as, logger, cl, token, key)
123 pkiClient := new(mocks.PKIClient)
124 pkiClient.On("RequestCertificate", mock.Anything, mock.Anything).Return(&protocol.CertificateResponse{
125 CertDer: der,
126 CertPem: []byte(cert),
127 }, nil).Once()
128 
129 m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) {
130 s.On("RegisterIdentity", mock.Anything, mock.Anything).Return(&protocol.RegisterIdentityResponse{
131 Apex: testApex,
132 }, nil)
133 
134 defaultNoHostnames(s)
135 transportHelper(t1, der)
136 
137 t1.Identify = cl
138 }
139 
140 client, t2, assertion := setupClient(t, as, ctx, logger, pkiClient, cfg, nil, m, false, 1)
141 defer assertion()
142 defer client.Close()
143 
144 client.Start(ctx)
145 
146 conn, err := t2.DialStream(ctx, cl, protocol.Stream_DIRECT)
147 as.NoError(err)
148 as.NoError(rpc.Send(conn, &protocol.Link{
149 Alpn: protocol.Link_TCP,
150 Hostname: testHostname,
151 }))
152 
153 status := &protocol.TunnelStatus{}
154 rpc.BoundedReceive(conn, status, 1024)
155 as.Equal(protocol.TunnelStatusCode_STATUS_OK, status.GetStatus())
156 
157 buf := make([]byte, 2)
158 n, err := io.ReadFull(conn, buf)
159 as.NoError(err)
160 as.Equal(2, n)
161 as.Equal("hi", string(buf))
162}
163 
164func TestUnpublishTunnel(t *testing.T) {
165 as := require.New(t)
166 logger := zaptest.NewLogger(t)
167 
168 file, err := os.CreateTemp("", "client")
169 as.NoError(err)
170 defer os.Remove(file.Name())
171 
172 ctx := t.Context()
173 
174 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
175 fmt.Fprintf(w, "ok")
176 }))
177 defer ts.Close()
178 
179 token := &protocol.ClientToken{
180 Token: []byte("test"),
181 }
182 cl := &protocol.Node{
183 Id: chord.Random(),
184 }
185 
186 der, cert, key := makeCertificate(as, logger, cl, token, nil)
187 cfg := &Config{
188 path: file.Name(),
189 router: skipmap.NewString[route](),
190 Apex: testApex,
191 Certificate: cert,
192 PrivKey: key,
193 Tunnels: []Tunnel{
194 {
195 Hostname: testHostname,
196 Target: ts.URL,
197 },
198 },
199 }
200 as.NoError(cfg.validate())
201 
202 m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) {
203 s.On("UnpublishTunnel", mock.Anything, mock.MatchedBy(func(req *protocol.UnpublishTunnelRequest) bool {
204 return req.GetHostname() == testHostname
205 })).Return(&protocol.UnpublishTunnelResponse{}, nil).NotBefore(publishCall)
206 
207 defaultNoHostnames(s)
208 transportHelper(t1, der)
209 }
210 
211 client, _, assertion := setupClient(t, as, ctx, logger, nil, cfg, nil, m, false, 1)
212 defer assertion()
213 defer client.Close()
214 
215 client.Start(ctx)
216 
217 err = client.UnpublishTunnel(ctx, Tunnel{Hostname: testHostname})
218 as.NoError(err)
219 
220 curr := client.GetCurrentConfig()
221 as.Len(curr.Tunnels, 0)
222}
223 
224func TestReleaseTunnel(t *testing.T) {
225 as := require.New(t)
226 logger := zaptest.NewLogger(t)
227 
228 file, err := os.CreateTemp("", "client")
229 as.NoError(err)
230 defer os.Remove(file.Name())
231 
232 ctx := t.Context()
233 
234 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
235 fmt.Fprintf(w, "ok")
236 }))
237 defer ts.Close()
238 
239 token := &protocol.ClientToken{
240 Token: []byte("test"),
241 }
242 cl := &protocol.Node{
243 Id: chord.Random(),
244 }
245 
246 der, cert, key := makeCertificate(as, logger, cl, token, nil)
247 cfg := &Config{
248 path: file.Name(),
249 router: skipmap.NewString[route](),
250 Apex: testApex,
251 Certificate: cert,
252 PrivKey: key,
253 Tunnels: []Tunnel{
254 {
255 Hostname: testHostname,
256 Target: ts.URL,
257 },
258 },
259 }
260 as.NoError(cfg.validate())
261 
262 m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) {
263 s.On("ReleaseTunnel", mock.Anything, mock.MatchedBy(func(req *protocol.ReleaseTunnelRequest) bool {
264 return req.GetHostname() == testHostname
265 })).Return(&protocol.ReleaseTunnelResponse{}, nil).NotBefore(publishCall)
266 
267 defaultNoHostnames(s)
268 transportHelper(t1, der)
269 }
270 
271 client, _, assertion := setupClient(t, as, ctx, logger, nil, cfg, nil, m, false, 1)
272 defer assertion()
273 defer client.Close()
274 
275 client.Start(ctx)
276 
277 err = client.ReleaseTunnel(ctx, Tunnel{Hostname: testHostname})
278 as.NoError(err)
279 
280 curr := client.GetCurrentConfig()
281 as.Len(curr.Tunnels, 0)
282}