refact: fix listener

rabbitprincess committed Mar 11, 2026 at 00:34 UTC dd13d0800cd171fc23ce36f40f31da6d7a726433
11 files changed +687 -605
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 no longer fails on a temporarily unavailable relay. Each configured relay listener keeps retrying until it connects or the tunnel is stopped.
34 +- Startup is fail-fast: if any configured relay cannot register, the tunnel exits instead of partially publishing.
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
+6 -6
@@ -54,10 +54,10 @@ That distinction matters because `/sdk/connect` stops being ordinary HTTP once h
54
55 ### SDK (`sdk/`)
56
57 -- `Listener`: validates one relay URL, registers one lease per relay, maintains per-entry `readyTarget` reverse sessions, renews lease TTLs, and yields accepted tenant TLS connections; listener failure is terminal and callers create a new listener instead of reactivating the old one
57 +- `Listener`: validates one relay URL, registers one lease per relay, maintains per-entry `readyTarget` reverse sessions, renews lease TTLs, and yields accepted tenant TLS connections
58 - `relayclient.go`: internal relay transport helper for control-plane requests and reverse session dialing
59 - Default app flow is `RelayURL -> NewListener -> PublicURLs -> http.Server.Serve(listener)` or `RelayURLs -> Expose -> PublicURLs -> http.Server.Serve(exposure)`
60 -- `expose.go`: optional `RunHTTPApp` helper for serving one handler on both a local HTTP port and the relay listener
60 +- `expose.go`: optional `RunHTTP` helper for serving one handler on both a local HTTP port and the relay listener
61 - Relay-aware entry inspection is reserved for advanced callers such as `portal-tunnel`
62 - Tenant TLS is created automatically through the relay keyless signer; callers do not provide a local self-signed fallback path
63
@@ -77,8 +77,8 @@ That distinction matters because `/sdk/connect` stops being ordinary HTTP once h
77 3. Each relay hijacks `/sdk/connect` requests and places the connection in the per-lease broker ready queue.
78 4. While idle, the relay writes `0x00` keepalive markers.
79 5. A browser connects to the relay SNI listener.
80 -6. Relay extracts SNI from ClientHello, resolves a lease, and claims one ready reverse session.
81 -7. Relay writes `0x02` to activate that session.
80 +6. Relay extracts SNI from ClientHello, resolves a lease, and waits up to `ClaimTimeout` for one reverse session from that lease broker.
81 +7. Relay writes `0x02` to activate the claimed session.
82 8. SDK/tunnel receives `0x02`, starts tenant TLS locally using the relay-backed keyless signer, and the relay bridges raw encrypted bytes end-to-end.
83
84 Result: the relay decides routing, but tenant TLS termination still happens at the SDK/tunnel side.
@@ -94,8 +94,9 @@ Result: the relay decides routing, but tenant TLS termination still happens at t
94 - `reverse_token`
95 - optional `hostnames`
96 - optional `metadata`
97 - - optional `ttl_seconds`
97 + - optional `ttl`
98 - If no hostname is supplied, relay derives one from `name + root host`
99 +- Registration reserves the hostname and publishes the route immediately; if no reverse session is ready yet, inbound SNI claims wait up to `ClaimTimeout`
100 - `PORTAL_URL` is normalized to its host component only; path/query segments are ignored for routing
101
102 ### 2. Reverse Connect
@@ -113,7 +114,6 @@ Result: the relay decides routing, but tenant TLS termination still happens at t
114 - `POST /sdk/renew`
115 - Requires `lease_id` + `reverse_token`
116 - Extends lease TTL
116 -- Resets a previously dropped broker back to active
117
118 ### 4. Unregister
119
docs/greenfield-raw-tcp-sni-keyless.md
+9 -14
@@ -129,26 +129,22 @@ If future work wants tenant HTTP/2, it must first remove cross-tenant certificat
129 ```text
130 LeaseBroker
131 lease_id
132 - state: active | dropped | stopped
132 + state: open | closed
133 ready queue: bounded FIFO of idle reverse sessions
134 - metrics: ready_count, claimed_count, dropped_count, last_claim_at
134 + metrics: ready_count, claim_wait_duration
135 ```
136
137 ### LeaseBroker API
138
139 - `Offer(session) error`
140 - `Claim(ctx) (*ReverseSession, error)`
141 -- `Drop()`
142 -- `Reset()`
143 -- `Stop()`
141 +- `Close()`
142
143 ### Rules
144
147 -- `Offer` rejects immediately if broker is dropped or stopped
145 +- `Offer` rejects immediately if broker is closed
146 - `Claim` blocks until a valid idle session is available or timeout/cancel fires
149 -- `Drop` drains ready queue and closes all idle sessions
150 -- `Reset` reopens a dropped lease after successful re-registration
151 -- `Stop` is terminal and used only for process shutdown
147 +- `Close` drains the ready queue, closes all idle sessions, and wakes blocked claimers
148
149 No separate global `dropped` map exists outside the broker.
150
@@ -324,9 +320,9 @@ Transient:
320
321 ### SDK Agent Rules
322
327 -- Fatal rejection closes the current listener
328 -- Recovery creates a fresh listener; the SDK does not reactivate a failed listener in place
329 -- Transient failure retries with bounded backoff inside one listener lifecycle
323 +- Fatal rejection pauses the agent
324 +- Successful renew/register clears pause
325 +- Transient failure retries with bounded backoff
326
327 ## Backpressure
328
@@ -356,8 +352,7 @@ Required structured logs:
352 - bridge started
353 - bridge ended
354 - session evicted
359 -- lease dropped
360 -- lease reset
355 +- broker closed
356 - fatal agent pause
357
358 Core metrics:
portal/broker.go
+37 -71
@@ -14,29 +14,18 @@ import (
14 )
15
16 var (
17 - errLeaseDropped = errors.New("lease dropped")
18 - errLeaseStopped = errors.New("lease stopped")
17 + errBrokerClosed = errors.New("lease broker closed")
18 errBrokerFull = errors.New("broker ready queue full")
20 - errNoSessions = errors.New("no reverse sessions available")
21 -)
22 -
23 -type brokerState int
24 -
25 -const (
26 - brokerStateActive brokerState = iota
27 - brokerStateDropped
28 - brokerStateStopped
19 )
20
21 type leaseBroker struct {
32 - notify chan struct{}
33 - leaseID string
34 - ready []*reverseSession
35 - idleInterval time.Duration
36 - readyLimit int
37 - totalSessions int
38 - state brokerState
39 - mu sync.Mutex
22 + notify chan struct{}
23 + leaseID string
24 + ready []*reverseSession
25 + idleInterval time.Duration
26 + readyLimit int
27 + closedErr error
28 + mu sync.Mutex
29 }
30
31 func newLeaseBroker(leaseID string, idleInterval time.Duration, readyLimit int) *leaseBroker {
@@ -54,23 +43,22 @@ func (b *leaseBroker) Offer(session *reverseSession) error {
43 }
44
45 b.mu.Lock()
57 - defer b.mu.Unlock()
58 -
59 - switch b.state {
60 - case brokerStateDropped:
61 - return errLeaseDropped
62 - case brokerStateStopped:
63 - return errLeaseStopped
46 + if b.closedErr != nil {
47 + err := b.closedErr
48 + b.mu.Unlock()
49 + return err
50 }
51
52 if b.readyLimit > 0 && len(b.ready) >= b.readyLimit {
53 + b.mu.Unlock()
54 return errBrokerFull
55 }
56
57 session.StartIdle()
58 b.ready = append(b.ready, session)
72 - b.totalSessions++
59 b.signalLocked()
60 + b.mu.Unlock()
61 +
62 go b.watchSession(session)
63 return nil
64 }
@@ -78,13 +66,10 @@ func (b *leaseBroker) Offer(session *reverseSession) error {
66 func (b *leaseBroker) Claim(ctx context.Context) (*reverseSession, error) {
67 for {
68 b.mu.Lock()
81 - switch b.state {
82 - case brokerStateDropped:
83 - b.mu.Unlock()
84 - return nil, errLeaseDropped
85 - case brokerStateStopped:
69 + if b.closedErr != nil {
70 + err := b.closedErr
71 b.mu.Unlock()
87 - return nil, errLeaseStopped
72 + return nil, err
73 }
74
75 if len(b.ready) > 0 {
@@ -101,10 +86,6 @@ func (b *leaseBroker) Claim(ctx context.Context) (*reverseSession, error) {
86 }
87 return session, nil
88 }
104 - if b.totalSessions == 0 {
105 - b.mu.Unlock()
106 - return nil, errNoSessions
107 - }
89 b.mu.Unlock()
90
91 select {
@@ -115,34 +96,13 @@ func (b *leaseBroker) Claim(ctx context.Context) (*reverseSession, error) {
96 }
97 }
98
118 -func (b *leaseBroker) Drop() {
119 - b.transition(brokerStateDropped)
120 -}
121 -
122 -func (b *leaseBroker) Reset() {
123 - b.mu.Lock()
124 - defer b.mu.Unlock()
125 - if b.state == brokerStateDropped {
126 - b.state = brokerStateActive
127 - b.signalLocked()
128 - }
129 -}
130 -
131 -func (b *leaseBroker) Stop() {
132 - b.transition(brokerStateStopped)
133 -}
134 -
135 -func (b *leaseBroker) ReadyCount() int {
136 - b.mu.Lock()
137 - defer b.mu.Unlock()
138 - return len(b.ready)
139 -}
140 -
141 -func (b *leaseBroker) transition(state brokerState) {
99 +func (b *leaseBroker) Close() {
100 b.mu.Lock()
101 sessions := b.ready
102 b.ready = nil
145 - b.state = state
103 + if b.closedErr == nil {
104 + b.closedErr = errBrokerClosed
105 + }
106 b.signalLocked()
107 b.mu.Unlock()
108
@@ -151,25 +111,33 @@ func (b *leaseBroker) transition(state brokerState) {
111 }
112 }
113
114 +func (b *leaseBroker) ReadyCount() int {
115 + b.mu.Lock()
116 + defer b.mu.Unlock()
117 + return len(b.ready)
118 +}
119 +
120 func (b *leaseBroker) watchSession(session *reverseSession) {
121 <-session.Done()
122 +
123 + var readyCount int
124 +
125 b.mu.Lock()
157 - defer b.mu.Unlock()
126 for i := range b.ready {
127 if b.ready[i] == session {
128 b.ready = append(b.ready[:i], b.ready[i+1:]...)
129 break
130 }
131 }
164 - b.totalSessions--
132 + readyCount = len(b.ready)
133 log.Info().
134 Str("component", "relay-server").
135 Str("lease_id", b.leaseID).
136 Str("remote_addr", session.RemoteAddr()).
169 - Int("ready", len(b.ready)).
170 - Int("total_sessions", b.totalSessions).
137 + Int("ready", readyCount).
138 Msg("sdk reverse disconnected")
139 b.signalLocked()
140 + b.mu.Unlock()
141 }
142
143 func (b *leaseBroker) signalLocked() {
@@ -182,8 +150,7 @@ func (b *leaseBroker) signalLocked() {
150 type reverseSessionState int
151
152 const (
185 - reverseSessionAdmitted reverseSessionState = iota
186 - reverseSessionIdle
153 + reverseSessionIdle reverseSessionState = iota
154 reverseSessionClaimed
155 reverseSessionClosed
156 )
@@ -203,7 +170,7 @@ func newReverseSession(conn net.Conn, idleInterval time.Duration) *reverseSessio
170 return &reverseSession{
171 conn: conn,
172 idleInterval: idleInterval,
206 - state: reverseSessionAdmitted,
173 + state: reverseSessionIdle,
174 done: make(chan struct{}),
175 }
176 }
@@ -234,11 +201,10 @@ func (s *reverseSession) IsClosed() bool {
201
202 func (s *reverseSession) StartIdle() {
203 s.mu.Lock()
237 - if s.state != reverseSessionAdmitted {
204 + if s.state != reverseSessionIdle || s.keepaliveStop != nil {
205 s.mu.Unlock()
206 return
207 }
241 - s.state = reverseSessionIdle
208 stop := make(chan struct{})
209 done := make(chan struct{})
210 s.keepaliveStop = stop
portal/broker_test.go
+127 -2
@@ -2,6 +2,7 @@ package portal
2
3 import (
4 "context"
5 + "errors"
6 "io"
7 "net"
8 "testing"
@@ -59,7 +60,7 @@ func TestLeaseBrokerClaimActivatesTLSMarker(t *testing.T) {
60 }
61 }
62
62 -func TestLeaseBrokerDropClosesIdleSessions(t *testing.T) {
63 +func TestLeaseBrokerCloseClosesIdleSessions(t *testing.T) {
64 t.Parallel()
65
66 serverConn, clientConn := net.Pipe()
@@ -69,7 +70,7 @@ func TestLeaseBrokerDropClosesIdleSessions(t *testing.T) {
70 t.Fatalf("Offer() error = %v", err)
71 }
72
72 - broker.Drop()
73 + broker.Close()
74
75 buf := make([]byte, 1)
76 _ = clientConn.SetReadDeadline(time.Now().Add(time.Second))
@@ -77,3 +78,127 @@ func TestLeaseBrokerDropClosesIdleSessions(t *testing.T) {
78 t.Fatal("Read() succeeded, want connection close")
79 }
80 }
81 +
82 +func TestLeaseBrokerCloseUnblocksClaim(t *testing.T) {
83 + t.Parallel()
84 +
85 + broker := newLeaseBroker("lease-test", time.Hour, 2)
86 + claimCtx, cancel := context.WithTimeout(context.Background(), time.Second)
87 + defer cancel()
88 +
89 + started := make(chan struct{})
90 + errCh := make(chan error, 1)
91 + go func() {
92 + close(started)
93 + _, err := broker.Claim(claimCtx)
94 + errCh <- err
95 + }()
96 +
97 + <-started
98 + select {
99 + case err := <-errCh:
100 + t.Fatalf("Claim() returned before Close(): %v", err)
101 + case <-time.After(50 * time.Millisecond):
102 + }
103 +
104 + broker.Close()
105 +
106 + select {
107 + case err := <-errCh:
108 + if !errors.Is(err, errBrokerClosed) {
109 + t.Fatalf("Claim() error = %v, want %v", err, errBrokerClosed)
110 + }
111 + case <-time.After(time.Second):
112 + t.Fatal("timed out waiting for closed claim")
113 + }
114 +}
115 +
116 +func TestLeaseBrokerClaimWaitsForLateOffer(t *testing.T) {
117 + t.Parallel()
118 +
119 + broker := newLeaseBroker("lease-test", time.Hour, 2)
120 +
121 + serverConn, clientConn := net.Pipe()
122 + t.Cleanup(func() {
123 + _ = serverConn.Close()
124 + _ = clientConn.Close()
125 + })
126 +
127 + session := newReverseSession(serverConn, time.Hour)
128 + claimCtx, cancel := context.WithTimeout(context.Background(), time.Second)
129 + defer cancel()
130 +
131 + markerCh := make(chan byte, 1)
132 + errCh := make(chan error, 1)
133 + go func() {
134 + var marker [1]byte
135 + if _, err := io.ReadFull(clientConn, marker[:]); err != nil {
136 + errCh <- err
137 + return
138 + }
139 + markerCh <- marker[0]
140 + }()
141 +
142 + type claimResult struct {
143 + session *reverseSession
144 + err error
145 + }
146 + started := make(chan struct{})
147 + resultCh := make(chan claimResult, 1)
148 + go func() {
149 + close(started)
150 + claimed, err := broker.Claim(claimCtx)
151 + resultCh <- claimResult{session: claimed, err: err}
152 + }()
153 +
154 + <-started
155 + select {
156 + case result := <-resultCh:
157 + t.Fatalf("Claim() returned before Offer(): %#v", result)
158 + case <-time.After(50 * time.Millisecond):
159 + }
160 +
161 + if err := broker.Offer(session); err != nil {
162 + t.Fatalf("Offer() error = %v", err)
163 + }
164 +
165 + select {
166 + case result := <-resultCh:
167 + if result.err != nil {
168 + t.Fatalf("Claim() error = %v", result.err)
169 + }
170 + if result.session != session {
171 + t.Fatalf("Claim() returned unexpected session")
172 + }
173 + select {
174 + case err := <-errCh:
175 + t.Fatalf("ReadFull() error = %v", err)
176 + case marker := <-markerCh:
177 + if marker != types.MarkerTLSStart {
178 + t.Fatalf("marker = 0x%02x, want 0x%02x", marker, types.MarkerTLSStart)
179 + }
180 + case <-time.After(time.Second):
181 + t.Fatal("timed out waiting for activation marker")
182 + }
183 + case <-time.After(time.Second):
184 + t.Fatal("timed out waiting for claim")
185 + }
186 +}
187 +
188 +func TestLeaseBrokerClaimTimesOutWithoutSessions(t *testing.T) {
189 + t.Parallel()
190 +
191 + broker := newLeaseBroker("lease-test", time.Hour, 2)
192 +
193 + claimCtx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
194 + defer cancel()
195 +
196 + start := time.Now()
197 + _, err := broker.Claim(claimCtx)
198 + if !errors.Is(err, context.DeadlineExceeded) {
199 + t.Fatalf("Claim() error = %v, want %v", err, context.DeadlineExceeded)
200 + }
201 + if time.Since(start) < 40*time.Millisecond {
202 + t.Fatalf("Claim() returned too early: %v", time.Since(start))
203 + }
204 +}
portal/server.go
+3 -4
@@ -189,7 +189,7 @@ func (s *Server) Shutdown(ctx context.Context) error {
189
190 s.mu.Lock()
191 for _, lease := range s.leases {
192 - lease.Broker.Stop()
192 + lease.Broker.Close()
193 }
194 s.mu.Unlock()
195
@@ -535,7 +535,6 @@ func (s *Server) renewLease(req types.RenewRequest, clientIP string) (types.Rene
535 record.ClientIP = clientIP
536 s.cfg.Policy.IPFilter().RegisterLeaseIP(record.ID, clientIP)
537 }
538 - record.Broker.Reset()
538 return types.RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt}, nil
539 }
540
@@ -555,7 +554,7 @@ func (s *Server) unregisterLease(req types.UnregisterRequest) error {
554
555 s.routes.DeleteLease(record.Hostnames)
556 s.cfg.Policy.ForgetLease(record.ID)
558 - record.Broker.Drop()
557 + record.Broker.Close()
558 return nil
559 }
560
@@ -708,7 +707,7 @@ func (s *Server) cleanupExpiredLeases() {
707 for _, lease := range expired {
708 s.routes.DeleteLease(lease.Hostnames)
709 s.cfg.Policy.ForgetLease(lease.ID)
711 - lease.Broker.Drop()
710 + lease.Broker.Close()
711 }
712 }
713
portal/utils.go renamed
sdk/expose.go
+3 -2
@@ -83,7 +83,8 @@ func Expose(ctx context.Context, relayUrls []string, name string, metadata types
83 logger.Info().
84 Int("relay_count", len(exposure.relays)).
85 Strs("relays", exposure.RelayURLs()).
86 - Msg("exposure starting")
86 + Strs("public_urls", exposure.PublicURLs()).
87 + Msg("exposure ready")
88
89 return exposure, nil
90 }
@@ -155,7 +156,7 @@ func (e *Exposure) PublicURLs() []string {
156 if relay.listener == nil {
157 continue
158 }
158 - for _, rawURL := range relay.listener.publicURLs() {
159 + for _, rawURL := range relay.listener.PublicURLs() {
160 if _, ok := seen[rawURL]; ok {
161 continue
162 }
sdk/listener.go
+285 -416
@@ -4,9 +4,9 @@ import (
4 "context"
5 "crypto/tls"
6 "errors"
7 + "fmt"
8 "io"
9 "net"
9 - "strings"
10 "sync"
11 "time"
12
@@ -16,407 +16,316 @@ import (
16 "github.com/gosuda/portal/v2/types"
17 )
18
19 -const (
20 - defaultListenerRetryDelay = 1 * time.Second
21 -)
22 -
23 -type listenerState uint8
24 -
25 -const (
26 - listenerStatePending listenerState = iota
27 - listenerStateReady
28 - listenerStateClosed
29 -)
30 -
19 type ListenerConfig struct {
32 - Name string
33 - ReverseToken string
34 - Hostnames []string
35 - Metadata types.LeaseMetadata
36 - RootCAPEM []byte
37 - RetryCount int
20 + Name string
21 + ReverseToken string
22 + Hostnames []string
23 + Metadata types.LeaseMetadata
24 + RootCAPEM []byte
25 + DialTimeout time.Duration
26 + RequestTimeout time.Duration
27 + HandshakeTimeout time.Duration
28 + LeaseTTL time.Duration
29 + RenewBefore time.Duration
30 + ReadyTarget int
31 }
32
33 type Listener struct {
41 - ctx context.Context
42 - cancel context.CancelFunc
43 - accepted chan net.Conn
44 - refill chan struct{}
45 -
34 + tlsCloser io.Closer
35 + tlsConfig *tls.Config
36 readyTarget int
37 leaseTTL time.Duration
48 - renewInterval time.Duration
38 + renewBefore time.Duration
39 handshakeTimeout time.Duration
50 - retryCount int
51 - retryDelay time.Duration
52 -
53 - mu sync.Mutex
54 - api *relayClient
55 - leaseID string
56 - hostnames []string
57 - tlsConfig *tls.Config
58 - tlsCloser io.Closer
59 - activeSessions int
60 - sessionFailures int
61 - state listenerState
62 -
63 - closeOnce sync.Once
64 - closeErr error
40 + ctx context.Context
41 + cancel context.CancelFunc
42 + api *relayClient
43 + signal chan struct{}
44 + accepted chan net.Conn
45 + leaseID string
46 + hostnames []string
47 + metadata types.LeaseMetadata
48 +
49 + activeSessions int
50 + closeOnce sync.Once
51 + mu sync.Mutex
52 }
53
54 // NewListener creates one relay listener and its dedicated relay transport for one relay URL.
55 func NewListener(ctx context.Context, relayURL string, cfg ListenerConfig) (*Listener, error) {
69 - api, err := newRelayClient(relayURL, cfg)
70 - if err != nil {
71 - return nil, err
72 - }
56 if ctx == nil {
57 ctx = context.Background()
58 }
59
77 - readyTarget := defaultReadyTarget
78 - leaseTTL := defaultLeaseTTL
79 - handshakeTimeout := defaultHandshakeTimeout
80 - renewBefore := defaultRenewBefore
81 - renewInterval := leaseTTL / 2
82 - if leaseTTL > renewBefore {
83 - renewInterval = leaseTTL - renewBefore
60 + listenerCtx, cancel := context.WithCancel(ctx)
61 + readyTarget := cfg.ReadyTarget
62 + if readyTarget <= 0 {
63 + readyTarget = defaultReadyTarget
64 }
85 - if renewInterval <= 0 {
86 - renewInterval = leaseTTL / 2
65 + leaseTTL := cfg.LeaseTTL
66 + if leaseTTL <= 0 {
67 + leaseTTL = defaultLeaseTTL
68 }
88 - if renewInterval <= 0 {
89 - renewInterval = time.Second
69 + handshakeTimeout := cfg.HandshakeTimeout
70 + if handshakeTimeout <= 0 {
71 + handshakeTimeout = defaultHandshakeTimeout
72 }
91 -
92 - retryCount := cfg.RetryCount
93 - retryDelay := defaultListenerRetryDelay
94 - if retryCount < 0 {
95 - retryCount = 0
73 + renewBefore := cfg.RenewBefore
74 + if renewBefore <= 0 {
75 + renewBefore = defaultRenewBefore
76 }
97 - if retryDelay <= 0 {
98 - retryDelay = defaultListenerRetryDelay
77 +
78 + api, err := newRelayClient(listenerCtx, relayURL, cfg)
79 + if err != nil {
80 + cancel()
81 + return nil, err
82 }
83
101 - listenerCtx, cancel := context.WithCancel(ctx)
84 l := &Listener{
85 ctx: listenerCtx,
86 cancel: cancel,
87 + api: api,
88 + signal: make(chan struct{}, 1),
89 accepted: make(chan net.Conn, max(readyTarget*2, 1)),
106 - refill: make(chan struct{}, 1),
90 readyTarget: readyTarget,
91 leaseTTL: leaseTTL,
109 - renewInterval: renewInterval,
92 + renewBefore: renewBefore,
93 handshakeTimeout: handshakeTimeout,
111 - retryCount: retryCount,
112 - retryDelay: retryDelay,
113 - api: api,
114 - state: listenerStatePending,
94 }
95
117 - go l.run(listenerCtx)
96 + resp, err := api.registerLease(listenerCtx, cfg.Hostnames, leaseTTL)
97 + if err != nil {
98 + api.close()
99 + cancel()
100 + return nil, err
101 + }
102 +
103 + tlsConf, tlsCloser, err := keyless.BuildClientTLSConfig(api.baseURL.String(), resp.Hostnames)
104 + if err != nil {
105 + _ = api.unregisterLease(context.Background(), resp.LeaseID)
106 + api.close()
107 + cancel()
108 + return nil, err
109 + }
110 +
111 + if listenerCtx.Err() != nil {
112 + _ = api.unregisterLease(context.Background(), resp.LeaseID)
113 + _ = tlsCloser.Close()
114 + api.close()
115 + cancel()
116 + return nil, listenerCtx.Err()
117 + }
118 +
119 + l.mu.Lock()
120 + l.leaseID = resp.LeaseID
121 + l.hostnames = append([]string(nil), resp.Hostnames...)
122 + l.metadata = cloneMetadata(resp.Metadata)
123 + l.tlsConfig = tlsConf
124 + l.tlsCloser = tlsCloser
125 + l.mu.Unlock()
126 +
127 + go l.runSupervisor()
128 + go l.runRenewLoop()
129 + l.notify()
130 return l, nil
131 }
132
133 func (l *Listener) Accept() (net.Conn, error) {
134 select {
135 case <-l.ctx.Done():
124 - select {
125 - case conn := <-l.accepted:
126 - if conn != nil {
127 - _ = conn.Close()
128 - }
129 - default:
130 - }
136 return nil, net.ErrClosed
137 case conn := <-l.accepted:
138 if conn == nil {
139 return nil, net.ErrClosed
140 }
136 - select {
137 - case <-l.ctx.Done():
138 - _ = conn.Close()
139 - return nil, net.ErrClosed
140 - default:
141 - return conn, nil
142 - }
141 + return conn, nil
142 }
143 }
144
145 func (l *Listener) Close() error {
146 + var closeErr error
147 l.closeOnce.Do(func() {
148 - l.closeErr = l.shutdown()
148 + if l.cancel != nil {
149 + l.cancel()
150 + }
151 +
152 + l.mu.Lock()
153 + leaseID := l.leaseID
154 + tlsCloser := l.tlsCloser
155 + api := l.api
156 + l.leaseID = ""
157 + l.tlsConfig = nil
158 + l.tlsCloser = nil
159 + l.activeSessions = 0
160 + l.mu.Unlock()
161 +
162 + l.drainAccepted()
163 +
164 + if api != nil && leaseID != "" {
165 + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
166 + closeErr = errors.Join(closeErr, api.unregisterLease(ctx, leaseID))
167 + cancel()
168 + }
169 + if tlsCloser != nil {
170 + closeErr = errors.Join(closeErr, tlsCloser.Close())
171 + }
172 + if api != nil {
173 + api.close()
174 + }
175 })
150 - return l.closeErr
176 + return closeErr
177 }
178
179 func (l *Listener) Addr() net.Addr {
180 l.mu.Lock()
181 defer l.mu.Unlock()
156 -
157 - if strings.TrimSpace(l.leaseID) == "" {
182 + if l.leaseID == "" {
183 return listenerAddr("portal:pending")
184 }
185 return listenerAddr("portal:" + l.leaseID)
186 }
187
163 -func (l *Listener) run(runCtx context.Context) {
164 - logger := log.With().
165 - Str("component", "sdk-listener").
166 - Str("name", l.api.name).
167 - Logger()
168 -
169 - var err error
170 - for attempt := 1; ; attempt++ {
171 - err = l.establish(runCtx)
172 - if err == nil {
173 - l.mu.Lock()
174 - if l.state == listenerStateClosed {
175 - l.mu.Unlock()
176 - return
177 - }
178 - leaseID := l.leaseID
179 - hostnames := append([]string(nil), l.hostnames...)
180 - l.mu.Unlock()
181 -
182 - logger.Info().
183 - Str("lease_id", leaseID).
184 - Strs("hostnames", hostnames).
185 - Msg("listener connected")
186 -
187 - go l.runSessionPool(runCtx)
188 - go l.runRenewLoop(runCtx)
189 - l.signalRefill()
190 - return
191 - }
192 -
193 - if errors.Is(runCtx.Err(), context.Canceled) {
194 - return
195 - }
196 -
197 - logger.Warn().
198 - Err(err).
199 - Int("attempt", attempt).
200 - Dur("retry_in", l.retryDelay).
201 - Msg("listener bootstrap failed")
202 -
203 - if l.retryLimitReached(attempt) {
204 - break
205 - }
206 - if !sleepOrDone(runCtx, l.retryDelay) {
207 - return
208 - }
209 - }
210 -
211 - l.fail(err, "listener bootstrap retry limit reached")
188 +func (l *Listener) LeaseID() string {
189 + l.mu.Lock()
190 + defer l.mu.Unlock()
191 + return l.leaseID
192 }
193
214 -func (l *Listener) establish(runCtx context.Context) error {
194 +func (l *Listener) Hostnames() []string {
195 l.mu.Lock()
216 - if l.state == listenerStateClosed {
217 - l.mu.Unlock()
218 - return context.Canceled
219 - }
220 - api := l.api
221 - l.mu.Unlock()
222 -
223 - resp, err := api.registerLease(runCtx, nil, l.leaseTTL)
224 - if err != nil {
225 - return err
226 - }
196 + defer l.mu.Unlock()
197 + return append([]string(nil), l.hostnames...)
198 +}
199
228 - tlsConfig, tlsCloser, err := keyless.BuildClientTLSConfig(api.baseURL.String(), resp.Hostnames)
229 - if err != nil {
230 - _ = api.unregisterLease(context.Background(), resp.LeaseID)
231 - return err
232 - }
200 +func (l *Listener) Metadata() types.LeaseMetadata {
201 + l.mu.Lock()
202 + defer l.mu.Unlock()
203 + return cloneMetadata(l.metadata)
204 +}
205
206 +func (l *Listener) PublicURLs() []string {
207 l.mu.Lock()
235 - if l.state == listenerStateClosed || runCtx.Err() != nil {
236 - l.mu.Unlock()
237 - _ = api.unregisterLease(context.Background(), resp.LeaseID)
238 - _ = tlsCloser.Close()
239 - return context.Canceled
240 - }
241 - oldCloser := l.tlsCloser
242 - l.leaseID = resp.LeaseID
243 - l.hostnames = append([]string(nil), resp.Hostnames...)
244 - l.tlsConfig = tlsConfig
245 - l.tlsCloser = tlsCloser
246 - l.activeSessions = 0
247 - l.sessionFailures = 0
248 - l.state = listenerStateReady
208 + hostnames := append([]string(nil), l.hostnames...)
209 l.mu.Unlock()
210
251 - if oldCloser != nil {
252 - _ = oldCloser.Close()
211 + urls := make([]string, 0, len(hostnames))
212 + for _, host := range hostnames {
213 + urls = append(urls, "https://"+host)
214 }
254 - return nil
215 + return urls
216 }
217
257 -func (l *Listener) runSessionPool(runCtx context.Context) {
218 +func (l *Listener) runSupervisor() {
219 for {
220 select {
260 - case <-runCtx.Done():
221 + case <-l.ctx.Done():
222 return
262 - case <-l.refill:
223 + case <-l.signal:
224 }
225
265 - for {
266 - l.mu.Lock()
267 - ready := l.state == listenerStateReady &&
268 - l.api != nil &&
269 - strings.TrimSpace(l.leaseID) != "" &&
270 - l.tlsConfig != nil &&
271 - l.activeSessions < l.readyTarget
272 - if !ready {
273 - l.mu.Unlock()
274 - break
275 - }
276 - l.activeSessions++
277 - l.mu.Unlock()
278 -
279 - go l.runSession(runCtx)
226 + for l.reserveSessionSlot() {
227 + go l.runSession()
228 }
229 }
230 }
231
284 -func (l *Listener) runRenewLoop(runCtx context.Context) {
285 - failures := 0
286 - wait := l.renewInterval
287 -
288 - for {
289 - if !sleepOrDone(runCtx, wait) {
290 - return
291 - }
292 -
293 - l.mu.Lock()
294 - api := l.api
295 - leaseID := l.leaseID
296 - ready := l.state == listenerStateReady
297 - l.mu.Unlock()
298 -
299 - if !ready || api == nil || strings.TrimSpace(leaseID) == "" {
300 - wait = l.renewInterval
301 - continue
302 - }
303 -
304 - ctx, cancel := context.WithTimeout(runCtx, 10*time.Second)
305 - err := api.renewLease(ctx, leaseID, l.leaseTTL)
306 - cancel()
232 +func (l *Listener) runRenewLoop() {
233 + interval := l.leaseTTL / 2
234 + if interval <= 0 {
235 + interval = 30 * time.Second
236 + }
237 + if l.renewBefore > 0 && l.leaseTTL > l.renewBefore {
238 + interval = l.leaseTTL - l.renewBefore
239 + }
240 + if interval <= 0 {
241 + interval = 30 * time.Second
242 + }
243
308 - if err == nil {
309 - failures = 0
310 - wait = l.renewInterval
311 - continue
312 - }
244 + ticker := time.NewTicker(interval)
245 + defer ticker.Stop()
246
314 - if isLeaseNotFound(err) {
315 - log.Warn().
316 - Err(err).
317 - Str("component", "sdk-listener").
318 - Str("lease_id", leaseID).
319 - Msg("lease not found on relay, attempting re-registration")
247 + var consecutiveFailures int
248
321 - if reregErr := l.reregister(runCtx); reregErr == nil {
322 - failures = 0
323 - wait = l.renewInterval
249 + for {
250 + select {
251 + case <-l.ctx.Done():
252 + return
253 + case <-ticker.C:
254 + l.mu.Lock()
255 + leaseID := l.leaseID
256 + l.mu.Unlock()
257
325 - l.mu.Lock()
326 - newLeaseID := l.leaseID
327 - newHostnames := append([]string(nil), l.hostnames...)
328 - l.mu.Unlock()
258 + ctx, cancel := context.WithTimeout(l.context(), 10*time.Second)
259 + err := l.api.renewLease(ctx, leaseID, l.leaseTTL)
260 + cancel()
261
330 - log.Info().
262 + if err != nil {
263 + if isLeaseNotFound(err) {
264 + log.Warn().
265 + Str("component", "sdk-listener").
266 + Str("lease_id", leaseID).
267 + Msg("lease not found on relay, attempting re-registration")
268 + if reregErr := l.reregister(); reregErr != nil {
269 + log.Error().Err(reregErr).
270 + Str("component", "sdk-listener").
271 + Msg("lease re-registration failed")
272 + } else {
273 + consecutiveFailures = 0
274 + log.Info().
275 + Str("component", "sdk-listener").
276 + Str("lease_id", l.LeaseID()).
277 + Strs("hostnames", l.Hostnames()).
278 + Msg("lease re-registered successfully")
279 + continue
280 + }
281 + }
282 +
283 + consecutiveFailures++
284 + event := log.Warn()
285 + if consecutiveFailures >= 3 {
286 + event = log.Error()
287 + }
288 + event.Err(err).
289 Str("component", "sdk-listener").
332 - Str("lease_id", newLeaseID).
333 - Strs("hostnames", newHostnames).
334 - Msg("lease re-registered successfully")
335 - continue
290 + Str("lease_id", l.LeaseID()).
291 + Int("consecutive_failures", consecutiveFailures).
292 + Msg("lease renewal failed")
293 } else {
337 - err = reregErr
338 - log.Error().
339 - Err(reregErr).
340 - Str("component", "sdk-listener").
341 - Str("lease_id", leaseID).
342 - Msg("lease re-registration failed")
294 + consecutiveFailures = 0
295 }
296 }
345 -
346 - failures++
347 - event := log.Warn()
348 - if l.retryLimitReached(failures) {
349 - event = log.Error()
350 - }
351 - event.Err(err).
352 - Str("component", "sdk-listener").
353 - Str("lease_id", leaseID).
354 - Int("consecutive_failures", failures).
355 - Msg("lease renewal failed")
356 -
357 - if l.retryLimitReached(failures) {
358 - l.fail(err, "listener renew retry limit reached")
359 - return
360 - }
361 - wait = l.retryDelay
297 }
298 }
299
365 -func (l *Listener) runSession(runCtx context.Context) {
366 - defer func() {
367 - l.mu.Lock()
368 - if l.activeSessions > 0 {
369 - l.activeSessions--
370 - }
371 - l.mu.Unlock()
372 - l.signalRefill()
373 - }()
374 -
375 - fail := func(err error) {
376 - if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
377 - return
378 - }
379 -
380 - l.mu.Lock()
381 - if l.state == listenerStateClosed {
382 - l.mu.Unlock()
383 - return
384 - }
385 - l.sessionFailures++
386 - failures := l.sessionFailures
387 - l.mu.Unlock()
388 -
389 - if l.retryLimitReached(failures) {
390 - l.fail(err, "listener session retry limit reached")
391 - return
392 - }
393 -
394 - _ = sleepOrDone(runCtx, l.retryDelay)
395 - }
300 +func (l *Listener) runSession() {
301 + defer l.releaseSessionSlot()
302
303 + sessionCtx := l.context()
304 l.mu.Lock()
398 - ready := l.state == listenerStateReady
399 - api := l.api
305 leaseID := l.leaseID
401 - tlsConfig := l.tlsConfig
306 l.mu.Unlock()
403 - if !ready || api == nil || strings.TrimSpace(leaseID) == "" || tlsConfig == nil {
307 +
308 + conn, err := l.api.openReverseSession(sessionCtx, leaseID)
309 + if err != nil {
310 + sleepOrDone(sessionCtx, time.Second)
311 return
312 }
313
407 - conn, err := api.openReverseSession(runCtx, leaseID)
314 + claimed, err := l.awaitActivation(conn)
315 if err != nil {
409 - fail(err)
410 - return
316 + _ = conn.Close()
317 + if !claimed && !errors.Is(err, context.Canceled) && !errors.Is(err, net.ErrClosed) {
318 + sleepOrDone(sessionCtx, time.Second)
319 + }
320 }
321 +}
322
323 +func (l *Listener) awaitActivation(conn net.Conn) (bool, error) {
324 var marker [1]byte
325 for {
326 _ = conn.SetReadDeadline(time.Now().Add(2 * l.handshakeTimeout))
327 if _, err := io.ReadFull(conn, marker[:]); err != nil {
417 - _ = conn.Close()
418 - fail(err)
419 - return
328 + return false, err
329 }
330 _ = conn.SetReadDeadline(time.Time{})
331
@@ -424,177 +333,142 @@ func (l *Listener) runSession(runCtx context.Context) {
333 case types.MarkerKeepalive:
334 continue
335 case types.MarkerTLSStart:
427 - tlsConn := tls.Server(conn, tlsConfig)
428 - handshakeCtx, cancel := context.WithTimeout(runCtx, l.handshakeTimeout)
429 - err := tlsConn.HandshakeContext(handshakeCtx)
430 - cancel()
431 - if err != nil {
432 - _ = tlsConn.Close()
433 - fail(err)
434 - return
435 - }
436 -
437 - l.mu.Lock()
438 - if l.state != listenerStateClosed {
439 - l.sessionFailures = 0
440 - }
441 - l.mu.Unlock()
442 -
443 - select {
444 - case <-runCtx.Done():
445 - _ = tlsConn.Close()
446 - case l.accepted <- tlsConn:
336 + if err := l.activate(conn); err != nil {
337 + return true, err
338 }
448 - return
339 + return true, nil
340 default:
450 - _ = conn.Close()
451 - fail(errors.New("unexpected reverse marker"))
452 - return
341 + return false, fmt.Errorf("unexpected reverse marker: 0x%02x", marker[0])
342 }
343 }
344 }
345
457 -func (l *Listener) publicURLs() []string {
346 +func (l *Listener) activate(conn net.Conn) error {
347 l.mu.Lock()
459 - defer l.mu.Unlock()
348 + tlsCfg := l.tlsConfig
349 + l.mu.Unlock()
350
461 - if l.state != listenerStateReady || len(l.hostnames) == 0 {
462 - return nil
351 + tlsConn := tls.Server(conn, tlsCfg)
352 + handshakeCtx, cancel := context.WithTimeout(l.context(), l.handshakeTimeout)
353 + defer cancel()
354 + if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
355 + return err
356 }
357
465 - urls := make([]string, 0, len(l.hostnames))
466 - for _, host := range l.hostnames {
467 - urls = append(urls, "https://"+host)
358 + select {
359 + case <-l.ctx.Done():
360 + _ = tlsConn.Close()
361 + return l.context().Err()
362 + case l.accepted <- tlsConn:
363 + return nil
364 }
469 - return urls
365 }
366
472 -func (l *Listener) reregister(runCtx context.Context) error {
367 +func (l *Listener) reregister() error {
368 + ctx, cancel := context.WithTimeout(l.context(), 10*time.Second)
369 + defer cancel()
370 +
371 l.mu.Lock()
474 - if l.state == listenerStateClosed {
475 - l.mu.Unlock()
476 - return context.Canceled
477 - }
478 - api := l.api
372 hostnames := append([]string(nil), l.hostnames...)
373 l.mu.Unlock()
374
482 - ctx, cancel := context.WithTimeout(runCtx, 10*time.Second)
483 - defer cancel()
484 -
485 - resp, err := api.registerLease(ctx, hostnames, l.leaseTTL)
375 + resp, err := l.api.registerLease(ctx, hostnames, l.leaseTTL)
376 if err != nil {
377 return err
378 }
379
490 - tlsConfig, tlsCloser, err := keyless.BuildClientTLSConfig(api.baseURL.String(), resp.Hostnames)
380 + tlsConf, tlsCloser, err := keyless.BuildClientTLSConfig(l.api.baseURL.String(), resp.Hostnames)
381 if err != nil {
492 - _ = api.unregisterLease(context.Background(), resp.LeaseID)
382 + _ = l.api.unregisterLease(ctx, resp.LeaseID)
383 return err
384 }
385
496 - l.mu.Lock()
497 - if l.state == listenerStateClosed || runCtx.Err() != nil {
498 - l.mu.Unlock()
499 - _ = api.unregisterLease(context.Background(), resp.LeaseID)
386 + if l.isClosed() {
387 + _ = l.api.unregisterLease(context.Background(), resp.LeaseID)
388 _ = tlsCloser.Close()
389 return context.Canceled
390 }
391 +
392 + l.mu.Lock()
393 oldCloser := l.tlsCloser
394 l.leaseID = resp.LeaseID
395 l.hostnames = append([]string(nil), resp.Hostnames...)
506 - l.tlsConfig = tlsConfig
396 + l.metadata = cloneMetadata(resp.Metadata)
397 + l.tlsConfig = tlsConf
398 l.tlsCloser = tlsCloser
508 - l.sessionFailures = 0
509 - l.state = listenerStateReady
399 l.mu.Unlock()
400
401 if oldCloser != nil {
402 _ = oldCloser.Close()
403 }
515 - l.signalRefill()
404 +
405 + l.notify()
406 return nil
407 }
408
519 -func (l *Listener) shutdown() error {
520 - l.mu.Lock()
521 - if l.state == listenerStateClosed {
522 - l.mu.Unlock()
523 - return nil
524 - }
525 - l.state = listenerStateClosed
526 - cancel := l.cancel
527 - api := l.api
528 - leaseID := l.leaseID
529 - tlsCloser := l.tlsCloser
530 - l.leaseID = ""
531 - l.tlsConfig = nil
532 - l.tlsCloser = nil
533 - l.activeSessions = 0
534 - l.sessionFailures = 0
535 - l.mu.Unlock()
536 -
537 - if cancel != nil {
538 - cancel()
539 - }
540 - l.drainAccepted()
409 +func isLeaseNotFound(err error) bool {
410 + return errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeLeaseNotFound})
411 +}
412
542 - var closeErr error
543 - if api != nil && strings.TrimSpace(leaseID) != "" {
544 - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
545 - closeErr = errors.Join(closeErr, api.unregisterLease(ctx, leaseID))
546 - cancel()
547 - }
548 - if tlsCloser != nil {
549 - closeErr = errors.Join(closeErr, tlsCloser.Close())
413 +func (l *Listener) reserveSessionSlot() bool {
414 + l.mu.Lock()
415 + defer l.mu.Unlock()
416 + if l.isClosed() {
417 + return false
418 }
551 - if api != nil {
552 - api.close()
419 + if l.activeSessions >= l.readyTarget {
420 + return false
421 }
554 - return closeErr
422 + l.activeSessions++
423 + return true
424 }
425
557 -func (l *Listener) fail(err error, message string) {
558 - closed := false
559 - l.closeOnce.Do(func() {
560 - closed = true
561 - l.closeErr = errors.Join(err, l.shutdown())
562 - })
563 - if !closed {
564 - return
426 +func (l *Listener) releaseSessionSlot() {
427 + l.mu.Lock()
428 + if l.activeSessions > 0 {
429 + l.activeSessions--
430 }
566 -
567 - log.Error().
568 - Str("component", "sdk-listener").
569 - Str("name", l.api.name).
570 - Err(l.closeErr).
571 - Msg(message)
431 + l.mu.Unlock()
432 + l.notify()
433 }
434
574 -func (l *Listener) signalRefill() {
435 +func (l *Listener) notify() {
436 select {
576 - case l.refill <- struct{}{}:
437 + case l.signal <- struct{}{}:
438 default:
439 }
440 }
441
581 -func (l *Listener) retryLimitReached(failures int) bool {
582 - return l.retryCount > 0 && failures >= l.retryCount
583 -}
584 -
585 -func isLeaseNotFound(err error) bool {
586 - return errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeLeaseNotFound})
587 -}
588 -
589 -func sleepOrDone(ctx context.Context, d time.Duration) bool {
442 +func sleepOrDone(ctx context.Context, d time.Duration) {
443 timer := time.NewTimer(d)
444 defer timer.Stop()
592 -
445 select {
446 case <-ctx.Done():
595 - return false
447 case <-timer.C:
448 + }
449 +}
450 +
451 +type listenerAddr string
452 +
453 +func (a listenerAddr) Network() string { return "portal" }
454 +func (a listenerAddr) String() string { return string(a) }
455 +
456 +func (l *Listener) context() context.Context {
457 + if l.ctx != nil {
458 + return l.ctx
459 + }
460 + return context.Background()
461 +}
462 +
463 +func (l *Listener) isClosed() bool {
464 + if l.ctx == nil {
465 + return false
466 + }
467 + select {
468 + case <-l.ctx.Done():
469 return true
470 + default:
471 + return false
472 }
473 }
474
@@ -610,8 +484,3 @@ func (l *Listener) drainAccepted() {
484 }
485 }
486 }
613 -
614 -type listenerAddr string
615 -
616 -func (a listenerAddr) Network() string { return "portal" }
617 -func (a listenerAddr) String() string { return string(a) }
sdk/relayclient.go
+33 -24
@@ -33,17 +33,17 @@ const (
33 )
34
35 type relayClient struct {
36 - baseURL *url.URL
37 - httpClient *http.Client
38 - rawTLSConfig *tls.Config
39 - dialTimeout time.Duration
40 - name string
41 - requestedHostnames []string
42 - reverseToken string
43 - metadata types.LeaseMetadata
36 + baseURL *url.URL
37 + httpClient *http.Client
38 + rawTLSConfig *tls.Config
39 + dialTimeout time.Duration
40 + name string
41 + hostnames []string
42 + reverseToken string
43 + metadata types.LeaseMetadata
44 }
45
46 -func newRelayClient(relayURL string, cfg ListenerConfig) (*relayClient, error) {
46 +func newRelayClient(ctx context.Context, relayURL string, cfg ListenerConfig) (*relayClient, error) {
47 name := strings.TrimSpace(cfg.Name)
48 if name == "" {
49 return nil, errors.New("listener name is required")
@@ -70,7 +70,11 @@ func newRelayClient(relayURL string, cfg ListenerConfig) (*relayClient, error) {
70
71 rootCAPEM := append([]byte(nil), cfg.RootCAPEM...)
72 if len(rootCAPEM) == 0 && isLocalRelayHost(baseURL.Hostname()) {
73 - bootstrapCtx, cancel := context.WithTimeout(context.Background(), defaultDialTimeout+defaultHandshakeTimeout)
73 + bootstrapParent := ctx
74 + if bootstrapParent == nil {
75 + bootstrapParent = context.Background()
76 + }
77 + bootstrapCtx, cancel := context.WithTimeout(bootstrapParent, defaultDialTimeout+defaultHandshakeTimeout)
78 defer cancel()
79
80 _, resolvedCAPEM, bootstrapErr := keyless.ResolveMaterials(bootstrapCtx, baseURL.String(), baseURL.Hostname())
@@ -85,6 +89,15 @@ func newRelayClient(relayURL string, cfg ListenerConfig) (*relayClient, error) {
89 return nil, err
90 }
91
92 + dialTimeout := cfg.DialTimeout
93 + if dialTimeout <= 0 {
94 + dialTimeout = defaultDialTimeout
95 + }
96 + requestTimeout := cfg.RequestTimeout
97 + if requestTimeout <= 0 {
98 + requestTimeout = defaultRequestTimeout
99 + }
100 +
101 baseTLS := &tls.Config{
102 MinVersion: tls.VersionTLS12,
103 ServerName: baseURL.Hostname(),
@@ -98,24 +111,20 @@ func newRelayClient(relayURL string, cfg ListenerConfig) (*relayClient, error) {
111 }
112
113 api := &relayClient{
101 - baseURL: baseURL,
102 - httpClient: &http.Client{Transport: transport, Timeout: defaultRequestTimeout},
103 - rawTLSConfig: baseTLS,
104 - dialTimeout: defaultDialTimeout,
105 - name: name,
106 - requestedHostnames: append([]string(nil), cfg.Hostnames...),
107 - reverseToken: reverseToken,
108 - metadata: cloneMetadata(cfg.Metadata),
114 + baseURL: baseURL,
115 + httpClient: &http.Client{Transport: transport, Timeout: requestTimeout},
116 + rawTLSConfig: baseTLS,
117 + dialTimeout: dialTimeout,
118 + name: name,
119 + hostnames: append([]string(nil), cfg.Hostnames...),
120 + reverseToken: reverseToken,
121 + metadata: cloneMetadata(cfg.Metadata),
122 }
123
111 - checkCtx, cancel := context.WithTimeout(context.Background(), defaultRequestTimeout)
112 - defer cancel()
113 -
114 - if err := api.ensureCompatible(checkCtx); err != nil {
124 + if err := api.ensureCompatible(ctx); err != nil {
125 api.close()
126 return nil, err
127 }
118 -
128 return api, nil
129 }
130
@@ -130,7 +139,7 @@ func (a *relayClient) close() {
139
140 func (a *relayClient) registerLease(ctx context.Context, hostnames []string, ttl time.Duration) (types.RegisterResponse, error) {
141 if len(hostnames) == 0 {
133 - hostnames = a.requestedHostnames
142 + hostnames = a.hostnames
143 }
144
145 var resp types.RegisterResponse
sdk/sdk_test.go
+183 -65
@@ -3,11 +3,9 @@ package sdk
3 import (
4 "context"
5 "encoding/json"
6 - "net"
6 "net/http"
7 "net/http/httptest"
8 "reflect"
10 - "strings"
9 "sync/atomic"
10 "testing"
11 "time"
@@ -15,34 +13,108 @@ import (
13 "github.com/gosuda/portal/v2/types"
14 )
15
18 -func TestNewRelayClientAcceptsMatchingVersion(t *testing.T) {
16 +func TestNewListenerFailsFastOnRegisterError(t *testing.T) {
17 t.Parallel()
18
21 - server := newDomainServer(t, types.SDKProtocolVersion)
19 + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
20 + switch r.URL.Path {
21 + case types.PathSDKDomain:
22 + writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.DomainResponse]{
23 + OK: true,
24 + Data: types.DomainResponse{
25 + RootHost: "localhost",
26 + Version: types.SDKProtocolVersion,
27 + },
28 + })
29 + case types.PathSDKRegister:
30 + writeSDKTestEnvelope(w, http.StatusConflict, types.APIEnvelope[any]{
31 + OK: false,
32 + Error: &types.APIError{Code: types.APIErrorCodeHostnameConflict, Message: "hostname already registered"},
33 + })
34 + default:
35 + http.NotFound(w, r)
36 + return
37 + }
38 + }))
39 defer server.Close()
40
24 - api, err := newRelayClient(server.URL, ListenerConfig{Name: "demo"})
25 - if err != nil {
26 - t.Fatalf("newRelayClient() error = %v", err)
41 + listener, err := NewListener(context.Background(), server.URL, ListenerConfig{Name: "demo"})
42 + if err == nil {
43 + t.Fatal("NewListener() error = nil, want register failure")
44 + }
45 + if listener != nil {
46 + t.Fatalf("NewListener() listener = %#v, want nil", listener)
47 }
28 - defer api.close()
48 }
49
31 -func TestNewListenerRejectsVersionMismatch(t *testing.T) {
50 +func TestNewListenerRegistersLeaseWithMainContract(t *testing.T) {
51 t.Parallel()
52
34 - server := newDomainServer(t, "999")
53 + var registerReq types.RegisterRequest
54 + server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
55 + switch r.URL.Path {
56 + case types.PathSDKDomain:
57 + writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.DomainResponse]{
58 + OK: true,
59 + Data: types.DomainResponse{
60 + RootHost: "localhost",
61 + Version: types.SDKProtocolVersion,
62 + },
63 + })
64 + case types.PathSDKRegister:
65 + if err := json.NewDecoder(r.Body).Decode(&registerReq); err != nil {
66 + t.Fatalf("decode register request: %v", err)
67 + }
68 + writeSDKTestEnvelope(w, http.StatusCreated, types.APIEnvelope[types.RegisterResponse]{
69 + OK: true,
70 + Data: types.RegisterResponse{
71 + LeaseID: "lease-1",
72 + Hostnames: []string{"127.0.0.1"},
73 + Metadata: registerReq.Metadata,
74 + },
75 + })
76 + case types.PathSDKConnect:
77 + writeSDKTestEnvelope(w, http.StatusForbidden, types.APIEnvelope[any]{
78 + OK: false,
79 + Error: &types.APIError{Code: types.APIErrorCodeUnauthorized, Message: "not used in test"},
80 + })
81 + case types.PathSDKRenew:
82 + writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.RenewResponse]{
83 + OK: true,
84 + Data: types.RenewResponse{LeaseID: "lease-1"},
85 + })
86 + case types.PathSDKUnregister:
87 + writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[any]{OK: true})
88 + default:
89 + http.NotFound(w, r)
90 + }
91 + }))
92 defer server.Close()
93
37 - listener, err := NewListener(context.Background(), server.URL, ListenerConfig{Name: "demo"})
38 - if err == nil {
39 - t.Fatal("NewListener() error = nil, want version mismatch")
94 + listener, err := NewListener(context.Background(), server.URL, ListenerConfig{
95 + Name: "demo",
96 + Metadata: types.LeaseMetadata{Owner: "alice"},
97 + LeaseTTL: 42 * time.Second,
98 + })
99 + if err != nil {
100 + t.Fatalf("NewListener() error = %v", err)
101 }
41 - if listener != nil {
42 - t.Fatalf("NewListener() listener = %#v, want nil", listener)
102 + defer listener.Close()
103 +
104 + if registerReq.TTL != 42 {
105 + t.Fatalf("register request TTL = %d, want 42", registerReq.TTL)
106 + }
107 + if listener.LeaseID() != "lease-1" {
108 + t.Fatalf("LeaseID() = %q, want %q", listener.LeaseID(), "lease-1")
109 + }
110 + if got := listener.Hostnames(); !reflect.DeepEqual(got, []string{"127.0.0.1"}) {
111 + t.Fatalf("Hostnames() = %v, want %v", got, []string{"127.0.0.1"})
112 + }
113 + if got := listener.PublicURLs(); !reflect.DeepEqual(got, []string{"https://127.0.0.1"}) {
114 + t.Fatalf("PublicURLs() = %v, want %v", got, []string{"https://127.0.0.1"})
115 }
44 - if !strings.Contains(err.Error(), "version mismatch") {
45 - t.Fatalf("NewListener() error = %v, want version mismatch", err)
116 + if got := listener.Metadata(); got.Owner != "alice" {
117 + t.Fatalf("Metadata().Owner = %q, want %q", got.Owner, "alice")
118 }
119 }
120
@@ -56,9 +128,8 @@ func TestNewListenerReregistersOnLeaseNotFound(t *testing.T) {
128 writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.DomainResponse]{
129 OK: true,
130 Data: types.DomainResponse{
59 - RootHost: "relay.example.com",
60 - SuggestedHostname: "demo.relay.example.com",
61 - Version: types.SDKProtocolVersion,
131 + RootHost: "localhost",
132 + Version: types.SDKProtocolVersion,
133 },
134 })
135 case types.PathSDKRegister:
@@ -67,10 +138,11 @@ func TestNewListenerReregistersOnLeaseNotFound(t *testing.T) {
138 if count > 1 {
139 leaseID = "lease-2"
140 }
70 - writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.RegisterResponse]{
141 + writeSDKTestEnvelope(w, http.StatusCreated, types.APIEnvelope[types.RegisterResponse]{
142 OK: true,
143 Data: types.RegisterResponse{
73 - LeaseID: leaseID,
144 + LeaseID: leaseID,
145 + Hostnames: []string{"127.0.0.1"},
146 },
147 })
148 case types.PathSDKRenew:
@@ -89,6 +161,11 @@ func TestNewListenerReregistersOnLeaseNotFound(t *testing.T) {
161 OK: true,
162 Data: types.RenewResponse{LeaseID: req.LeaseID},
163 })
164 + case types.PathSDKConnect:
165 + writeSDKTestEnvelope(w, http.StatusForbidden, types.APIEnvelope[any]{
166 + OK: false,
167 + Error: &types.APIError{Code: types.APIErrorCodeUnauthorized, Message: "not used in test"},
168 + })
169 case types.PathSDKUnregister:
170 writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[any]{OK: true})
171 default:
@@ -97,33 +174,94 @@ func TestNewListenerReregistersOnLeaseNotFound(t *testing.T) {
174 }))
175 defer server.Close()
176
100 - api, err := newRelayClient(server.URL, ListenerConfig{Name: "demo"})
177 + listener, err := NewListener(context.Background(), server.URL, ListenerConfig{
178 + Name: "demo",
179 + LeaseTTL: 80 * time.Millisecond,
180 + RenewBefore: 40 * time.Millisecond,
181 + })
182 if err != nil {
102 - t.Fatalf("newRelayClient() error = %v", err)
103 - }
104 -
105 - listenerCtx, cancel := context.WithCancel(context.Background())
106 - listener := &Listener{
107 - ctx: listenerCtx,
108 - cancel: cancel,
109 - accepted: make(chan net.Conn, 1),
110 - refill: make(chan struct{}, 1),
111 - readyTarget: 0,
112 - leaseTTL: defaultLeaseTTL,
113 - renewInterval: 10 * time.Millisecond,
114 - handshakeTimeout: defaultHandshakeTimeout,
115 - retryCount: 0,
116 - retryDelay: 10 * time.Millisecond,
117 - api: api,
118 - state: listenerStatePending,
119 - }
120 - go listener.run(listenerCtx)
183 + t.Fatalf("NewListener() error = %v", err)
184 + }
185 defer listener.Close()
186
187 waitForSDKTest(t, func() bool {
124 - listener.mu.Lock()
125 - defer listener.mu.Unlock()
126 - return listener.leaseID == "lease-2"
188 + return listener.LeaseID() == "lease-2"
189 + })
190 +}
191 +
192 +func TestExposeFailsFastWhenAnyRelayCannotRegister(t *testing.T) {
193 + t.Parallel()
194 +
195 + var unregisterCount atomic.Int32
196 + goodServer := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
197 + switch r.URL.Path {
198 + case types.PathSDKDomain:
199 + writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.DomainResponse]{
200 + OK: true,
201 + Data: types.DomainResponse{
202 + RootHost: "localhost",
203 + Version: types.SDKProtocolVersion,
204 + },
205 + })
206 + case types.PathSDKRegister:
207 + writeSDKTestEnvelope(w, http.StatusCreated, types.APIEnvelope[types.RegisterResponse]{
208 + OK: true,
209 + Data: types.RegisterResponse{
210 + LeaseID: "lease-good",
211 + Hostnames: []string{"127.0.0.1"},
212 + },
213 + })
214 + case types.PathSDKConnect:
215 + writeSDKTestEnvelope(w, http.StatusForbidden, types.APIEnvelope[any]{
216 + OK: false,
217 + Error: &types.APIError{Code: types.APIErrorCodeUnauthorized, Message: "not used in test"},
218 + })
219 + case types.PathSDKRenew:
220 + writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.RenewResponse]{
221 + OK: true,
222 + Data: types.RenewResponse{LeaseID: "lease-good"},
223 + })
224 + case types.PathSDKUnregister:
225 + unregisterCount.Add(1)
226 + writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[any]{OK: true})
227 + default:
228 + http.NotFound(w, r)
229 + }
230 + }))
231 + defer goodServer.Close()
232 +
233 + badServer := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
234 + switch r.URL.Path {
235 + case types.PathSDKDomain:
236 + writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.DomainResponse]{
237 + OK: true,
238 + Data: types.DomainResponse{
239 + RootHost: "localhost",
240 + Version: types.SDKProtocolVersion,
241 + },
242 + })
243 + case types.PathSDKRegister:
244 + writeSDKTestEnvelope(w, http.StatusConflict, types.APIEnvelope[any]{
245 + OK: false,
246 + Error: &types.APIError{Code: types.APIErrorCodeHostnameConflict, Message: "hostname already registered"},
247 + })
248 + default:
249 + http.NotFound(w, r)
250 + return
251 + }
252 + }))
253 + defer badServer.Close()
254 +
255 + exposure, err := Expose(context.Background(), []string{goodServer.URL, badServer.URL}, "demo", types.LeaseMetadata{})
256 + if err == nil {
257 + t.Fatal("Expose() error = nil, want register failure")
258 + }
259 + if exposure != nil {
260 + t.Fatalf("Expose() exposure = %#v, want nil", exposure)
261 + }
262 +
263 + waitForSDKTest(t, func() bool {
264 + return unregisterCount.Load() > 0
265 })
266 }
267
@@ -159,26 +297,6 @@ func TestNormalizeRelayURLs(t *testing.T) {
297 }
298 }
299
162 -func newDomainServer(t *testing.T, version string) *httptest.Server {
163 - t.Helper()
164 -
165 - return httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
166 - if r.URL.Path != types.PathSDKDomain {
167 - http.NotFound(w, r)
168 - return
169 - }
170 -
171 - writeSDKTestEnvelope(w, http.StatusOK, types.APIEnvelope[types.DomainResponse]{
172 - OK: true,
173 - Data: types.DomainResponse{
174 - RootHost: "relay.example.com",
175 - SuggestedHostname: "demo.relay.example.com",
176 - Version: version,
177 - },
178 - })
179 - }))
180 -}
181 -
300 func writeSDKTestEnvelope[T any](w http.ResponseWriter, status int, envelope types.APIEnvelope[T]) {
301 w.Header().Set("Content-Type", "application/json")
302 w.WriteHeader(status)