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 {