File
Blob: tun/client/outcomes_test.go
| 1 | package client |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "errors" |
| 7 | "net/http" |
| 8 | "net/http/httptest" |
| 9 | "os" |
| 10 | "path/filepath" |
| 11 | "testing" |
| 12 | "time" |
| 13 | |
| 14 | "go.miragespace.co/specter/spec/protocol" |
| 15 | "go.miragespace.co/specter/spec/rpc" |
| 16 | |
| 17 | "github.com/stretchr/testify/mock" |
| 18 | "github.com/stretchr/testify/require" |
| 19 | "github.com/zhangyunhao116/skipmap" |
| 20 | "go.uber.org/atomic" |
| 21 | "go.uber.org/zap/zaptest" |
| 22 | ) |
| 23 | |
| 24 | func newOutcomeTestClient(t *testing.T, tunnels []Tunnel) (*Client, *mockTunnelClient) { |
| 25 | t.Helper() |
| 26 | |
| 27 | cfg := &Config{ |
| 28 | path: filepath.Join(t.TempDir(), "client.yml"), |
| 29 | router: skipmap.NewString[route](), |
| 30 | Version: 2, |
| 31 | Apex: testApex, |
| 32 | Tunnels: tunnels, |
| 33 | } |
| 34 | require.NoError(t, cfg.validate()) |
| 35 | cfg.buildRouter() |
| 36 | require.NoError(t, cfg.writeFile()) |
| 37 | |
| 38 | tunnelClient := new(mockTunnelClient) |
| 39 | c := &Client{ |
| 40 | ClientConfig: ClientConfig{ |
| 41 | Logger: zaptest.NewLogger(t), |
| 42 | Configuration: cfg, |
| 43 | }, |
| 44 | tunnelClient: tunnelClient, |
| 45 | forwarder: &forwarder{ |
| 46 | logger: zaptest.NewLogger(t), |
| 47 | rootDomain: atomic.NewString(testApex), |
| 48 | proxies: skipmap.NewString[*httpProxy](), |
| 49 | }, |
| 50 | connections: skipmap.NewString[*protocol.Node](), |
| 51 | } |
| 52 | c.connections.Store("gateway.example.com", &protocol.Node{ |
| 53 | Id: 1, |
| 54 | Address: "gateway.example.com", |
| 55 | }) |
| 56 | t.Cleanup(func() { tunnelClient.TunnelService.AssertExpectations(t) }) |
| 57 | return c, tunnelClient |
| 58 | } |
| 59 | |
| 60 | func TestLocalReloadRejectsInvalidConfiguration(t *testing.T) { |
| 61 | for _, tc := range []struct { |
| 62 | name string |
| 63 | content string |
| 64 | remove bool |
| 65 | }{ |
| 66 | {name: "invalid YAML", content: "tunnels: ["}, |
| 67 | {name: "invalid target", content: "version: 2\ntunnels:\n - hostname: changed\n target: unsupported://localhost\n"}, |
| 68 | {name: "unreadable file", remove: true}, |
| 69 | } { |
| 70 | t.Run(tc.name, func(t *testing.T) { |
| 71 | original := Tunnel{Hostname: "existing", Target: "http://localhost:8080"} |
| 72 | c, tunnelClient := newOutcomeTestClient(t, []Tunnel{original}) |
| 73 | if tc.remove { |
| 74 | require.NoError(t, os.Remove(c.Configuration.path)) |
| 75 | } else { |
| 76 | require.NoError(t, os.WriteFile(c.Configuration.path, []byte(tc.content), 0600)) |
| 77 | } |
| 78 | |
| 79 | response := httptest.NewRecorder() |
| 80 | c.localHandler().ServeHTTP(response, httptest.NewRequest(http.MethodPost, "/api/reload", nil)) |
| 81 | |
| 82 | require.Equal(t, http.StatusBadRequest, response.Code) |
| 83 | require.Contains(t, response.Header().Get("Content-Type"), "application/json") |
| 84 | var result SyncResult |
| 85 | require.NoError(t, json.Unmarshal(response.Body.Bytes(), &result)) |
| 86 | require.NotEmpty(t, result.Error) |
| 87 | require.False(t, result.Applied) |
| 88 | require.False(t, result.Saved) |
| 89 | require.Len(t, c.Configuration.Tunnels, 1) |
| 90 | require.Equal(t, original.Hostname, c.Configuration.Tunnels[0].Hostname) |
| 91 | require.Equal(t, original.Target, c.Configuration.Tunnels[0].Target) |
| 92 | route, found := c.Configuration.router.Load(original.Hostname) |
| 93 | require.True(t, found) |
| 94 | require.Equal(t, original.Target, route.parsed.String()) |
| 95 | tunnelClient.TunnelService.AssertNotCalled(t, "PublishTunnel", mock.Anything, mock.Anything) |
| 96 | }) |
| 97 | } |
| 98 | } |
| 99 | |
| 100 | func TestLocalReloadReportsPartialPublishFailure(t *testing.T) { |
| 101 | c, tunnelClient := newOutcomeTestClient(t, []Tunnel{ |
| 102 | {Hostname: "published", Target: "http://localhost:8080"}, |
| 103 | {Hostname: "failed", Target: "tcp://localhost:5432"}, |
| 104 | }) |
| 105 | |
| 106 | tunnelClient.TunnelService.On("RegisteredHostnames", mock.Anything, mock.Anything). |
| 107 | Return(&protocol.RegisteredHostnamesResponse{}, nil) |
| 108 | tunnelClient.TunnelService.On("PublishTunnel", mock.Anything, mock.MatchedBy(func(req *protocol.PublishTunnelRequest) bool { |
| 109 | return req.GetHostname() == "published" |
| 110 | })).Return(&protocol.PublishTunnelResponse{Published: c.getConnectedNodes()}, nil).Once() |
| 111 | tunnelClient.TunnelService.On("PublishTunnel", mock.Anything, mock.MatchedBy(func(req *protocol.PublishTunnelRequest) bool { |
| 112 | return req.GetHostname() == "failed" |
| 113 | })).Return(nil, errors.New("publication rejected")).Once() |
| 114 | |
| 115 | response := httptest.NewRecorder() |
| 116 | c.localHandler().ServeHTTP(response, httptest.NewRequest(http.MethodPost, "/api/reload", nil)) |
| 117 | |
| 118 | require.Equal(t, http.StatusInternalServerError, response.Code) |
| 119 | require.Contains(t, response.Header().Get("Content-Type"), "application/json") |
| 120 | var result SyncResult |
| 121 | require.NoError(t, json.Unmarshal(response.Body.Bytes(), &result)) |
| 122 | require.True(t, result.Applied) |
| 123 | require.True(t, result.Saved) |
| 124 | require.Contains(t, result.Error, "publication rejected") |
| 125 | require.Len(t, result.Tunnels, 2) |
| 126 | |
| 127 | byHostname := make(map[string]TunnelSyncResult, len(result.Tunnels)) |
| 128 | for _, tunnel := range result.Tunnels { |
| 129 | byHostname[tunnel.Hostname] = tunnel |
| 130 | } |
| 131 | published := byHostname["published"] |
| 132 | require.Equal(t, "http://localhost:8080", published.Target) |
| 133 | require.True(t, published.Published) |
| 134 | require.Equal(t, 1, published.PublishedEndpoints) |
| 135 | require.Empty(t, published.Error) |
| 136 | failed := byHostname["failed"] |
| 137 | require.Equal(t, "tcp://localhost:5432", failed.Target) |
| 138 | require.False(t, failed.Published) |
| 139 | require.Zero(t, failed.PublishedEndpoints) |
| 140 | require.Contains(t, failed.Error, "publication rejected") |
| 141 | |
| 142 | onDisk, err := NewConfig(c.Configuration.path) |
| 143 | require.NoError(t, err) |
| 144 | require.Len(t, onDisk.Tunnels, 2) |
| 145 | } |
| 146 | |
| 147 | func TestLocalTunnelRemovalReportsPersistenceFailure(t *testing.T) { |
| 148 | for _, tc := range []struct { |
| 149 | name string |
| 150 | method string |
| 151 | result any |
| 152 | }{ |
| 153 | {name: "unpublish", method: "UnpublishTunnel", result: &protocol.UnpublishTunnelResponse{}}, |
| 154 | {name: "release", method: "ReleaseTunnel", result: &protocol.ReleaseTunnelResponse{}}, |
| 155 | } { |
| 156 | t.Run(tc.name, func(t *testing.T) { |
| 157 | c, tunnelClient := newOutcomeTestClient(t, []Tunnel{ |
| 158 | {Hostname: "existing", Target: "http://localhost:8080"}, |
| 159 | }) |
| 160 | tunnelClient.TunnelService.On(tc.method, mock.Anything, mock.MatchedBy(func(req any) bool { |
| 161 | if req, ok := req.(*protocol.UnpublishTunnelRequest); ok { |
| 162 | return req.GetHostname() == "existing" |
| 163 | } |
| 164 | if req, ok := req.(*protocol.ReleaseTunnelRequest); ok { |
| 165 | return req.GetHostname() == "existing" |
| 166 | } |
| 167 | return false |
| 168 | })).Return(tc.result, nil).Once() |
| 169 | |
| 170 | // A regular file cannot contain the config path, even when tests run as root. |
| 171 | originalPath := c.Configuration.path |
| 172 | originalContents, err := os.ReadFile(originalPath) |
| 173 | require.NoError(t, err) |
| 174 | c.Configuration.path = filepath.Join(originalPath, "client.yml") |
| 175 | |
| 176 | response := httptest.NewRecorder() |
| 177 | c.localHandler().ServeHTTP(response, httptest.NewRequest(http.MethodPost, "/api/"+tc.name+"/existing", nil)) |
| 178 | |
| 179 | require.Equal(t, http.StatusInternalServerError, response.Code) |
| 180 | require.Contains(t, response.Header().Get("Content-Type"), "application/json") |
| 181 | var result SyncResult |
| 182 | require.NoError(t, json.Unmarshal(response.Body.Bytes(), &result)) |
| 183 | require.NotEmpty(t, result.Error) |
| 184 | require.True(t, result.Applied, "successful network removal must be distinguished from persistence failure") |
| 185 | require.False(t, result.Saved) |
| 186 | require.Empty(t, c.GetCurrentConfig().Tunnels) |
| 187 | _, found := c.Configuration.router.Load("existing") |
| 188 | require.False(t, found, "a successful network removal must also remove the live route") |
| 189 | contents, err := os.ReadFile(originalPath) |
| 190 | require.NoError(t, err) |
| 191 | require.Equal(t, originalContents, contents) |
| 192 | |
| 193 | require.False(t, c.getStatus().Pending, "a save failure must not schedule publication") |
| 194 | require.Nil(t, c.getStatus().RetryAt) |
| 195 | |
| 196 | // An explicit reload accepts the file's contents, including the old entry. |
| 197 | c.Configuration.path = originalPath |
| 198 | tunnelClient.TunnelService.On("RegisteredHostnames", mock.Anything, mock.Anything). |
| 199 | Return(&protocol.RegisteredHostnamesResponse{}, nil).Once() |
| 200 | tunnelClient.TunnelService.On("PublishTunnel", mock.Anything, mock.Anything). |
| 201 | Return(&protocol.PublishTunnelResponse{Published: c.getConnectedNodes()}, nil).Once() |
| 202 | reloadResponse := httptest.NewRecorder() |
| 203 | c.localHandler().ServeHTTP(reloadResponse, httptest.NewRequest(http.MethodPost, "/api/reload", nil)) |
| 204 | require.Equal(t, http.StatusNoContent, reloadResponse.Code) |
| 205 | require.Equal(t, "existing", c.GetCurrentConfig().Tunnels[0].Hostname) |
| 206 | require.True(t, c.getStatus().Synchronization.Saved) |
| 207 | _, found = c.Configuration.router.Load("existing") |
| 208 | require.True(t, found) |
| 209 | }) |
| 210 | } |
| 211 | } |
| 212 | |
| 213 | func TestLocalStatusRemainsAvailableDuringPublication(t *testing.T) { |
| 214 | c, tunnelClient := newOutcomeTestClient(t, []Tunnel{ |
| 215 | {Hostname: "existing", Target: "http://localhost:8080"}, |
| 216 | }) |
| 217 | c.Configuration.Certificate = "private-certificate-test-value" |
| 218 | c.Configuration.PrivKey = "private-key-test-value" |
| 219 | c.syncStateMu.Lock() |
| 220 | c.lastSync = SyncResult{Tunnels: []TunnelSyncResult{ |
| 221 | {Hostname: "removed", Target: "http://localhost:9090", Published: true, PublishedEndpoints: 1}, |
| 222 | {Hostname: "existing", Target: "http://localhost:7070", Published: true, PublishedEndpoints: 1}, |
| 223 | }} |
| 224 | c.syncStateMu.Unlock() |
| 225 | |
| 226 | publicationStarted := make(chan struct{}) |
| 227 | publicationRelease := make(chan struct{}) |
| 228 | publicationDone := make(chan SyncResult, 1) |
| 229 | tunnelClient.TunnelService.On("RegisteredHostnames", mock.Anything, mock.Anything). |
| 230 | Return(&protocol.RegisteredHostnamesResponse{}, nil).Once() |
| 231 | tunnelClient.TunnelService.On("PublishTunnel", mock.Anything, mock.Anything). |
| 232 | Run(func(mock.Arguments) { |
| 233 | close(publicationStarted) |
| 234 | <-publicationRelease |
| 235 | }).Return(&protocol.PublishTunnelResponse{Published: c.getConnectedNodes()}, nil).Once() |
| 236 | go func() { publicationDone <- c.SyncConfigTunnels(t.Context()) }() |
| 237 | defer func() { |
| 238 | close(publicationRelease) |
| 239 | result := <-publicationDone |
| 240 | require.Empty(t, result.Error) |
| 241 | }() |
| 242 | select { |
| 243 | case <-publicationStarted: |
| 244 | case <-time.After(2 * time.Second): |
| 245 | t.Fatal("publication did not start") |
| 246 | } |
| 247 | |
| 248 | statusDone := make(chan *httptest.ResponseRecorder, 1) |
| 249 | go func() { |
| 250 | response := httptest.NewRecorder() |
| 251 | c.localHandler().ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/api/status", nil)) |
| 252 | statusDone <- response |
| 253 | }() |
| 254 | var response *httptest.ResponseRecorder |
| 255 | select { |
| 256 | case response = <-statusDone: |
| 257 | case <-time.After(2 * time.Second): |
| 258 | t.Fatal("status endpoint blocked on publication") |
| 259 | } |
| 260 | |
| 261 | require.Equal(t, http.StatusOK, response.Code) |
| 262 | require.Equal(t, "no-store", response.Header().Get("Cache-Control")) |
| 263 | require.Contains(t, response.Header().Get("Content-Type"), "application/json") |
| 264 | var status ClientStatus |
| 265 | require.NoError(t, json.Unmarshal(response.Body.Bytes(), &status)) |
| 266 | require.Equal(t, testApex, status.Apex) |
| 267 | require.Len(t, status.ConnectedNodes, 1) |
| 268 | require.Equal(t, "gateway.example.com", status.ConnectedNodes[0].GetAddress()) |
| 269 | require.Equal(t, []TunnelSyncResult{ |
| 270 | {Hostname: "existing", Target: "http://localhost:8080"}, |
| 271 | }, status.Synchronization.Tunnels, "status must reflect configured targets, not stale acknowledgements") |
| 272 | require.NotContains(t, response.Body.String(), "private-certificate-test-value") |
| 273 | require.NotContains(t, response.Body.String(), "private-key-test-value") |
| 274 | require.NotContains(t, response.Body.String(), `"certificate"`) |
| 275 | require.NotContains(t, response.Body.String(), `"privKey"`) |
| 276 | } |
| 277 | |
| 278 | func TestLocalUnpublishClearsSynchronization(t *testing.T) { |
| 279 | for _, tc := range []struct { |
| 280 | name string |
| 281 | err error |
| 282 | }{ |
| 283 | {name: "published"}, |
| 284 | {name: "publication failed", err: errors.New("publication rejected")}, |
| 285 | } { |
| 286 | t.Run(tc.name, func(t *testing.T) { |
| 287 | c, tunnelClient := newOutcomeTestClient(t, []Tunnel{{Hostname: "existing", Target: "http://localhost:8080"}}) |
| 288 | tunnelClient.TunnelService.On("RegisteredHostnames", mock.Anything, mock.Anything). |
| 289 | Return(&protocol.RegisteredHostnamesResponse{}, nil).Once() |
| 290 | tunnelClient.TunnelService.On("PublishTunnel", mock.Anything, mock.Anything). |
| 291 | Return(&protocol.PublishTunnelResponse{Published: c.getConnectedNodes()}, tc.err).Once() |
| 292 | tunnelClient.TunnelService.On("UnpublishTunnel", mock.Anything, mock.MatchedBy(func(req *protocol.UnpublishTunnelRequest) bool { |
| 293 | return req.GetHostname() == "existing" |
| 294 | })).Return(&protocol.UnpublishTunnelResponse{}, nil).Once() |
| 295 | initial := c.SyncConfigTunnels(t.Context()) |
| 296 | require.Equal(t, tc.err != nil, c.getStatus().Pending) |
| 297 | require.Equal(t, tc.err == nil, initial.Tunnels[0].Published) |
| 298 | |
| 299 | handler := c.localHandler() |
| 300 | removalResponse := httptest.NewRecorder() |
| 301 | handler.ServeHTTP(removalResponse, httptest.NewRequest(http.MethodPost, "/api/unpublish/existing", nil)) |
| 302 | require.Equal(t, http.StatusOK, removalResponse.Code) |
| 303 | statusResponse := httptest.NewRecorder() |
| 304 | handler.ServeHTTP(statusResponse, httptest.NewRequest(http.MethodGet, "/api/status", nil)) |
| 305 | require.Equal(t, http.StatusOK, statusResponse.Code) |
| 306 | var status ClientStatus |
| 307 | require.NoError(t, json.Unmarshal(statusResponse.Body.Bytes(), &status)) |
| 308 | require.Equal(t, []TunnelSyncResult{}, status.Synchronization.Tunnels, "empty configuration must serialize as an empty array") |
| 309 | require.True(t, status.Synchronization.Applied) |
| 310 | require.True(t, status.Synchronization.Saved) |
| 311 | require.False(t, status.Pending) |
| 312 | require.Nil(t, status.RetryAt) |
| 313 | require.Empty(t, status.Synchronization.Error) |
| 314 | }) |
| 315 | } |
| 316 | } |
| 317 | |
| 318 | func TestGatewayChangePreservesUnreloadedConfigurationEdits(t *testing.T) { |
| 319 | const originalTarget = "http://localhost:8080" |
| 320 | const editedTarget = "http://localhost:9090" |
| 321 | c, tunnelClient := newOutcomeTestClient(t, []Tunnel{ |
| 322 | {Hostname: "existing", Target: originalTarget}, |
| 323 | }) |
| 324 | tunnelClient.TunnelService.On("RegisteredHostnames", mock.Anything, mock.Anything). |
| 325 | Return(&protocol.RegisteredHostnamesResponse{}, nil).Once() |
| 326 | tunnelClient.TunnelService.On("PublishTunnel", mock.Anything, mock.Anything). |
| 327 | Return(&protocol.PublishTunnelResponse{Published: c.getConnectedNodes()}, nil).Once() |
| 328 | |
| 329 | initial := c.SyncConfigTunnels(t.Context()) |
| 330 | require.Empty(t, initial.Error) |
| 331 | require.True(t, initial.Saved) |
| 332 | |
| 333 | // Discovering another gateway must not rewrite edits awaiting explicit reload. |
| 334 | nodes := append(c.getConnectedNodes(), &protocol.Node{Id: 2, Address: "second.example.com"}) |
| 335 | tunnelClient.TunnelService.On("GetNodes", mock.Anything, mock.Anything). |
| 336 | Return(&protocol.GetNodesResponse{Nodes: nodes}, nil).Once() |
| 337 | for _, node := range nodes { |
| 338 | tunnelClient.TunnelService.On("Ping", mock.MatchedBy(func(ctx context.Context) bool { |
| 339 | return rpc.GetNode(ctx).GetId() == node.GetId() |
| 340 | }), mock.Anything).Return(&protocol.ClientPingResponse{Node: node}, nil) |
| 341 | } |
| 342 | tunnelClient.TunnelService.On("PublishTunnel", mock.Anything, mock.Anything). |
| 343 | Return(&protocol.PublishTunnelResponse{Published: nodes}, nil).Once() |
| 344 | edited := c.GetCurrentConfig() |
| 345 | edited.Tunnels[0].Target = editedTarget |
| 346 | require.NoError(t, edited.writeFile()) |
| 347 | c.reconcileConnections(t.Context()) |
| 348 | |
| 349 | status := c.getStatus() |
| 350 | require.False(t, status.Pending) |
| 351 | require.Empty(t, status.Synchronization.Error) |
| 352 | require.True(t, status.Synchronization.Tunnels[0].Published) |
| 353 | require.Equal(t, 2, status.Synchronization.Tunnels[0].PublishedEndpoints) |
| 354 | require.Equal(t, originalTarget, status.Synchronization.Tunnels[0].Target) |
| 355 | require.Equal(t, originalTarget, c.GetCurrentConfig().Tunnels[0].Target) |
| 356 | route, found := c.Configuration.router.Load("existing") |
| 357 | require.True(t, found) |
| 358 | require.Equal(t, originalTarget, route.parsed.String()) |
| 359 | onDisk, err := NewConfig(edited.path) |
| 360 | require.NoError(t, err) |
| 361 | require.Equal(t, editedTarget, onDisk.Tunnels[0].Target) |
| 362 | } |