Skip to content
File

Blob: tun/server/dial_recovery_test.go

go245 lines
1package server
2 
3import (
4 "context"
5 "errors"
6 "sync/atomic"
7 "testing"
8 "time"
9 
10 "go.miragespace.co/specter/spec/protocol"
11 "go.miragespace.co/specter/spec/rpc"
12 "go.miragespace.co/specter/spec/transport"
13 "go.miragespace.co/specter/spec/tun"
14 "go.miragespace.co/specter/util/bufconn"
15 
16 "github.com/stretchr/testify/mock"
17 "github.com/stretchr/testify/require"
18)
19 
20func TestDialClientRefreshesExhaustedRoutes(t *testing.T) {
21 as := require.New(t)
22 _, node, clientT, _, serv := getFixture(t, as)
23 t.Cleanup(serv.routeCache.Close)
24 t.Cleanup(serv.keylessCache.Close)
25 link := &protocol.Link{Hostname: "recovered.example.com", Alpn: protocol.Link_HTTP}
26 oldClient, chordNode, tunnelNode := getIdentities()
27 newClient := &protocol.Node{Address: "new-client:123", Id: oldClient.GetId() + 1}
28 clientT.On("Identity").Return(tunnelNode)
29 
30 for _, client := range []*protocol.Node{oldClient, newClient} {
31 route := &protocol.TunnelRoute{
32 Hostname: link.GetHostname(), ClientDestination: client,
33 ChordDestination: chordNode, TunnelDestination: tunnelNode,
34 }
35 wire, err := route.MarshalVT()
36 as.NoError(err)
37 node.On("Get", mock.Anything, []byte(tun.RoutingKey(link.GetHostname(), 1))).Return(wire, nil).Once()
38 for i := 2; i <= tun.NumRedundantLinks; i++ {
39 node.On("Get", mock.Anything, []byte(tun.RoutingKey(link.GetHostname(), i))).Return(nil, nil).Once()
40 }
41 }
42 serv.RoutesPreload(link.GetHostname())
43 clientT.On("DialStream", mock.Anything, oldClient, protocol.Stream_DIRECT).Return(nil, transport.ErrNoDirect).Once()
44 client, remote := bufconn.BufferedPipe(8192)
45 t.Cleanup(func() { client.Close(); remote.Close() })
46 clientT.On("DialStream", mock.Anything, newClient, protocol.Stream_DIRECT).Return(client, nil).Once()
47 
48 conn, err := serv.DialClient(t.Context(), link)
49 as.NoError(err)
50 as.Same(client, conn)
51 received := &protocol.Link{}
52 as.NoError(rpc.BoundedReceive(remote, received, 1024))
53 as.Equal(link.GetHostname(), received.GetHostname())
54 node.AssertExpectations(t)
55 clientT.AssertExpectations(t)
56}
57 
58func TestDialClientKeepsUsableFallback(t *testing.T) {
59 as := require.New(t)
60 _, node, clientT, _, serv := getFixture(t, as)
61 t.Cleanup(serv.routeCache.Close)
62 t.Cleanup(serv.keylessCache.Close)
63 link := &protocol.Link{Hostname: "fallback.example.com"}
64 client, chordNode, tunnelNode := getIdentities()
65 fallback := &protocol.Node{Address: "fallback:123", Id: client.GetId() + 1}
66 clientT.On("Identity").Return(tunnelNode)
67 for i, destination := range []*protocol.Node{client, fallback} {
68 wire, err := (&protocol.TunnelRoute{
69 Hostname: link.GetHostname(), ClientDestination: destination,
70 ChordDestination: chordNode, TunnelDestination: tunnelNode,
71 }).MarshalVT()
72 as.NoError(err)
73 node.On("Get", mock.Anything, []byte(tun.RoutingKey(link.GetHostname(), i+1))).Return(wire, nil).Once()
74 }
75 node.On("Get", mock.Anything, []byte(tun.RoutingKey(link.GetHostname(), 3))).Return(nil, errors.New("temporary lookup failure")).Once()
76 clientT.On("DialStream", mock.Anything, client, protocol.Stream_DIRECT).Return(nil, transport.ErrNoDirect).Once()
77 conn, remote := bufconn.BufferedPipe(8192)
78 t.Cleanup(func() { conn.Close(); remote.Close() })
79 clientT.On("DialStream", mock.Anything, fallback, protocol.Stream_DIRECT).Return(conn, nil).Once()
80 
81 got, err := serv.DialClient(t.Context(), link)
82 as.NoError(err)
83 as.Same(conn, got)
84 node.AssertExpectations(t)
85 clientT.AssertExpectations(t)
86}
87 
88func TestDialClientRefreshCooldown(t *testing.T) {
89 as := require.New(t)
90 _, node, clientT, _, serv := getFixture(t, as)
91 t.Cleanup(serv.routeCache.Close)
92 t.Cleanup(serv.keylessCache.Close)
93 link := &protocol.Link{Hostname: "still-offline.example.com"}
94 client, chordNode, tunnelNode := getIdentities()
95 route := &protocol.TunnelRoute{
96 Hostname: link.GetHostname(), ClientDestination: client,
97 ChordDestination: chordNode, TunnelDestination: tunnelNode,
98 }
99 as.True(serv.routeCache.SetWithTTL(link.GetHostname(), &routesResult{routes: []*protocol.TunnelRoute{route}}, 128, routePositiveTTL))
100 wire, err := route.MarshalVT()
101 as.NoError(err)
102 // The first and third requests refresh; the second is inside the cooldown.
103 node.On("Get", mock.Anything, []byte(tun.RoutingKey(link.GetHostname(), 1))).Return(wire, nil).Twice()
104 for i := 2; i <= tun.NumRedundantLinks; i++ {
105 node.On("Get", mock.Anything, []byte(tun.RoutingKey(link.GetHostname(), i))).Return(nil, nil).Twice()
106 }
107 clientT.On("Identity").Return(tunnelNode)
108 clientT.On("DialStream", mock.Anything, client, protocol.Stream_DIRECT).Return(nil, transport.ErrNoDirect).Times(5)
109 for i := 0; i < 2; i++ {
110 conn, err := serv.DialClient(t.Context(), link)
111 as.Nil(conn)
112 as.ErrorIs(err, tun.ErrTunnelClientNotConnected)
113 }
114 cached, ok := serv.routeCache.Get(link.GetHostname())
115 as.True(ok)
116 expired := *cached
117 expired.refreshAfter = time.Now().Add(-time.Second)
118 as.True(serv.routeCache.SetWithTTL(link.GetHostname(), &expired, 128, routePositiveTTL))
119 conn, err := serv.DialClient(t.Context(), link)
120 as.Nil(conn)
121 as.ErrorIs(err, tun.ErrTunnelClientNotConnected)
122 node.AssertExpectations(t)
123 clientT.AssertExpectations(t)
124}
125 
126func TestDialClientSharesRefreshWithConcurrentAndDelayedFailures(t *testing.T) {
127 as := require.New(t)
128 _, node, clientT, _, serv := getFixture(t, as)
129 t.Cleanup(serv.routeCache.Close)
130 t.Cleanup(serv.keylessCache.Close)
131 link := &protocol.Link{Hostname: "concurrent.example.com"}
132 oldClient, chordNode, tunnelNode := getIdentities()
133 newClient := &protocol.Node{Address: "new-client:123", Id: oldClient.GetId() + 1}
134 route := &protocol.TunnelRoute{
135 Hostname: link.GetHostname(), ClientDestination: oldClient,
136 ChordDestination: chordNode, TunnelDestination: tunnelNode,
137 }
138 as.True(serv.routeCache.SetWithTTL(link.GetHostname(), &routesResult{routes: []*protocol.TunnelRoute{route}}, 128, routePositiveTTL))
139 fresh := &protocol.TunnelRoute{
140 Hostname: link.GetHostname(), ClientDestination: newClient,
141 ChordDestination: chordNode, TunnelDestination: tunnelNode,
142 }
143 wire, err := fresh.MarshalVT()
144 as.NoError(err)
145 lookupStarted := make(chan struct{}, 1)
146 finishLookup := make(chan struct{})
147 node.On("Get", mock.Anything, []byte(tun.RoutingKey(link.GetHostname(), 1))).Return(wire, nil).Once().Run(func(mock.Arguments) {
148 lookupStarted <- struct{}{}
149 <-finishLookup
150 })
151 for i := 2; i <= tun.NumRedundantLinks; i++ {
152 node.On("Get", mock.Anything, []byte(tun.RoutingKey(link.GetHostname(), i))).Return(nil, nil).Once()
153 }
154 const callers = 12
155 started := make(chan struct{}, callers)
156 release := make(chan struct{})
157 releaseDelayed := make(chan struct{})
158 var attempts atomic.Int32
159 clientT.On("Identity").Return(tunnelNode)
160 clientT.On("DialStream", mock.Anything, oldClient, protocol.Stream_DIRECT).Return(nil, transport.ErrNoDirect).Times(callers).Run(func(mock.Arguments) {
161 delayed := attempts.Add(1) == callers
162 started <- struct{}{}
163 if delayed {
164 <-releaseDelayed
165 } else {
166 <-release
167 }
168 })
169 clientT.On("DialStream", mock.Anything, newClient, protocol.Stream_DIRECT).Return(nil, transport.ErrNoDirect).Times(callers)
170 results := make(chan error, callers)
171 for range callers {
172 go func() {
173 _, err := serv.DialClient(t.Context(), link)
174 results <- err
175 }()
176 }
177 for range callers {
178 awaitRouteResult(t, started)
179 }
180 close(release)
181 awaitRouteResult(t, lookupStarted)
182 close(finishLookup)
183 for range callers - 1 {
184 as.ErrorIs(awaitRouteResult(t, results), tun.ErrTunnelClientNotConnected)
185 }
186 // This caller exhausted the old generation after its replacement was cached.
187 close(releaseDelayed)
188 as.ErrorIs(awaitRouteResult(t, results), tun.ErrTunnelClientNotConnected)
189 node.AssertExpectations(t)
190 clientT.AssertExpectations(t)
191}
192 
193func TestRouteLookupCancelledWaiterDoesNotCancelSharedLoad(t *testing.T) {
194 as := require.New(t)
195 _, node, clientT, _, serv := getFixture(t, as)
196 t.Cleanup(serv.routeCache.Close)
197 t.Cleanup(serv.keylessCache.Close)
198 hostname := "cancelled.example.com"
199 client, chordNode, tunnelNode := getIdentities()
200 clientT.On("Identity").Return(tunnelNode)
201 wire, err := (&protocol.TunnelRoute{
202 Hostname: hostname, ClientDestination: client,
203 ChordDestination: chordNode, TunnelDestination: tunnelNode,
204 }).MarshalVT()
205 as.NoError(err)
206 started := make(chan context.Context, 1)
207 release := make(chan struct{})
208 node.On("Get", mock.Anything, []byte(tun.RoutingKey(hostname, 1))).Return(wire, nil).Once().Run(func(args mock.Arguments) {
209 started <- args.Get(0).(context.Context)
210 <-release
211 })
212 for i := 2; i <= tun.NumRedundantLinks; i++ {
213 node.On("Get", mock.Anything, []byte(tun.RoutingKey(hostname, i))).Return(nil, nil).Once()
214 }
215 ctx, cancel := context.WithCancel(t.Context())
216 defer cancel()
217 result := make(chan error, 1)
218 go func() {
219 _, err := serv.lookupRoutes(ctx, hostname, nil)
220 result <- err
221 }()
222 lookupCtx := awaitRouteResult(t, started)
223 cancel()
224 as.ErrorIs(awaitRouteResult(t, result), context.Canceled)
225 as.NoError(lookupCtx.Err())
226 close(release)
227 ret, err := serv.lookupRoutes(t.Context(), hostname, nil)
228 as.NoError(err)
229 as.NoError(ret.err)
230 as.Len(ret.routes, 1)
231 node.AssertExpectations(t)
232}
233 
234func awaitRouteResult[T any](t *testing.T, ch <-chan T) T {
235 t.Helper()
236 select {
237 case ret := <-ch:
238 return ret
239 case <-time.After(2 * time.Second):
240 t.Fatal("timed out waiting for route operation")
241 var zero T
242 return zero
243 }
244}