sdk: refactor listener
Kim committed
Mar 10, 2026 at 19:26 UTC
f6ec06e2d891930417985f9c2fd3fd5520b5e1ac
8 files changed
+923
-373
cmd/portal-tunnel/README.md
+1
-1
@@ -31,6 +31,6 @@ Portal-tunnel connects a local service to a Portal relay with the legacy CLI sha
31
32
- Multiple relay URLs are registered independently. Each relay gets its own lease ID and public URLs.
33
- Portal-tunnel now consumes one aggregate SDK listener, so the CLI no longer manages per-relay listener loops itself.
34
-- Startup is fail-fast: if any configured relay cannot register, the tunnel exits instead of partially publishing.
34
+- Startup no longer fails on a temporarily unavailable relay. Each configured relay listener keeps retrying until it connects or the tunnel is stopped.
35
- Tenant TLS is provisioned automatically through the relay keyless signer. The SDK fetches the relay certificate chain and uses `/v1/sign` for remote signing.
36
- When the local service is unreachable, the tunnel returns an HTTP 503 page.
docs/architecture.md
+2
-2
@@ -54,9 +54,9 @@ That distinction matters because `/sdk/connect` stops being ordinary HTTP once h
54
55
### SDK (`sdk/`)
56
57
-- `Client`: validates one or more relay URLs and owns per-relay HTTP client and raw TLS dial config
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 -> NewClient -> Listen -> PublicURLs -> http.Server.Serve(listener)`
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
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
sdk/client.go
+26
-95
@@ -32,70 +32,70 @@ const (
32
defaultHTTPShutdownTimeout = 5 * time.Second
33
)
34
35
-type ClientOption func(*Client)
35
+type RelayClientOption func(*RelayClient)
36
37
-func WithRootCAPEM(rootCAPEM []byte) ClientOption {
37
+func WithRootCAPEM(rootCAPEM []byte) RelayClientOption {
38
rootCAPEM = append([]byte(nil), rootCAPEM...)
39
- return func(client *Client) {
39
+ return func(client *RelayClient) {
40
client.rootCAPEM = append([]byte(nil), rootCAPEM...)
41
}
42
}
43
44
-func WithInsecureSkipVerify(skip bool) ClientOption {
45
- return func(client *Client) {
44
+func WithInsecureSkipVerify(skip bool) RelayClientOption {
45
+ return func(client *RelayClient) {
46
client.insecureSkipVerify = skip
47
}
48
}
49
50
-func WithDialTimeout(timeout time.Duration) ClientOption {
51
- return func(client *Client) {
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) ClientOption {
59
- return func(client *Client) {
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) ClientOption {
67
- return func(client *Client) {
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) ClientOption {
75
- return func(client *Client) {
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) ClientOption {
83
- return func(client *Client) {
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) ClientOption {
91
- return func(client *Client) {
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 Client struct {
98
+type RelayClient struct {
99
baseURL *url.URL
100
httpClient *http.Client
101
rawTLSConfig *tls.Config
@@ -109,7 +109,7 @@ type Client struct {
109
readyTarget int
110
}
111
112
-func NewClient(relayURL string, options ...ClientOption) (*Client, error) {
112
+func NewRelayClient(relayURL string, options ...RelayClientOption) (*RelayClient, error) {
113
baseURL, err := url.Parse(strings.TrimSpace(relayURL))
114
if err != nil {
115
return nil, fmt.Errorf("parse relay url: %w", err)
@@ -124,7 +124,7 @@ func NewClient(relayURL string, options ...ClientOption) (*Client, error) {
124
baseURL.RawQuery = ""
125
baseURL.Fragment = ""
126
127
- client := &Client{
127
+ client := &RelayClient{
128
baseURL: baseURL,
129
dialTimeout: defaultDialTimeout,
130
requestTimeout: defaultRequestTimeout,
@@ -176,7 +176,7 @@ func NewClient(relayURL string, options ...ClientOption) (*Client, error) {
176
return client, nil
177
}
178
179
-func (c *Client) Close() {
179
+func (c *RelayClient) Close() {
180
if c == nil || c.httpClient == nil {
181
return
182
}
@@ -184,76 +184,7 @@ func (c *Client) Close() {
184
transport.CloseIdleConnections()
185
}
186
}
187
-
188
-func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, error) {
189
- if strings.TrimSpace(req.Name) == "" {
190
- return nil, errors.New("listener name is required")
191
- }
192
- if ctx == nil {
193
- ctx = context.Background()
194
- }
195
-
196
- reverseToken := strings.TrimSpace(req.ReverseToken)
197
- if reverseToken == "" {
198
- reverseToken = randomToken()
199
- }
200
-
201
- readyTarget := req.ReadyTarget
202
- if readyTarget <= 0 {
203
- readyTarget = c.readyTarget
204
- }
205
- leaseTTL := req.LeaseTTL
206
- if leaseTTL <= 0 {
207
- leaseTTL = c.leaseTTL
208
- }
209
- acceptedCap := max(readyTarget*2, 1)
210
-
211
- registerReq := types.RegisterRequest{
212
- Name: req.Name,
213
- Hostnames: req.Hostnames,
214
- Metadata: req.Metadata,
215
- ReverseToken: reverseToken,
216
- TLS: true,
217
- TTLSeconds: int(leaseTTL / time.Second),
218
- }
219
-
220
- var registerResp types.RegisterResponse
221
- if err := c.doJSON(ctx, http.MethodPost, types.PathSDKRegister, registerReq, ®isterResp); err != nil {
222
- return nil, err
223
- }
224
-
225
- tlsConf, tlsCloser, err := keyless.BuildClientTLSConfig(c.baseURL.String(), registerResp.Hostnames)
226
- if err != nil {
227
- _ = c.unregisterLease(context.Background(), registerResp.LeaseID, reverseToken)
228
- return nil, err
229
- }
230
-
231
- listenerCtx, cancel := context.WithCancel(ctx)
232
- listener := &Listener{
233
- client: c,
234
- baseContext: func() context.Context { return listenerCtx },
235
- ctxDone: listenerCtx.Done(),
236
- cancel: cancel,
237
- name: strings.TrimSpace(req.Name),
238
- leaseID: registerResp.LeaseID,
239
- hostnames: registerResp.Hostnames,
240
- metadata: registerResp.Metadata,
241
- reverseToken: reverseToken,
242
- leaseTTL: leaseTTL,
243
- readyTarget: readyTarget,
244
- tlsConfig: tlsConf,
245
- tlsCloser: tlsCloser,
246
- accepted: make(chan net.Conn, acceptedCap),
247
- signal: make(chan struct{}, 1),
248
- }
249
-
250
- go listener.runSupervisor()
251
- go listener.runRenewLoop()
252
- listener.notify()
253
- return listener, nil
254
-}
255
-
256
-func (c *Client) doJSON(ctx context.Context, method, path string, payload any, out any) error {
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)
@@ -299,7 +230,7 @@ func (c *Client) doJSON(ctx context.Context, method, path string, payload any, o
230
return json.Unmarshal(envelope.Data, out)
231
}
232
302
-func (c *Client) registerLease(ctx context.Context, req types.RegisterRequest) (types.RegisterResponse, error) {
233
+func (c *RelayClient) registerLease(ctx context.Context, req types.RegisterRequest) (types.RegisterResponse, error) {
234
var resp types.RegisterResponse
235
if err := c.doJSON(ctx, http.MethodPost, types.PathSDKRegister, req, &resp); err != nil {
236
return types.RegisterResponse{}, err
@@ -307,7 +238,7 @@ func (c *Client) registerLease(ctx context.Context, req types.RegisterRequest) (
238
return resp, nil
239
}
240
310
-func (c *Client) renewLease(ctx context.Context, leaseID, reverseToken string, ttl time.Duration) error {
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{
243
LeaseID: leaseID,
244
ReverseToken: reverseToken,
@@ -315,14 +246,14 @@ func (c *Client) renewLease(ctx context.Context, leaseID, reverseToken string, t
246
}, &types.RenewResponse{})
247
}
248
318
-func (c *Client) unregisterLease(ctx context.Context, leaseID, reverseToken string) error {
249
+func (c *RelayClient) unregisterLease(ctx context.Context, leaseID, reverseToken string) error {
250
return c.doJSON(ctx, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
251
LeaseID: leaseID,
252
ReverseToken: reverseToken,
253
}, nil)
254
}
255
325
-func (c *Client) openReverseSession(ctx context.Context, leaseID, reverseToken string) (net.Conn, error) {
256
+func (c *RelayClient) openReverseSession(ctx context.Context, leaseID, reverseToken string) (net.Conn, error) {
257
dialer := &tls.Dialer{
258
NetDialer: &net.Dialer{Timeout: c.dialTimeout},
259
Config: c.rawTLSConfig.Clone(),
sdk/client_test.go
+7
-7
@@ -15,7 +15,7 @@ import (
15
"github.com/gosuda/portal/v2/types"
16
)
17
18
-func TestNewClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
18
+func TestNewRelayClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
19
t.Parallel()
20
21
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
@@ -23,9 +23,9 @@ func TestNewClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
23
}))
24
defer server.Close()
25
26
- client, err := NewClient(server.URL)
26
+ client, err := NewRelayClient(server.URL)
27
if err != nil {
28
- t.Fatalf("NewClient() error = %v", err)
28
+ t.Fatalf("NewRelayClient() error = %v", err)
29
}
30
defer client.Close()
31
@@ -40,7 +40,7 @@ func TestNewClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
40
}
41
}
42
43
-func TestNewClientAppliesOptions(t *testing.T) {
43
+func TestNewRelayClientAppliesOptions(t *testing.T) {
44
t.Parallel()
45
46
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
@@ -52,7 +52,7 @@ func TestNewClientAppliesOptions(t *testing.T) {
52
Type: "CERTIFICATE",
53
Bytes: server.Certificate().Raw,
54
})
55
- client, err := NewClient(
55
+ client, err := NewRelayClient(
56
"https://relay.example.com/base/",
57
WithRootCAPEM(rootCAPEM),
58
WithInsecureSkipVerify(true),
@@ -64,7 +64,7 @@ func TestNewClientAppliesOptions(t *testing.T) {
64
WithReadyTarget(3),
65
)
66
if err != nil {
67
- t.Fatalf("NewClient() error = %v", err)
67
+ t.Fatalf("NewRelayClient() error = %v", err)
68
}
69
defer client.Close()
70
@@ -147,7 +147,7 @@ func TestOpenReverseSessionPreservesAPIErrorCode(t *testing.T) {
147
t.Fatalf("server client transport type = %T, want *http.Transport", server.Client().Transport)
148
}
149
150
- client := &Client{
150
+ client := &RelayClient{
151
baseURL: baseURL,
152
httpClient: &http.Client{
153
Transport: transport.Clone(),
sdk/helper.go
+24
-16
@@ -16,8 +16,8 @@ import (
16
"github.com/gosuda/portal/v2/types"
17
)
18
19
-// Exposure owns the lifecycle of one or more relay listeners plus their
20
-// clients and accepts traffic from all of them through one net.Listener.
19
+// Exposure owns the lifecycle of one or more relay listeners and accepts
20
+// traffic from all of them through one net.Listener.
21
type Exposure struct {
22
listener net.Listener
23
relays []exposureRelay
@@ -53,14 +53,22 @@ func Expose(ctx context.Context, relayUrls []string, name string, metadata types
53
}
54
55
for _, relayURL := range relayURLs {
56
- client, err := NewClient(relayURL)
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 := client.Listen(ctx, ListenRequest{
61
+ listener, err := NewListener(ctx, ListenRequest{
62
Name: name,
63
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,
72
})
73
if err != nil {
74
client.Close()
@@ -68,10 +76,9 @@ func Expose(ctx context.Context, relayUrls []string, name string, metadata types
76
}
77
78
relays = append(relays, exposureRelay{
71
- relayURL: relayURL,
72
- publicURLs: append([]string(nil), listener.PublicURLs()...),
73
- client: client,
74
- listener: listener,
79
+ relayURL: relayURL,
80
+ client: client,
81
+ listener: listener,
82
})
83
}
84
@@ -94,8 +101,7 @@ func Expose(ctx context.Context, relayUrls []string, name string, metadata types
101
logger.Info().
102
Int("relay_count", len(exposure.relays)).
103
Strs("relays", exposure.RelayURLs()).
97
- Strs("public_urls", exposure.PublicURLs()).
98
- Msg("exposure ready")
104
+ Msg("exposure starting")
105
106
return exposure, nil
107
}
@@ -164,7 +170,10 @@ func (e *Exposure) PublicURLs() []string {
170
out := make([]string, 0, len(e.relays))
171
seen := make(map[string]struct{})
172
for _, relay := range e.relays {
167
- for _, rawURL := range relay.publicURLs {
173
+ if relay.listener == nil {
174
+ continue
175
+ }
176
+ for _, rawURL := range relay.listener.publicURLs() {
177
if _, ok := seen[rawURL]; ok {
178
continue
179
}
@@ -189,7 +198,7 @@ func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr
198
return RunHTTP(ctx, relayListener, handler, localAddr)
199
}
200
192
-// Close closes the merged listener and all underlying SDK clients.
201
+// Close closes the merged listener and all underlying relay listeners and clients.
202
func (e *Exposure) Close() error {
203
if e == nil {
204
return nil
@@ -328,10 +337,9 @@ func RunHTTP(ctx context.Context, relayListener net.Listener, handler http.Handl
337
}
338
339
type exposureRelay struct {
331
- relayURL string
332
- publicURLs []string
333
- client *Client
334
- listener *Listener
340
+ relayURL string
341
+ client *RelayClient
342
+ listener *Listener
343
}
344
345
type exposureConn struct {
sdk/helper_test.go
+12
-6
@@ -531,12 +531,18 @@ func TestExposureAccessorsReturnCopies(t *testing.T) {
531
exposure := &Exposure{
532
relays: []exposureRelay{
533
{
534
- relayURL: "https://relay-1.example.com",
535
- publicURLs: []string{"https://app.example.com"},
534
+ relayURL: "https://relay-1.example.com",
535
+ listener: &Listener{
536
+ hostnames: []string{"app.example.com"},
537
+ state: listenerStateReady,
538
+ },
539
},
540
{
538
- relayURL: "https://relay-2.example.com",
539
- publicURLs: []string{"https://app.example.com"},
541
+ relayURL: "https://relay-2.example.com",
542
+ listener: &Listener{
543
+ hostnames: []string{"app.example.com"},
544
+ state: listenerStateReady,
545
+ },
546
},
547
},
548
}
@@ -556,8 +562,8 @@ func TestExposureAccessorsReturnCopies(t *testing.T) {
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
}
559
- if got, want := exposure.relays[0].publicURLs, []string{"https://app.example.com"}; !reflect.DeepEqual(got, want) {
560
- t.Fatalf("relays[0].publicURLs = %v, want %v", got, want)
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
sdk/listener.go
+537
-219
@@ -4,9 +4,9 @@ import (
4
"context"
5
"crypto/tls"
6
"errors"
7
- "fmt"
7
"io"
8
"net"
9
+ "strings"
10
"sync"
11
"time"
12
@@ -16,6 +16,20 @@ import (
16
"github.com/gosuda/portal/v2/types"
17
)
18
19
+const (
20
+ defaultListenerRetryCount = 30
21
+ defaultListenerRetryDelay = time.Second
22
+)
23
+
24
+type listenerState uint8
25
+
26
+const (
27
+ listenerStatePending listenerState = iota
28
+ listenerStateReady
29
+ listenerStateStale
30
+ listenerStateClosed
31
+)
32
+
33
type ListenRequest struct {
34
Name string
35
ReverseToken string
@@ -25,205 +39,468 @@ type ListenRequest struct {
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
50
+}
51
+
52
type Listener struct {
29
- tlsCloser io.Closer
30
- tlsConfig *tls.Config
31
- baseContext func() context.Context
32
- ctxDone <-chan struct{}
33
- cancel context.CancelFunc
34
- client *Client
35
- signal chan struct{}
36
- accepted chan net.Conn
37
- name string
38
- leaseID string
39
- reverseToken string
40
- hostnames []string
41
- metadata types.LeaseMetadata
42
- readyTarget int
43
- leaseTTL time.Duration
44
-
45
- activeSessions int
46
- closeOnce sync.Once
47
- mu sync.Mutex
53
+ ctx context.Context
54
+ cancel context.CancelFunc
55
+ accepted chan net.Conn
56
+ refill chan struct{}
57
+
58
+ name string
59
+ reverseToken string
60
+ metadata types.LeaseMetadata
61
+ readyTarget int
62
+ leaseTTL time.Duration
63
+ renewInterval time.Duration
64
+ handshakeTimeout time.Duration
65
+ retryCount int
66
+ retryDelay time.Duration
67
+
68
+ mu sync.Mutex
69
+ client *RelayClient
70
+ leaseID string
71
+ hostnames []string
72
+ tlsConfig *tls.Config
73
+ tlsCloser io.Closer
74
+ activeSessions int
75
+ sessionFailures int
76
+ state listenerState
77
+ runID uint64
78
+
79
+ closeOnce sync.Once
80
+ closeErr error
81
+}
82
+
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")
89
+ }
90
+ if ctx == nil {
91
+ ctx = context.Background()
92
+ }
93
+
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
+ }
124
+ renewInterval := leaseTTL / 2
125
+ if leaseTTL > renewBefore {
126
+ renewInterval = leaseTTL - renewBefore
127
+ }
128
+ if renewInterval <= 0 {
129
+ renewInterval = leaseTTL / 2
130
+ }
131
+ if renewInterval <= 0 {
132
+ renewInterval = time.Second
133
+ }
134
+
135
+ retryCount := opts.retryCount
136
+ if retryCount <= 0 {
137
+ retryCount = defaultListenerRetryCount
138
+ }
139
+ retryDelay := opts.retryDelay
140
+ if retryDelay <= 0 {
141
+ retryDelay = defaultListenerRetryDelay
142
+ }
143
+
144
+ listenerCtx, cancel := context.WithCancel(ctx)
145
+ l := &Listener{
146
+ ctx: listenerCtx,
147
+ cancel: cancel,
148
+ accepted: make(chan net.Conn, max(readyTarget*2, 1)),
149
+ refill: make(chan struct{}, 1),
150
+ name: strings.TrimSpace(req.Name),
151
+ reverseToken: reverseToken,
152
+ metadata: cloneMetadata(req.Metadata),
153
+ readyTarget: readyTarget,
154
+ leaseTTL: leaseTTL,
155
+ renewInterval: renewInterval,
156
+ handshakeTimeout: handshakeTimeout,
157
+ retryCount: retryCount,
158
+ retryDelay: retryDelay,
159
+ client: opts.client,
160
+ hostnames: append([]string(nil), req.Hostnames...),
161
+ state: listenerStatePending,
162
+ runID: 1,
163
+ }
164
+
165
+ go l.run(listenerCtx, l.runID)
166
+ return l, nil
167
}
168
169
func (l *Listener) Accept() (net.Conn, error) {
170
select {
52
- case <-l.ctxDone:
171
+ case <-l.ctx.Done():
172
+ select {
173
+ case conn := <-l.accepted:
174
+ if conn != nil {
175
+ _ = conn.Close()
176
+ }
177
+ default:
178
+ }
179
return nil, net.ErrClosed
180
case conn := <-l.accepted:
181
if conn == nil {
182
return nil, net.ErrClosed
183
}
58
- return conn, nil
184
+ select {
185
+ case <-l.ctx.Done():
186
+ _ = conn.Close()
187
+ return nil, net.ErrClosed
188
+ default:
189
+ return conn, nil
190
+ }
191
}
192
}
193
194
func (l *Listener) Close() error {
195
var closeErr error
196
l.closeOnce.Do(func() {
65
- l.cancel()
66
-
67
- l.mu.Lock()
68
- leaseID := l.leaseID
69
- tlsCloser := l.tlsCloser
70
- l.mu.Unlock()
71
-
72
- ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
73
- defer cancel()
74
- if err := l.client.unregisterLease(ctx, leaseID, l.reverseToken); err != nil {
75
- closeErr = err
76
- }
77
- if tlsCloser != nil {
78
- closeErr = errors.Join(closeErr, tlsCloser.Close())
79
- }
197
+ closeErr = l.closeCurrent()
198
})
199
return closeErr
200
}
201
202
func (l *Listener) Addr() net.Addr {
85
- return listenerAddr("portal:" + l.leaseID)
86
-}
87
-
88
-func (l *Listener) LeaseID() string {
203
l.mu.Lock()
204
defer l.mu.Unlock()
91
- return l.leaseID
92
-}
205
94
-func (l *Listener) Hostnames() []string {
95
- l.mu.Lock()
96
- defer l.mu.Unlock()
97
- return l.hostnames
206
+ if strings.TrimSpace(l.leaseID) == "" {
207
+ return listenerAddr("portal:pending")
208
+ }
209
+ return listenerAddr("portal:" + l.leaseID)
210
}
211
100
-func (l *Listener) Metadata() types.LeaseMetadata {
101
- l.mu.Lock()
102
- defer l.mu.Unlock()
103
- return l.metadata
104
-}
212
+func (l *Listener) Reactivate(ctx context.Context) error {
213
+ if ctx == nil {
214
+ ctx = context.Background()
215
+ }
216
106
-func (l *Listener) PublicURLs() []string {
217
l.mu.Lock()
108
- hostnames := l.hostnames
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
111
- urls := make([]string, 0, len(hostnames))
112
- for _, host := range hostnames {
113
- urls = append(urls, "https://"+host)
114
- }
115
- return urls
237
+ l.drainAccepted()
238
+ go l.run(runCtx, runID)
239
+ return nil
240
}
241
118
-func (l *Listener) runSupervisor() {
119
- for {
120
- select {
121
- case <-l.ctxDone:
242
+func (l *Listener) run(runCtx context.Context, runID uint64) {
243
+ logger := log.With().
244
+ Str("component", "sdk-listener").
245
+ Str("name", l.name).
246
+ Logger()
247
+
248
+ var err error
249
+ for attempt := 1; attempt <= l.retryCount; attempt++ {
250
+ err = l.establish(runID, runCtx)
251
+ if err == nil {
252
+ l.mu.Lock()
253
+ if l.runID != runID {
254
+ l.mu.Unlock()
255
+ return
256
+ }
257
+ leaseID := l.leaseID
258
+ hostnames := append([]string(nil), l.hostnames...)
259
+ l.mu.Unlock()
260
+
261
+ logger.Info().
262
+ Str("lease_id", leaseID).
263
+ Strs("hostnames", hostnames).
264
+ Msg("listener connected")
265
+
266
+ go l.runSessionPool(runCtx, runID)
267
+ go l.runRenewLoop(runCtx, runID)
268
+ l.signalRefill()
269
return
123
- case <-l.signal:
270
}
271
126
- for l.reserveSessionSlot() {
127
- go l.runSession()
272
+ if errors.Is(runCtx.Err(), context.Canceled) {
273
+ return
274
+ }
275
+
276
+ logger.Warn().
277
+ Err(err).
278
+ Int("attempt", attempt).
279
+ Dur("retry_in", l.retryDelay).
280
+ Msg("listener bootstrap failed")
281
+
282
+ if attempt == l.retryCount {
283
+ break
284
+ }
285
+ if !sleepOrDone(runCtx, l.retryDelay) {
286
+ return
287
}
288
}
289
+
290
+ l.fail(runID, err, "listener bootstrap retry limit reached")
291
}
292
132
-func (l *Listener) runRenewLoop() {
133
- interval := l.leaseTTL / 2
134
- if interval <= 0 {
135
- interval = 30 * time.Second
293
+func (l *Listener) establish(runID uint64, runCtx context.Context) error {
294
+ l.mu.Lock()
295
+ if l.runID != runID {
296
+ l.mu.Unlock()
297
+ return context.Canceled
298
}
137
- if l.client.renewBefore > 0 && l.leaseTTL > l.client.renewBefore {
138
- interval = l.leaseTTL - l.client.renewBefore
299
+ client := l.client
300
+ hostnames := append([]string(nil), l.hostnames...)
301
+ l.mu.Unlock()
302
+
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
+ })
311
+ if err != nil {
312
+ return err
313
}
140
- if interval <= 0 {
141
- interval = 30 * time.Second
314
+
315
+ tlsConfig, tlsCloser, err := keyless.BuildClientTLSConfig(l.client.baseURL.String(), resp.Hostnames)
316
+ if err != nil {
317
+ _ = client.unregisterLease(context.Background(), resp.LeaseID, l.reverseToken)
318
+ return err
319
}
320
144
- ticker := time.NewTicker(interval)
145
- defer ticker.Stop()
321
+ l.mu.Lock()
322
+ if l.runID != runID || l.state == listenerStateClosed {
323
+ l.mu.Unlock()
324
+ _ = client.unregisterLease(context.Background(), resp.LeaseID, l.reverseToken)
325
+ _ = tlsCloser.Close()
326
+ return context.Canceled
327
+ }
328
+ oldCloser := l.tlsCloser
329
+ l.client = client
330
+ l.leaseID = resp.LeaseID
331
+ l.hostnames = append([]string(nil), resp.Hostnames...)
332
+ l.tlsConfig = tlsConfig
333
+ l.tlsCloser = tlsCloser
334
+ l.activeSessions = 0
335
+ l.sessionFailures = 0
336
+ l.state = listenerStateReady
337
+ l.mu.Unlock()
338
147
- var consecutiveFailures int
339
+ if oldCloser != nil {
340
+ _ = oldCloser.Close()
341
+ }
342
+ return nil
343
+}
344
345
+func (l *Listener) runSessionPool(runCtx context.Context, runID uint64) {
346
for {
347
select {
151
- case <-l.ctxDone:
348
+ case <-runCtx.Done():
349
return
153
- case <-ticker.C:
350
+ case <-l.refill:
351
+ }
352
+
353
+ for {
354
l.mu.Lock()
155
- leaseID := l.leaseID
355
+ ready := l.runID == runID &&
356
+ l.state == listenerStateReady &&
357
+ l.client != nil &&
358
+ strings.TrimSpace(l.leaseID) != "" &&
359
+ l.tlsConfig != nil &&
360
+ l.activeSessions < l.readyTarget
361
+ if !ready {
362
+ l.mu.Unlock()
363
+ break
364
+ }
365
+ l.activeSessions++
366
l.mu.Unlock()
367
158
- ctx, cancel := context.WithTimeout(l.context(), 10*time.Second)
159
- err := l.client.renewLease(ctx, leaseID, l.reverseToken, l.leaseTTL)
160
- cancel()
368
+ go l.runSession(runCtx, runID)
369
+ }
370
+ }
371
+}
372
162
- if err != nil {
163
- if isLeaseNotFound(err) {
164
- log.Warn().
165
- Str("component", "sdk-listener").
166
- Str("lease_id", leaseID).
167
- Msg("lease not found on relay, attempting re-registration")
168
- if reregErr := l.reregister(); reregErr != nil {
169
- log.Error().Err(reregErr).
170
- Str("component", "sdk-listener").
171
- Msg("lease re-registration failed")
172
- } else {
173
- consecutiveFailures = 0
174
- log.Info().
175
- Str("component", "sdk-listener").
176
- Str("lease_id", l.LeaseID()).
177
- Strs("hostnames", l.Hostnames()).
178
- Msg("lease re-registered successfully")
179
- continue
180
- }
181
- }
182
-
183
- consecutiveFailures++
184
- event := log.Warn()
185
- if consecutiveFailures >= 3 {
186
- event = log.Error()
187
- }
188
- event.Err(err).
373
+func (l *Listener) runRenewLoop(runCtx context.Context, runID uint64) {
374
+ failures := 0
375
+ wait := l.renewInterval
376
+
377
+ for {
378
+ if !sleepOrDone(runCtx, wait) {
379
+ return
380
+ }
381
+
382
+ l.mu.Lock()
383
+ current := l.runID == runID
384
+ client := l.client
385
+ leaseID := l.leaseID
386
+ ready := l.state == listenerStateReady
387
+ l.mu.Unlock()
388
+
389
+ if !current || !ready || client == nil || strings.TrimSpace(leaseID) == "" {
390
+ wait = l.renewInterval
391
+ continue
392
+ }
393
+
394
+ ctx, cancel := context.WithTimeout(runCtx, 10*time.Second)
395
+ err := client.renewLease(ctx, leaseID, l.reverseToken, l.leaseTTL)
396
+ cancel()
397
+
398
+ if err == nil {
399
+ failures = 0
400
+ wait = l.renewInterval
401
+ continue
402
+ }
403
+
404
+ if isLeaseNotFound(err) {
405
+ log.Warn().
406
+ Str("component", "sdk-listener").
407
+ Str("lease_id", leaseID).
408
+ Msg("lease not found on relay, attempting re-registration")
409
+
410
+ err = l.establish(runID, runCtx)
411
+ if err == nil {
412
+ failures = 0
413
+ wait = l.renewInterval
414
+ l.signalRefill()
415
+
416
+ l.mu.Lock()
417
+ leaseID = l.leaseID
418
+ hostnames := append([]string(nil), l.hostnames...)
419
+ l.mu.Unlock()
420
+
421
+ log.Info().
422
Str("component", "sdk-listener").
190
- Str("lease_id", l.LeaseID()).
191
- Int("consecutive_failures", consecutiveFailures).
192
- Msg("lease renewal failed")
193
- } else {
194
- consecutiveFailures = 0
423
+ Str("lease_id", leaseID).
424
+ Strs("hostnames", hostnames).
425
+ Msg("lease re-registered successfully")
426
+ continue
427
}
428
}
429
+
430
+ failures++
431
+ event := log.Warn()
432
+ if failures >= l.retryCount {
433
+ event = log.Error()
434
+ }
435
+ event.Err(err).
436
+ Str("component", "sdk-listener").
437
+ Str("lease_id", leaseID).
438
+ Int("consecutive_failures", failures).
439
+ Msg("lease renewal failed")
440
+
441
+ if failures >= l.retryCount {
442
+ l.fail(runID, err, "listener renew retry limit reached")
443
+ return
444
+ }
445
+ wait = l.retryDelay
446
}
447
}
448
200
-func (l *Listener) runSession() {
201
- defer l.releaseSessionSlot()
449
+func (l *Listener) runSession(runCtx context.Context, runID uint64) {
450
+ defer func() {
451
+ l.mu.Lock()
452
+ if l.runID == runID && l.activeSessions > 0 {
453
+ l.activeSessions--
454
+ }
455
+ l.mu.Unlock()
456
+ l.signalRefill()
457
+ }()
458
+
459
+ fail := func(err error) {
460
+ if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
461
+ return
462
+ }
463
+
464
+ l.mu.Lock()
465
+ if l.runID != runID {
466
+ l.mu.Unlock()
467
+ return
468
+ }
469
+ l.sessionFailures++
470
+ failures := l.sessionFailures
471
+ l.mu.Unlock()
472
+
473
+ if failures >= l.retryCount {
474
+ l.fail(runID, err, "listener session retry limit reached")
475
+ return
476
+ }
477
+
478
+ _ = sleepOrDone(runCtx, l.retryDelay)
479
+ }
480
203
- sessionCtx := l.context()
481
l.mu.Lock()
482
+ ready := l.runID == runID && l.state == listenerStateReady
483
+ client := l.client
484
leaseID := l.leaseID
485
+ tlsConfig := l.tlsConfig
486
l.mu.Unlock()
207
- conn, err := l.client.openReverseSession(sessionCtx, leaseID, l.reverseToken)
208
- if err != nil {
209
- sleepOrDone(sessionCtx, time.Second)
487
+ if !ready || client == nil || strings.TrimSpace(leaseID) == "" || tlsConfig == nil {
488
return
489
}
490
213
- if err := l.awaitActivation(conn); err != nil {
214
- _ = conn.Close()
215
- if !errors.Is(err, context.Canceled) && !errors.Is(err, net.ErrClosed) {
216
- sleepOrDone(sessionCtx, time.Second)
217
- }
491
+ conn, err := client.openReverseSession(runCtx, leaseID, l.reverseToken)
492
+ if err != nil {
493
+ fail(err)
494
+ return
495
}
219
-}
496
221
-func (l *Listener) awaitActivation(conn net.Conn) error {
497
var marker [1]byte
498
for {
224
- _ = conn.SetReadDeadline(time.Now().Add(2 * l.client.handshakeTimeout))
499
+ _ = conn.SetReadDeadline(time.Now().Add(2 * l.handshakeTimeout))
500
if _, err := io.ReadFull(conn, marker[:]); err != nil {
226
- return err
501
+ _ = conn.Close()
502
+ fail(err)
503
+ return
504
}
505
_ = conn.SetReadDeadline(time.Time{})
506
@@ -231,139 +508,180 @@ func (l *Listener) awaitActivation(conn net.Conn) error {
508
case types.MarkerKeepalive:
509
continue
510
case types.MarkerTLSStart:
234
- return l.activate(conn)
511
+ tlsConn := tls.Server(conn, tlsConfig)
512
+ handshakeCtx, cancel := context.WithTimeout(runCtx, l.handshakeTimeout)
513
+ err := tlsConn.HandshakeContext(handshakeCtx)
514
+ cancel()
515
+ if err != nil {
516
+ _ = tlsConn.Close()
517
+ fail(err)
518
+ return
519
+ }
520
+
521
+ l.mu.Lock()
522
+ l.sessionFailures = 0
523
+ l.mu.Unlock()
524
+
525
+ select {
526
+ case <-runCtx.Done():
527
+ _ = tlsConn.Close()
528
+ case l.accepted <- tlsConn:
529
+ }
530
+ return
531
default:
236
- return fmt.Errorf("unexpected reverse marker: 0x%02x", marker[0])
532
+ _ = conn.Close()
533
+ fail(errors.New("unexpected reverse marker"))
534
+ return
535
}
536
}
537
}
538
241
-func (l *Listener) activate(conn net.Conn) error {
539
+func (l *Listener) publicURLs() []string {
540
l.mu.Lock()
243
- tlsCfg := l.tlsConfig
244
- l.mu.Unlock()
245
- // Reuse the shared config so session ticket state survives across connections.
246
- tlsConn := tls.Server(conn, tlsCfg)
247
- handshakeCtx, cancel := context.WithTimeout(l.context(), l.client.handshakeTimeout)
248
- defer cancel()
249
- if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
250
- return err
251
- }
541
+ defer l.mu.Unlock()
542
253
- select {
254
- case <-l.ctxDone:
255
- _ = tlsConn.Close()
256
- return l.context().Err()
257
- case l.accepted <- tlsConn:
543
+ if l.state != listenerStateReady || len(l.hostnames) == 0 {
544
return nil
545
}
260
-}
546
262
-func (l *Listener) reregister() error {
263
- ctx, cancel := context.WithTimeout(l.context(), 10*time.Second)
264
- defer cancel()
547
+ urls := make([]string, 0, len(l.hostnames))
548
+ for _, host := range l.hostnames {
549
+ urls = append(urls, "https://"+host)
550
+ }
551
+ return urls
552
+}
553
554
+func (l *Listener) closeCurrent() error {
555
l.mu.Lock()
267
- hostnames := l.hostnames
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++
567
l.mu.Unlock()
568
270
- resp, err := l.client.registerLease(ctx, types.RegisterRequest{
271
- Name: l.name,
272
- Hostnames: hostnames,
273
- Metadata: l.metadata,
274
- ReverseToken: l.reverseToken,
275
- TLS: true,
276
- TTLSeconds: int(l.leaseTTL / time.Second),
277
- })
278
- if err != nil {
279
- return err
569
+ if cancel != nil {
570
+ cancel()
571
}
572
+ l.drainAccepted()
573
282
- tlsConf, tlsCloser, err := keyless.BuildClientTLSConfig(l.client.baseURL.String(), resp.Hostnames)
283
- if err != nil {
284
- _ = l.client.unregisterLease(ctx, resp.LeaseID, l.reverseToken)
285
- return err
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()
579
}
287
-
288
- l.mu.Lock()
289
- oldCloser := l.tlsCloser
290
- l.leaseID = resp.LeaseID
291
- l.hostnames = resp.Hostnames
292
- l.metadata = resp.Metadata
293
- l.tlsConfig = tlsConf
294
- l.tlsCloser = tlsCloser
295
- l.mu.Unlock()
296
-
297
- if oldCloser != nil {
298
- _ = oldCloser.Close()
580
+ if tlsCloser != nil {
581
+ closeErr = errors.Join(closeErr, tlsCloser.Close())
582
}
300
-
301
- l.notify()
302
- return nil
303
-}
304
-
305
-func isLeaseNotFound(err error) bool {
306
- return errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeLeaseNotFound})
583
+ l.closeErr = closeErr
584
+ return closeErr
585
}
586
309
-func (l *Listener) reserveSessionSlot() bool {
310
- l.mu.Lock()
311
- defer l.mu.Unlock()
312
- if l.isClosed() {
313
- return false
587
+func (l *Listener) fail(runID uint64, err error, message string) {
588
+ closeErr, changed := l.markStale(runID)
589
+ if !changed {
590
+ return
591
}
315
- if l.activeSessions >= l.readyTarget {
316
- return false
592
+ if closeErr != nil {
593
+ err = errors.Join(err, closeErr)
594
}
318
- l.activeSessions++
319
- return true
595
+ log.Error().
596
+ Str("component", "sdk-listener").
597
+ Str("name", l.name).
598
+ Err(err).
599
+ Msg(message)
600
}
601
322
-func (l *Listener) releaseSessionSlot() {
602
+func (l *Listener) markStale(runID uint64) (error, bool) {
603
l.mu.Lock()
324
- l.activeSessions--
604
+ if l.runID != runID || l.state == listenerStateClosed || l.state == listenerStateStale {
605
+ l.mu.Unlock()
606
+ return nil, false
607
+ }
608
+
609
+ l.state = listenerStateStale
610
+ cancel := l.cancel
611
+ client := l.client
612
+ leaseID := l.leaseID
613
+ tlsCloser := l.tlsCloser
614
+ l.leaseID = ""
615
+ l.tlsConfig = nil
616
+ l.tlsCloser = nil
617
+ l.activeSessions = 0
618
+ l.sessionFailures = 0
619
l.mu.Unlock()
326
- l.notify()
620
+
621
+ if cancel != nil {
622
+ cancel()
623
+ }
624
+ l.drainAccepted()
625
+
626
+ var closeErr error
627
+ if client != nil && strings.TrimSpace(leaseID) != "" {
628
+ ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
629
+ closeErr = errors.Join(closeErr, client.unregisterLease(ctx, leaseID, l.reverseToken))
630
+ cancel()
631
+ }
632
+ if tlsCloser != nil {
633
+ closeErr = errors.Join(closeErr, tlsCloser.Close())
634
+ }
635
+ return closeErr, true
636
}
637
329
-func (l *Listener) notify() {
638
+func (l *Listener) signalRefill() {
639
select {
331
- case l.signal <- struct{}{}:
640
+ case l.refill <- struct{}{}:
641
default:
642
}
643
}
644
336
-func sleepOrDone(ctx context.Context, d time.Duration) {
645
+func isLeaseNotFound(err error) bool {
646
+ return errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeLeaseNotFound})
647
+}
648
+
649
+func sleepOrDone(ctx context.Context, d time.Duration) bool {
650
timer := time.NewTimer(d)
651
defer timer.Stop()
652
+
653
select {
654
case <-ctx.Done():
655
+ return false
656
case <-timer.C:
657
+ return true
658
}
659
}
660
345
-type listenerAddr string
346
-
347
-func (a listenerAddr) Network() string { return "portal" }
348
-func (a listenerAddr) String() string { return string(a) }
349
-
350
-func (l *Listener) context() context.Context {
351
- if l.baseContext != nil {
352
- if ctx := l.baseContext(); ctx != nil {
353
- return ctx
661
+func (l *Listener) drainAccepted() {
662
+ for {
663
+ select {
664
+ case conn := <-l.accepted:
665
+ if conn != nil {
666
+ _ = conn.Close()
667
+ }
668
+ default:
669
+ return
670
}
671
}
356
- return context.Background()
672
}
673
359
-func (l *Listener) isClosed() bool {
360
- if l.ctxDone == nil {
361
- return false
362
- }
363
- select {
364
- case <-l.ctxDone:
365
- return true
366
- default:
367
- return false
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
+
684
+type listenerAddr string
685
+
686
+func (a listenerAddr) Network() string { return "portal" }
687
+func (a listenerAddr) String() string { return string(a) }
sdk/listener_test.go
+314
-27
@@ -1,63 +1,54 @@
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
10
-func TestListenerAccessors(t *testing.T) {
18
+func TestListenerSnapshotAndAddr(t *testing.T) {
19
t.Parallel()
20
21
listener := &Listener{
22
leaseID: "lease-1",
23
hostnames: []string{"app.relay.example.com"},
16
- metadata: types.LeaseMetadata{
17
- Owner: "alice",
18
- Tags: []string{"one", "two"},
19
- },
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
}
25
- if listener.LeaseID() != "lease-1" {
26
- t.Fatalf("LeaseID() = %q, want %q", listener.LeaseID(), "lease-1")
27
- }
28
-
29
- hostnames := listener.Hostnames()
30
- if len(hostnames) != 1 || hostnames[0] != "app.relay.example.com" {
31
- t.Fatalf("Hostnames() = %#v, want [app.relay.example.com]", hostnames)
32
- }
33
-
34
- metadata := listener.Metadata()
35
- if metadata.Owner != "alice" {
36
- t.Fatalf("Metadata().Owner = %q, want %q", metadata.Owner, "alice")
37
- }
38
- if len(metadata.Tags) != 2 {
39
- t.Fatalf("Metadata().Tags len = %d, want 2", len(metadata.Tags))
40
- }
30
42
- publicURLs := listener.PublicURLs()
31
+ publicURLs := listener.publicURLs()
32
if len(publicURLs) != 1 || publicURLs[0] != "https://app.relay.example.com" {
44
- t.Fatalf("PublicURLs() = %#v, want [https://app.relay.example.com]", publicURLs)
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
51
- done := make(chan struct{})
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{
58
- ctxDone: done,
59
- accepted: make(chan net.Conn, 2),
60
- hostnames: []string{"app.relay.example.com"},
49
+ ctx: ctx,
50
+ cancel: cancel,
51
+ accepted: make(chan net.Conn, 2),
52
}
53
listener.accepted <- serverConn1
54
listener.accepted <- serverConn2
@@ -81,3 +72,299 @@ func TestListenerAccept(t *testing.T) {
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
+}