Skip to content
File

Blob: tun/server/handler_test.go

go487 lines
1package server
2 
3import (
4 "context"
5 "encoding/json"
6 "errors"
7 "net"
8 "net/http"
9 "net/http/httptest"
10 "net/url"
11 "strings"
12 "testing"
13 "testing/synctest"
14 "time"
15 
16 "go.miragespace.co/specter/spec/chord"
17 "go.miragespace.co/specter/spec/mocks"
18 "go.miragespace.co/specter/spec/protocol"
19 "go.miragespace.co/specter/spec/transport"
20 "go.miragespace.co/specter/spec/tun"
21 
22 "github.com/go-chi/chi/v5"
23 "github.com/stretchr/testify/mock"
24 "github.com/stretchr/testify/require"
25)
26 
27type handlerVNode struct {
28 chord.VNode
29 prefixList func(context.Context, []byte) ([][]byte, error)
30}
31 
32func (n handlerVNode) PrefixList(ctx context.Context, prefix []byte) ([][]byte, error) {
33 return n.prefixList(ctx, prefix)
34}
35 
36type handlerTransport struct {
37 transport.Transport
38 dialStream func(context.Context, *protocol.Node, protocol.Stream_Type) (net.Conn, error)
39}
40 
41func (tr handlerTransport) DialStream(ctx context.Context, peer *protocol.Node, kind protocol.Stream_Type) (net.Conn, error) {
42 return tr.dialStream(ctx, peer, kind)
43}
44 
45type handlerQueryService func(context.Context, *protocol.ListTunnelsRequest) (*protocol.ListTunnelsResponse, error)
46 
47func (f handlerQueryService) ListTunnels(ctx context.Context, req *protocol.ListTunnelsRequest) (*protocol.ListTunnelsResponse, error) {
48 return f(ctx, req)
49}
50 
51func handlerRouter(node chord.VNode, tr transport.Transport) http.Handler {
52 serv := &Server{Config: Config{Chord: node, TunnelTransport: tr}}
53 router := chi.NewRouter()
54 router.Mount("/clients", TunnelServerHandler(serv))
55 return router
56}
57 
58func handlerClientTransport(t *testing.T, service handlerQueryService, onDial func(context.Context)) transport.Transport {
59 t.Helper()
60 client := httptest.NewServer(protocol.NewClientQueryServiceServer(service))
61 t.Cleanup(client.Close)
62 return handlerTransport{dialStream: func(ctx context.Context, peer *protocol.Node, kind protocol.Stream_Type) (net.Conn, error) {
63 if peer.GetId() != 111111 || peer.GetAddress() != "fake-address" || !peer.GetRendezvous() || kind != protocol.Stream_RPC {
64 t.Errorf("unexpected client RPC destination: %v, stream %v", peer, kind)
65 }
66 if onDial != nil {
67 onDial(ctx)
68 }
69 return (&net.Dialer{}).DialContext(ctx, "tcp", client.Listener.Addr().String())
70 }}
71}
72 
73func handlerTunnelInfo(t *testing.T, w *httptest.ResponseRecorder) tunnelsInfo {
74 t.Helper()
75 require.Equal(t, http.StatusOK, w.Code)
76 require.Equal(t, "application/json; charset=utf-8", w.Header().Get("Content-Type"))
77 require.Equal(t, "no-store", w.Header().Get("Cache-Control"))
78 var info tunnelsInfo
79 require.NoError(t, json.Unmarshal(w.Body.Bytes(), &info))
80 require.NotNil(t, info.Tunnels, "empty tunnel results must be arrays")
81 return info
82}
83 
84func handlerTunnelRows(t *testing.T, info tunnelsInfo) map[string]clientTunnel {
85 t.Helper()
86 rows := make(map[string]clientTunnel, len(info.Tunnels))
87 for i, tunnel := range info.Tunnels {
88 if i > 0 {
89 require.Less(t, info.Tunnels[i-1].Hostname, tunnel.Hostname, "tunnels must be sorted and unique")
90 }
91 rows[tunnel.Hostname] = tunnel
92 }
93 return rows
94}
95 
96func TestHandlerListConnectedClients(t *testing.T) {
97 node := new(mocks.VNode)
98 clientT := new(mocks.Transport)
99 clientT.On("ListConnected").Return([]transport.ConnectedPeer{
100 {
101 Identity: &protocol.Node{Id: 222222, Address: "second-client"},
102 Addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.2"), Port: 4200},
103 Version: "v2.0.0",
104 },
105 {
106 Identity: &protocol.Node{Id: 111111, Address: "fake-address"},
107 Addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 4200},
108 Version: "v1.0.0",
109 },
110 }).Once()
111 clientT.On("Identity").Return(&protocol.Node{Address: "local-node"}).Once()
112 
113 w := httptest.NewRecorder()
114 handlerRouter(node, clientT).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/clients", nil))
115 
116 require.Equal(t, http.StatusOK, w.Code)
117 require.Equal(t, "application/json; charset=utf-8", w.Header().Get("Content-Type"))
118 require.Equal(t, "no-store", w.Header().Get("Cache-Control"))
119 var body map[string]json.RawMessage
120 require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
121 var observedAt string
122 require.NoError(t, json.Unmarshal(body["observedAt"], &observedAt))
123 _, err := time.Parse(time.RFC3339, observedAt)
124 require.NoError(t, err)
125 delete(body, "observedAt")
126 payload, err := json.Marshal(body)
127 require.NoError(t, err)
128 require.JSONEq(t, `{
129 "node": "local-node",
130 "clients": [
131 {"clientId":"111111","identity":"111111/fake-address","address":"192.0.2.1:4200","version":"v1.0.0","url":"/_internal/tun/111111/fake-address"},
132 {"clientId":"222222","identity":"222222/second-client","address":"192.0.2.2:4200","version":"v2.0.0","url":"/_internal/tun/222222/second-client"}
133 ]
134 }`, string(payload))
135 require.Empty(t, node.Calls, "the connected-client list must use only the local transport snapshot")
136 node.AssertExpectations(t)
137 clientT.AssertExpectations(t)
138}
139 
140func TestHandlerEmptyConnectedClientsIsArray(t *testing.T) {
141 clientT := new(mocks.Transport)
142 clientT.On("ListConnected").Return([]transport.ConnectedPeer(nil)).Once()
143 clientT.On("Identity").Return(&protocol.Node{Address: "local-node"}).Once()
144 w := httptest.NewRecorder()
145 handlerRouter(new(mocks.VNode), clientT).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/clients", nil))
146 var body map[string]json.RawMessage
147 require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
148 require.JSONEq(t, `[]`, string(body["clients"]))
149 clientT.AssertExpectations(t)
150}
151 
152func TestHandlerClientIdentityURLRoundTrip(t *testing.T) {
153 for _, address := range []string{`client/<script>alert("address")</script>?x=1&y=2`, "client%2Fescaped"} {
154 t.Run(address, func(t *testing.T) {
155 clientT := new(mocks.Transport)
156 clientT.On("ListConnected").Return([]transport.ConnectedPeer{{
157 Identity: &protocol.Node{Id: 111111, Address: address},
158 Addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 4200},
159 Version: `<script>alert("version")</script>`,
160 }}).Once()
161 clientT.On("Identity").Return(&protocol.Node{Address: `<script>alert("node")</script>`}).Once()
162 w := httptest.NewRecorder()
163 handlerRouter(new(mocks.VNode), clientT).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/clients", nil))
164 var list connectedInfo
165 require.NoError(t, json.Unmarshal(w.Body.Bytes(), &list))
166 require.Len(t, list.Clients, 1)
167 require.Equal(t, `<script>alert("node")</script>`, list.Node)
168 require.Equal(t, `<script>alert("version")</script>`, list.Clients[0].Version)
169 require.Equal(t, "111111/"+address, list.Clients[0].Identity)
170 require.Equal(t, "/_internal/tun/111111/"+url.PathEscape(address), list.Clients[0].URL)
171 require.NotContains(t, w.Body.String(), "<script>")
172 
173 // Follow the listed browser URL against the corresponding API route.
174 clientT.On("DialStream", mock.Anything, mock.MatchedBy(func(peer *protocol.Node) bool {
175 return peer.GetId() == 111111 && peer.GetAddress() == address && peer.GetRendezvous()
176 }), protocol.Stream_RPC).Return(nil, errors.New("client unavailable")).Once()
177 node := handlerVNode{prefixList: func(_ context.Context, prefix []byte) ([][]byte, error) {
178 expected := tun.ClientHostnamesPrefix(&protocol.ClientToken{Token: []byte(address)})
179 if string(prefix) != expected {
180 t.Errorf("registration prefix = %q, want %q", prefix, expected)
181 }
182 return nil, nil
183 }}
184 w = httptest.NewRecorder()
185 path := strings.Replace(list.Clients[0].URL, "/_internal/tun", "/clients", 1)
186 handlerRouter(node, clientT).ServeHTTP(w, httptest.NewRequest(http.MethodGet, path, nil))
187 info := handlerTunnelInfo(t, w)
188 require.Equal(t, address, info.Address)
189 require.Equal(t, "111111/"+address, info.Identity)
190 clientT.AssertExpectations(t)
191 })
192 }
193}
194 
195func TestHandlerListClientTunnels(t *testing.T) {
196 for _, tc := range []struct {
197 name string
198 configurationErr error
199 registrationErr error
200 want map[string][2]string
201 }{
202 {
203 name: "union of configured and registered hostnames",
204 want: map[string][2]string{
205 "configured-only": {"Yes", "No"},
206 "registered-only": {"No", "Yes"},
207 "shared": {"Yes", "Yes"},
208 },
209 },
210 {
211 name: "client failure preserves registrations",
212 configurationErr: errors.New("client unavailable"),
213 want: map[string][2]string{
214 "registered-only": {"Unknown", "Yes"},
215 "shared": {"Unknown", "Yes"},
216 },
217 },
218 {
219 name: "registration failure preserves client configuration",
220 registrationErr: errors.New("ring unavailable"),
221 want: map[string][2]string{
222 "configured-only": {"Yes", "Unknown"},
223 "shared": {"Yes", "Unknown"},
224 },
225 },
226 {
227 name: "both sources unavailable",
228 configurationErr: errors.New("client unavailable"),
229 registrationErr: errors.New("ring unavailable"),
230 want: map[string][2]string{},
231 },
232 } {
233 t.Run(tc.name, func(t *testing.T) {
234 clientT := handlerClientTransport(t, func(context.Context, *protocol.ListTunnelsRequest) (*protocol.ListTunnelsResponse, error) {
235 return &protocol.ListTunnelsResponse{Tunnels: []*protocol.ClientTunnel{
236 {Hostname: "shared", Target: "http://shared-target"},
237 {Hostname: "configured-only", Target: "http://configured-target"},
238 }}, tc.configurationErr
239 }, nil)
240 node := handlerVNode{prefixList: func(_ context.Context, prefix []byte) ([][]byte, error) {
241 expected := tun.ClientHostnamesPrefix(&protocol.ClientToken{Token: []byte("fake-address")})
242 if string(prefix) != expected {
243 t.Errorf("registration prefix = %q, want %q", prefix, expected)
244 }
245 return [][]byte{[]byte("shared"), []byte("registered-only"), []byte("shared")}, tc.registrationErr
246 }}
247 
248 w := httptest.NewRecorder()
249 handlerRouter(node, clientT).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/clients/111111/fake-address", nil))
250 
251 info := handlerTunnelInfo(t, w)
252 require.Equal(t, "111111/fake-address", info.Identity)
253 require.Equal(t, "fake-address", info.Address)
254 rows := handlerTunnelRows(t, info)
255 require.Len(t, rows, len(tc.want))
256 for hostname, statuses := range tc.want {
257 require.Contains(t, rows, hostname)
258 require.Equal(t, statuses[0], rows[hostname].Configured, hostname)
259 require.Equal(t, statuses[1], rows[hostname].Registered, hostname)
260 }
261 if tc.configurationErr != nil {
262 require.Contains(t, info.ConfigurationError, tc.configurationErr.Error())
263 require.NotContains(t, w.Body.String(), "http://configured-target")
264 } else {
265 require.Empty(t, info.ConfigurationError)
266 require.Equal(t, "http://configured-target", rows["configured-only"].Target)
267 }
268 if tc.registrationErr != nil {
269 require.Equal(t, tc.registrationErr.Error(), info.RegistrationError)
270 } else {
271 require.Empty(t, info.RegistrationError)
272 }
273 })
274 }
275}
276 
277func TestHandlerClientTunnelsLookupsShareDeadlineAndRunConcurrently(t *testing.T) {
278 configurationStarted := make(chan struct{})
279 registrationStarted := make(chan struct{})
280 dialContexts := make(chan context.Context, 1)
281 registrationContexts := make(chan context.Context, 1)
282 clientT := handlerClientTransport(t, func(ctx context.Context, _ *protocol.ListTunnelsRequest) (*protocol.ListTunnelsResponse, error) {
283 close(configurationStarted)
284 select {
285 case <-registrationStarted:
286 return &protocol.ListTunnelsResponse{Tunnels: []*protocol.ClientTunnel{{Hostname: "shared", Target: "http://target"}}}, nil
287 case <-ctx.Done():
288 return nil, ctx.Err()
289 }
290 }, func(ctx context.Context) { dialContexts <- ctx })
291 node := handlerVNode{prefixList: func(ctx context.Context, _ []byte) ([][]byte, error) {
292 registrationContexts <- ctx
293 close(registrationStarted)
294 select {
295 case <-configurationStarted:
296 return [][]byte{[]byte("shared")}, nil
297 case <-ctx.Done():
298 return nil, ctx.Err()
299 }
300 }}
301 
302 start := time.Now()
303 w := httptest.NewRecorder()
304 req := httptest.NewRequest(http.MethodGet, "/clients/111111/fake-address", nil).WithContext(t.Context())
305 handlerRouter(node, clientT).ServeHTTP(w, req)
306 
307 require.Equal(t, http.StatusOK, w.Code)
308 rows := handlerTunnelRows(t, handlerTunnelInfo(t, w))
309 require.Equal(t, "Yes", rows["shared"].Configured)
310 require.Equal(t, "Yes", rows["shared"].Registered)
311 dialDeadline, ok := (<-dialContexts).Deadline()
312 require.True(t, ok, "the tunnel dial needs a bounded request context")
313 registrationDeadline, ok := (<-registrationContexts).Deadline()
314 require.True(t, ok)
315 require.Equal(t, registrationDeadline, dialDeadline)
316 require.WithinDuration(t, start.Add(lookupTimeout), registrationDeadline, time.Second)
317}
318 
319func TestHandlerClientTunnelsCancellation(t *testing.T) {
320 for _, completed := range []bool{false, true} {
321 name := "stalled registration"
322 if completed {
323 name = "completed registration"
324 }
325 t.Run(name, func(t *testing.T) {
326 synctest.Test(t, func(t *testing.T) {
327 release := make(chan struct{})
328 defer close(release)
329 var dialContext, registrationContext context.Context
330 dialStopped := make(chan struct{})
331 clientT := handlerTransport{dialStream: func(ctx context.Context, _ *protocol.Node, _ protocol.Stream_Type) (net.Conn, error) {
332 dialContext = ctx
333 defer close(dialStopped)
334 <-ctx.Done()
335 return nil, ctx.Err()
336 }}
337 node := handlerVNode{prefixList: func(ctx context.Context, _ []byte) ([][]byte, error) {
338 registrationContext = ctx
339 if completed {
340 return [][]byte{[]byte("registered-only")}, nil
341 }
342 <-release // Simulate storage that does not honor cancellation.
343 return nil, ctx.Err()
344 }}
345 ctx, cancel := context.WithCancel(t.Context())
346 defer cancel()
347 w := httptest.NewRecorder()
348 done := make(chan struct{})
349 go func() {
350 defer close(done)
351 req := httptest.NewRequest(http.MethodGet, "/clients/111111/fake-address", nil).WithContext(ctx)
352 handlerRouter(node, clientT).ServeHTTP(w, req)
353 }()
354 synctest.Wait()
355 require.NotNil(t, dialContext)
356 require.NotNil(t, registrationContext)
357 cancel()
358 synctest.Wait()
359 select {
360 case <-done:
361 default:
362 t.Fatal("handler did not finish after cancellation")
363 }
364 select {
365 case <-dialStopped:
366 default:
367 t.Fatal("client dial did not stop after cancellation")
368 }
369 require.ErrorIs(t, dialContext.Err(), context.Canceled)
370 require.ErrorIs(t, registrationContext.Err(), context.Canceled)
371 info := handlerTunnelInfo(t, w)
372 require.Contains(t, info.ConfigurationError, context.Canceled.Error())
373 if completed {
374 require.Empty(t, info.RegistrationError)
375 require.Equal(t, []clientTunnel{{Hostname: "registered-only", Configured: "Unknown", Registered: "Yes"}}, info.Tunnels)
376 } else {
377 require.Equal(t, context.Canceled.Error(), info.RegistrationError)
378 require.Empty(t, info.Tunnels)
379 }
380 })
381 })
382 }
383}
384 
385func TestHandlerConnectedSessionMetadata(t *testing.T) {
386 s, node, clientT, _ := sessionFixture(t)
387 physical := func() *mocks.PhysicalConn {
388 conn := mocks.NewPhysicalConn(nil)
389 t.Cleanup(func() { conn.Close("test cleanup") })
390 return conn
391 }
392 ephemeralConn := physical()
393 tokenConn := physical()
394 replacementConn := physical()
395 owner := "v2:123:public-owner-identity"
396 for i, item := range []struct {
397 conn transport.PhysicalConn
398 mode sessionMode
399 hostname string
400 owner *protocol.ClientToken
401 }{
402 {
403 conn: ephemeralConn,
404 mode: ephemeral,
405 hostname: "test",
406 },
407 {
408 conn: tokenConn,
409 mode: delegated,
410 hostname: "custom.example.com",
411 owner: &protocol.ClientToken{Token: []byte(owner)},
412 },
413 } {
414 sess := &session{
415 alias: tun.SessionAlias([16]byte{byte(i)}),
416 hostname: item.hostname,
417 conn: item.conn,
418 mode: item.mode,
419 owner: item.owner,
420 }
421 require.NoError(t, s.sessions.reserve(sess))
422 require.True(t, sess.activate(t.Context(), time.Now().Add(time.Minute)))
423 }
424 clientT.On("ListConnected").Return([]transport.ConnectedPeer{
425 {
426 Identity: &protocol.Node{
427 Id: 1,
428 Address: "ephemeral-client",
429 },
430 Addr: &net.UDPAddr{Port: 1},
431 Physical: ephemeralConn,
432 },
433 {
434 Identity: &protocol.Node{
435 Id: 2,
436 Address: "token-client",
437 },
438 Addr: &net.UDPAddr{Port: 2},
439 Physical: tokenConn,
440 },
441 {
442 Identity: &protocol.Node{
443 Id: 2,
444 Address: "token-client",
445 },
446 Addr: &net.UDPAddr{Port: 3},
447 Physical: replacementConn,
448 },
449 }).Twice()
450 read := func() connectedInfo {
451 w := httptest.NewRecorder()
452 TunnelServerHandler(s).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/", nil))
453 require.Equal(t, http.StatusOK, w.Code)
454 var info connectedInfo
455 require.NoError(t, json.Unmarshal(w.Body.Bytes(), &info))
456 require.NotContains(t, w.Body.String(), "tg1_")
457 return info
458 }
459 info := read()
460 require.Len(t, info.Clients, 3)
461 require.Equal(t, "ephemeral", info.Clients[0].SessionMode)
462 require.Equal(t, "test."+s.Apex, info.Clients[0].Hostname)
463 require.Equal(t, "1", info.Clients[0].ClientID)
464 require.Empty(t, info.Clients[0].OwnerIdentity)
465 for _, client := range info.Clients[1:] {
466 if client.Address == ":2" {
467 require.Equal(t, "token", client.SessionMode)
468 require.Equal(t, owner, client.OwnerIdentity)
469 require.Equal(t, "123", client.OwnerLabel)
470 require.Empty(t, client.OwnerURL, "an owner absent from the local snapshot has no local link")
471 require.Equal(t, "custom.example.com", client.Hostname)
472 } else {
473 require.Empty(t, client.SessionMode, "a replacement connection must not inherit session metadata")
474 require.Empty(t, client.OwnerIdentity)
475 require.Empty(t, client.Hostname)
476 }
477 }
478 tokenConn.Close("disconnected")
479 for _, client := range read().Clients[1:] {
480 require.Empty(t, client.SessionMode)
481 require.Empty(t, client.OwnerIdentity)
482 require.Empty(t, client.Hostname)
483 }
484 require.Empty(t, node.Calls, "session metadata must not query persistent registration")
485 clientT.AssertExpectations(t)
486}