refactor(sdk): simplify listener lifecycle and rename internal relay helpers
Kim committed
Mar 10, 2026 at 21:46 UTC
f36d1ee2d4e06f77dd07b35971150f12973d2099
12 files changed
+512
-1624
AGENTS.md
+1
-1
@@ -22,7 +22,7 @@ Descriptive docs under `docs/` should match current code paths.
22
4. **Keep explicit root-domain fallback behavior through SNI no-route handling to the admin/API listener.**
23
- Why: preserves the intended split between root-host control-plane traffic and tenant subdomain traffic.
24
25
-5. **All leases require `TLS=true`.** The register endpoint rejects `TLS=false`.
25
+5. **All leases are TLS-only.** The register endpoint does not accept a non-TLS mode.
26
- Why: all tenant routes are expected to stay on the TLS passthrough path.
27
28
## TLS and Identity Invariants
docs/architecture.md
+6
-7
@@ -54,23 +54,23 @@ That distinction matters because `/sdk/connect` stops being ordinary HTTP once h
54
55
### SDK (`sdk/`)
56
57
-- `RelayClient`: validates one or more relay URLs and owns per-relay HTTP client and raw TLS dial config
58
-- `Listener`: registers one lease per relay, maintains per-entry `readyTarget` reverse sessions, renews lease TTLs, and yields accepted tenant TLS connections through one aggregate listener surface
59
-- Default app flow is `RelayURLs -> NewRelayClient -> Listener -> PublicURLs -> http.Server.Serve(listener)`
60
-- `helper.go`: optional `RunHTTPApp` helper for serving one handler on both a local HTTP port and the relay listener
57
+- `Listener`: validates one relay URL, registers one lease per relay, maintains per-entry `readyTarget` reverse sessions, renews lease TTLs, and yields accepted tenant TLS connections; listener failure is terminal and callers create a new listener instead of reactivating the old one
58
+- `relayclient.go`: internal relay transport helper for control-plane requests and reverse session dialing
59
+- Default app flow is `RelayURL -> NewListener -> PublicURLs -> http.Server.Serve(listener)` or `RelayURLs -> Expose -> PublicURLs -> http.Server.Serve(exposure)`
60
+- `expose.go`: optional `RunHTTPApp` helper for serving one handler on both a local HTTP port and the relay listener
61
- Relay-aware entry inspection is reserved for advanced callers such as `portal-tunnel`
62
- Tenant TLS is created automatically through the relay keyless signer; callers do not provide a local self-signed fallback path
63
64
### Tunnel (`cmd/portal-tunnel`)
65
66
-- Creates one SDK client, registers one lease per relay through the SDK, and consumes one aggregate listener
66
+- Creates one SDK listener per relay through the SDK and consumes one aggregate listener
67
- Accepts claimed tenant connections from the relay
68
- Proxies raw TCP to a local `--host`
69
- Returns an HTTP 503 response when the local target is unavailable
70
71
## Transport Model
72
73
-### Raw reverse transport (`TLS=true` only)
73
+### Raw reverse transport (TLS only)
74
75
1. SDK/tunnel registers one lease per relay with `POST /sdk/register`.
76
2. SDK opens one or more reverse sessions per registered lease with `GET /sdk/connect?lease_id=...`.
@@ -92,7 +92,6 @@ Result: the relay decides routing, but tenant TLS termination still happens at t
92
- Caller provides:
93
- `name`
94
- `reverse_token`
95
- - `tls=true`
95
- optional `hostnames`
96
- optional `metadata`
97
- optional `ttl_seconds`
docs/greenfield-raw-tcp-sni-keyless.md
+4
-4
@@ -248,7 +248,7 @@ All HTTP endpoints in this list are HTTP/1.1 only.
248
249
### Register
250
251
-- Requires `lease_id`, `name`, `reverse_token`, `tls=true`
251
+- Requires `name` and `reverse_token`
252
- Creates or resets lease broker
253
- Registers route in `RouteTable`
254
@@ -324,9 +324,9 @@ Transient:
324
325
### SDK Agent Rules
326
327
-- Fatal rejection pauses the agent
328
-- Successful renew/register clears pause
329
-- Transient failure retries with bounded backoff
327
+- Fatal rejection closes the current listener
328
+- Recovery creates a fresh listener; the SDK does not reactivate a failed listener in place
329
+- Transient failure retries with bounded backoff inside one listener lifecycle
330
331
## Backpressure
332
portal/server.go
+5
-7
@@ -280,6 +280,7 @@ func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
280
writeAPIData(w, http.StatusOK, types.DomainResponse{
281
RootHost: s.cfg.RootHost,
282
SuggestedHostname: suggestHostname(name, s.cfg.RootHost),
283
+ Version: types.SDKProtocolVersion,
284
})
285
}
286
@@ -452,9 +453,6 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
453
if strings.TrimSpace(req.ReverseToken) == "" {
454
return types.RegisterResponse{}, errors.New("reverse token is required")
455
}
455
- if !req.TLS {
456
- return types.RegisterResponse{}, errors.New("tls must be true")
457
- }
456
if s.isClientIPBanned(clientIP) {
457
return types.RegisterResponse{}, errIPBanned
458
}
@@ -474,8 +472,8 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
472
}
473
474
ttl := s.cfg.LeaseTTL
477
- if req.TTLSeconds > 0 {
478
- ttl = time.Duration(req.TTLSeconds) * time.Second
475
+ if req.TTL > 0 {
476
+ ttl = time.Duration(req.TTL) * time.Second
477
}
478
479
leaseID := randomID("lease_")
@@ -528,8 +526,8 @@ func (s *Server) renewLease(req types.RenewRequest, clientIP string) (types.Rene
526
}
527
528
ttl := s.cfg.LeaseTTL
531
- if req.TTLSeconds > 0 {
532
- ttl = time.Duration(req.TTLSeconds) * time.Second
529
+ if req.TTL > 0 {
530
+ ttl = time.Duration(req.TTL) * time.Second
531
}
532
record.ExpiresAt = time.Now().Add(ttl)
533
record.LastSeenAt = time.Now()
sdk/client_test.go
deleted
-183
@@ -1,183 +0,0 @@
1
-package sdk
2
-
3
-import (
4
- "context"
5
- "crypto/tls"
6
- "encoding/json"
7
- "encoding/pem"
8
- "errors"
9
- "net/http"
10
- "net/http/httptest"
11
- "net/url"
12
- "testing"
13
- "time"
14
-
15
- "github.com/gosuda/portal/v2/types"
16
-)
17
-
18
-func TestNewRelayClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
19
- t.Parallel()
20
-
21
- server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
22
- w.WriteHeader(http.StatusOK)
23
- }))
24
- defer server.Close()
25
-
26
- client, err := NewRelayClient(server.URL)
27
- if err != nil {
28
- t.Fatalf("NewRelayClient() error = %v", err)
29
- }
30
- defer client.Close()
31
-
32
- resp, err := client.httpClient.Get(server.URL)
33
- if err != nil {
34
- t.Fatalf("httpClient.Get() error = %v", err)
35
- }
36
- _ = resp.Body.Close()
37
-
38
- if resp.StatusCode != http.StatusOK {
39
- t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusOK)
40
- }
41
-}
42
-
43
-func TestNewRelayClientAppliesOptions(t *testing.T) {
44
- t.Parallel()
45
-
46
- server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
47
- w.WriteHeader(http.StatusOK)
48
- }))
49
- defer server.Close()
50
-
51
- rootCAPEM := pem.EncodeToMemory(&pem.Block{
52
- Type: "CERTIFICATE",
53
- Bytes: server.Certificate().Raw,
54
- })
55
- client, err := NewRelayClient(
56
- "https://relay.example.com/base/",
57
- WithRootCAPEM(rootCAPEM),
58
- WithInsecureSkipVerify(true),
59
- WithDialTimeout(2*time.Second),
60
- WithRequestTimeout(3*time.Second),
61
- WithHandshakeTimeout(4*time.Second),
62
- WithLeaseTTL(5*time.Minute),
63
- WithRenewBefore(45*time.Second),
64
- WithReadyTarget(3),
65
- )
66
- if err != nil {
67
- t.Fatalf("NewRelayClient() error = %v", err)
68
- }
69
- defer client.Close()
70
-
71
- if got := client.baseURL.String(); got != "https://relay.example.com/base" {
72
- t.Fatalf("baseURL.String() = %q, want %q", got, "https://relay.example.com/base")
73
- }
74
- if !client.insecureSkipVerify {
75
- t.Fatal("insecureSkipVerify = false, want true")
76
- }
77
- if client.dialTimeout != 2*time.Second {
78
- t.Fatalf("dialTimeout = %v, want %v", client.dialTimeout, 2*time.Second)
79
- }
80
- if client.requestTimeout != 3*time.Second {
81
- t.Fatalf("requestTimeout = %v, want %v", client.requestTimeout, 3*time.Second)
82
- }
83
- if client.handshakeTimeout != 4*time.Second {
84
- t.Fatalf("handshakeTimeout = %v, want %v", client.handshakeTimeout, 4*time.Second)
85
- }
86
- if client.leaseTTL != 5*time.Minute {
87
- t.Fatalf("leaseTTL = %v, want %v", client.leaseTTL, 5*time.Minute)
88
- }
89
- if client.renewBefore != 45*time.Second {
90
- t.Fatalf("renewBefore = %v, want %v", client.renewBefore, 45*time.Second)
91
- }
92
- if client.readyTarget != 3 {
93
- t.Fatalf("readyTarget = %d, want %d", client.readyTarget, 3)
94
- }
95
- if client.httpClient.Timeout != 3*time.Second {
96
- t.Fatalf("httpClient.Timeout = %v, want %v", client.httpClient.Timeout, 3*time.Second)
97
- }
98
- if !client.rawTLSConfig.InsecureSkipVerify {
99
- t.Fatal("rawTLSConfig.InsecureSkipVerify = false, want true")
100
- }
101
- if string(client.rootCAPEM) != string(rootCAPEM) {
102
- t.Fatalf("rootCAPEM = %q, want copied PEM input", string(client.rootCAPEM))
103
- }
104
-
105
- rootCAPEM[0] = 'X'
106
- if string(client.rootCAPEM) == string(rootCAPEM) {
107
- t.Fatalf("rootCAPEM changed with caller slice mutation: %q", string(client.rootCAPEM))
108
- }
109
- if len(client.rootCAPEM) == 0 || client.rootCAPEM[0] != '-' {
110
- t.Fatalf("rootCAPEM changed with caller slice mutation: %q", string(client.rootCAPEM))
111
- }
112
-}
113
-
114
-func TestOpenReverseSessionPreservesAPIErrorCode(t *testing.T) {
115
- t.Parallel()
116
-
117
- server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
118
- if r.URL.Path != types.PathSDKConnect {
119
- t.Fatalf("request path = %q, want %q", r.URL.Path, types.PathSDKConnect)
120
- }
121
- if got := r.URL.Query().Get("lease_id"); got != "lease-123" {
122
- t.Fatalf("lease_id = %q, want %q", got, "lease-123")
123
- }
124
-
125
- w.Header().Set("Content-Type", "application/json")
126
- w.WriteHeader(http.StatusForbidden)
127
- if flusher, ok := w.(http.Flusher); ok {
128
- flusher.Flush()
129
- }
130
- time.Sleep(25 * time.Millisecond)
131
- _ = json.NewEncoder(w).Encode(types.APIEnvelope[any]{
132
- OK: false,
133
- Error: &types.APIError{Code: "unauthorized", Message: "bad reverse token"},
134
- })
135
- }))
136
- server.EnableHTTP2 = false
137
- server.StartTLS()
138
- defer server.Close()
139
-
140
- baseURL, err := url.Parse(server.URL)
141
- if err != nil {
142
- t.Fatalf("url.Parse() error = %v", err)
143
- }
144
-
145
- transport, ok := server.Client().Transport.(*http.Transport)
146
- if !ok {
147
- t.Fatalf("server client transport type = %T, want *http.Transport", server.Client().Transport)
148
- }
149
-
150
- client := &RelayClient{
151
- baseURL: baseURL,
152
- httpClient: &http.Client{
153
- Transport: transport.Clone(),
154
- Timeout: defaultRequestTimeout,
155
- },
156
- rawTLSConfig: &tls.Config{
157
- MinVersion: tls.VersionTLS12,
158
- ServerName: baseURL.Hostname(),
159
- RootCAs: transport.TLSClientConfig.RootCAs,
160
- NextProtos: []string{"http/1.1"},
161
- },
162
- dialTimeout: defaultDialTimeout,
163
- }
164
-
165
- _, err = client.openReverseSession(context.Background(), "lease-123", "tok_123")
166
- if err == nil {
167
- t.Fatal("openReverseSession() error = nil, want APIRequestError")
168
- }
169
-
170
- var apiErr *types.APIRequestError
171
- if !errors.As(err, &apiErr) {
172
- t.Fatalf("openReverseSession() error = %T, want *types.APIRequestError", err)
173
- }
174
- if apiErr.StatusCode != http.StatusForbidden {
175
- t.Fatalf("APIRequestError.StatusCode = %d, want %d", apiErr.StatusCode, http.StatusForbidden)
176
- }
177
- if apiErr.Code != "unauthorized" {
178
- t.Fatalf("APIRequestError.Code = %q, want %q", apiErr.Code, "unauthorized")
179
- }
180
- if apiErr.Message != "bad reverse token" {
181
- t.Fatalf("APIRequestError.Message = %q, want %q", apiErr.Message, "bad reverse token")
182
- }
183
-}
sdk/expose.go
renamed
+2
-26
@@ -45,39 +45,21 @@ func Expose(ctx context.Context, relayUrls []string, name string, metadata types
45
if relay.listener != nil {
46
closeErr = errors.Join(closeErr, relay.listener.Close())
47
}
48
- if relay.client != nil {
49
- relay.client.Close()
50
- }
48
}
49
return closeErr
50
}
51
52
for _, relayURL := range relayURLs {
56
- client, err := NewRelayClient(relayURL)
57
- if err != nil {
58
- return nil, errors.Join(fmt.Errorf("new client %q: %w", relayURL, err), cleanup())
59
- }
60
-
61
- listener, err := NewListener(ctx, ListenRequest{
53
+ listener, err := NewListener(ctx, relayURL, ListenerConfig{
54
Name: name,
55
Metadata: metadata,
64
- }, listenerOptions{
65
- client: client,
66
- handshakeTimeout: client.handshakeTimeout,
67
- renewBefore: client.renewBefore,
68
- retryCount: defaultListenerRetryCount,
69
- retryDelay: defaultListenerRetryDelay,
70
- defaultLeaseTTL: client.leaseTTL,
71
- defaultReadyTarget: client.readyTarget,
56
})
57
if err != nil {
74
- client.Close()
58
return nil, errors.Join(fmt.Errorf("listen %q: %w", relayURL, err), cleanup())
59
}
60
61
relays = append(relays, exposureRelay{
62
relayURL: relayURL,
80
- client: client,
63
listener: listener,
64
})
65
}
@@ -198,7 +180,7 @@ func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr
180
return RunHTTP(ctx, relayListener, handler, localAddr)
181
}
182
201
-// Close closes the merged listener and all underlying relay listeners and clients.
183
+// Close closes the merged listener and all underlying relay listeners.
184
func (e *Exposure) Close() error {
185
if e == nil {
186
return nil
@@ -209,11 +191,6 @@ func (e *Exposure) Close() error {
191
if e.listener != nil {
192
closeErr = errors.Join(closeErr, e.listener.Close())
193
}
212
- for _, relay := range e.relays {
213
- if relay.client != nil {
214
- relay.client.Close()
215
- }
216
- }
194
195
logger := log.With().Str("component", "sdk-exposure").Logger()
196
event := logger.Info().
@@ -338,7 +315,6 @@ func RunHTTP(ctx context.Context, relayListener net.Listener, handler http.Handl
315
316
type exposureRelay struct {
317
relayURL string
341
- client *RelayClient
318
listener *Listener
319
}
320
sdk/helper_test.go
deleted
-638
@@ -1,638 +0,0 @@
1
-package sdk
2
-
3
-import (
4
- "context"
5
- "errors"
6
- "io"
7
- "net"
8
- "net/http"
9
- "reflect"
10
- "sort"
11
- "testing"
12
- "time"
13
-
14
- "github.com/gosuda/portal/v2/types"
15
-)
16
-
17
-func TestRunHTTPRelayOnly(t *testing.T) {
18
- t.Parallel()
19
-
20
- listener, err := net.Listen("tcp", "127.0.0.1:0")
21
- if err != nil {
22
- t.Fatalf("Listen() error = %v", err)
23
- }
24
- defer listener.Close()
25
-
26
- ctx, cancel := context.WithCancel(context.Background())
27
- defer cancel()
28
-
29
- errCh := make(chan error, 1)
30
- go func() {
31
- errCh <- RunHTTP(ctx, listener, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
32
- _, _ = io.WriteString(w, "ok")
33
- }), "")
34
- }()
35
-
36
- waitForHTTP(t, "http://"+listener.Addr().String())
37
- cancel()
38
-
39
- select {
40
- case err := <-errCh:
41
- if err != nil {
42
- t.Fatalf("RunHTTP() error = %v", err)
43
- }
44
- case <-time.After(3 * time.Second):
45
- t.Fatal("RunHTTP() did not exit after context cancellation")
46
- }
47
-}
48
-
49
-func TestExposureRunHTTPLocalOnly(t *testing.T) {
50
- t.Parallel()
51
-
52
- localListener, err := net.Listen("tcp", "127.0.0.1:0")
53
- if err != nil {
54
- t.Fatalf("Listen() error = %v", err)
55
- }
56
- localAddr := localListener.Addr().String()
57
- _ = localListener.Close()
58
-
59
- ctx, cancel := context.WithCancel(context.Background())
60
- defer cancel()
61
-
62
- var exposure *Exposure
63
- errCh := make(chan error, 1)
64
- go func() {
65
- errCh <- exposure.RunHTTP(ctx, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
66
- _, _ = io.WriteString(w, "ok")
67
- }), localAddr)
68
- }()
69
-
70
- waitForHTTP(t, "http://"+localAddr)
71
- cancel()
72
-
73
- select {
74
- case err := <-errCh:
75
- if err != nil {
76
- t.Fatalf("RunHTTP() error = %v", err)
77
- }
78
- case <-time.After(3 * time.Second):
79
- t.Fatal("RunHTTP() did not exit after context cancellation")
80
- }
81
-}
82
-
83
-func TestRunHTTPRequiresRelayOrLocal(t *testing.T) {
84
- t.Parallel()
85
-
86
- err := RunHTTP(context.Background(), nil, http.HandlerFunc(func(http.ResponseWriter, *http.Request) {}), "")
87
- if err == nil {
88
- t.Fatal("RunHTTP() error = nil, want error")
89
- }
90
-}
91
-
92
-func TestRunHTTPLocalAndRelay(t *testing.T) {
93
- t.Parallel()
94
-
95
- relayListener, err := net.Listen("tcp", "127.0.0.1:0")
96
- if err != nil {
97
- t.Fatalf("Listen() error = %v", err)
98
- }
99
- defer relayListener.Close()
100
-
101
- localListener, err := net.Listen("tcp", "127.0.0.1:0")
102
- if err != nil {
103
- t.Fatalf("Listen() error = %v", err)
104
- }
105
- localAddr := localListener.Addr().String()
106
- _ = localListener.Close()
107
-
108
- ctx, cancel := context.WithCancel(context.Background())
109
- defer cancel()
110
-
111
- errCh := make(chan error, 1)
112
- go func() {
113
- errCh <- RunHTTP(ctx, relayListener, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
114
- _, _ = io.WriteString(w, "ok")
115
- }), localAddr)
116
- }()
117
-
118
- waitForHTTP(t, "http://"+relayListener.Addr().String())
119
- waitForHTTP(t, "http://"+localAddr)
120
- cancel()
121
-
122
- select {
123
- case err := <-errCh:
124
- if err != nil {
125
- t.Fatalf("RunHTTP() error = %v", err)
126
- }
127
- case <-time.After(3 * time.Second):
128
- t.Fatal("RunHTTP() did not exit after context cancellation")
129
- }
130
-}
131
-
132
-func TestRunHTTPRelayListenerCloseIsNormal(t *testing.T) {
133
- t.Parallel()
134
-
135
- listener, err := net.Listen("tcp", "127.0.0.1:0")
136
- if err != nil {
137
- t.Fatalf("Listen() error = %v", err)
138
- }
139
-
140
- ctx, cancel := context.WithCancel(context.Background())
141
- defer cancel()
142
-
143
- errCh := make(chan error, 1)
144
- go func() {
145
- errCh <- RunHTTP(ctx, listener, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
146
- _, _ = io.WriteString(w, "ok")
147
- }), "")
148
- }()
149
-
150
- waitForHTTP(t, "http://"+listener.Addr().String())
151
-
152
- if err := listener.Close(); err != nil {
153
- t.Fatalf("listener.Close() error = %v", err)
154
- }
155
-
156
- select {
157
- case err := <-errCh:
158
- if err != nil {
159
- t.Fatalf("RunHTTP() error = %v", err)
160
- }
161
- case <-time.After(3 * time.Second):
162
- t.Fatal("RunHTTP() did not exit after listener close")
163
- }
164
-}
165
-
166
-func TestRunHTTPRelayListenerCloseKeepsLocalRunning(t *testing.T) {
167
- t.Parallel()
168
-
169
- relayListener, err := net.Listen("tcp", "127.0.0.1:0")
170
- if err != nil {
171
- t.Fatalf("Listen() error = %v", err)
172
- }
173
-
174
- localListener, err := net.Listen("tcp", "127.0.0.1:0")
175
- if err != nil {
176
- t.Fatalf("Listen() error = %v", err)
177
- }
178
- localAddr := localListener.Addr().String()
179
- _ = localListener.Close()
180
-
181
- ctx, cancel := context.WithCancel(context.Background())
182
- defer cancel()
183
-
184
- errCh := make(chan error, 1)
185
- go func() {
186
- errCh <- RunHTTP(ctx, relayListener, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
187
- _, _ = io.WriteString(w, "ok")
188
- }), localAddr)
189
- }()
190
-
191
- waitForHTTP(t, "http://"+relayListener.Addr().String())
192
- waitForHTTP(t, "http://"+localAddr)
193
-
194
- if err := relayListener.Close(); err != nil {
195
- t.Fatalf("relayListener.Close() error = %v", err)
196
- }
197
-
198
- waitForHTTP(t, "http://"+localAddr)
199
-
200
- select {
201
- case err := <-errCh:
202
- t.Fatalf("RunHTTP() exited early with %v", err)
203
- case <-time.After(200 * time.Millisecond):
204
- }
205
-
206
- cancel()
207
-
208
- select {
209
- case err := <-errCh:
210
- if err != nil {
211
- t.Fatalf("RunHTTP() error = %v", err)
212
- }
213
- case <-time.After(3 * time.Second):
214
- t.Fatal("RunHTTP() did not exit after context cancellation")
215
- }
216
-}
217
-
218
-func waitForHTTP(t *testing.T, rawURL string) {
219
- t.Helper()
220
-
221
- client := &http.Client{Timeout: 200 * time.Millisecond}
222
- deadline := time.Now().Add(3 * time.Second)
223
- for time.Now().Before(deadline) {
224
- resp, err := client.Get(rawURL)
225
- if err == nil {
226
- _ = resp.Body.Close()
227
- if resp.StatusCode == http.StatusOK {
228
- return
229
- }
230
- }
231
- time.Sleep(50 * time.Millisecond)
232
- }
233
-
234
- t.Fatalf("timed out waiting for %s", rawURL)
235
-}
236
-
237
-func TestMergeListenersRequiresInput(t *testing.T) {
238
- t.Parallel()
239
-
240
- listener, err := mergeListeners()
241
- if err == nil {
242
- t.Fatal("mergeListeners() error = nil, want error")
243
- }
244
- if listener != nil {
245
- t.Fatalf("mergeListeners() listener = %#v, want nil", listener)
246
- }
247
-}
248
-
249
-func TestMergeListenersAcceptsFromAllSources(t *testing.T) {
250
- t.Parallel()
251
-
252
- listener1, err := net.Listen("tcp", "127.0.0.1:0")
253
- if err != nil {
254
- t.Fatalf("Listen() error = %v", err)
255
- }
256
- defer listener1.Close()
257
-
258
- listener2, err := net.Listen("tcp", "127.0.0.1:0")
259
- if err != nil {
260
- t.Fatalf("Listen() error = %v", err)
261
- }
262
- defer listener2.Close()
263
-
264
- merged, err := mergeListeners(listener1, listener2)
265
- if err != nil {
266
- t.Fatalf("mergeListeners() error = %v", err)
267
- }
268
- defer merged.Close()
269
-
270
- client1, err := net.Dial("tcp", listener1.Addr().String())
271
- if err != nil {
272
- t.Fatalf("Dial() error = %v", err)
273
- }
274
- defer client1.Close()
275
-
276
- client2, err := net.Dial("tcp", listener2.Addr().String())
277
- if err != nil {
278
- t.Fatalf("Dial() error = %v", err)
279
- }
280
- defer client2.Close()
281
-
282
- got := make([]string, 0, 2)
283
- for range 2 {
284
- conn, err := merged.Accept()
285
- if err != nil {
286
- t.Fatalf("Accept() error = %v", err)
287
- }
288
- got = append(got, conn.LocalAddr().String())
289
- _ = conn.Close()
290
- }
291
-
292
- want := []string{listener1.Addr().String(), listener2.Addr().String()}
293
- sort.Strings(got)
294
- sort.Strings(want)
295
- if len(got) != len(want) || got[0] != want[0] || got[1] != want[1] {
296
- t.Fatalf("accepted local addrs = %v, want %v", got, want)
297
- }
298
-}
299
-
300
-func TestMergeListenersContinuesAfterOneSourceCloses(t *testing.T) {
301
- t.Parallel()
302
-
303
- listener1, err := net.Listen("tcp", "127.0.0.1:0")
304
- if err != nil {
305
- t.Fatalf("Listen() error = %v", err)
306
- }
307
- defer listener1.Close()
308
-
309
- listener2, err := net.Listen("tcp", "127.0.0.1:0")
310
- if err != nil {
311
- t.Fatalf("Listen() error = %v", err)
312
- }
313
- defer listener2.Close()
314
-
315
- merged, err := mergeListeners(listener1, listener2)
316
- if err != nil {
317
- t.Fatalf("mergeListeners() error = %v", err)
318
- }
319
- defer merged.Close()
320
-
321
- if err := listener1.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
322
- t.Fatalf("listener1.Close() error = %v", err)
323
- }
324
-
325
- acceptCh := make(chan net.Conn, 1)
326
- errCh := make(chan error, 1)
327
- go func() {
328
- conn, acceptErr := merged.Accept()
329
- if acceptErr != nil {
330
- errCh <- acceptErr
331
- return
332
- }
333
- acceptCh <- conn
334
- }()
335
-
336
- client, err := net.Dial("tcp", listener2.Addr().String())
337
- if err != nil {
338
- t.Fatalf("Dial() error = %v", err)
339
- }
340
- defer client.Close()
341
-
342
- select {
343
- case acceptErr := <-errCh:
344
- t.Fatalf("Accept() error = %v", acceptErr)
345
- case conn := <-acceptCh:
346
- if conn.LocalAddr().String() != listener2.Addr().String() {
347
- t.Fatalf("conn.LocalAddr() = %q, want %q", conn.LocalAddr().String(), listener2.Addr().String())
348
- }
349
- _ = conn.Close()
350
- case <-time.After(3 * time.Second):
351
- t.Fatal("Accept() did not return connection from surviving listener")
352
- }
353
-}
354
-
355
-func TestMergeListenersCloseUnblocksAccept(t *testing.T) {
356
- t.Parallel()
357
-
358
- listener, err := net.Listen("tcp", "127.0.0.1:0")
359
- if err != nil {
360
- t.Fatalf("Listen() error = %v", err)
361
- }
362
- defer listener.Close()
363
-
364
- merged, err := mergeListeners(listener)
365
- if err != nil {
366
- t.Fatalf("mergeListeners() error = %v", err)
367
- }
368
-
369
- errCh := make(chan error, 1)
370
- go func() {
371
- _, acceptErr := merged.Accept()
372
- errCh <- acceptErr
373
- }()
374
-
375
- if err := merged.Close(); err != nil {
376
- t.Fatalf("Close() error = %v", err)
377
- }
378
-
379
- select {
380
- case acceptErr := <-errCh:
381
- if !errors.Is(acceptErr, net.ErrClosed) {
382
- t.Fatalf("Accept() error = %v, want net.ErrClosed", acceptErr)
383
- }
384
- case <-time.After(3 * time.Second):
385
- t.Fatal("Accept() did not unblock after Close()")
386
- }
387
-}
388
-
389
-func TestMergedListenerAcceptDoesNotReturnQueuedConnAfterClose(t *testing.T) {
390
- t.Parallel()
391
-
392
- conn := &stubConn{}
393
- closed := make(chan struct{})
394
- close(closed)
395
-
396
- merged := &mergedListener{
397
- accepted: make(chan net.Conn, 1),
398
- closed: closed,
399
- }
400
- merged.accepted <- conn
401
-
402
- gotConn, err := merged.Accept()
403
- if gotConn != nil {
404
- t.Fatalf("Accept() conn = %#v, want nil", gotConn)
405
- }
406
- if !errors.Is(err, net.ErrClosed) {
407
- t.Fatalf("Accept() error = %v, want net.ErrClosed", err)
408
- }
409
- if conn.closeCount != 1 {
410
- t.Fatalf("conn close count = %d, want 1", conn.closeCount)
411
- }
412
-}
413
-
414
-func TestNormalizeRelayURLs(t *testing.T) {
415
- t.Parallel()
416
-
417
- got, err := NormalizeRelayURLs([]string{
418
- " localhost:4017 , https://relay.example.com/base/relay?x=1#frag ",
419
- "https://relay.example.com/base",
420
- })
421
- if err != nil {
422
- t.Fatalf("NormalizeRelayURLs() error = %v", err)
423
- }
424
-
425
- want := []string{
426
- "https://localhost:4017",
427
- "https://relay.example.com/base",
428
- }
429
- if !reflect.DeepEqual(got, want) {
430
- t.Fatalf("NormalizeRelayURLs() = %v, want %v", got, want)
431
- }
432
-}
433
-
434
-func TestNormalizeRelayURLRejectsNonHTTPS(t *testing.T) {
435
- t.Parallel()
436
-
437
- _, err := NormalizeRelayURL("http://relay.example.com")
438
- if err == nil {
439
- t.Fatal("NormalizeRelayURL() error = nil, want error")
440
- }
441
-}
442
-
443
-func TestNormalizeTargetAddr(t *testing.T) {
444
- t.Parallel()
445
-
446
- tests := []struct {
447
- name string
448
- input string
449
- want string
450
- }{
451
- {
452
- name: "host only",
453
- input: "localhost",
454
- want: "localhost:80",
455
- },
456
- {
457
- name: "host and port",
458
- input: "127.0.0.1:8080",
459
- want: "127.0.0.1:8080",
460
- },
461
- {
462
- name: "http url",
463
- input: "http://localhost:3000",
464
- want: "localhost:3000",
465
- },
466
- {
467
- name: "https url default port preserved by host parsing",
468
- input: "https://example.com",
469
- want: "example.com:80",
470
- },
471
- {
472
- name: "ipv6 host",
473
- input: "::1",
474
- want: "[::1]:80",
475
- },
476
- }
477
-
478
- for _, tt := range tests {
479
- t.Run(tt.name, func(t *testing.T) {
480
- t.Parallel()
481
-
482
- got, err := NormalizeTargetAddr(tt.input)
483
- if err != nil {
484
- t.Fatalf("NormalizeTargetAddr() error = %v", err)
485
- }
486
- if got != tt.want {
487
- t.Fatalf("NormalizeTargetAddr() = %q, want %q", got, tt.want)
488
- }
489
- })
490
- }
491
-}
492
-
493
-func TestNormalizeTargetAddrRejectsInvalidInput(t *testing.T) {
494
- t.Parallel()
495
-
496
- inputs := []string{
497
- "",
498
- "ftp://example.com",
499
- "http://example.com/path",
500
- "http://example.com?a=1",
501
- "http://example.com#frag",
502
- "host:port:extra",
503
- }
504
-
505
- for _, input := range inputs {
506
- t.Run(input, func(t *testing.T) {
507
- t.Parallel()
508
-
509
- if _, err := NormalizeTargetAddr(input); err == nil {
510
- t.Fatalf("NormalizeTargetAddr(%q) error = nil, want error", input)
511
- }
512
- })
513
- }
514
-}
515
-
516
-func TestExposeNoRelayInputs(t *testing.T) {
517
- t.Parallel()
518
-
519
- exposure, err := Expose(context.Background(), nil, "demo", types.LeaseMetadata{})
520
- if err != nil {
521
- t.Fatalf("Expose() error = %v", err)
522
- }
523
- if exposure != nil {
524
- t.Fatalf("Expose() exposure = %#v, want nil", exposure)
525
- }
526
-}
527
-
528
-func TestExposureAccessorsReturnCopies(t *testing.T) {
529
- t.Parallel()
530
-
531
- exposure := &Exposure{
532
- relays: []exposureRelay{
533
- {
534
- relayURL: "https://relay-1.example.com",
535
- listener: &Listener{
536
- hostnames: []string{"app.example.com"},
537
- state: listenerStateReady,
538
- },
539
- },
540
- {
541
- relayURL: "https://relay-2.example.com",
542
- listener: &Listener{
543
- hostnames: []string{"app.example.com"},
544
- state: listenerStateReady,
545
- },
546
- },
547
- },
548
- }
549
-
550
- relayURLs := exposure.RelayURLs()
551
- publicURLs := exposure.PublicURLs()
552
-
553
- relayURLs[0] = "changed"
554
- publicURLs[0] = "changed"
555
-
556
- if got, want := exposure.RelayURLs(), []string{"https://relay-1.example.com", "https://relay-2.example.com"}; !reflect.DeepEqual(got, want) {
557
- t.Fatalf("RelayURLs() = %v, want %v", got, want)
558
- }
559
- if got, want := exposure.PublicURLs(), []string{"https://app.example.com"}; !reflect.DeepEqual(got, want) {
560
- t.Fatalf("PublicURLs() = %v, want %v", got, want)
561
- }
562
- if got, want := exposure.relays[0].relayURL, "https://relay-1.example.com"; got != want {
563
- t.Fatalf("relays[0].relayURL = %q, want %q", got, want)
564
- }
565
- if got, want := exposure.relays[0].listener.publicURLs(), []string{"https://app.example.com"}; !reflect.DeepEqual(got, want) {
566
- t.Fatalf("relays[0].listener.publicURLs() = %v, want %v", got, want)
567
- }
568
-}
569
-
570
-func TestExposureCloseIsIdempotent(t *testing.T) {
571
- t.Parallel()
572
-
573
- listener := &stubListener{}
574
- exposure := &Exposure{listener: listener}
575
-
576
- if err := exposure.Close(); err != nil {
577
- t.Fatalf("Close() error = %v", err)
578
- }
579
- if err := exposure.Close(); err != nil {
580
- t.Fatalf("second Close() error = %v", err)
581
- }
582
- if listener.closeCount != 1 {
583
- t.Fatalf("listener close count = %d, want 1", listener.closeCount)
584
- }
585
-}
586
-
587
-func TestExposureImplementsListener(t *testing.T) {
588
- t.Parallel()
589
-
590
- listener := &stubListener{addr: listenerAddr("merged:test")}
591
- exposure := &Exposure{listener: listener}
592
-
593
- if got := exposure.Addr().String(); got != "merged:test" {
594
- t.Fatalf("Addr().String() = %q, want %q", got, "merged:test")
595
- }
596
-
597
- _, err := exposure.Accept()
598
- if !errors.Is(err, net.ErrClosed) {
599
- t.Fatalf("Accept() error = %v, want net.ErrClosed", err)
600
- }
601
-}
602
-
603
-type stubListener struct {
604
- closeCount int
605
- addr net.Addr
606
-}
607
-
608
-func (l *stubListener) Accept() (net.Conn, error) {
609
- return nil, net.ErrClosed
610
-}
611
-
612
-func (l *stubListener) Close() error {
613
- l.closeCount++
614
- if l.closeCount > 1 {
615
- return errors.New("listener closed more than once")
616
- }
617
- return nil
618
-}
619
-
620
-func (l *stubListener) Addr() net.Addr {
621
- if l.addr != nil {
622
- return l.addr
623
- }
624
- return listenerAddr("stub")
625
-}
626
-
627
-type stubConn struct {
628
- closeCount int
629
-}
630
-
631
-func (c *stubConn) Read(_ []byte) (int, error) { return 0, io.EOF }
632
-func (c *stubConn) Write(b []byte) (int, error) { return len(b), nil }
633
-func (c *stubConn) Close() error { c.closeCount++; return nil }
634
-func (c *stubConn) LocalAddr() net.Addr { return listenerAddr("stub-local") }
635
-func (c *stubConn) RemoteAddr() net.Addr { return listenerAddr("stub-remote") }
636
-func (c *stubConn) SetDeadline(_ time.Time) error { return nil }
637
-func (c *stubConn) SetReadDeadline(_ time.Time) error { return nil }
638
-func (c *stubConn) SetWriteDeadline(_ time.Time) error { return nil }
sdk/listener.go
+143
-213
@@ -17,8 +17,7 @@ import (
17
)
18
19
const (
20
- defaultListenerRetryCount = 30
21
- defaultListenerRetryDelay = time.Second
20
+ defaultListenerRetryDelay = 1 * time.Second
21
)
22
23
type listenerState uint8
@@ -26,27 +25,16 @@ type listenerState uint8
25
const (
26
listenerStatePending listenerState = iota
27
listenerStateReady
29
- listenerStateStale
28
listenerStateClosed
29
)
30
33
-type ListenRequest struct {
31
+type ListenerConfig struct {
32
Name string
33
ReverseToken string
34
Hostnames []string
35
Metadata types.LeaseMetadata
38
- ReadyTarget int
39
- LeaseTTL time.Duration
40
-}
41
-
42
-type listenerOptions struct {
43
- client *RelayClient
44
- handshakeTimeout time.Duration
45
- renewBefore time.Duration
46
- retryCount int
47
- retryDelay time.Duration
48
- defaultLeaseTTL time.Duration
49
- defaultReadyTarget int
36
+ RootCAPEM []byte
37
+ RetryCount int
38
}
39
40
type Listener struct {
@@ -55,9 +43,6 @@ type Listener struct {
43
accepted chan net.Conn
44
refill chan struct{}
45
58
- name string
59
- reverseToken string
60
- metadata types.LeaseMetadata
46
readyTarget int
47
leaseTTL time.Duration
48
renewInterval time.Duration
@@ -66,7 +51,7 @@ type Listener struct {
51
retryDelay time.Duration
52
53
mu sync.Mutex
69
- client *RelayClient
54
+ api *relayClient
55
leaseID string
56
hostnames []string
57
tlsConfig *tls.Config
@@ -74,53 +59,25 @@ type Listener struct {
59
activeSessions int
60
sessionFailures int
61
state listenerState
77
- runID uint64
62
63
closeOnce sync.Once
64
closeErr error
65
}
66
83
-func NewListener(ctx context.Context, req ListenRequest, opts listenerOptions) (*Listener, error) {
84
- if strings.TrimSpace(req.Name) == "" {
85
- return nil, errors.New("listener name is required")
86
- }
87
- if opts.client == nil {
88
- return nil, errors.New("listener client is required")
67
+// NewListener creates one relay listener and its dedicated relay transport for one relay URL.
68
+func NewListener(ctx context.Context, relayURL string, cfg ListenerConfig) (*Listener, error) {
69
+ api, err := newRelayClient(relayURL, cfg)
70
+ if err != nil {
71
+ return nil, err
72
}
73
if ctx == nil {
74
ctx = context.Background()
75
}
76
94
- reverseToken := strings.TrimSpace(req.ReverseToken)
95
- if reverseToken == "" {
96
- reverseToken = randomToken()
97
- }
98
-
99
- readyTarget := req.ReadyTarget
100
- if readyTarget <= 0 {
101
- readyTarget = opts.defaultReadyTarget
102
- }
103
- if readyTarget <= 0 {
104
- readyTarget = defaultReadyTarget
105
- }
106
-
107
- leaseTTL := req.LeaseTTL
108
- if leaseTTL <= 0 {
109
- leaseTTL = opts.defaultLeaseTTL
110
- }
111
- if leaseTTL <= 0 {
112
- leaseTTL = defaultLeaseTTL
113
- }
114
-
115
- handshakeTimeout := opts.handshakeTimeout
116
- if handshakeTimeout <= 0 {
117
- handshakeTimeout = defaultHandshakeTimeout
118
- }
119
-
120
- renewBefore := opts.renewBefore
121
- if renewBefore <= 0 {
122
- renewBefore = defaultRenewBefore
123
- }
77
+ readyTarget := defaultReadyTarget
78
+ leaseTTL := defaultLeaseTTL
79
+ handshakeTimeout := defaultHandshakeTimeout
80
+ renewBefore := defaultRenewBefore
81
renewInterval := leaseTTL / 2
82
if leaseTTL > renewBefore {
83
renewInterval = leaseTTL - renewBefore
@@ -132,11 +89,11 @@ func NewListener(ctx context.Context, req ListenRequest, opts listenerOptions) (
89
renewInterval = time.Second
90
}
91
135
- retryCount := opts.retryCount
136
- if retryCount <= 0 {
137
- retryCount = defaultListenerRetryCount
92
+ retryCount := cfg.RetryCount
93
+ retryDelay := defaultListenerRetryDelay
94
+ if retryCount < 0 {
95
+ retryCount = 0
96
}
139
- retryDelay := opts.retryDelay
97
if retryDelay <= 0 {
98
retryDelay = defaultListenerRetryDelay
99
}
@@ -147,22 +104,17 @@ func NewListener(ctx context.Context, req ListenRequest, opts listenerOptions) (
104
cancel: cancel,
105
accepted: make(chan net.Conn, max(readyTarget*2, 1)),
106
refill: make(chan struct{}, 1),
150
- name: strings.TrimSpace(req.Name),
151
- reverseToken: reverseToken,
152
- metadata: cloneMetadata(req.Metadata),
107
readyTarget: readyTarget,
108
leaseTTL: leaseTTL,
109
renewInterval: renewInterval,
110
handshakeTimeout: handshakeTimeout,
111
retryCount: retryCount,
112
retryDelay: retryDelay,
159
- client: opts.client,
160
- hostnames: append([]string(nil), req.Hostnames...),
113
+ api: api,
114
state: listenerStatePending,
162
- runID: 1,
115
}
116
165
- go l.run(listenerCtx, l.runID)
117
+ go l.run(listenerCtx)
118
return l, nil
119
}
120
@@ -192,11 +144,10 @@ func (l *Listener) Accept() (net.Conn, error) {
144
}
145
146
func (l *Listener) Close() error {
195
- var closeErr error
147
l.closeOnce.Do(func() {
197
- closeErr = l.closeCurrent()
148
+ l.closeErr = l.shutdown()
149
})
199
- return closeErr
150
+ return l.closeErr
151
}
152
153
func (l *Listener) Addr() net.Addr {
@@ -209,48 +160,18 @@ func (l *Listener) Addr() net.Addr {
160
return listenerAddr("portal:" + l.leaseID)
161
}
162
212
-func (l *Listener) Reactivate(ctx context.Context) error {
213
- if ctx == nil {
214
- ctx = context.Background()
215
- }
216
-
217
- l.mu.Lock()
218
- if l.state == listenerStateClosed {
219
- l.mu.Unlock()
220
- return net.ErrClosed
221
- }
222
- if l.state != listenerStateStale {
223
- l.mu.Unlock()
224
- return errors.New("listener is not stale")
225
- }
226
-
227
- runCtx, cancel := context.WithCancel(ctx)
228
- l.ctx = runCtx
229
- l.cancel = cancel
230
- l.state = listenerStatePending
231
- l.activeSessions = 0
232
- l.sessionFailures = 0
233
- l.runID++
234
- runID := l.runID
235
- l.mu.Unlock()
236
-
237
- l.drainAccepted()
238
- go l.run(runCtx, runID)
239
- return nil
240
-}
241
-
242
-func (l *Listener) run(runCtx context.Context, runID uint64) {
163
+func (l *Listener) run(runCtx context.Context) {
164
logger := log.With().
165
Str("component", "sdk-listener").
245
- Str("name", l.name).
166
+ Str("name", l.api.name).
167
Logger()
168
169
var err error
249
- for attempt := 1; attempt <= l.retryCount; attempt++ {
250
- err = l.establish(runID, runCtx)
170
+ for attempt := 1; ; attempt++ {
171
+ err = l.establish(runCtx)
172
if err == nil {
173
l.mu.Lock()
253
- if l.runID != runID {
174
+ if l.state == listenerStateClosed {
175
l.mu.Unlock()
176
return
177
}
@@ -263,8 +184,8 @@ func (l *Listener) run(runCtx context.Context, runID uint64) {
184
Strs("hostnames", hostnames).
185
Msg("listener connected")
186
266
- go l.runSessionPool(runCtx, runID)
267
- go l.runRenewLoop(runCtx, runID)
187
+ go l.runSessionPool(runCtx)
188
+ go l.runRenewLoop(runCtx)
189
l.signalRefill()
190
return
191
}
@@ -279,7 +200,7 @@ func (l *Listener) run(runCtx context.Context, runID uint64) {
200
Dur("retry_in", l.retryDelay).
201
Msg("listener bootstrap failed")
202
282
- if attempt == l.retryCount {
203
+ if l.retryLimitReached(attempt) {
204
break
205
}
206
if !sleepOrDone(runCtx, l.retryDelay) {
@@ -287,46 +208,37 @@ func (l *Listener) run(runCtx context.Context, runID uint64) {
208
}
209
}
210
290
- l.fail(runID, err, "listener bootstrap retry limit reached")
211
+ l.fail(err, "listener bootstrap retry limit reached")
212
}
213
293
-func (l *Listener) establish(runID uint64, runCtx context.Context) error {
214
+func (l *Listener) establish(runCtx context.Context) error {
215
l.mu.Lock()
295
- if l.runID != runID {
216
+ if l.state == listenerStateClosed {
217
l.mu.Unlock()
218
return context.Canceled
219
}
299
- client := l.client
300
- hostnames := append([]string(nil), l.hostnames...)
220
+ api := l.api
221
l.mu.Unlock()
222
303
- resp, err := client.registerLease(runCtx, types.RegisterRequest{
304
- Name: l.name,
305
- Hostnames: hostnames,
306
- Metadata: cloneMetadata(l.metadata),
307
- ReverseToken: l.reverseToken,
308
- TLS: true,
309
- TTLSeconds: int(l.leaseTTL / time.Second),
310
- })
223
+ resp, err := api.registerLease(runCtx, nil, l.leaseTTL)
224
if err != nil {
225
return err
226
}
227
315
- tlsConfig, tlsCloser, err := keyless.BuildClientTLSConfig(l.client.baseURL.String(), resp.Hostnames)
228
+ tlsConfig, tlsCloser, err := keyless.BuildClientTLSConfig(api.baseURL.String(), resp.Hostnames)
229
if err != nil {
317
- _ = client.unregisterLease(context.Background(), resp.LeaseID, l.reverseToken)
230
+ _ = api.unregisterLease(context.Background(), resp.LeaseID)
231
return err
232
}
233
234
l.mu.Lock()
322
- if l.runID != runID || l.state == listenerStateClosed {
235
+ if l.state == listenerStateClosed || runCtx.Err() != nil {
236
l.mu.Unlock()
324
- _ = client.unregisterLease(context.Background(), resp.LeaseID, l.reverseToken)
237
+ _ = api.unregisterLease(context.Background(), resp.LeaseID)
238
_ = tlsCloser.Close()
239
return context.Canceled
240
}
241
oldCloser := l.tlsCloser
329
- l.client = client
242
l.leaseID = resp.LeaseID
243
l.hostnames = append([]string(nil), resp.Hostnames...)
244
l.tlsConfig = tlsConfig
@@ -342,7 +254,7 @@ func (l *Listener) establish(runID uint64, runCtx context.Context) error {
254
return nil
255
}
256
345
-func (l *Listener) runSessionPool(runCtx context.Context, runID uint64) {
257
+func (l *Listener) runSessionPool(runCtx context.Context) {
258
for {
259
select {
260
case <-runCtx.Done():
@@ -352,9 +264,8 @@ func (l *Listener) runSessionPool(runCtx context.Context, runID uint64) {
264
265
for {
266
l.mu.Lock()
355
- ready := l.runID == runID &&
356
- l.state == listenerStateReady &&
357
- l.client != nil &&
267
+ ready := l.state == listenerStateReady &&
268
+ l.api != nil &&
269
strings.TrimSpace(l.leaseID) != "" &&
270
l.tlsConfig != nil &&
271
l.activeSessions < l.readyTarget
@@ -365,12 +276,12 @@ func (l *Listener) runSessionPool(runCtx context.Context, runID uint64) {
276
l.activeSessions++
277
l.mu.Unlock()
278
368
- go l.runSession(runCtx, runID)
279
+ go l.runSession(runCtx)
280
}
281
}
282
}
283
373
-func (l *Listener) runRenewLoop(runCtx context.Context, runID uint64) {
284
+func (l *Listener) runRenewLoop(runCtx context.Context) {
285
failures := 0
286
wait := l.renewInterval
287
@@ -380,19 +291,18 @@ func (l *Listener) runRenewLoop(runCtx context.Context, runID uint64) {
291
}
292
293
l.mu.Lock()
383
- current := l.runID == runID
384
- client := l.client
294
+ api := l.api
295
leaseID := l.leaseID
296
ready := l.state == listenerStateReady
297
l.mu.Unlock()
298
389
- if !current || !ready || client == nil || strings.TrimSpace(leaseID) == "" {
299
+ if !ready || api == nil || strings.TrimSpace(leaseID) == "" {
300
wait = l.renewInterval
301
continue
302
}
303
304
ctx, cancel := context.WithTimeout(runCtx, 10*time.Second)
395
- err := client.renewLease(ctx, leaseID, l.reverseToken, l.leaseTTL)
305
+ err := api.renewLease(ctx, leaseID, l.leaseTTL)
306
cancel()
307
308
if err == nil {
@@ -403,33 +313,39 @@ func (l *Listener) runRenewLoop(runCtx context.Context, runID uint64) {
313
314
if isLeaseNotFound(err) {
315
log.Warn().
316
+ Err(err).
317
Str("component", "sdk-listener").
318
Str("lease_id", leaseID).
319
Msg("lease not found on relay, attempting re-registration")
320
410
- err = l.establish(runID, runCtx)
411
- if err == nil {
321
+ if reregErr := l.reregister(runCtx); reregErr == nil {
322
failures = 0
323
wait = l.renewInterval
414
- l.signalRefill()
324
325
l.mu.Lock()
417
- leaseID = l.leaseID
418
- hostnames := append([]string(nil), l.hostnames...)
326
+ newLeaseID := l.leaseID
327
+ newHostnames := append([]string(nil), l.hostnames...)
328
l.mu.Unlock()
329
330
log.Info().
331
Str("component", "sdk-listener").
423
- Str("lease_id", leaseID).
424
- Strs("hostnames", hostnames).
332
+ Str("lease_id", newLeaseID).
333
+ Strs("hostnames", newHostnames).
334
Msg("lease re-registered successfully")
335
continue
336
+ } else {
337
+ err = reregErr
338
+ log.Error().
339
+ Err(reregErr).
340
+ Str("component", "sdk-listener").
341
+ Str("lease_id", leaseID).
342
+ Msg("lease re-registration failed")
343
}
344
}
345
346
failures++
347
event := log.Warn()
432
- if failures >= l.retryCount {
348
+ if l.retryLimitReached(failures) {
349
event = log.Error()
350
}
351
event.Err(err).
@@ -438,18 +354,18 @@ func (l *Listener) runRenewLoop(runCtx context.Context, runID uint64) {
354
Int("consecutive_failures", failures).
355
Msg("lease renewal failed")
356
441
- if failures >= l.retryCount {
442
- l.fail(runID, err, "listener renew retry limit reached")
357
+ if l.retryLimitReached(failures) {
358
+ l.fail(err, "listener renew retry limit reached")
359
return
360
}
361
wait = l.retryDelay
362
}
363
}
364
449
-func (l *Listener) runSession(runCtx context.Context, runID uint64) {
365
+func (l *Listener) runSession(runCtx context.Context) {
366
defer func() {
367
l.mu.Lock()
452
- if l.runID == runID && l.activeSessions > 0 {
368
+ if l.activeSessions > 0 {
369
l.activeSessions--
370
}
371
l.mu.Unlock()
@@ -462,7 +378,7 @@ func (l *Listener) runSession(runCtx context.Context, runID uint64) {
378
}
379
380
l.mu.Lock()
465
- if l.runID != runID {
381
+ if l.state == listenerStateClosed {
382
l.mu.Unlock()
383
return
384
}
@@ -470,8 +386,8 @@ func (l *Listener) runSession(runCtx context.Context, runID uint64) {
386
failures := l.sessionFailures
387
l.mu.Unlock()
388
473
- if failures >= l.retryCount {
474
- l.fail(runID, err, "listener session retry limit reached")
389
+ if l.retryLimitReached(failures) {
390
+ l.fail(err, "listener session retry limit reached")
391
return
392
}
393
@@ -479,16 +395,16 @@ func (l *Listener) runSession(runCtx context.Context, runID uint64) {
395
}
396
397
l.mu.Lock()
482
- ready := l.runID == runID && l.state == listenerStateReady
483
- client := l.client
398
+ ready := l.state == listenerStateReady
399
+ api := l.api
400
leaseID := l.leaseID
401
tlsConfig := l.tlsConfig
402
l.mu.Unlock()
487
- if !ready || client == nil || strings.TrimSpace(leaseID) == "" || tlsConfig == nil {
403
+ if !ready || api == nil || strings.TrimSpace(leaseID) == "" || tlsConfig == nil {
404
return
405
}
406
491
- conn, err := client.openReverseSession(runCtx, leaseID, l.reverseToken)
407
+ conn, err := api.openReverseSession(runCtx, leaseID)
408
if err != nil {
409
fail(err)
410
return
@@ -519,7 +435,9 @@ func (l *Listener) runSession(runCtx context.Context, runID uint64) {
435
}
436
437
l.mu.Lock()
522
- l.sessionFailures = 0
438
+ if l.state != listenerStateClosed {
439
+ l.sessionFailures = 0
440
+ }
441
l.mu.Unlock()
442
443
select {
@@ -551,64 +469,62 @@ func (l *Listener) publicURLs() []string {
469
return urls
470
}
471
554
-func (l *Listener) closeCurrent() error {
472
+func (l *Listener) reregister(runCtx context.Context) error {
473
l.mu.Lock()
556
- l.state = listenerStateClosed
557
- cancel := l.cancel
558
- client := l.client
559
- leaseID := l.leaseID
560
- tlsCloser := l.tlsCloser
561
- l.leaseID = ""
562
- l.tlsConfig = nil
563
- l.tlsCloser = nil
564
- l.activeSessions = 0
565
- l.sessionFailures = 0
566
- l.runID++
474
+ if l.state == listenerStateClosed {
475
+ l.mu.Unlock()
476
+ return context.Canceled
477
+ }
478
+ api := l.api
479
+ hostnames := append([]string(nil), l.hostnames...)
480
l.mu.Unlock()
481
569
- if cancel != nil {
570
- cancel()
571
- }
572
- l.drainAccepted()
482
+ ctx, cancel := context.WithTimeout(runCtx, 10*time.Second)
483
+ defer cancel()
484
574
- var closeErr error
575
- if client != nil && strings.TrimSpace(leaseID) != "" {
576
- ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
577
- closeErr = errors.Join(closeErr, client.unregisterLease(ctx, leaseID, l.reverseToken))
578
- cancel()
485
+ resp, err := api.registerLease(ctx, hostnames, l.leaseTTL)
486
+ if err != nil {
487
+ return err
488
}
580
- if tlsCloser != nil {
581
- closeErr = errors.Join(closeErr, tlsCloser.Close())
489
+
490
+ tlsConfig, tlsCloser, err := keyless.BuildClientTLSConfig(api.baseURL.String(), resp.Hostnames)
491
+ if err != nil {
492
+ _ = api.unregisterLease(context.Background(), resp.LeaseID)
493
+ return err
494
}
583
- l.closeErr = closeErr
584
- return closeErr
585
-}
495
587
-func (l *Listener) fail(runID uint64, err error, message string) {
588
- closeErr, changed := l.markStale(runID)
589
- if !changed {
590
- return
496
+ l.mu.Lock()
497
+ if l.state == listenerStateClosed || runCtx.Err() != nil {
498
+ l.mu.Unlock()
499
+ _ = api.unregisterLease(context.Background(), resp.LeaseID)
500
+ _ = tlsCloser.Close()
501
+ return context.Canceled
502
}
592
- if closeErr != nil {
593
- err = errors.Join(err, closeErr)
503
+ oldCloser := l.tlsCloser
504
+ l.leaseID = resp.LeaseID
505
+ l.hostnames = append([]string(nil), resp.Hostnames...)
506
+ l.tlsConfig = tlsConfig
507
+ l.tlsCloser = tlsCloser
508
+ l.sessionFailures = 0
509
+ l.state = listenerStateReady
510
+ l.mu.Unlock()
511
+
512
+ if oldCloser != nil {
513
+ _ = oldCloser.Close()
514
}
595
- log.Error().
596
- Str("component", "sdk-listener").
597
- Str("name", l.name).
598
- Err(err).
599
- Msg(message)
515
+ l.signalRefill()
516
+ return nil
517
}
518
602
-func (l *Listener) markStale(runID uint64) (error, bool) {
519
+func (l *Listener) shutdown() error {
520
l.mu.Lock()
604
- if l.runID != runID || l.state == listenerStateClosed || l.state == listenerStateStale {
521
+ if l.state == listenerStateClosed {
522
l.mu.Unlock()
606
- return nil, false
523
+ return nil
524
}
608
-
609
- l.state = listenerStateStale
525
+ l.state = listenerStateClosed
526
cancel := l.cancel
611
- client := l.client
527
+ api := l.api
528
leaseID := l.leaseID
529
tlsCloser := l.tlsCloser
530
l.leaseID = ""
@@ -624,15 +540,35 @@ func (l *Listener) markStale(runID uint64) (error, bool) {
540
l.drainAccepted()
541
542
var closeErr error
627
- if client != nil && strings.TrimSpace(leaseID) != "" {
543
+ if api != nil && strings.TrimSpace(leaseID) != "" {
544
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
629
- closeErr = errors.Join(closeErr, client.unregisterLease(ctx, leaseID, l.reverseToken))
545
+ closeErr = errors.Join(closeErr, api.unregisterLease(ctx, leaseID))
546
cancel()
547
}
548
if tlsCloser != nil {
549
closeErr = errors.Join(closeErr, tlsCloser.Close())
550
}
635
- return closeErr, true
551
+ if api != nil {
552
+ api.close()
553
+ }
554
+ return closeErr
555
+}
556
+
557
+func (l *Listener) fail(err error, message string) {
558
+ closed := false
559
+ l.closeOnce.Do(func() {
560
+ closed = true
561
+ l.closeErr = errors.Join(err, l.shutdown())
562
+ })
563
+ if !closed {
564
+ return
565
+ }
566
+
567
+ log.Error().
568
+ Str("component", "sdk-listener").
569
+ Str("name", l.api.name).
570
+ Err(l.closeErr).
571
+ Msg(message)
572
}
573
574
func (l *Listener) signalRefill() {
@@ -642,6 +578,10 @@ func (l *Listener) signalRefill() {
578
}
579
}
580
581
+func (l *Listener) retryLimitReached(failures int) bool {
582
+ return l.retryCount > 0 && failures >= l.retryCount
583
+}
584
+
585
func isLeaseNotFound(err error) bool {
586
return errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeLeaseNotFound})
587
}
@@ -671,16 +611,6 @@ func (l *Listener) drainAccepted() {
611
}
612
}
613
674
-func cloneMetadata(metadata types.LeaseMetadata) types.LeaseMetadata {
675
- return types.LeaseMetadata{
676
- Description: metadata.Description,
677
- Owner: metadata.Owner,
678
- Thumbnail: metadata.Thumbnail,
679
- Tags: append([]string(nil), metadata.Tags...),
680
- Hide: metadata.Hide,
681
- }
682
-}
683
-
614
type listenerAddr string
615
616
func (a listenerAddr) Network() string { return "portal" }
sdk/listener_test.go
deleted
-370
@@ -1,370 +0,0 @@
1
-package sdk
2
-
3
-import (
4
- "context"
5
- "encoding/json"
6
- "encoding/pem"
7
- "errors"
8
- "net"
9
- "net/http"
10
- "net/http/httptest"
11
- "sync/atomic"
12
- "testing"
13
- "time"
14
-
15
- "github.com/gosuda/portal/v2/types"
16
-)
17
-
18
-func TestListenerSnapshotAndAddr(t *testing.T) {
19
- t.Parallel()
20
-
21
- listener := &Listener{
22
- leaseID: "lease-1",
23
- hostnames: []string{"app.relay.example.com"},
24
- state: listenerStateReady,
25
- }
26
-
27
- if listener.Addr().String() != "portal:lease-1" {
28
- t.Fatalf("Addr().String() = %q, want %q", listener.Addr().String(), "portal:lease-1")
29
- }
30
-
31
- publicURLs := listener.publicURLs()
32
- if len(publicURLs) != 1 || publicURLs[0] != "https://app.relay.example.com" {
33
- t.Fatalf("publicURLs() = %#v, want [https://app.relay.example.com]", publicURLs)
34
- }
35
-}
36
-
37
-func TestListenerAccept(t *testing.T) {
38
- t.Parallel()
39
-
40
- ctx, cancel := context.WithCancel(context.Background())
41
- defer cancel()
42
-
43
- serverConn1, clientConn1 := net.Pipe()
44
- defer clientConn1.Close()
45
- serverConn2, clientConn2 := net.Pipe()
46
- defer clientConn2.Close()
47
-
48
- listener := &Listener{
49
- ctx: ctx,
50
- cancel: cancel,
51
- accepted: make(chan net.Conn, 2),
52
- }
53
- listener.accepted <- serverConn1
54
- listener.accepted <- serverConn2
55
-
56
- conn, err := listener.Accept()
57
- if err != nil {
58
- t.Fatalf("Accept() error = %v", err)
59
- }
60
- defer conn.Close()
61
-
62
- if conn != serverConn1 {
63
- t.Fatal("Accept() did not return the original connection")
64
- }
65
-
66
- plainConn, err := listener.Accept()
67
- if err != nil {
68
- t.Fatalf("Accept() error = %v", err)
69
- }
70
- defer plainConn.Close()
71
- if plainConn != serverConn2 {
72
- t.Fatal("Accept() did not return the original connection")
73
- }
74
-}
75
-
76
-func TestNewListenerRetriesBootstrapUntilSuccess(t *testing.T) {
77
- t.Parallel()
78
-
79
- var registerCalls atomic.Int32
80
- server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
81
- switch r.URL.Path {
82
- case types.PathSDKRegister:
83
- if registerCalls.Add(1) < 3 {
84
- writeTestAPIEnvelope(w, http.StatusServiceUnavailable, types.APIEnvelope[struct{}]{
85
- OK: false,
86
- Error: &types.APIError{
87
- Code: types.APIErrorCodeLeaseRejected,
88
- Message: "relay unavailable",
89
- },
90
- })
91
- return
92
- }
93
-
94
- writeTestAPIEnvelope(w, http.StatusOK, types.APIEnvelope[types.RegisterResponse]{
95
- OK: true,
96
- Data: types.RegisterResponse{
97
- LeaseID: "lease-1",
98
- Metadata: types.LeaseMetadata{
99
- Owner: "alice",
100
- },
101
- },
102
- })
103
- case types.PathSDKRenew:
104
- writeTestAPIEnvelope(w, http.StatusOK, types.APIEnvelope[types.RenewResponse]{
105
- OK: true,
106
- Data: types.RenewResponse{LeaseID: "lease-1"},
107
- })
108
- case types.PathSDKUnregister:
109
- writeTestAPIEnvelope(w, http.StatusOK, types.APIEnvelope[struct{}]{OK: true})
110
- case types.PathSDKConnect:
111
- writeTestAPIEnvelope(w, http.StatusServiceUnavailable, types.APIEnvelope[struct{}]{
112
- OK: false,
113
- Error: &types.APIError{
114
- Code: types.APIErrorCodeSessionCreateFailed,
115
- Message: "reverse sessions unavailable",
116
- },
117
- })
118
- default:
119
- http.NotFound(w, r)
120
- }
121
- }))
122
- defer server.Close()
123
-
124
- rootCAPEM := pem.EncodeToMemory(&pem.Block{
125
- Type: "CERTIFICATE",
126
- Bytes: server.Certificate().Raw,
127
- })
128
-
129
- client, err := NewRelayClient(server.URL, WithRootCAPEM(rootCAPEM))
130
- if err != nil {
131
- t.Fatalf("NewRelayClient() error = %v", err)
132
- }
133
-
134
- listener, err := NewListener(context.Background(), ListenRequest{
135
- Name: "demo",
136
- }, listenerOptions{
137
- client: client,
138
- handshakeTimeout: time.Second,
139
- renewBefore: time.Second,
140
- retryCount: defaultListenerRetryCount,
141
- retryDelay: 10 * time.Millisecond,
142
- defaultLeaseTTL: time.Minute,
143
- defaultReadyTarget: 1,
144
- })
145
- if err != nil {
146
- t.Fatalf("newListener() error = %v", err)
147
- }
148
- defer listener.Close()
149
-
150
- waitForListenerCondition(t, func() bool {
151
- listener.mu.Lock()
152
- defer listener.mu.Unlock()
153
- return listener.leaseID == "lease-1"
154
- })
155
-
156
- if got := registerCalls.Load(); got < 3 {
157
- t.Fatalf("register call count = %d, want at least 3", got)
158
- }
159
-}
160
-
161
-func TestNewListenerBecomesStaleAfterRetryLimit(t *testing.T) {
162
- t.Parallel()
163
-
164
- server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
165
- if r.URL.Path == types.PathSDKRegister {
166
- writeTestAPIEnvelope(w, http.StatusServiceUnavailable, types.APIEnvelope[struct{}]{
167
- OK: false,
168
- Error: &types.APIError{
169
- Code: types.APIErrorCodeLeaseRejected,
170
- Message: "relay unavailable",
171
- },
172
- })
173
- return
174
- }
175
- http.NotFound(w, r)
176
- }))
177
- defer server.Close()
178
-
179
- rootCAPEM := pem.EncodeToMemory(&pem.Block{
180
- Type: "CERTIFICATE",
181
- Bytes: server.Certificate().Raw,
182
- })
183
-
184
- client, err := NewRelayClient(server.URL, WithRootCAPEM(rootCAPEM))
185
- if err != nil {
186
- t.Fatalf("NewRelayClient() error = %v", err)
187
- }
188
-
189
- listener, err := NewListener(context.Background(), ListenRequest{
190
- Name: "demo",
191
- }, listenerOptions{
192
- client: client,
193
- handshakeTimeout: time.Second,
194
- renewBefore: time.Second,
195
- retryCount: 2,
196
- retryDelay: 10 * time.Millisecond,
197
- defaultLeaseTTL: time.Minute,
198
- defaultReadyTarget: 1,
199
- })
200
- if err != nil {
201
- t.Fatalf("newListener() error = %v", err)
202
- }
203
- defer listener.Close()
204
-
205
- waitForListenerCondition(t, func() bool {
206
- listener.mu.Lock()
207
- defer listener.mu.Unlock()
208
- return listener.state == listenerStateStale
209
- })
210
-
211
- if _, err := listener.Accept(); !errors.Is(err, net.ErrClosed) {
212
- t.Fatalf("Accept() error = %v, want net.ErrClosed", err)
213
- }
214
-}
215
-
216
-func TestListenerReactivateAfterStale(t *testing.T) {
217
- t.Parallel()
218
-
219
- var allowRegister atomic.Int32
220
- server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
221
- switch r.URL.Path {
222
- case types.PathSDKRegister:
223
- if allowRegister.Load() == 0 {
224
- writeTestAPIEnvelope(w, http.StatusServiceUnavailable, types.APIEnvelope[struct{}]{
225
- OK: false,
226
- Error: &types.APIError{
227
- Code: types.APIErrorCodeLeaseRejected,
228
- Message: "relay unavailable",
229
- },
230
- })
231
- return
232
- }
233
-
234
- writeTestAPIEnvelope(w, http.StatusOK, types.APIEnvelope[types.RegisterResponse]{
235
- OK: true,
236
- Data: types.RegisterResponse{
237
- LeaseID: "lease-2",
238
- Metadata: types.LeaseMetadata{
239
- Owner: "alice",
240
- },
241
- },
242
- })
243
- case types.PathSDKRenew:
244
- writeTestAPIEnvelope(w, http.StatusOK, types.APIEnvelope[types.RenewResponse]{
245
- OK: true,
246
- Data: types.RenewResponse{LeaseID: "lease-2"},
247
- })
248
- case types.PathSDKUnregister:
249
- writeTestAPIEnvelope(w, http.StatusOK, types.APIEnvelope[struct{}]{OK: true})
250
- case types.PathSDKConnect:
251
- hijacker, ok := w.(http.Hijacker)
252
- if !ok {
253
- t.Fatal("response writer does not support hijacking")
254
- }
255
- conn, rw, err := hijacker.Hijack()
256
- if err != nil {
257
- t.Fatalf("Hijack() error = %v", err)
258
- }
259
- _, _ = rw.WriteString("HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n")
260
- _ = rw.Flush()
261
- time.Sleep(2 * time.Second)
262
- _ = conn.Close()
263
- default:
264
- http.NotFound(w, r)
265
- }
266
- }))
267
- defer server.Close()
268
-
269
- rootCAPEM := pem.EncodeToMemory(&pem.Block{
270
- Type: "CERTIFICATE",
271
- Bytes: server.Certificate().Raw,
272
- })
273
-
274
- client, err := NewRelayClient(server.URL, WithRootCAPEM(rootCAPEM))
275
- if err != nil {
276
- t.Fatalf("NewRelayClient() error = %v", err)
277
- }
278
-
279
- listener, err := NewListener(context.Background(), ListenRequest{
280
- Name: "demo",
281
- }, listenerOptions{
282
- client: client,
283
- handshakeTimeout: time.Second,
284
- renewBefore: time.Second,
285
- retryCount: 1,
286
- retryDelay: 10 * time.Millisecond,
287
- defaultLeaseTTL: time.Minute,
288
- defaultReadyTarget: 1,
289
- })
290
- if err != nil {
291
- t.Fatalf("newListener() error = %v", err)
292
- }
293
- defer listener.Close()
294
-
295
- waitForListenerCondition(t, func() bool {
296
- listener.mu.Lock()
297
- defer listener.mu.Unlock()
298
- return listener.state == listenerStateStale
299
- })
300
-
301
- allowRegister.Store(1)
302
- if err := listener.Reactivate(context.Background()); err != nil {
303
- t.Fatalf("Reactivate() error = %v", err)
304
- }
305
-
306
- waitForListenerCondition(t, func() bool {
307
- listener.mu.Lock()
308
- defer listener.mu.Unlock()
309
- return listener.state == listenerStateReady && listener.leaseID == "lease-2"
310
- })
311
-}
312
-
313
-func TestListenerAcceptDoesNotReturnQueuedConnAfterClose(t *testing.T) {
314
- t.Parallel()
315
-
316
- conn := &listenerStubConn{}
317
- ctx, cancel := context.WithCancel(context.Background())
318
- cancel()
319
-
320
- listener := &Listener{
321
- ctx: ctx,
322
- cancel: cancel,
323
- accepted: make(chan net.Conn, 1),
324
- }
325
- listener.accepted <- conn
326
-
327
- gotConn, err := listener.Accept()
328
- if gotConn != nil {
329
- t.Fatalf("Accept() conn = %#v, want nil", gotConn)
330
- }
331
- if !errors.Is(err, net.ErrClosed) {
332
- t.Fatalf("Accept() error = %v, want net.ErrClosed", err)
333
- }
334
- if conn.closeCount != 1 {
335
- t.Fatalf("conn close count = %d, want 1", conn.closeCount)
336
- }
337
-}
338
-
339
-type listenerStubConn struct {
340
- closeCount int
341
-}
342
-
343
-func (c *listenerStubConn) Read(_ []byte) (int, error) { return 0, nil }
344
-func (c *listenerStubConn) Write(b []byte) (int, error) { return len(b), nil }
345
-func (c *listenerStubConn) Close() error { c.closeCount++; return nil }
346
-func (c *listenerStubConn) LocalAddr() net.Addr { return listenerAddr("stub-local") }
347
-func (c *listenerStubConn) RemoteAddr() net.Addr { return listenerAddr("stub-remote") }
348
-func (c *listenerStubConn) SetDeadline(_ time.Time) error { return nil }
349
-func (c *listenerStubConn) SetReadDeadline(_ time.Time) error { return nil }
350
-func (c *listenerStubConn) SetWriteDeadline(_ time.Time) error { return nil }
351
-
352
-func waitForListenerCondition(t *testing.T, fn func() bool) {
353
- t.Helper()
354
-
355
- deadline := time.Now().Add(5 * time.Second)
356
- for time.Now().Before(deadline) {
357
- if fn() {
358
- return
359
- }
360
- time.Sleep(10 * time.Millisecond)
361
- }
362
-
363
- t.Fatal("timed out waiting for listener condition")
364
-}
365
-
366
-func writeTestAPIEnvelope[T any](w http.ResponseWriter, status int, envelope types.APIEnvelope[T]) {
367
- w.Header().Set("Content-Type", "application/json")
368
- w.WriteHeader(status)
369
- _ = json.NewEncoder(w).Encode(envelope)
370
-}
sdk/relayclient.go
renamed
+147
-172
@@ -32,84 +32,28 @@ const (
32
defaultHTTPShutdownTimeout = 5 * time.Second
33
)
34
35
-type RelayClientOption func(*RelayClient)
36
-
37
-func WithRootCAPEM(rootCAPEM []byte) RelayClientOption {
38
- rootCAPEM = append([]byte(nil), rootCAPEM...)
39
- return func(client *RelayClient) {
40
- client.rootCAPEM = append([]byte(nil), rootCAPEM...)
41
- }
42
-}
43
-
44
-func WithInsecureSkipVerify(skip bool) RelayClientOption {
45
- return func(client *RelayClient) {
46
- client.insecureSkipVerify = skip
47
- }
48
-}
49
-
50
-func WithDialTimeout(timeout time.Duration) RelayClientOption {
51
- return func(client *RelayClient) {
52
- if timeout > 0 {
53
- client.dialTimeout = timeout
54
- }
55
- }
56
-}
57
-
58
-func WithRequestTimeout(timeout time.Duration) RelayClientOption {
59
- return func(client *RelayClient) {
60
- if timeout > 0 {
61
- client.requestTimeout = timeout
62
- }
63
- }
64
-}
65
-
66
-func WithHandshakeTimeout(timeout time.Duration) RelayClientOption {
67
- return func(client *RelayClient) {
68
- if timeout > 0 {
69
- client.handshakeTimeout = timeout
70
- }
71
- }
72
-}
73
-
74
-func WithLeaseTTL(ttl time.Duration) RelayClientOption {
75
- return func(client *RelayClient) {
76
- if ttl > 0 {
77
- client.leaseTTL = ttl
78
- }
79
- }
80
-}
81
-
82
-func WithRenewBefore(d time.Duration) RelayClientOption {
83
- return func(client *RelayClient) {
84
- if d > 0 {
85
- client.renewBefore = d
86
- }
87
- }
88
-}
89
-
90
-func WithReadyTarget(n int) RelayClientOption {
91
- return func(client *RelayClient) {
92
- if n > 0 {
93
- client.readyTarget = n
94
- }
95
- }
96
-}
97
-
98
-type RelayClient struct {
35
+type relayClient struct {
36
baseURL *url.URL
37
httpClient *http.Client
38
rawTLSConfig *tls.Config
102
- rootCAPEM []byte
103
- insecureSkipVerify bool
39
dialTimeout time.Duration
105
- requestTimeout time.Duration
106
- handshakeTimeout time.Duration
107
- leaseTTL time.Duration
108
- renewBefore time.Duration
109
- readyTarget int
40
+ name string
41
+ requestedHostnames []string
42
+ reverseToken string
43
+ metadata types.LeaseMetadata
44
}
45
112
-func NewRelayClient(relayURL string, options ...RelayClientOption) (*RelayClient, error) {
46
+func newRelayClient(relayURL string, cfg ListenerConfig) (*relayClient, error) {
47
+ name := strings.TrimSpace(cfg.Name)
48
+ if name == "" {
49
+ return nil, errors.New("listener name is required")
50
+ }
51
+
52
+ reverseToken := strings.TrimSpace(cfg.ReverseToken)
53
+ if reverseToken == "" {
54
+ reverseToken = randomToken()
55
+ }
56
+
57
baseURL, err := url.Parse(strings.TrimSpace(relayURL))
58
if err != nil {
59
return nil, fmt.Errorf("parse relay url: %w", err)
@@ -124,43 +68,28 @@ func NewRelayClient(relayURL string, options ...RelayClientOption) (*RelayClient
68
baseURL.RawQuery = ""
69
baseURL.Fragment = ""
70
127
- client := &RelayClient{
128
- baseURL: baseURL,
129
- dialTimeout: defaultDialTimeout,
130
- requestTimeout: defaultRequestTimeout,
131
- handshakeTimeout: defaultHandshakeTimeout,
132
- leaseTTL: defaultLeaseTTL,
133
- renewBefore: defaultRenewBefore,
134
- readyTarget: defaultReadyTarget,
135
- }
136
- for _, option := range options {
137
- if option != nil {
138
- option(client)
139
- }
140
- }
141
-
142
- if len(client.rootCAPEM) == 0 && !client.insecureSkipVerify && isLocalRelayHost(baseURL.Hostname()) {
71
+ rootCAPEM := append([]byte(nil), cfg.RootCAPEM...)
72
+ if len(rootCAPEM) == 0 && isLocalRelayHost(baseURL.Hostname()) {
73
bootstrapCtx, cancel := context.WithTimeout(context.Background(), defaultDialTimeout+defaultHandshakeTimeout)
74
defer cancel()
75
146
- _, rootCAPEM, bootstrapErr := keyless.ResolveMaterials(bootstrapCtx, baseURL.String(), baseURL.Hostname())
76
+ _, resolvedCAPEM, bootstrapErr := keyless.ResolveMaterials(bootstrapCtx, baseURL.String(), baseURL.Hostname())
77
if bootstrapErr != nil {
78
return nil, fmt.Errorf("bootstrap localhost relay trust: %w", bootstrapErr)
79
}
150
- client.rootCAPEM = rootCAPEM
80
+ rootCAPEM = resolvedCAPEM
81
}
82
153
- rootCAs, err := buildRootCAs(client.rootCAPEM)
83
+ rootCAs, err := buildRootCAs(rootCAPEM)
84
if err != nil {
85
return nil, err
86
}
87
88
baseTLS := &tls.Config{
159
- MinVersion: tls.VersionTLS12,
160
- ServerName: baseURL.Hostname(),
161
- RootCAs: rootCAs,
162
- InsecureSkipVerify: client.insecureSkipVerify,
163
- NextProtos: []string{"http/1.1"},
89
+ MinVersion: tls.VersionTLS12,
90
+ ServerName: baseURL.Hostname(),
91
+ RootCAs: rootCAs,
92
+ NextProtos: []string{"http/1.1"},
93
}
94
95
transport := &http.Transport{
@@ -168,104 +97,94 @@ func NewRelayClient(relayURL string, options ...RelayClientOption) (*RelayClient
97
ForceAttemptHTTP2: false,
98
}
99
171
- client.httpClient = &http.Client{
172
- Transport: transport,
173
- Timeout: client.requestTimeout,
100
+ api := &relayClient{
101
+ baseURL: baseURL,
102
+ httpClient: &http.Client{Transport: transport, Timeout: defaultRequestTimeout},
103
+ rawTLSConfig: baseTLS,
104
+ dialTimeout: defaultDialTimeout,
105
+ name: name,
106
+ requestedHostnames: append([]string(nil), cfg.Hostnames...),
107
+ reverseToken: reverseToken,
108
+ metadata: cloneMetadata(cfg.Metadata),
109
}
175
- client.rawTLSConfig = baseTLS
176
- return client, nil
110
+
111
+ checkCtx, cancel := context.WithTimeout(context.Background(), defaultRequestTimeout)
112
+ defer cancel()
113
+
114
+ if err := api.ensureCompatible(checkCtx); err != nil {
115
+ api.close()
116
+ return nil, err
117
+ }
118
+
119
+ return api, nil
120
}
121
179
-func (c *RelayClient) Close() {
180
- if c == nil || c.httpClient == nil {
122
+func (a *relayClient) close() {
123
+ if a == nil || a.httpClient == nil {
124
return
125
}
183
- if transport, ok := c.httpClient.Transport.(*http.Transport); ok {
126
+ if transport, ok := a.httpClient.Transport.(*http.Transport); ok {
127
transport.CloseIdleConnections()
128
}
129
}
187
-func (c *RelayClient) doJSON(ctx context.Context, method, path string, payload any, out any) error {
188
- var body io.Reader
189
- if payload != nil {
190
- buf, err := json.Marshal(payload)
191
- if err != nil {
192
- return fmt.Errorf("marshal payload: %w", err)
193
- }
194
- body = bytes.NewReader(buf)
195
- }
196
-
197
- ref, _ := url.Parse(path)
198
- req, err := http.NewRequestWithContext(ctx, method, c.baseURL.ResolveReference(ref).String(), body)
199
- if err != nil {
200
- return err
201
- }
202
- req.Header.Set("Content-Type", "application/json")
130
204
- resp, err := c.httpClient.Do(req)
205
- if err != nil {
206
- return err
131
+func (a *relayClient) registerLease(ctx context.Context, hostnames []string, ttl time.Duration) (types.RegisterResponse, error) {
132
+ if len(hostnames) == 0 {
133
+ hostnames = a.requestedHostnames
134
}
208
- defer resp.Body.Close()
135
210
- var envelope types.APIEnvelope[json.RawMessage]
211
- if err := json.NewDecoder(resp.Body).Decode(&envelope); err != nil {
212
- return fmt.Errorf("decode response: %w", err)
213
- }
214
- if !envelope.OK {
215
- if envelope.Error == nil {
216
- return &types.APIRequestError{
217
- StatusCode: resp.StatusCode,
218
- Message: fmt.Sprintf("api request failed with status %d", resp.StatusCode),
219
- }
220
- }
221
- return &types.APIRequestError{
222
- StatusCode: resp.StatusCode,
223
- Code: envelope.Error.Code,
224
- Message: envelope.Error.Message,
225
- }
226
- }
227
- if out == nil {
228
- return nil
229
- }
230
- return json.Unmarshal(envelope.Data, out)
231
-}
232
-
233
-func (c *RelayClient) registerLease(ctx context.Context, req types.RegisterRequest) (types.RegisterResponse, error) {
136
var resp types.RegisterResponse
235
- if err := c.doJSON(ctx, http.MethodPost, types.PathSDKRegister, req, &resp); err != nil {
137
+ if err := a.doJSON(ctx, http.MethodPost, types.PathSDKRegister, types.RegisterRequest{
138
+ Name: a.name,
139
+ Hostnames: append([]string(nil), hostnames...),
140
+ Metadata: cloneMetadata(a.metadata),
141
+ ReverseToken: a.reverseToken,
142
+ TTL: int(ttl / time.Second),
143
+ }, &resp); err != nil {
144
return types.RegisterResponse{}, err
145
}
146
return resp, nil
147
}
148
241
-func (c *RelayClient) renewLease(ctx context.Context, leaseID, reverseToken string, ttl time.Duration) error {
242
- return c.doJSON(ctx, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
149
+func (a *relayClient) ensureCompatible(ctx context.Context) error {
150
+ var resp types.DomainResponse
151
+ if err := a.doJSON(ctx, http.MethodGet, types.PathSDKDomain, nil, &resp); err != nil {
152
+ return fmt.Errorf("check relay compatibility: %w", err)
153
+ }
154
+ if strings.TrimSpace(resp.Version) != types.SDKProtocolVersion {
155
+ return fmt.Errorf("relay sdk version mismatch: relay=%q client=%q", strings.TrimSpace(resp.Version), types.SDKProtocolVersion)
156
+ }
157
+ return nil
158
+}
159
+
160
+func (a *relayClient) renewLease(ctx context.Context, leaseID string, ttl time.Duration) error {
161
+ return a.doJSON(ctx, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
162
LeaseID: leaseID,
244
- ReverseToken: reverseToken,
245
- TTLSeconds: int(ttl / time.Second),
163
+ ReverseToken: a.reverseToken,
164
+ TTL: int(ttl / time.Second),
165
}, &types.RenewResponse{})
166
}
167
249
-func (c *RelayClient) unregisterLease(ctx context.Context, leaseID, reverseToken string) error {
250
- return c.doJSON(ctx, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
168
+func (a *relayClient) unregisterLease(ctx context.Context, leaseID string) error {
169
+ return a.doJSON(ctx, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
170
LeaseID: leaseID,
252
- ReverseToken: reverseToken,
171
+ ReverseToken: a.reverseToken,
172
}, nil)
173
}
174
256
-func (c *RelayClient) openReverseSession(ctx context.Context, leaseID, reverseToken string) (net.Conn, error) {
175
+func (a *relayClient) openReverseSession(ctx context.Context, leaseID string) (net.Conn, error) {
176
dialer := &tls.Dialer{
258
- NetDialer: &net.Dialer{Timeout: c.dialTimeout},
259
- Config: c.rawTLSConfig.Clone(),
177
+ NetDialer: &net.Dialer{Timeout: a.dialTimeout},
178
+ Config: a.rawTLSConfig.Clone(),
179
}
180
262
- conn, err := dialer.DialContext(ctx, "tcp", ensurePort(c.baseURL.Host))
181
+ conn, err := dialer.DialContext(ctx, "tcp", ensurePort(a.baseURL.Host))
182
if err != nil {
183
return nil, err
184
}
185
186
connectRef, _ := url.Parse(types.PathSDKConnect)
268
- connectURL := c.baseURL.ResolveReference(connectRef)
187
+ connectURL := a.baseURL.ResolveReference(connectRef)
188
query := connectURL.Query()
189
query.Set("lease_id", leaseID)
190
connectURL.RawQuery = query.Encode()
@@ -273,10 +192,10 @@ func (c *RelayClient) openReverseSession(ctx context.Context, leaseID, reverseTo
192
req := &http.Request{
193
Method: http.MethodGet,
194
URL: connectURL,
276
- Host: c.baseURL.Host,
195
+ Host: a.baseURL.Host,
196
Header: make(http.Header),
197
}
279
- req.Header.Set(types.HeaderReverseToken, reverseToken)
198
+ req.Header.Set(types.HeaderReverseToken, a.reverseToken)
199
req.Header.Set("Connection", "keep-alive")
200
201
if writeErr := req.Write(conn); writeErr != nil {
@@ -301,25 +220,50 @@ func (c *RelayClient) openReverseSession(ctx context.Context, leaseID, reverseTo
220
return wrapBufferedConn(conn, reader), nil
221
}
222
304
-func decodeAPIResponseError(resp *http.Response) error {
305
- if resp == nil {
306
- return &types.APIRequestError{Message: "empty api response"}
223
+func (a *relayClient) doJSON(ctx context.Context, method, path string, payload any, out any) error {
224
+ var body io.Reader
225
+ if payload != nil {
226
+ buf, err := json.Marshal(payload)
227
+ if err != nil {
228
+ return fmt.Errorf("marshal payload: %w", err)
229
+ }
230
+ body = bytes.NewReader(buf)
231
}
232
309
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 8<<10))
233
+ ref, _ := url.Parse(path)
234
+ req, err := http.NewRequestWithContext(ctx, method, a.baseURL.ResolveReference(ref).String(), body)
235
+ if err != nil {
236
+ return err
237
+ }
238
+ req.Header.Set("Content-Type", "application/json")
239
+
240
+ resp, err := a.httpClient.Do(req)
241
+ if err != nil {
242
+ return err
243
+ }
244
+ defer resp.Body.Close()
245
+
246
var envelope types.APIEnvelope[json.RawMessage]
311
- if err := json.Unmarshal(body, &envelope); err == nil && envelope.Error != nil {
247
+ if err := json.NewDecoder(resp.Body).Decode(&envelope); err != nil {
248
+ return fmt.Errorf("decode response: %w", err)
249
+ }
250
+ if !envelope.OK {
251
+ if envelope.Error == nil {
252
+ return &types.APIRequestError{
253
+ StatusCode: resp.StatusCode,
254
+ Message: fmt.Sprintf("api request failed with status %d", resp.StatusCode),
255
+ }
256
+ }
257
return &types.APIRequestError{
258
StatusCode: resp.StatusCode,
259
Code: envelope.Error.Code,
260
Message: envelope.Error.Message,
261
}
262
}
318
-
319
- return &types.APIRequestError{
320
- StatusCode: resp.StatusCode,
321
- Message: strings.TrimSpace(string(body)),
263
+ if out == nil {
264
+ return nil
265
}
266
+ return json.Unmarshal(envelope.Data, out)
267
}
268
269
func buildRootCAs(rootCAPEM []byte) (*x509.CertPool, error) {
@@ -333,6 +277,27 @@ func buildRootCAs(rootCAPEM []byte) (*x509.CertPool, error) {
277
return pool, nil
278
}
279
280
+func decodeAPIResponseError(resp *http.Response) error {
281
+ if resp == nil {
282
+ return &types.APIRequestError{Message: "empty api response"}
283
+ }
284
+
285
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, 8<<10))
286
+ var envelope types.APIEnvelope[json.RawMessage]
287
+ if err := json.Unmarshal(body, &envelope); err == nil && envelope.Error != nil {
288
+ return &types.APIRequestError{
289
+ StatusCode: resp.StatusCode,
290
+ Code: envelope.Error.Code,
291
+ Message: envelope.Error.Message,
292
+ }
293
+ }
294
+
295
+ return &types.APIRequestError{
296
+ StatusCode: resp.StatusCode,
297
+ Message: strings.TrimSpace(string(body)),
298
+ }
299
+}
300
+
301
func randomToken() string {
302
buf := make([]byte, 8)
303
if _, err := rand.Read(buf); err != nil {
@@ -360,6 +325,16 @@ func isLocalRelayHost(host string) bool {
325
return strings.HasSuffix(host, ".localhost")
326
}
327
328
+func cloneMetadata(metadata types.LeaseMetadata) types.LeaseMetadata {
329
+ return types.LeaseMetadata{
330
+ Description: metadata.Description,
331
+ Owner: metadata.Owner,
332
+ Thumbnail: metadata.Thumbnail,
333
+ Tags: append([]string(nil), metadata.Tags...),
334
+ Hide: metadata.Hide,
335
+ }
336
+}
337
+
338
type bufferedConn struct {
339
net.Conn
340
reader *bytes.Reader
sdk/sdk_test.go
new
+200
@@ -0,0 +1,200 @@
1
+package sdk
2
+
3
+import (
4
+ "context"
5
+ "encoding/json"
6
+ "net"
7
+ "net/http"
8
+ "net/http/httptest"
9
+ "reflect"
10
+ "strings"
11
+ "sync/atomic"
12
+ "testing"
13
+ "time"
14
+
15
+ "github.com/gosuda/portal/v2/types"
16
+)
17
+
18
+func TestNewRelayClientAcceptsMatchingVersion(t *testing.T) {
19
+ t.Parallel()
20
+
21
+ server := newDomainServer(t, types.SDKProtocolVersion)
22
+ defer server.Close()
23
+
24
+ api, err := newRelayClient(server.URL, ListenerConfig{Name: "demo"})
25
+ if err != nil {
26
+ t.Fatalf("newRelayClient() error = %v", err)
27
+ }
28
+ defer api.close()
29
+}
30
+
31
+func TestNewListenerRejectsVersionMismatch(t *testing.T) {
32
+ t.Parallel()
33
+
34
+ server := newDomainServer(t, "999")
35
+ defer server.Close()
36
+
37
+ listener, err := NewListener(context.Background(), server.URL, ListenerConfig{Name: "demo"})
38
+ if err == nil {
39
+ t.Fatal("NewListener() error = nil, want version mismatch")
40
+ }
41
+ if listener != nil {
42
+ t.Fatalf("NewListener() listener = %#v, want nil", listener)
43
+ }
44
+ if !strings.Contains(err.Error(), "version mismatch") {
45
+ t.Fatalf("NewListener() error = %v, want version mismatch", err)
46
+ }
47
+}
48
+
49
+func TestNewListenerReregistersOnLeaseNotFound(t *testing.T) {
50
+ t.Parallel()
51
+
52
+ var registerCount atomic.Int32
53
+ server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
54
+ switch r.URL.Path {
55
+ case types.PathSDKDomain:
56
+ writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.DomainResponse]{
57
+ OK: true,
58
+ Data: types.DomainResponse{
59
+ RootHost: "relay.example.com",
60
+ SuggestedHostname: "demo.relay.example.com",
61
+ Version: types.SDKProtocolVersion,
62
+ },
63
+ })
64
+ case types.PathSDKRegister:
65
+ count := registerCount.Add(1)
66
+ leaseID := "lease-1"
67
+ if count > 1 {
68
+ leaseID = "lease-2"
69
+ }
70
+ writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.RegisterResponse]{
71
+ OK: true,
72
+ Data: types.RegisterResponse{
73
+ LeaseID: leaseID,
74
+ },
75
+ })
76
+ case types.PathSDKRenew:
77
+ var req types.RenewRequest
78
+ if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
79
+ t.Fatalf("decode renew request: %v", err)
80
+ }
81
+ if req.LeaseID == "lease-1" {
82
+ writeSDKTestEnvelope(w, http.StatusNotFound, types.APIEnvelope[any]{
83
+ OK: false,
84
+ Error: &types.APIError{Code: types.APIErrorCodeLeaseNotFound, Message: "lease not found"},
85
+ })
86
+ return
87
+ }
88
+ writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.RenewResponse]{
89
+ OK: true,
90
+ Data: types.RenewResponse{LeaseID: req.LeaseID},
91
+ })
92
+ case types.PathSDKUnregister:
93
+ writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[any]{OK: true})
94
+ default:
95
+ http.NotFound(w, r)
96
+ }
97
+ }))
98
+ defer server.Close()
99
+
100
+ api, err := newRelayClient(server.URL, ListenerConfig{Name: "demo"})
101
+ if err != nil {
102
+ t.Fatalf("newRelayClient() error = %v", err)
103
+ }
104
+
105
+ listenerCtx, cancel := context.WithCancel(context.Background())
106
+ listener := &Listener{
107
+ ctx: listenerCtx,
108
+ cancel: cancel,
109
+ accepted: make(chan net.Conn, 1),
110
+ refill: make(chan struct{}, 1),
111
+ readyTarget: 0,
112
+ leaseTTL: defaultLeaseTTL,
113
+ renewInterval: 10 * time.Millisecond,
114
+ handshakeTimeout: defaultHandshakeTimeout,
115
+ retryCount: 0,
116
+ retryDelay: 10 * time.Millisecond,
117
+ api: api,
118
+ state: listenerStatePending,
119
+ }
120
+ go listener.run(listenerCtx)
121
+ defer listener.Close()
122
+
123
+ waitForSDKTest(t, func() bool {
124
+ listener.mu.Lock()
125
+ defer listener.mu.Unlock()
126
+ return listener.leaseID == "lease-2"
127
+ })
128
+}
129
+
130
+func TestExposeNoRelayInputs(t *testing.T) {
131
+ t.Parallel()
132
+
133
+ exposure, err := Expose(context.Background(), nil, "demo", types.LeaseMetadata{})
134
+ if err != nil {
135
+ t.Fatalf("Expose() error = %v", err)
136
+ }
137
+ if exposure != nil {
138
+ t.Fatalf("Expose() exposure = %#v, want nil", exposure)
139
+ }
140
+}
141
+
142
+func TestNormalizeRelayURLs(t *testing.T) {
143
+ t.Parallel()
144
+
145
+ got, err := NormalizeRelayURLs([]string{
146
+ " localhost:4017 , https://relay.example.com/base/relay?x=1#frag ",
147
+ "https://relay.example.com/base",
148
+ })
149
+ if err != nil {
150
+ t.Fatalf("NormalizeRelayURLs() error = %v", err)
151
+ }
152
+
153
+ want := []string{
154
+ "https://localhost:4017",
155
+ "https://relay.example.com/base",
156
+ }
157
+ if !reflect.DeepEqual(got, want) {
158
+ t.Fatalf("NormalizeRelayURLs() = %v, want %v", got, want)
159
+ }
160
+}
161
+
162
+func newDomainServer(t *testing.T, version string) *httptest.Server {
163
+ t.Helper()
164
+
165
+ return httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
166
+ if r.URL.Path != types.PathSDKDomain {
167
+ http.NotFound(w, r)
168
+ return
169
+ }
170
+
171
+ writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.DomainResponse]{
172
+ OK: true,
173
+ Data: types.DomainResponse{
174
+ RootHost: "relay.example.com",
175
+ SuggestedHostname: "demo.relay.example.com",
176
+ Version: version,
177
+ },
178
+ })
179
+ }))
180
+}
181
+
182
+func writeSDKTestEnvelope[T any](w http.ResponseWriter, status int, envelope types.APIEnvelope[T]) {
183
+ w.Header().Set("Content-Type", "application/json")
184
+ w.WriteHeader(status)
185
+ _ = json.NewEncoder(w).Encode(envelope)
186
+}
187
+
188
+func waitForSDKTest(t *testing.T, fn func() bool) {
189
+ t.Helper()
190
+
191
+ deadline := time.Now().Add(5 * time.Second)
192
+ for time.Now().Before(deadline) {
193
+ if fn() {
194
+ return
195
+ }
196
+ time.Sleep(10 * time.Millisecond)
197
+ }
198
+
199
+ t.Fatal("timed out waiting for condition")
200
+}
types/api.go
+4
-3
@@ -10,6 +10,7 @@ const (
10
HeaderReverseToken = "X-Portal-Token"
11
MarkerKeepalive = byte(0x00)
12
MarkerTLSStart = byte(0x02)
13
+ SDKProtocolVersion = "1"
14
)
15
16
type APIEnvelope[T any] struct {
@@ -72,8 +73,7 @@ type RegisterRequest struct {
73
ReverseToken string `json:"reverse_token"`
74
Hostnames []string `json:"hostnames,omitempty"`
75
Metadata LeaseMetadata `json:"metadata"`
75
- TTLSeconds int `json:"ttl_seconds,omitempty"`
76
- TLS bool `json:"tls"`
76
+ TTL int `json:"ttl,omitempty"`
77
}
78
79
type RegisterResponse struct {
@@ -87,7 +87,7 @@ type RegisterResponse struct {
87
type RenewRequest struct {
88
LeaseID string `json:"lease_id"`
89
ReverseToken string `json:"reverse_token"`
90
- TTLSeconds int `json:"ttl_seconds,omitempty"`
90
+ TTL int `json:"ttl,omitempty"`
91
}
92
93
type RenewResponse struct {
@@ -103,6 +103,7 @@ type UnregisterRequest struct {
103
type DomainResponse struct {
104
RootHost string `json:"root_host"`
105
SuggestedHostname string `json:"suggested_hostname"`
106
+ Version string `json:"version"`
107
}
108
109
type AdminLoginRequest struct {