File
Blob: tun/server/dial_recovery_test.go
| 1 | package server |
| 2 | |
| 3 | import ( |
| 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 | |
| 20 | func 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 | |
| 58 | func 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 | |
| 88 | func 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 | |
| 126 | func 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 | |
| 193 | func 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 | |
| 234 | func 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 | } |