Skip to content
File

Blob: tun/client/outcomes_test.go

go363 lines
1package client
2 
3import (
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 
24func 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 
60func 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 
100func 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 
147func 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 
213func 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 
278func 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 
318func 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}