Skip to content
File

Blob: tun/client/lightweight_test.go

go430 lines
1package client
2 
3import (
4 "bufio"
5 "context"
6 "crypto/ed25519"
7 "crypto/rand"
8 "crypto/tls"
9 "crypto/x509"
10 "fmt"
11 "io"
12 "net"
13 "net/http"
14 "net/http/httptest"
15 "os"
16 "strings"
17 "sync"
18 "testing"
19 "time"
20 
21 "go.miragespace.co/specter/spec/mocks"
22 "go.miragespace.co/specter/spec/protocol"
23 "go.miragespace.co/specter/spec/rpc"
24 "go.miragespace.co/specter/spec/transport"
25 "go.miragespace.co/specter/tun/client/dialer"
26 "go.miragespace.co/specter/util/acceptor"
27 
28 "github.com/stretchr/testify/mock"
29 "github.com/stretchr/testify/require"
30 "github.com/twitchtv/twirp"
31 "go.uber.org/zap"
32 "go.uber.org/zap/zaptest/observer"
33)
34 
35type lightweightTransport struct {
36 *mocks.MemoryTransport
37 cert *x509.Certificate
38 attempts chan *mocks.PhysicalConn
39 addresses chan string
40 mu sync.Mutex
41 aliases map[string]string
42 connections map[string]*mocks.PhysicalConn
43}
44 
45func (t *lightweightTransport) DialStream(ctx context.Context, node *protocol.Node, kind protocol.Stream_Type) (net.Conn, error) {
46 address := node.GetAddress()
47 t.mu.Lock()
48 if canonical, ok := t.aliases[address]; ok {
49 node = &protocol.Node{Address: canonical}
50 }
51 pc := t.connections[node.GetAddress()]
52 if pc == nil || pc.Err() != nil {
53 pc = mocks.NewPhysicalConn(func(d *transport.StreamDelegate) {
54 d.Certificate, d.Identity = t.cert, node
55 if d.Kind == protocol.Stream_DIRECT {
56 t.Self <- d
57 } else {
58 t.Other <- d
59 }
60 })
61 if t.connections != nil {
62 t.connections[node.GetAddress()] = pc
63 }
64 }
65 t.mu.Unlock()
66 if t.attempts != nil {
67 t.attempts <- pc
68 }
69 if t.addresses != nil {
70 t.addresses <- address
71 }
72 return pc.OpenStream(kind)
73}
74 
75func lightweightUpstream(t *testing.T, pc *mocks.PhysicalConn, hostname string) string {
76 t.Helper()
77 direct, err := pc.OpenStream(protocol.Stream_DIRECT)
78 require.NoError(t, err)
79 require.NoError(t, rpc.Send(direct, &protocol.Link{
80 Alpn: protocol.Link_HTTP,
81 Hostname: hostname,
82 }))
83 tp := &http.Transport{DialContext: func(context.Context, string, string) (net.Conn, error) { return direct, nil }}
84 defer tp.CloseIdleConnections()
85 hc := &http.Client{
86 Transport: tp,
87 Timeout: time.Second,
88 }
89 resp, err := hc.Get("http://" + hostname + "/")
90 require.NoError(t, err)
91 defer resp.Body.Close()
92 data, err := io.ReadAll(resp.Body)
93 require.NoError(t, err)
94 return string(data)
95}
96 
97func TestLightweightRun(t *testing.T) {
98 dir := t.TempDir()
99 t.Chdir(dir)
100 ctx, cancel := context.WithTimeout(t.Context(), 12*time.Second)
101 defer cancel()
102 core, logs := observer.New(zap.DebugLevel)
103 logger := zap.New(core)
104 target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, "upstream") }))
105 defer target.Close()
106 _, key, err := ed25519.GenerateKey(rand.Reader)
107 require.NoError(t, err)
108 der, _, _ := makeCertificate(require.New(t), logger, &protocol.Node{Id: 1}, &protocol.ClientToken{Token: []byte("test")}, key)
109 cert, err := x509.ParseCertificate(der)
110 require.NoError(t, err)
111 pki := &mocks.PKIClient{}
112 pki.On("RequestCertificate", mock.Anything, mock.Anything).Return(&protocol.CertificateResponse{CertDer: der}, nil).Once()
113 left, right := mocks.PipeTransport()
114 tp := &lightweightTransport{
115 MemoryTransport: left,
116 cert: cert,
117 attempts: make(chan *mocks.PhysicalConn, 4),
118 addresses: make(chan string, 4),
119 }
120 left.On("WithClientCertificate", mock.AnythingOfType("tls.Certificate")).Run(func(args mock.Arguments) { require.Len(t, args.Get(0).(tls.Certificate).Certificate, 1) }).Return(nil).Once()
121 service := &mocks.TunnelService{}
122 home := &protocol.Node{Address: "home.example.com:443"}
123 response := &protocol.OpenSessionResponse{
124 Hostname: "temporary",
125 Node: home,
126 Apex: "example.com",
127 }
128 service.On("OpenEphemeralSession", mock.Anything, mock.Anything).Return(nil, twirp.Unavailable.Error("try again")).Once()
129 opened := make(chan context.Context, 1)
130 release := make(chan struct{})
131 service.On("OpenEphemeralSession", mock.Anything, mock.Anything).Run(func(args mock.Arguments) { opened <- args.Get(0).(context.Context); <-release }).Return(response, nil).Once()
132 service.On("OpenEphemeralSession", mock.Anything, mock.Anything).Return(response, nil).Once()
133 service.On("OpenEphemeralSession", mock.Anything, mock.Anything).Return(nil, twirp.PermissionDenied.Error("denied")).Once()
134 router := transport.NewStreamRouter(logger, nil, right)
135 acc := acceptor.NewH2Acceptor(nil)
136 defer acc.Close()
137 setupRPC(ctx, logger, service, &mocks.KeylessService{}, router, acc)
138 go router.Accept(ctx)
139 outputR, outputW := io.Pipe()
140 defer outputR.Close()
141 l, err := NewLightweightClient(LightweightConfig{
142 Logger: logger,
143 Transport: tp,
144 PKIClient: pki,
145 Apex: &dialer.ParsedApex{
146 Host: "example.com",
147 Port: 443,
148 },
149 Target: target.URL,
150 Output: outputW,
151 })
152 require.NoError(t, err)
153 result := make(chan error, 1)
154 go func() { result <- l.Run(ctx); outputW.Close() }()
155 first := <-tp.attempts
156 require.Equal(t, "example.com:443", <-tp.addresses)
157 require.Eventually(t, func() bool { return first.Err() != nil }, 3*time.Second, time.Millisecond)
158 second := <-tp.attempts
159 require.NotSame(t, first, second)
160 require.Equal(t, "example.com:443", <-tp.addresses)
161 select {
162 case <-opened:
163 case <-ctx.Done():
164 t.Fatal("Open not reached")
165 }
166 require.Equal(t, "upstream", lightweightUpstream(t, second, "temporary"))
167 close(release)
168 printed := make(chan string, 1)
169 go func() { data, _ := io.ReadAll(outputR); printed <- string(data) }()
170 require.Eventually(t, func() bool { return logs.FilterMessage("Tunnel ready").Len() == 1 }, time.Second, time.Millisecond)
171 second.Close("disconnect")
172 third := <-tp.attempts
173 require.Equal(t, home.Address, <-tp.addresses)
174 require.Eventually(t, func() bool { return logs.FilterMessage("Tunnel recovered").Len() == 1 }, time.Second, time.Millisecond)
175 third.Close("disconnect again")
176 fourth := <-tp.attempts
177 require.Equal(t, home.Address, <-tp.addresses)
178 select {
179 case err := <-result:
180 require.True(t, definite(err))
181 require.ErrorContains(t, err, "denied")
182 case <-ctx.Done():
183 t.Fatal("client did not stop")
184 }
185 require.Error(t, fourth.Err())
186 require.Equal(t, "https://temporary.example.com\n", <-printed)
187 entries, err := os.ReadDir(dir)
188 require.NoError(t, err)
189 require.Empty(t, entries)
190 service.AssertExpectations(t)
191 service.AssertNotCalled(t, "Ping", mock.Anything, mock.Anything)
192 pki.AssertExpectations(t)
193 require.NotContains(t, strings.Join([]string{l.URL()}, ""), "tg1_")
194}
195 
196func TestTokenConnections(t *testing.T) {
197 for _, tc := range []struct {
198 name string
199 expiresAt int64
200 expiryLog string
201 revoke bool
202 }{
203 {
204 name: "cancel",
205 expiryLog: "none",
206 },
207 {
208 name: "revoked",
209 expiresAt: time.Date(2030, 2, 3, 4, 5, 6, 0, time.FixedZone("test", -7*60*60)).Unix(),
210 expiryLog: "2030-02-03T11:05:06Z",
211 revoke: true,
212 },
213 } {
214 t.Run(tc.name, func(t *testing.T) {
215 dir := t.TempDir()
216 t.Chdir(dir)
217 ctx, cancel := context.WithTimeout(t.Context(), 8*time.Second)
218 defer cancel()
219 core, logs := observer.New(zap.DebugLevel)
220 logger := zap.New(core)
221 target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { fmt.Fprint(w, "upstream") }))
222 defer target.Close()
223 der, _, _ := makeCertificate(require.New(t), logger, &protocol.Node{Id: 1}, &protocol.ClientToken{Token: []byte("certificate")}, nil)
224 cert, err := x509.ParseCertificate(der)
225 require.NoError(t, err)
226 pki := &mocks.PKIClient{}
227 pki.On("RequestCertificate", mock.Anything, mock.Anything).Return(&protocol.CertificateResponse{CertDer: der}, nil).Once()
228 left, right := mocks.PipeTransport()
229 tp := &lightweightTransport{
230 MemoryTransport: left,
231 cert: cert,
232 addresses: make(chan string, 32),
233 aliases: map[string]string{"alias-a.example.com:443": "a.example.com:443"},
234 connections: make(map[string]*mocks.PhysicalConn),
235 }
236 left.On("WithClientCertificate", mock.AnythingOfType("tls.Certificate")).Return(nil).Once()
237 service := &mocks.TunnelService{}
238 apex := &protocol.Node{Address: "example.com:443"}
239 nodes := []*protocol.Node{
240 {Address: "a.example.com:443"},
241 {Address: "b.example.com:443"},
242 {Address: "c.example.com:443"},
243 }
244 at := func(address string) any {
245 return mock.MatchedBy(func(ctx context.Context) bool {
246 return rpc.GetDelegation(ctx).Identity.GetAddress() == address
247 })
248 }
249 service.On("GetNodes", at(apex.Address), mock.Anything).Return(&protocol.GetNodesResponse{
250 Nodes: []*protocol.Node{
251 nodes[0],
252 {Address: "alias-a.example.com:443"},
253 nodes[1],
254 },
255 }, nil).Once()
256 service.On("GetNodes", mock.Anything, mock.Anything).Return(&protocol.GetNodesResponse{Nodes: nodes}, nil)
257 const token = "tg1_secret-that-must-not-appear"
258 request := func(slot uint32) any {
259 return mock.MatchedBy(func(req *protocol.OpenDelegatedSessionRequest) bool {
260 return req.Token == token && req.RouteSlot == slot
261 })
262 }
263 service.On("OpenDelegatedSession", at(apex.Address), request(1)).Return(nil, twirp.ResourceExhausted.Error("apex is full")).Once()
264 opened := [3]chan *mocks.PhysicalConn{
265 make(chan *mocks.PhysicalConn, 2),
266 make(chan *mocks.PhysicalConn, 2),
267 make(chan *mocks.PhysicalConn, 2),
268 }
269 open := func(slot uint32, gate <-chan struct{}) {
270 service.On("OpenDelegatedSession", at(nodes[slot-1].Address), request(slot)).
271 Run(func(args mock.Arguments) {
272 delegation := rpc.GetDelegation(args.Get(0).(context.Context))
273 pc := delegation.Conn.(transport.PhysicalConnProvider).PhysicalConn().(*mocks.PhysicalConn)
274 opened[slot-1] <- pc
275 if gate != nil {
276 select {
277 case <-gate:
278 case <-ctx.Done():
279 }
280 }
281 }).Return(&protocol.OpenSessionResponse{
282 Hostname: "delegated",
283 Node: nodes[slot-1],
284 Apex: "example.com",
285 GrantId: "grant-for-client-test",
286 ExpiresAt: tc.expiresAt,
287 }, nil).Once()
288 }
289 allowSecond := make(chan struct{})
290 allowRepair := make(chan struct{})
291 open(1, nil)
292 open(2, allowSecond)
293 open(3, nil)
294 open(3, allowRepair)
295 if tc.revoke {
296 service.On("OpenDelegatedSession", at(nodes[0].Address), request(1)).Return(nil, twirp.PermissionDenied.Error("grant revoked")).Once()
297 }
298 router := transport.NewStreamRouter(logger, nil, right)
299 acc := acceptor.NewH2Acceptor(nil)
300 defer acc.Close()
301 setupRPC(ctx, logger, service, &mocks.KeylessService{}, router, acc)
302 go router.Accept(ctx)
303 outputR, outputW := io.Pipe()
304 defer outputR.Close()
305 printed := make(chan string, 8)
306 go func() {
307 defer close(printed)
308 scanner := bufio.NewScanner(outputR)
309 for scanner.Scan() {
310 printed <- scanner.Text()
311 }
312 }()
313 l, err := NewLightweightClient(LightweightConfig{
314 Logger: logger,
315 Transport: tp,
316 PKIClient: pki,
317 Apex: &dialer.ParsedApex{
318 Host: "example.com",
319 Port: 443,
320 },
321 Target: target.URL,
322 Token: token,
323 Output: outputW,
324 })
325 require.NoError(t, err)
326 result := make(chan error, 1)
327 finished := make(chan struct{})
328 go func() {
329 err := l.Run(ctx)
330 outputW.Close()
331 result <- err
332 close(finished)
333 }()
334 t.Cleanup(func() {
335 cancel()
336 select {
337 case <-finished:
338 case <-time.After(time.Second):
339 t.Error("token client did not stop")
340 }
341 })
342 receive := func(slot int) *mocks.PhysicalConn {
343 t.Helper()
344 select {
345 case pc := <-opened[slot-1]:
346 return pc
347 case <-ctx.Done():
348 t.Fatalf("slot %d did not open: %v", slot, logs.All())
349 return nil
350 }
351 }
352 first, second := receive(1), receive(2)
353 select {
354 case line := <-printed:
355 require.Equal(t, "https://delegated.example.com", line)
356 case <-ctx.Done():
357 t.Fatal("URL was not printed with the first attachment")
358 }
359 require.Equal(t, 1, logs.FilterMessage("Tunnel connection ready").Len())
360 var dialed []string
361 for len(tp.addresses) > 0 {
362 dialed = append(dialed, <-tp.addresses)
363 }
364 require.Equal(t, apex.Address, dialed[0])
365 require.Contains(t, dialed, "alias-a.example.com:443")
366 require.NoError(t, first.Err(), "a duplicate physical handle must stay open")
367 require.Equal(t, "upstream", lightweightUpstream(t, first, "delegated"))
368 close(allowSecond)
369 third := receive(3)
370 require.Eventually(t, func() bool { return logs.FilterMessage("Tunnel connection ready").Len() == 3 }, time.Second, time.Millisecond)
371 require.NotSame(t, first, second)
372 require.NotSame(t, first, third)
373 require.NotSame(t, second, third)
374 third.Close("connection lost")
375 replacement := receive(3)
376 require.NotSame(t, third, replacement)
377 require.NoError(t, first.Err())
378 require.NoError(t, second.Err())
379 require.Equal(t, "upstream", lightweightUpstream(t, second, "delegated"))
380 close(allowRepair)
381 require.Eventually(t, func() bool { return logs.FilterMessage("Tunnel connection ready").Len() == 4 }, time.Second, time.Millisecond)
382 if tc.revoke {
383 first.Close("reconnect after revocation")
384 } else {
385 cancel()
386 }
387 select {
388 case err := <-result:
389 if tc.revoke {
390 require.True(t, tokenFatal(err))
391 require.ErrorContains(t, err, "grant revoked")
392 } else {
393 require.NoError(t, err)
394 }
395 case <-time.After(2 * time.Second):
396 t.Fatal("client did not stop after cancellation or grant rejection")
397 }
398 for line := range printed {
399 t.Errorf("URL printed again: %q", line)
400 }
401 require.Equal(t, 1, logs.FilterMessage("Tunnel ready").Len())
402 slots := map[string]int{
403 nodes[0].Address: 1,
404 nodes[1].Address: 2,
405 nodes[2].Address: 3,
406 }
407 for _, entry := range logs.FilterMessage("Tunnel connection ready").All() {
408 fields := entry.ContextMap()
409 require.Equal(t, "grant-for-client-test", fields["grantId"])
410 require.Equal(t, tc.expiryLog, fields["expiresAt"])
411 require.EqualValues(t, slots[fields["server"].(string)], fields["slot"])
412 }
413 for _, entry := range logs.All() {
414 require.NotContains(t, fmt.Sprint(entry.Message, entry.ContextMap()), token)
415 }
416 for _, pc := range tp.connections {
417 require.Error(t, pc.Err(), "shutdown must close every physical connection")
418 }
419 entries, err := os.ReadDir(dir)
420 require.NoError(t, err)
421 require.Empty(t, entries)
422 service.AssertExpectations(t)
423 service.AssertNotCalled(t, "Ping", mock.Anything, mock.Anything)
424 service.AssertNotCalled(t, "RegisterIdentity", mock.Anything, mock.Anything)
425 left.AssertExpectations(t)
426 pki.AssertExpectations(t)
427 })
428 }
429}