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, &registerResp); 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 +}