Skip to content
File

Blob: integrations/lightweight_test.go

go300 lines
1package integrations
2 
3import (
4 "bytes"
5 "context"
6 "crypto/sha256"
7 "crypto/tls"
8 "encoding/json"
9 "errors"
10 "fmt"
11 "io"
12 "net"
13 "net/http"
14 "net/http/httptest"
15 "net/url"
16 "os"
17 "path/filepath"
18 "testing"
19 "time"
20 
21 clientcmd "go.miragespace.co/specter/cmd/client"
22 servercmd "go.miragespace.co/specter/cmd/server"
23 
24 "github.com/stretchr/testify/require"
25 "github.com/twitchtv/twirp"
26 "github.com/urfave/cli/v3"
27 "go.uber.org/zap/zaptest/observer"
28)
29 
30func startIntegrationApp(t *testing.T, command *cli.Command, args ...string) (*observer.ObservedLogs, func() error) {
31 t.Helper()
32 app, logs := compileApp(command)
33 app.Metadata["apexOverride"] = serverApex
34 app.Writer = io.Discard
35 // Keep peers available while cleanup stops each app in reverse start order.
36 ctx, cancel := context.WithCancel(context.WithoutCancel(t.Context()))
37 done := make(chan error, 1)
38 go func() { done <- app.Run(ctx, append([]string{"specter"}, args...)) }()
39 stopped := false
40 stop := func() error {
41 if stopped {
42 return nil
43 }
44 stopped = true
45 cancel()
46 // Match the existing in-process server fixture: cancel its serving context.
47 // Its Chord leave path uses that same canceled transport context, so waiting
48 // for graceful ring departure here would test unrelated server shutdown.
49 if command.Name == "server" {
50 return nil
51 }
52 select {
53 case err := <-done:
54 return err
55 case <-time.After(30 * time.Second):
56 return fmt.Errorf("app did not stop")
57 }
58 }
59 t.Cleanup(func() { require.NoError(t, stop()) })
60 t.Cleanup(func() {
61 if t.Failed() {
62 for _, entry := range logs.All() {
63 t.Log(entry.Message, entry.ContextMap())
64 }
65 }
66 })
67 return logs, stop
68}
69 
70func waitIntegrationLog(t *testing.T, logs *observer.ObservedLogs, message string) {
71 t.Helper()
72 require.Eventually(t, func() bool { return logs.FilterMessage(message).Len() > 0 }, 30*time.Second, 20*time.Millisecond, message)
73}
74 
75func startLightweightServers(t *testing.T, ports, httpPorts []int) {
76 t.Helper()
77 for i, port := range ports {
78 // Keep physical peers visible in the existing bounded successor discovery;
79 // virtual-node placement is covered by Chord's own integration tests.
80 args := []string{"server", "--cert-dir", "../certs", "--data-dir", t.TempDir(), "--listen", fmt.Sprintf("127.0.0.1:%d", port), "--listen-http", fmt.Sprint(httpPorts[i]), "--apex", serverApex, "--virtual", "1"}
81 if i > 0 {
82 args = append(args, "--join", fmt.Sprintf("127.0.0.1:%d", ports[0]))
83 }
84 logs, _ := startIntegrationApp(t, servercmd.Generate(), args...)
85 waitIntegrationLog(t, logs, "specter server started")
86 waitIntegrationLog(t, logs, "gateway server started")
87 }
88}
89 
90func fetchLightweight(t *testing.T, authority string, port int) (int, string) {
91 t.Helper()
92 cfg := &tls.Config{
93 ServerName: authority,
94 InsecureSkipVerify: true,
95 NextProtos: []string{"h2"},
96 }
97 tp := &http.Transport{
98 ForceAttemptHTTP2: true,
99 TLSClientConfig: cfg,
100 DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
101 return (&tls.Dialer{Config: cfg}).DialContext(ctx, "tcp", fmt.Sprintf("127.0.0.1:%d", port))
102 },
103 }
104 defer tp.CloseIdleConnections()
105 c := &http.Client{
106 Transport: tp,
107 Timeout: 5 * time.Second,
108 }
109 resp, err := c.Get("https://" + authority + "/")
110 if err != nil {
111 return 0, err.Error()
112 }
113 defer resp.Body.Close()
114 require.Equal(t, 2, resp.ProtoMajor)
115 data, err := io.ReadAll(resp.Body)
116 require.NoError(t, err)
117 return resp.StatusCode, string(data)
118}
119 
120func lightweightAuthority(t *testing.T, logs *observer.ObservedLogs) string {
121 t.Helper()
122 waitIntegrationLog(t, logs, "Tunnel ready")
123 parsed, err := url.Parse(logs.FilterMessage("Tunnel ready").All()[0].ContextMap()["url"].(string))
124 require.NoError(t, err)
125 return parsed.Hostname()
126}
127 
128func localTokenRequest(t *testing.T, port int, method, path string, body any, output any) int {
129 t.Helper()
130 var data []byte
131 var err error
132 if body != nil {
133 data, err = json.Marshal(body)
134 require.NoError(t, err)
135 }
136 req, err := http.NewRequest(method, fmt.Sprintf("http://127.0.0.1:%d/api%s", port, path), bytes.NewReader(data))
137 require.NoError(t, err)
138 req.Header.Set("Content-Type", "application/json")
139 resp, err := (&http.Client{Timeout: 10 * time.Second}).Do(req)
140 require.NoError(t, err)
141 defer resp.Body.Close()
142 if len(path) >= 7 && path[:7] == "/tokens" {
143 require.Equal(t, "no-store", resp.Header.Get("Cache-Control"))
144 }
145 if output != nil {
146 require.NoError(t, json.NewDecoder(resp.Body).Decode(output))
147 }
148 return resp.StatusCode
149}
150 
151func runTokenCLI(t *testing.T, args ...string) (string, error) {
152 t.Helper()
153 // Each stopped-client command bootstraps again; respect the existing per-IP RPC limit.
154 time.Sleep(1100 * time.Millisecond)
155 app, _ := compileApp(clientcmd.Generate())
156 app.Metadata["apexOverride"] = serverApex
157 var out bytes.Buffer
158 app.Writer = &out
159 ctx, cancel := context.WithTimeout(t.Context(), 20*time.Second)
160 defer cancel()
161 err := app.Run(ctx, append([]string{"specter", "client", "--insecure", "token"}, args...))
162 return out.String(), err
163}
164 
165func TestIntegrationLightweight(t *testing.T) {
166 if os.Getenv("GO_INTEGRATION_TUNNEL") == "" {
167 t.Skip("set GO_INTEGRATION_TUNNEL=1")
168 }
169 ports := []int{21958, 21959, 21960}
170 startLightweightServers(t, ports, []int{21858, 21859, 21860})
171 target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, "lightweight") }))
172 defer target.Close()
173 t.Run("ephemeral", func(t *testing.T) {
174 logs, stop := startIntegrationApp(t, clientcmd.Generate(), "client", "--insecure", "expose", "--apex", "127.0.0.1:21958", target.URL)
175 authority := lightweightAuthority(t, logs)
176 for _, port := range ports {
177 code, body := fetchLightweight(t, authority, port)
178 require.Equal(t, 200, code, body)
179 require.Equal(t, "lightweight", body)
180 }
181 require.NoError(t, stop())
182 require.Eventually(t, func() bool { code, _ := fetchLightweight(t, authority, ports[0]); return code != 200 }, 5*time.Second, 20*time.Millisecond)
183 })
184 t.Run("late upstream", func(t *testing.T) {
185 listener, err := net.Listen("tcp", "127.0.0.1:0")
186 require.NoError(t, err)
187 addr := listener.Addr().String()
188 listener.Close()
189 logs, _ := startIntegrationApp(t, clientcmd.Generate(), "client", "--insecure", "expose", "--apex", "127.0.0.1:21958", "http://"+addr)
190 authority := lightweightAuthority(t, logs)
191 code, _ := fetchLightweight(t, authority, ports[0])
192 require.Equal(t, 502, code)
193 listener, err = net.Listen("tcp", addr)
194 require.NoError(t, err)
195 upstream := &http.Server{Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, "late") })}
196 defer upstream.Close()
197 go upstream.Serve(listener)
198 code, body := fetchLightweight(t, authority, ports[0])
199 require.Equal(t, 200, code, body)
200 require.Equal(t, "late", body)
201 })
202 t.Run("delegated", func(t *testing.T) {
203 path := filepath.Join(t.TempDir(), "owner.yaml")
204 require.NoError(t, os.WriteFile(path, []byte(fmt.Sprintf("version: 2\napex: 127.0.0.1:21958\ntunnels:\n - target: %s\n", target.URL)), 0600))
205 ownerLogs, stopOwner := startIntegrationApp(t, clientcmd.Generate(), "client", "--insecure", "tunnel", "--config", path, "--server", "127.0.0.1:21881")
206 waitIntegrationLog(t, ownerLogs, "Local server started")
207 var hostnames []struct {
208 Hostname string `json:"hostname"`
209 }
210 require.Equal(t, 200, localTokenRequest(t, 21881, "GET", "/ls", nil, &hostnames))
211 require.Len(t, hostnames, 1)
212 hostname := hostnames[0].Hostname
213 require.Equal(t, 200, localTokenRequest(t, 21881, "POST", "/unpublish/"+hostname, nil, nil))
214 var minted struct {
215 Token string `json:"token"`
216 Grant struct {
217 ID string `json:"id"`
218 } `json:"grant"`
219 }
220 require.Equal(t, 201, localTokenRequest(t, 21881, "POST", "/tokens", map[string]string{"hostname": hostname}, &minted))
221 require.Contains(t, minted.Token, "tg1_")
222 tokenPath := filepath.Join(t.TempDir(), "token")
223 require.NoError(t, os.WriteFile(tokenPath, []byte(minted.Token), 0600))
224 logs, stopServe := startIntegrationApp(t, clientcmd.Generate(), "client", "--insecure", "serve", "--apex", "127.0.0.1:21959", "--token-file", tokenPath, target.URL)
225 authority := lightweightAuthority(t, logs)
226 activeConnections := func() map[string]string {
227 active := make(map[string]string)
228 for _, entry := range logs.All() {
229 fields := entry.ContextMap()
230 slot := fmt.Sprint(fields["slot"])
231 server, _ := fields["server"].(string)
232 switch entry.Message {
233 case "Tunnel connection ready":
234 active[slot] = server
235 case "Tunnel disconnected":
236 if active[slot] == server {
237 delete(active, slot)
238 }
239 }
240 }
241 return active
242 }
243 require.Eventually(t, func() bool {
244 active := activeConnections()
245 servers := make(map[string]bool)
246 for _, slot := range []string{"1", "2", "3"} {
247 if active[slot] == "" {
248 return false
249 }
250 servers[active[slot]] = true
251 }
252 return len(active) == 3 && len(servers) == 3
253 }, 45*time.Second, 20*time.Millisecond, "discover and activate three distinct servers from the apex")
254 require.Equal(t, 1, logs.FilterMessage("Tunnel ready").Len())
255 for _, entry := range logs.FilterMessage("Tunnel connection ready").All() {
256 fields := entry.ContextMap()
257 require.Equal(t, minted.Grant.ID, fields["grantId"])
258 require.Equal(t, "none", fields["expiresAt"])
259 }
260 for _, port := range ports {
261 code, body := fetchLightweight(t, authority, port)
262 require.Equal(t, 200, code, body)
263 require.Equal(t, "lightweight", body)
264 }
265 require.NoError(t, stopOwner())
266 before, err := os.ReadFile(path)
267 require.NoError(t, err)
268 out, err := runTokenCLI(t, "list", "--config", path)
269 require.NoError(t, err)
270 require.Contains(t, out, minted.Grant.ID)
271 require.NotContains(t, out, "tg1_")
272 after, err := os.ReadFile(path)
273 require.NoError(t, err)
274 require.Equal(t, sha256.Sum256(before), sha256.Sum256(after))
275 out, err = runTokenCLI(t, "revoke", "--config", path, minted.Grant.ID)
276 require.NoError(t, err)
277 require.Contains(t, out, `"revoked": true`)
278 app, _ := compileApp(clientcmd.Generate())
279 app.Metadata["apexOverride"] = serverApex
280 app.Writer = io.Discard
281 ctx, cancel := context.WithTimeout(t.Context(), 15*time.Second)
282 defer cancel()
283 err = app.Run(ctx, []string{"specter", "client", "--insecure", "serve", "--apex", "127.0.0.1:21959", "--token-file", tokenPath, target.URL})
284 require.ErrorContains(t, err, "unknown or revoked grant")
285 require.Eventually(t, func() bool { return len(activeConnections()) == 0 }, 45*time.Second, 100*time.Millisecond, "revocation closes all token connections")
286 for _, port := range ports {
287 code, body := fetchLightweight(t, authority, port)
288 require.Equal(t, 503, code, body)
289 }
290 // Cancellation may beat the client's next Open after its last connection
291 // closes; otherwise that Open must fail with the revoked-grant error.
292 if err := stopServe(); err != nil {
293 var rpcError twirp.Error
294 require.True(t, errors.As(err, &rpcError), "%v", err)
295 require.Equal(t, twirp.Unauthenticated, rpcError.Code())
296 require.Equal(t, "unknown or revoked grant", rpcError.Msg())
297 }
298 })
299}